已合并
fix(threshold_grad_v2_d): align comparison semantics #9773
fix(threshold_grad_v2_d): align comparison semantics #9773
已合并
babeiee创建于 9月3日
共 4 个文件变更+32-24
@@ -14,6 +14,8 @@
14 */14 */
15 15 
16#include "threshold_grad_v2_d_tiling.h"16#include "threshold_grad_v2_d_tiling.h"
17+#include <cmath>
18+#include <limits>
17#include <graph/utils/type_utils.h>19#include <graph/utils/type_utils.h>
18#include "log/log.h"20#include "log/log.h"
19#include "atvoss/broadcast/broadcast_tiling.h"21#include "atvoss/broadcast/broadcast_tiling.h"
@@ -87,13 +89,19 @@ ge::graphStatus ThresholdGradV2DTiling::DoOpTiling()
87 tilingKey = GET_TPL_TILING_KEY(brcBaseTiling.GetSchMode(), THRESHOLD_GRAD_V2_D_TPL_FP32);89 tilingKey = GET_TPL_TILING_KEY(brcBaseTiling.GetSchMode(), THRESHOLD_GRAD_V2_D_TPL_FP32);
88 brcBaseTiling.SetScalar<float>(thresHold);90 brcBaseTiling.SetScalar<float>(thresHold);
89 } else if (input0DType == ge::DT_INT32) {91 } else if (input0DType == ge::DT_INT32) {
90- BroadcastBaseTiling<ThresholdGradV2DInt32Dag<int32_t>::OpDag> brcBaseTiling(context_);92+ constexpr float int32MinAsFloat = static_cast<float>(std::numeric_limits<int32_t>::min());
93+ constexpr float int32MaxAsFloat = static_cast<float>(std::numeric_limits<int32_t>::max());
94+ OP_CHECK_IF(!std::isfinite(thresHold) || thresHold < int32MinAsFloat || thresHold >= int32MaxAsFloat,
95+ OP_LOGE(context_->GetNodeName(), "threshold must be a finite value representable as int32"),
96+ return ge::GRAPH_FAILED);
97+ const int32_t int32Threshold = static_cast<int32_t>(thresHold);
98+ BroadcastBaseTiling<ThresholdGradV2DInt32Dag::OpDag> brcBaseTiling(context_);
91 baseTilingResult = brcBaseTiling.DoTiling();99 baseTilingResult = brcBaseTiling.DoTiling();
92 OP_CHECK_IF(baseTilingResult == ge::GRAPH_FAILED,100 OP_CHECK_IF(baseTilingResult == ge::GRAPH_FAILED,
93- OP_LOGE(context_->GetNodeName(), "BroadcastBaseTiling<ThresholdGradV2DDag<int32_t>::OpDag> failed"),101+ OP_LOGE(context_->GetNodeName(), "BroadcastBaseTiling<ThresholdGradV2DInt32Dag::OpDag> failed"),
94 return ge::GRAPH_FAILED);102 return ge::GRAPH_FAILED);
95 tilingKey = GET_TPL_TILING_KEY(brcBaseTiling.GetSchMode(), THRESHOLD_GRAD_V2_D_TPL_INT32);103 tilingKey = GET_TPL_TILING_KEY(brcBaseTiling.GetSchMode(), THRESHOLD_GRAD_V2_D_TPL_INT32);
96- brcBaseTiling.SetScalar<float>(thresHold);104+ brcBaseTiling.SetScalar<int32_t>(int32Threshold);
97 } else if (input0DType == ge::DT_INT8) {105 } else if (input0DType == ge::DT_INT8) {
98 BroadcastBaseTiling<ThresholdGradV2D8BDag<int8_t>::OpDag> brcBaseTiling(context_);106 BroadcastBaseTiling<ThresholdGradV2D8BDag<int8_t>::OpDag> brcBaseTiling(context_);
99 baseTilingResult = brcBaseTiling.DoTiling();107 baseTilingResult = brcBaseTiling.DoTiling();
@@ -22,7 +22,7 @@
22namespace ThresholdGradV2DOp {22namespace ThresholdGradV2DOp {
23using namespace Ops::Base;23using namespace Ops::Base;
24 24 
25-constexpr int COMPARE_MODE_GT = 1;25+constexpr int COMPARE_MODE_LE = 3;
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_GT>, OpCopyInSelfCast, data_threshold>;39+ using Compare = Bind<Vec::Compare<uint8_t, float, COMPARE_MODE_LE>, OpCopyInSelfCast, data_threshold>;
40- using Select = Bind<Vec::Select<uint8_t, float, SELECT_MODE_TENSOR>, Compare, OpCopyInGradCast, data_zero>;40+ using Select = Bind<Vec::Select<uint8_t, float, SELECT_MODE_TENSOR>, Compare, data_zero, OpCopyInGradCast>;
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>;
@@ -47,19 +47,15 @@ struct ThresholdGradV2D8BDag {
47 using OpDag = DAGSch<Outputs, void, MemCfg>;47 using OpDag = DAGSch<Outputs, void, MemCfg>;
48};48};
49 49 
50-template <typename U>
51struct ThresholdGradV2DInt32Dag {50struct ThresholdGradV2DInt32Dag {
52- using const_zero = MAKE_CONST(float, 0.0);51+ using const_zero = MAKE_CONST(int32_t, 0);
53- using data_threshold = Bind<Vec::Duplicate<float>, Placeholder::Var<float, 0>>;52+ using data_threshold = Bind<Vec::Duplicate<int32_t>, Placeholder::Var<int32_t, 0>>;
54- using data_zero = Bind<Vec::Duplicate<float>, const_zero>;53+ using data_zero = Bind<Vec::Duplicate<int32_t>, const_zero>;
55- using OpCopyInGrad = Bind<Vec::CopyInBrc<U>, Placeholder::In0<U>>;54+ using OpCopyInGrad = Bind<Vec::CopyInBrc<int32_t>, Placeholder::In0<int32_t>>;
56- using OpCopyInGradCast = Bind<Vec::Cast<float, U, 1>, OpCopyInGrad>;55+ using OpCopyInSelf = Bind<Vec::CopyInBrc<int32_t>, Placeholder::In1<int32_t>>;
57- using OpCopyInSelf = Bind<Vec::CopyInBrc<U>, Placeholder::In1<U>>;56+ using Compare = Bind<Vec::Compare<uint8_t, int32_t, COMPARE_MODE_LE>, OpCopyInSelf, data_threshold>;
58- using OpCopyInSelfCast = Bind<Vec::Cast<float, U, 1>, OpCopyInSelf>;57+ using Select = Bind<Vec::Select<uint8_t, int32_t, SELECT_MODE_TENSOR>, Compare, data_zero, OpCopyInGrad>;
59- using Compare = Bind<Vec::Compare<uint8_t, float, COMPARE_MODE_GT>, OpCopyInSelfCast, data_threshold>;58+ using OpCopyOut = Bind<Vec::CopyOut<int32_t>, Placeholder::Out0<int32_t>, Select>;
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>;
62- using OpCopyOut = Bind<Vec::CopyOut<U>, Placeholder::Out0<U>, SelectCast>;
63 // 指定输出节点59 // 指定输出节点
64 using Outputs = Elems<OpCopyOut>;60 using Outputs = Elems<OpCopyOut>;
65 using MemCfg = MemOptCfg<MemLevel::LEVEL_2>;61 using MemCfg = MemOptCfg<MemLevel::LEVEL_2>;
@@ -75,8 +71,8 @@ struct ThresholdGradV2DDag {
75 using OpCopyInGradCast = Bind<Vec::Cast<float, U, 0>, OpCopyInGrad>;71 using OpCopyInGradCast = Bind<Vec::Cast<float, U, 0>, OpCopyInGrad>;
76 using OpCopyInSelf = Bind<Vec::CopyInBrc<U>, Placeholder::In1<U>>;72 using OpCopyInSelf = Bind<Vec::CopyInBrc<U>, Placeholder::In1<U>>;
77 using OpCopyInSelfCast = Bind<Vec::Cast<float, U, 0>, OpCopyInSelf>;73 using OpCopyInSelfCast = Bind<Vec::Cast<float, U, 0>, OpCopyInSelf>;
78- using Compare = Bind<Vec::Compare<uint8_t, float, COMPARE_MODE_GT>, OpCopyInSelfCast, data_threshold>;74+ using Compare = Bind<Vec::Compare<uint8_t, float, COMPARE_MODE_LE>, OpCopyInSelfCast, data_threshold>;
79- using Select = Bind<Vec::Select<uint8_t, float, SELECT_MODE_TENSOR>, Compare, OpCopyInGradCast, data_zero>;75+ using Select = Bind<Vec::Select<uint8_t, float, SELECT_MODE_TENSOR>, Compare, data_zero, OpCopyInGradCast>;
80 using SelectCast = Bind<Vec::Cast<U, float, 1>, Select>;76 using SelectCast = Bind<Vec::Cast<U, float, 1>, Select>;
81 using OpCopyOut = Bind<Vec::CopyOut<U>, Placeholder::Out0<U>, SelectCast>;77 using OpCopyOut = Bind<Vec::CopyOut<U>, Placeholder::Out0<U>, SelectCast>;
82 // 指定输出节点78 // 指定输出节点
@@ -36,7 +36,7 @@ __global__ __aicore__ void threshold_grad_v2_d(GM_ADDR gradOutput, GM_ADDR self,
36 BroadcastSch<schMode, ThresholdGradV2DDag<float>::OpDag> sch(tiling);36 BroadcastSch<schMode, ThresholdGradV2DDag<float>::OpDag> sch(tiling);
37 sch.Process(gradOutput, self, out);37 sch.Process(gradOutput, self, out);
38 } else if constexpr (dtype == static_cast<uint64_t>(THRESHOLD_GRAD_V2_D_TPL_INT32)) {38 } else if constexpr (dtype == static_cast<uint64_t>(THRESHOLD_GRAD_V2_D_TPL_INT32)) {
39- BroadcastSch<schMode, ThresholdGradV2DInt32Dag<int32_t>::OpDag> sch(tiling);39+ BroadcastSch<schMode, ThresholdGradV2DInt32Dag::OpDag> sch(tiling);
40 sch.Process(gradOutput, self, out);40 sch.Process(gradOutput, self, out);
41 } else if constexpr (dtype == static_cast<uint64_t>(THRESHOLD_GRAD_V2_D_TPL_INT8)) {41 } else if constexpr (dtype == static_cast<uint64_t>(THRESHOLD_GRAD_V2_D_TPL_INT8)) {
42 BroadcastSch<schMode, ThresholdGradV2D8BDag<int8_t>::OpDag> sch(tiling);42 BroadcastSch<schMode, ThresholdGradV2D8BDag<int8_t>::OpDag> sch(tiling);
@@ -34,9 +34,13 @@ def threshold_grad_v2_d_golden(grad_output, self_tensor, *, threshold=1.0, **kwa
34 grad_output_t = numpy_to_torch_tensor(grad_output)34 grad_output_t = numpy_to_torch_tensor(grad_output)
35 self_t = numpy_to_torch_tensor(self_tensor)35 self_t = numpy_to_torch_tensor(self_tensor)
36 grad_output_t, self_t = torch.broadcast_tensors(grad_output_t, self_t)36 grad_output_t, self_t = torch.broadcast_tensors(grad_output_t, self_t)
37- mask = self_t.to(torch.float32) > float(threshold)37+ output_dtype = grad_output_t.dtype
38- result = torch.where(mask, grad_output_t, torch.zeros_like(grad_output_t))38+ # The kernel promotes every supported dtype to float32 for comparison and
39- return torch_to_numpy_tensor(result.cpu())39+ # selection, then casts the selected gradient back to the output dtype.
40+ result = torch.ops.aten.threshold_backward(
41+ grad_output_t.to(torch.float32), self_t.to(torch.float32), threshold
陈
陈陈展熹29 天前

[P1] INT32 golden 仍会丢失精度

本 PR 的 kernel 已改为 int32 直接比较,但这里仍将 self_t 和 grad_output_t 无条件转换为 float32。以 self=16777217、threshold=16777216 为例,self_t.to(torch.float32) 会舍入为 16777216,golden 返回 0,而 kernel 的 int32 比较应透传 grad;grad 本身超过 2^24 时也会在转为 float32 后被改值。因此 INT32 边界用例会被错误判定,本 PR 的精确整数语义无法由该 golden 验证。请对 INT32 保持整数比较,并补充 2^24 邻域边界用例。

likedislike
wangweidong
29 天前 评论:
42+ )
43+ return torch_to_numpy_tensor(result.to(output_dtype).cpu())
40 44 
41 45 
42def aclnn_threshold_backward_golden(gradOutput, self, threshold, out, **kwargs):46def aclnn_threshold_backward_golden(gradOutput, self, threshold, out, **kwargs):