已合并
fix(threshold_grad_v2_d): 对齐接口及比较语义 #9150
fix(threshold_grad_v2_d): 对齐接口及比较语义 #9150
已合并
babeiee创建于 8月25日
共 3 个文件变更+17-20
@@ -12,8 +12,8 @@
12 12 
13/*!13/*!
14 * \file threshold_grad_v2_d_def.cpp14 * \file threshold_grad_v2_d_def.cpp
15- * \brief ThresholdGradV2D 算子定义:gradOutput/self -> out,fp16/fp32/bf16/int32/int8/uint815+ * \brief ThresholdGradV2D 算子定义:gradients/features -> backprops,fp16/fp32/bf16/int32/int8/uint8
16- * 属性 threshold(Float, 默认 1.0)。out = self>threshold ? gradOutput : 016+ * 属性 threshold(Float)。backprops = features>threshold ? gradients : 0
17 */17 */
18#include "register/op_def_registry.h"18#include "register/op_def_registry.h"
19 19 
@@ -22,21 +22,21 @@ class ThresholdGradV2D : public OpDef {
22public:22public:
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 @@
22namespace ThresholdGradV2DOp {22namespace ThresholdGradV2DOp {
23using namespace Ops::Base;23using namespace Ops::Base;
24 24 
25-constexpr int COMPARE_MODE_LE = 3;25+constexpr int COMPARE_MODE_GT = 1;
26constexpr int SELECT_MODE_TENSOR = 2;26constexpr int SELECT_MODE_TENSOR = 2;
27 27 
28template <typename U>28template <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 kwargs28 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.dtype31+ grad_output_t, self_t = torch.broadcast_tensors(grad_output_t, self_t)
32- # The kernel promotes every supported dtype to float32 for comparison and32+ 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())