Pull Request已成功合入, 合并人@CANN-robot
(感谢 xchu42 的贡献)变更摘要
本次 PR 主要修复 Concat 算子在 AscendC 后端上的精度问题与分组策略缺陷。核心改动包括三方面:在 Concat16MultipleColumns 和 MultipleInputsConcat16Rows 的转置流程中增加 AscendC::PipeBarrier<PIPE_V>() 同步指令以防止流水线竞争导致的精度错误;将临时缓冲区大小计算中的对齐判断逻辑从基于 concat 维度之后的维度乘积,改为基于 vectorized_axis/vectorized_strides 向量化信息判断是否全对齐;以及在 A3 场景下分组时对尾轴 stride 做 32B 对齐,避免产生过大的 group。
主要改动
-
ConcatExtend 增加 PipeBarrier 同步: 在
autofuse/ascendc/api/concat.h中,Concat16MultipleColumns和MultipleInputsConcat16Rows函数的第一次转置前、第二次转置前后分别插入AscendC::PipeBarrier<PIPE_V>(),消除流水线中的数据竞争,解决部分场景下的精度问题。 -
对齐判断改为基于向量化轴信息:
autofuse/ascir/reg_func/concat.cpp中IsAllStaticAligned函数重构,不再通过 concat 维度之后所有维度的乘积取模判断对齐,而是基于output_attr.vectorized_axis和vectorized_strides找到 concat 维度在向量化轴中的位置,再对各输入的repeats[pos] * vectorized_strides[concat_dim]做对齐检查;CalcForDefaultKernel和CalcConcatTmpSizeV2同步适配新接口。 -
分组 stride 计算尾轴对齐到 32B:
autofuse/optimize/task_generator/concat_group_partitioner.cpp的ParseConcatNode中,在concat_by_transpose_为 true 时,对最后一个轴(尾轴)的value向上对齐到kAlignment / dtype_size_(即 32B),防止 A3 场景下因未对齐产生过大的 concat group。 -
concat_by_transpose_初始化时机后移: 将ConcatGroupPartitioner::Initialize()中对concat_by_transpose_的赋值移至PartitionGroups()中,确保在需要时再获取后端配置。 -
测试用例补充向量化字段: 在
autofuse/tests/ut/ascir/reg_func/test_reg_func_concat.cpp的CreateStaticGraph和CreateStaticGraphNotLastAxis中为所有 tensor 节点补全vectorized_axis和vectorized_strides字段,使单测与新对齐判断逻辑匹配。


代码审查
审查总结
变更文件审查结果
| 文件 | 审查结果 |
|---|---|
autofuse/ascendc/api/concat.h |
✅ 无问题。新增的 PipeBarrier<PIPE_V>() 同步指令位置正确,与已有同步模式一致,能解决 PR 所述的精度问题。 |
autofuse/ascir/reg_func/concat.cpp |
✅ 无问题。IsAllStaticAligned 改用 vectorized_axis/vectorized_strides 进行对齐判断,逻辑正确;CalcForDefaultKernel / CalcConcatTmpSize / CalcConcatTmpSizeV2 的签名和调用适配正确。 |
autofuse/optimize/task_generator/concat_group_partitioner.cpp |
⚠️ P3 低风险:align_size 理论上可能为 0 导致除零(极端异常 dtype),已报告防御性检查建议。 |
autofuse/optimize/task_generator/concat_group_partitioner.h |
🔴 P0 编译错误:第 65 行 static bool InputHasTransposeOrReduce(...) const; 将 static 与 const 同时用于成员函数,为非法 C++,将导致编译失败。 |
autofuse/tests/ut/ascir/reg_func/test_reg_func_concat.cpp |
✅ 无问题。测试新增的 vectorized_axis / vectorized_strides 字段与生产代码一致,测试预期正确。 |
统计
- P0: 1 个(编译错误)
- P3: 1 个(防御性检查建议)
- 总计: 2 个发现
总体风险判断
高风险 — P0 编译错误将阻断构建,必须修复后才能合入。其余变更(同步指令补齐、对齐判断重构、分组 stride 对齐)逻辑正确且意图清晰。
| 类型 | 数量 |
|---|---|
| 🔴 阻塞 | 1 |
| 🟡 建议 | 0 |
⛔ 需要修改


/approve


Pull Request
描述
变更类型
请选择本次引入的变更类型:
关联的Issue
如何测试
描述测试此变更的步骤和前提条件:
NA
核对清单
其他信息
在此添加任何其他关于本次 PR 的说明。