已合并
fix(threshold_grad_v2_d): 对齐接口及比较语义 #9150
babeiee创建于 8月25日
fix(threshold_grad_v2_d): 对齐接口及比较语义 #9150
已合并
共 3 个文件变更+17-20
| @@ -12,8 +12,8 @@ | |||
| 12 | 12 | ||
| 13 | /*! | 13 | /*! |
| 14 | * \file threshold_grad_v2_d_def.cpp | 14 | * \file threshold_grad_v2_d_def.cpp |
| 15 | - * \brief ThresholdGradV2D 算子定义:gradOutput/self -> out,fp16/fp32/bf16/int32/int8/uint8 | 15 | + * \brief ThresholdGradV2D 算子定义:gradients/features -> backprops,fp16/fp32/bf16/int32/int8/uint8 |
| 16 | - * 属性 threshold(Float, 默认 1.0)。out = self>threshold ? gradOutput : 0 | 16 | + * 属性 threshold(Float)。backprops = features>threshold ? gradients : 0 |
| 17 | */ | 17 | */ |
| 18 | 18 | ||
| 19 | 19 | ||
| @@ -22,21 +22,21 @@ class ThresholdGradV2D : public OpDef { | |||
| 22 | public: | 22 | public: |
| 23 | explicit ThresholdGradV2D(const char* name) : OpDef(name) | 23 | explicit ThresholdGradV2D(const char* name) : OpDef(name) |
| 24 | { | 24 | { |
| 25 | - this->Input("gradOutput") | 25 | + this->Input("gradients") |
| 26 | .ParamType(REQUIRED) | 26 | .ParamType(REQUIRED) |
| 27 | .DataType({ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_BF16, ge::DT_INT32, ge::DT_INT8, ge::DT_UINT8}) | 27 | .DataType({ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_BF16, ge::DT_INT32, ge::DT_INT8, ge::DT_UINT8}) |
| 28 | .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND}) | 28 | .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND}) |
| 29 | .UnknownShapeFormat( | 29 | .UnknownShapeFormat( |
| 30 | {ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND}) | 30 | {ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND}) |
| 31 | .AutoContiguous(); | 31 | .AutoContiguous(); |
| 32 | - this->Input("self") | 32 | + this->Input("features") |
| 33 | .ParamType(REQUIRED) | 33 | .ParamType(REQUIRED) |
| 34 | .DataType({ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_BF16, ge::DT_INT32, ge::DT_INT8, ge::DT_UINT8}) | 34 | .DataType({ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_BF16, ge::DT_INT32, ge::DT_INT8, ge::DT_UINT8}) |
| 35 | .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND}) | 35 | .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND}) |
| 36 | .UnknownShapeFormat( | 36 | .UnknownShapeFormat( |
| 37 | {ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND}) | 37 | {ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND}) |
| 38 | .AutoContiguous(); | 38 | .AutoContiguous(); |
| 39 | - this->Output("out") | 39 | + this->Output("backprops") |
| 40 | .ParamType(REQUIRED) | 40 | .ParamType(REQUIRED) |
| 41 | .DataType({ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_BF16, ge::DT_INT32, ge::DT_INT8, ge::DT_UINT8}) | 41 | .DataType({ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_BF16, ge::DT_INT32, ge::DT_INT8, ge::DT_UINT8}) |
| 42 | .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND}) | 42 | .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND}) |
| @@ -44,7 +44,7 @@ public: | |||
| 44 | {ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND}) | 44 | {ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND}) |
| 45 | .AutoContiguous(); | 45 | .AutoContiguous(); |
| 46 | 46 | ||
| 47 | - this->Attr("threshold").AttrType(OPTIONAL).Float(1.0); | 47 | + this->Attr("threshold").AttrType(REQUIRED).Float(); |
| 48 | 48 | ||
| 49 | // 目标芯片为 Ascend950PR/DT(arch35)。6 dtype(含 int8/uint8)依赖 arch35 矢量 ISA,arch22 不支持。 | 49 | // 目标芯片为 Ascend950PR/DT(arch35)。6 dtype(含 int8/uint8)依赖 arch35 矢量 ISA,arch22 不支持。 |
| 50 | OpAICoreConfig aiCoreConfig; | 50 | OpAICoreConfig aiCoreConfig; |
| @@ -22,7 +22,7 @@ | |||
| 22 | namespace ThresholdGradV2DOp { | 22 | namespace ThresholdGradV2DOp { |
| 23 | using namespace Ops::Base; | 23 | using namespace Ops::Base; |
| 24 | 24 | ||
| 25 | -constexpr int COMPARE_MODE_LE = 3; | 25 | +constexpr int COMPARE_MODE_GT = 1; |
| 26 | constexpr int SELECT_MODE_TENSOR = 2; | 26 | constexpr int SELECT_MODE_TENSOR = 2; |
| 27 | 27 | ||
| 28 | template <typename U> | 28 | template <typename U> |
| @@ -36,8 +36,8 @@ struct ThresholdGradV2D8BDag { | |||
| 36 | using OpCopyInSelf = Bind<Vec::CopyInBrc<U>, Placeholder::In1<U>>; | 36 | using OpCopyInSelf = Bind<Vec::CopyInBrc<U>, Placeholder::In1<U>>; |
| 37 | using OpCopyInSelfHalf = Bind<Vec::Cast<half, U, 0>, OpCopyInSelf>; | 37 | using OpCopyInSelfHalf = Bind<Vec::Cast<half, U, 0>, OpCopyInSelf>; |
| 38 | using OpCopyInSelfCast = Bind<Vec::Cast<float, half, 0>, OpCopyInSelfHalf>; | 38 | using OpCopyInSelfCast = Bind<Vec::Cast<float, half, 0>, OpCopyInSelfHalf>; |
| 39 | - using Compare = Bind<Vec::Compare<uint8_t, float, COMPARE_MODE_LE>, OpCopyInSelfCast, data_threshold>; | 39 | + using Compare = Bind<Vec::Compare<uint8_t, float, COMPARE_MODE_GT>, OpCopyInSelfCast, data_threshold>; |
| 40 | - using Select = Bind<Vec::Select<uint8_t, float, SELECT_MODE_TENSOR>, Compare, data_zero, OpCopyInGradCast>; | 40 | + using Select = Bind<Vec::Select<uint8_t, float, SELECT_MODE_TENSOR>, Compare, OpCopyInGradCast, data_zero>; |
| 41 | using SelectHalf = Bind<Vec::Cast<half, float, 1>, Select>; | 41 | using SelectHalf = Bind<Vec::Cast<half, float, 1>, Select>; |
| 42 | using SelectCast = Bind<Vec::Cast<U, half, 1>, SelectHalf>; | 42 | using SelectCast = Bind<Vec::Cast<U, half, 1>, SelectHalf>; |
| 43 | using OpCopyOut = Bind<Vec::CopyOut<U>, Placeholder::Out0<U>, SelectCast>; | 43 | using OpCopyOut = Bind<Vec::CopyOut<U>, Placeholder::Out0<U>, SelectCast>; |
| @@ -56,8 +56,8 @@ struct ThresholdGradV2DInt32Dag { | |||
| 56 | using OpCopyInGradCast = Bind<Vec::Cast<float, U, 1>, OpCopyInGrad>; | 56 | using OpCopyInGradCast = Bind<Vec::Cast<float, U, 1>, OpCopyInGrad>; |
| 57 | using OpCopyInSelf = Bind<Vec::CopyInBrc<U>, Placeholder::In1<U>>; | 57 | using OpCopyInSelf = Bind<Vec::CopyInBrc<U>, Placeholder::In1<U>>; |
| 58 | using OpCopyInSelfCast = Bind<Vec::Cast<float, U, 1>, OpCopyInSelf>; | 58 | using OpCopyInSelfCast = Bind<Vec::Cast<float, U, 1>, OpCopyInSelf>; |
| 59 | - using Compare = Bind<Vec::Compare<uint8_t, float, COMPARE_MODE_LE>, OpCopyInSelfCast, data_threshold>; | 59 | + using Compare = Bind<Vec::Compare<uint8_t, float, COMPARE_MODE_GT>, OpCopyInSelfCast, data_threshold>; |
| 60 | - using Select = Bind<Vec::Select<uint8_t, float, SELECT_MODE_TENSOR>, Compare, data_zero, OpCopyInGradCast>; | 60 | + using Select = Bind<Vec::Select<uint8_t, float, SELECT_MODE_TENSOR>, Compare, OpCopyInGradCast, data_zero>; |
| 61 | using SelectCast = Bind<Vec::Cast<U, float, 1>, Select>; | 61 | using SelectCast = Bind<Vec::Cast<U, float, 1>, Select>; |
| 62 | using OpCopyOut = Bind<Vec::CopyOut<U>, Placeholder::Out0<U>, SelectCast>; | 62 | using OpCopyOut = Bind<Vec::CopyOut<U>, Placeholder::Out0<U>, SelectCast>; |
| 63 | // 指定输出节点 | 63 | // 指定输出节点 |
| @@ -75,8 +75,8 @@ struct ThresholdGradV2DDag { | |||
| 75 | using OpCopyInGradCast = Bind<Vec::Cast<float, U, 0>, OpCopyInGrad>; | 75 | using OpCopyInGradCast = Bind<Vec::Cast<float, U, 0>, OpCopyInGrad>; |
| 76 | using OpCopyInSelf = Bind<Vec::CopyInBrc<U>, Placeholder::In1<U>>; | 76 | using OpCopyInSelf = Bind<Vec::CopyInBrc<U>, Placeholder::In1<U>>; |
| 77 | using OpCopyInSelfCast = Bind<Vec::Cast<float, U, 0>, OpCopyInSelf>; | 77 | using OpCopyInSelfCast = Bind<Vec::Cast<float, U, 0>, OpCopyInSelf>; |
| 78 | - using Compare = Bind<Vec::Compare<uint8_t, float, COMPARE_MODE_LE>, OpCopyInSelfCast, data_threshold>; | 78 | + using Compare = Bind<Vec::Compare<uint8_t, float, COMPARE_MODE_GT>, OpCopyInSelfCast, data_threshold>; |
| 79 | - using Select = Bind<Vec::Select<uint8_t, float, SELECT_MODE_TENSOR>, Compare, data_zero, OpCopyInGradCast>; | 79 | + using Select = Bind<Vec::Select<uint8_t, float, SELECT_MODE_TENSOR>, Compare, OpCopyInGradCast, data_zero>; |
| 80 | using SelectCast = Bind<Vec::Cast<U, float, 1>, Select>; | 80 | using SelectCast = Bind<Vec::Cast<U, float, 1>, Select>; |
| 81 | using OpCopyOut = Bind<Vec::CopyOut<U>, Placeholder::Out0<U>, SelectCast>; | 81 | using OpCopyOut = Bind<Vec::CopyOut<U>, Placeholder::Out0<U>, SelectCast>; |
| 82 | // 指定输出节点 | 82 | // 指定输出节点 |
| @@ -28,10 +28,7 @@ def threshold_grad_v2_d_golden(grad_output, self_tensor, *, threshold=1.0, **kwa | |||
| 28 | del kwargs | 28 | del kwargs |
| 29 | grad_output_t = numpy_to_torch_tensor(grad_output) | 29 | grad_output_t = numpy_to_torch_tensor(grad_output) |
| 30 | self_t = numpy_to_torch_tensor(self_tensor) | 30 | self_t = numpy_to_torch_tensor(self_tensor) |
| 31 | - output_dtype = grad_output_t.dtype | 31 | + grad_output_t, self_t = torch.broadcast_tensors(grad_output_t, self_t) |
| 32 | - # The kernel promotes every supported dtype to float32 for comparison and | 32 | + mask = self_t.to(torch.float32) > float(threshold) |
| 33 | - # selection, then casts the selected gradient back to the output dtype. | 33 | + result = torch.where(mask, grad_output_t, torch.zeros_like(grad_output_t)) |
| 34 | - result = torch.ops.aten.threshold_backward( | 34 | + return torch_to_numpy_tensor(result.cpu()) |
| 35 | - grad_output_t.to(torch.float32), self_t.to(torch.float32), threshold | ||
| 36 | - ) | ||
| 37 | - return torch_to_numpy_tensor(result.to(output_dtype).cpu()) | ||