文件最后提交记录最后更新时间
4 个月前
4 个月前
2 个月前
4 个月前
4 个月前
4 个月前
3 个月前
README

BroadcastGradientArgs

产品支持情况

产品 是否支持
Ascend 950PR/Ascend 950DT
Atlas A3 训练系列产品/Atlas A3 推理系列产品 ×
Atlas A2 训练系列产品/Atlas A2 推理系列产品 ×
Atlas 200I/500 A2 推理产品 ×
Atlas 推理系列产品 ×
Atlas 训练系列产品 ×

功能说明

  • 算子功能:

    BroadcastGradientArgs是TensorFlow中用于计算梯度传播所需的广播维度索引的算子。

    它的核心功能是:在反向传播过程中,根据两个张量在正向传播时的原始形状,自动识别出它们因广播机制而扩展的维度,并输出需要在哪些维度上对梯度进行约简,以便将梯度从广播后的形状还原为每个原始张量的形状。

  • 示例:

    例1:
      原始张量a的shape为[2, 1, 4, 1, 6]
      原始张量b的shape为[2, 3, 1, 5, 1]
      x1_data: [2, 1, 4, 1, 6]   # x1_shape=[5]
      x2_data: [2, 3, 1, 5, 1]   # x2_shape=[5]
      
      y1_data: [1, 3]            # y1_shape=[2]
      y2_data: [2, 4]            # y2_shape=[2]
    
    例2:
      原始张量a的shape为[4, 1, 6]
      原始张量b的shape为[2, 3, 1, 5, 1]
      x1_data: [4, 1, 6]         # x1_shape=[3]
      x2_data: [2, 3, 1, 5, 1]   # x2_shape=[5]
      
      y1_data: [0, 1, 3]         # y1_shape=[3]
      y2_data: [2, 4]            # y2_shape=[2]
    
    例3:
      原始张量a的shape为[2, 1, 4, 1, 6]
      原始张量b的shape为[2, 1, 4, 1, 6]
      x1_data: [2, 1, 4, 1, 6]   # x1_shape=[5]
      x2_data: [2, 1, 4, 1, 6]   # x2_shape=[5]
      
      y1_data: []            # y1_shape=[0]
      y2_data: []            # y2_shape=[0]
    

参数说明

参数名 输入/输出/属性 描述 数据类型 数据格式
x1 输入 shape必须是1维,data是示例中原始张量a的shape。 INT32、INT64 ND
x2 输入 shape必须是1维,data是示例中原始张量b的shape,数据类型与x1的数据类型保持一致。 INT32、INT64 ND
y1 输出 shape必须是1维,表示x1对应的张量shape中需要广播的索引,数据类型与x1的数据类型保持一致。 INT32、INT64 ND
y2 输出 shape必须是1维,表示x2对应的张量shape中需要广播的索引,数据类型与x1的数据类型保持一致。 INT32、INT64 ND

约束说明

调用说明

调用方式 调用样例 说明
图模式调用 test_geir_broadcast_gradient_args 通过算子IR构图方式调用BroadcastGradientArgs算子。