已合并
fix(threshold_grad_v2_d): align comparison semantics #9773
babeiee创建于 9月3日
fix(threshold_grad_v2_d): align comparison semantics #9773
已合并
共 4 个文件变更+32-24
| @@ -14,6 +14,8 @@ | |||
| 14 | */ | 14 | */ |
| 15 | 15 | ||
| 16 | 16 | ||
| 17 | + | ||
| 18 | + | ||
| 17 | 19 | ||
| 18 | 20 | ||
| 19 | 21 | ||
| @@ -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 @@ | |||
| 22 | namespace ThresholdGradV2DOp { | 22 | namespace ThresholdGradV2DOp { |
| 23 | using namespace Ops::Base; | 23 | using namespace Ops::Base; |
| 24 | 24 | ||
| 25 | -constexpr int COMPARE_MODE_GT = 1; | 25 | +constexpr int COMPARE_MODE_LE = 3; |
| 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_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> | ||
| 51 | struct ThresholdGradV2DInt32Dag { | 50 | struct 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 | ||
陈 | |||
| 42 | + ) | ||
| 43 | + return torch_to_numpy_tensor(result.to(output_dtype).cpu()) | ||
| 40 | 44 | ||
| 41 | 45 | ||
| 42 | def aclnn_threshold_backward_golden(gradOutput, self, threshold, out, **kwargs): | 46 | def aclnn_threshold_backward_golden(gradOutput, self, threshold, out, **kwargs): |
[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邻域边界用例。