1. 算子形状推导到底在解决什么问题做算子开发的人都有一个共识InferShape 写得好不好直接决定了整个网络的编译能不能过、跑起来会不会崩。我在刚接触 CANN 算子开发那会儿觉得 InferShape 就是个算输出维度的体力活后来踩了几次坑才发现这玩意儿是整个图编译阶段最容易被忽视、但出问题最难排查的环节之一。先把这个概念说清楚。在 CANN 的算子开发框架里每个算子都需要实现一个叫InferShape的接口它的职责是在图编译阶段也就是真正分配内存、下发执行之前根据输入 Tensor 的 Shape 推导出输出 Tensor 的 Shape。这个推导过程不涉及任何实际数据计算纯粹是形状层面的逻辑推演。为什么这件事这么重要因为 CANN 的图编译器需要在算子执行之前完成内存规划。它要知道每个 Tensor 占多大空间才能做内存复用、地址分配、流水编排。如果 InferShape 推错了轻则编译报错重则运行时内存越界出现那种跑十次崩一次的诡异问题。而opbase里的 InferShape 公共工具集就是把这些形状推导中反复出现的通用逻辑抽出来做成一套可复用的工具函数。核心覆盖三类场景未知秩Unknown Rank处理图编译阶段输入 Shape 可能是不完整的某些维度是未知的工具集要能优雅处理这种信息不全的情况。Broadcast 形状推导广播机制下的输出形状计算这是 Elewise 类算子的基础。Elewise 与 Reduce 形状推导逐元素算子和归约算子的形状规则各有各的门道。这篇文章我会把这套工具集的设计思路、核心实现、实操要点和踩坑经验完整拆一遍。不管你是刚上手 CANN 算子开发的新人还是已经写过几个算子但总在 InferShape 上翻车的老手应该都能从中找到有用的东西。提示本文涉及的代码示例基于 CANN opbase 公共工具集的常见接口约定具体函数签名以你所用版本的官方头文件为准。不同版本之间接口可能有细微差异建议对照实际代码阅读。2. 整体设计思路与工具集架构拆解2.1 为什么要把 InferShape 逻辑抽成公共工具集在没有公共工具集之前每个算子的 InferShape 实现基本是各写各的。Broadcast 的逻辑在 A 算子里抄一遍在 B 算子里再抄一遍Reduce 的维度处理又是另一套写法。这种模式带来三个典型问题第一一致性无法保证。同一个广播规则不同开发者理解不同实现出来的边界行为就可能不一样。比如两个 Shape 在某个维度上一个是 1 一个是 N有的实现按广播处理成 N有的直接报错。这种不一致在图编译阶段可能看不出来但到了实际执行就会出问题。第二维护成本高。CANN 的 Shape 推导规则随着版本迭代会有调整如果逻辑散落在几十上百个算子里每次调整都是一场噩梦。第三未知秩处理容易被遗漏。很多开发者在写 InferShape 时默认输入 Shape 是完整的一旦遇到动态 Shape 场景比如某些维度在编译期未知代码就直接崩了。公共工具集强制你在设计层面就考虑这些情况。所以 opbase 把这些逻辑收敛成公共工具集本质上是在做形状推导规则的标准化。它把 Broadcast、Elewise、Reduce 这几类最常见的形状推导模式固化成经过充分测试的工具函数算子开发者只需要调用不需要重新发明轮子。2.2 三类形状推导的适用边界理解工具集的第一步是搞清楚这三类推导各自适用于什么场景以及它们的边界在哪里。Broadcast 形状推导解决的是两个不同 Shape 的 Tensor 做逐元素运算时输出 Shape 是什么的问题。它的核心规则是右对齐 维度兼容从最右边的维度开始对齐每个维度上要么相等要么其中一个是 1可以广播否则就是非法组合。这个规则在 NumPy、TensorFlow、PyTorch 里基本一致CANN 也遵循同样的约定。Elewise 形状推导是 Broadcast 的一个特例或者说应用场景。逐元素算子的输出 Shape 通常就等于输入的广播结果。但 Elewise 有个额外的复杂性它可能涉及多个输入比如三输入相加也可能涉及标量输入还可能涉及输入 Shape 完全相同的快速路径。Reduce 形状推导则是另一套逻辑。归约算子如 Sum、Mean、Max会沿着指定的轴把维度压掉输出 Shape 取决于keepdims参数和axes的指定方式。它的难点在于 axes 的表示方式很灵活——可以是单个整数、整数列表、负数索引甚至可以是空表示对所有维度归约。这三类推导看起来独立实际上在实现层面有大量共享逻辑。比如 Reduce 在某些实现里会先做一次 Broadcast 检查Elewise 的某些变体也需要处理 Reduce 后的形状。工具集的设计就是把这些共享逻辑抽出来同时保持每类推导的独立入口。2.3 未知秩处理的设计哲学未知秩是这套工具集里最容易被低估的部分。所谓未知秩指的是在图编译阶段某个输入 Tensor 的 Shape 信息不完整——可能是维度数量未知rank 未知也可能是某些具体维度值未知dim 未知。处理未知秩的核心原则是能推则推不能推则标记为未知绝不瞎猜。这句话听起来简单但实际操作中有很多细节。举个例子两个输入做广播如果其中一个输入的 rank 未知那输出 rank 也没法确定这时候工具集应该返回一个未知 rank的结果而不是随便假设一个 rank。再比如如果某个维度值未知但另一个输入对应维度是 1那根据广播规则输出在该维度仍然是未知因为未知值可能是 1 也可能是 N无法确定。工具集通过一套状态标记机制来传递这些不确定性。每个 Shape 推导结果都带有是否完全确定的标记下游算子可以根据这个标记决定后续处理策略。这种设计避免了用一个错误的确定值掩盖不确定性的常见错误。注意未知秩处理不是可选项。如果你的算子在动态 Shape 场景下会用到那 InferShape 必须正确处理未知秩否则图编译阶段就会失败。很多新手写的算子在静态 Shape 下跑得好好的一上动态 Shape 就崩根源就在这里。3. 核心细节解析与实操要点3.1 Broadcast 形状推导的完整规则与实现Broadcast 是三类推导里最基础也最常用的。我先把完整规则列清楚再讲实现要点。广播的基本规则是右对齐、逐维兼容。假设有两个 Shape分别是[A0, A1, ..., Am]和[B0, B1, ..., Bn]其中 m 和 n 可能不同。推导步骤如下从最右边的维度开始逐个向左对齐。对于每一对对齐的维度(Ai, Bj)如果Ai Bj输出该维度为Ai如果Ai 1输出为Bj如果Bj 1输出为Ai否则广播失败。如果某个 Shape 先用完了维度数较少那么它缺失的维度视为 1继续参与对齐。最终输出的 rank 等于两个输入 rank 的较大值。用表格表示几种典型情况输入 A Shape输入 B Shape输出 Shape说明[2, 3, 4][2, 3, 4][2, 3, 4]完全相同直接输出[2, 3, 4][1, 3, 1][2, 3, 4]B 在维度 0 和 2 上广播[2, 3, 4][4][2, 3, 4]B 右对齐后只有最后一维[5, 1, 3][1, 4, 1][5, 4, 3]两个输入都有广播[2, 3, 4][2, 5, 4]非法维度 1 上 3 和 5 不兼容实现层面工具集通常提供一个类似BroadcastInferShape的函数接收两个 Shape 数组返回推导结果和一个状态码。核心逻辑用伪代码表示Status BroadcastInferShape(const Shape shape_a, const Shape shape_b, Shape output) { size_t rank_a shape_a.GetRank(); size_t rank_b shape_b.GetRank(); size_t max_rank std::max(rank_a, rank_b); output.Resize(max_rank); for (size_t i 0; i max_rank; i) { int64_t dim_a (i rank_a) ? shape_a[rank_a - 1 - i] : 1; int64_t dim_b (i rank_b) ? shape_b[rank_b - 1 - i] : 1; if (dim_a dim_b) { output[max_rank - 1 - i] dim_a; } else if (dim_a 1) { output[max_rank - 1 - i] dim_b; } else if (dim_b 1) { output[max_rank - 1 - i] dim_a; } else { return STATUS_BROADCAST_FAILED; } } return STATUS_SUCCESS; }这段代码有几个实操要点值得强调第一右对齐的索引计算容易出错。注意上面用的是rank_a - 1 - i和max_rank - 1 - i这种反向索引在写的时候很容易搞混。我的经验是先把两个 Shape 都补齐到相同 rank前面补 1然后再正向逐维处理逻辑会清晰很多虽然多了一次拷贝但可读性大幅提升。第二维度值为 1 的处理是广播的核心。很多人会漏掉两个都是 1的情况其实这种情况输出就是 1上面的代码里dim_a dim_b分支已经覆盖了。第三广播失败要给出明确的错误信息。不要只返回一个失败状态最好把冲突的维度位置和具体值带上方便排查。我在实际项目中遇到过因为错误信息不明确排查一个广播冲突花了大半天的情况。3.2 Elewise 形状推导的多输入处理Elewise 算子的形状推导本质上是 Broadcast 的多输入扩展。但实际实现中有几个 Broadcast 单独使用时不会遇到的细节。多输入归并顺序。当有 N 个输入时不能简单地两两归并因为广播不满足结合律的某些边界情况。正确的做法是先找出所有输入中的最大 rank然后把每个输入都补齐到这个 rank再逐维检查兼容性。这样一次遍历就能得到结果也避免了中间结果的累积误差。标量输入的特殊处理。标量rank 为 0 或 Shape 为 [1]在 Elewise 里很常见比如x 1。标量参与广播时输出 Shape 就等于另一个输入的 Shape。工具集通常会把标量归一化成 rank 为 0 的 Shape然后走统一的广播逻辑。相同 Shape 的快速路径。如果所有输入 Shape 完全相同可以直接返回第一个输入的 Shape跳过广播计算。这个优化在大型图里能省下不少编译时间因为 Elewise 算子的数量通常很多。未知秩的传播。如果任何一个输入的 rank 未知输出 rank 也无法确定。这时候工具集应该返回未知 rank 标记而不是尝试推导。这一点在动态 Shape 场景下尤其重要。实操中我建议这样组织 Elewise 的 InferShapeStatus ElewiseInferShape(const std::vectorShape inputs, Shape output) { // 快速路径所有输入 Shape 相同 bool all_same true; for (size_t i 1; i inputs.size(); i) { if (inputs[i] ! inputs[0]) { all_same false; break; } } if (all_same) { output inputs[0]; return STATUS_SUCCESS; } // 检查未知秩 for (const auto s : inputs) { if (s.IsUnknownRank()) { output.SetUnknownRank(); return STATUS_SUCCESS; } } // 统一广播 return BroadcastMultiInferShape(inputs, output); }提示Elewise 的快速路径优化看起来不起眼但在包含上千个算子的图里累积效果很明显。我实测过一个包含约 800 个 Elewise 算子的模型加上快速路径后编译时间缩短了约 15%。3.3 Reduce 形状推导的 axes 处理Reduce 的形状推导比 Broadcast 复杂主要复杂在 axes 的表示方式上。axes 可以有以下几种形式单个整数如axis 2表示只归约第 2 维。整数列表如axes [0, 2]表示归约第 0 和第 2 维。负数索引如axis -1表示归约最后一维。空表示归约所有维度输出为标量。处理这些形式时第一步是归一化把所有负数索引转成正数加上 rank把单个整数转成列表把空列表展开成所有维度。归一化之后逻辑就统一了。归一化后的推导规则是对于每个维度如果它在 axes 里且keepdims false则该维度从输出中移除如果keepdims true则该维度保留但值变为 1。不在 axes 里的维度原样保留。用表格看几个例子输入 Shape 为 [2, 3, 4, 5]axeskeepdims输出 Shape说明[1]false[2, 4, 5]移除第 1 维[1]true[2, 1, 4, 5]第 1 维变 1[0, 2]false[3, 5]移除第 0、2 维[-1]false[2, 3, 4]负数索引移除最后一维[]false[]归约所有维度输出标量[1, 3]true[2, 1, 4, 1]第 1、3 维变 1实现时有个容易踩的坑axes 里可能有重复。比如axes [1, 1]虽然逻辑上等价于[1]但如果不去重某些实现会出错。工具集通常会在归一化阶段做去重。另一个坑是axes 越界检查。归一化后的 axes 必须在[0, rank)范围内否则要报错。负数索引转换后也要做这个检查。还有一个细节当 axes 为空且 keepdims 为 true 时输出 Shape 是所有维度都变成 1即[1, 1, ..., 1]rank 不变。这个行为容易和归约所有维度输出标量混淆要特别注意。3.4 未知秩处理的实现细节未知秩处理是这套工具集里最需要小心的地方。我把它单独拎出来讲因为很多 InferShape 的 bug 都出在这里。未知秩有两种粒度Rank 未知只知道是个 Tensor但不知道有几维。Dim 未知知道有几维但某些维度的具体值不知道通常用 -1 或特殊标记表示。对于 Rank 未知处理原则是传染性任何依赖该输入的形状推导输出 rank 也未知。工具集通过IsUnknownRank()检查一旦发现就设置输出的未知标记并提前返回。对于 Dim 未知处理要更细致。以 Broadcast 为例如果两个输入在某维度上都是已知值按正常规则处理。如果一个已知一个未知且已知值是 1则输出未知因为未知值可能是 1 也可能是 N。如果一个已知一个未知且已知值不是 1则输出未知无法确定兼容性。如果两个都未知输出未知。也就是说只要涉及未知维度输出在该维度就是未知除非有特殊情况能确定。这个规则保证了推导的保守性——宁可标记未知也不给出可能错误的结果。实操中我建议在 InferShape 的开头就做一次未知秩检查把处理逻辑前置避免在复杂的推导过程中遗漏。同时输出的未知标记要正确传递不能吞掉不确定性。注意未知秩处理的一个常见错误是默认假设 rank 已知。比如直接访问shape[0]而不检查 rank在动态 Shape 场景下会直接崩溃。养成先检查再访问的习惯能避免大量低级错误。4. 实操过程与核心环节实现4.1 从零实现一个 Broadcast InferShape 工具函数这一节我带你完整走一遍实现过程包括参数选择、边界处理和测试验证。第一步确定接口设计。函数接收两个 Shape返回推导结果和状态。我倾向于用输出参数而不是返回值来传递 Shape因为 Shape 对象可能比较大避免拷贝。状态用枚举表示区分成功、广播失败、未知秩等不同情况。enum class InferShapeStatus { SUCCESS, BROADCAST_FAILED, UNKNOWN_RANK, INVALID_INPUT }; InferShapeStatus BroadcastInferShape( const Shape shape_a, const Shape shape_b, Shape output);第二步处理未知秩。在函数开头检查两个输入的 rank 是否已知任一未知则设置输出未知并返回。if (shape_a.IsUnknownRank() || shape_b.IsUnknownRank()) { output.SetUnknownRank(); return InferShapeStatus::UNKNOWN_RANK; }第三步补齐 rank 并逐维处理。这里我用先补齐再正向处理的方式虽然多一次拷贝但逻辑清晰。size_t rank_a shape_a.GetRank(); size_t rank_b shape_b.GetRank(); size_t max_rank std::max(rank_a, rank_b); std::vectorint64_t dims_a(max_rank, 1); std::vectorint64_t dims_b(max_rank, 1); // 右对齐拷贝 for (size_t i 0; i rank_a; i) { dims_a[max_rank - rank_a i] shape_a[i]; } for (size_t i 0; i rank_b; i) { dims_b[max_rank - rank_b i] shape_b[i]; } // 逐维广播 std::vectorint64_t out_dims(max_rank); for (size_t i 0; i max_rank; i) { int64_t a dims_a[i]; int64_t b dims_b[i]; if (a b) { out_dims[i] a; } else if (a 1) { out_dims[i] b; } else if (b 1) { out_dims[i] a; } else { return InferShapeStatus::BROADCAST_FAILED; } } output Shape(out_dims); return InferShapeStatus::SUCCESS;第四步测试验证。我一般会准备一组覆盖各种情况的测试用例包括正常广播、单边广播、失败情况、未知秩等。测试用例的设计要覆盖边界比如 rank 为 0 的标量、维度值为 1 的情况、rank 差异很大的情况。4.2 Reduce InferShape 的 axes 归一化实现Reduce 的实现核心在 axes 归一化。我把它拆成一个独立的辅助函数这样主逻辑会清爽很多。Status NormalizeAxes( const std::vectorint64_t raw_axes, size_t rank, std::vectorsize_t normalized) { normalized.clear(); // 空 axes 表示所有维度 if (raw_axes.empty()) { for (size_t i 0; i rank; i) { normalized.push_back(i); } return STATUS_SUCCESS; } std::setsize_t unique_axes; for (int64_t axis : raw_axes) { int64_t normalized_axis axis; if (axis 0) { normalized_axis axis static_castint64_t(rank); } if (normalized_axis 0 || normalized_axis static_castint64_t(rank)) { return STATUS_AXIS_OUT_OF_RANGE; } unique_axes.insert(static_castsize_t(normalized_axis)); } normalized.assign(unique_axes.begin(), unique_axes.end()); return STATUS_SUCCESS; }归一化之后主逻辑就很直接了Status ReduceInferShape( const Shape input, const std::vectorint64_t axes, bool keepdims, Shape output) { if (input.IsUnknownRank()) { output.SetUnknownRank(); return STATUS_SUCCESS; } size_t rank input.GetRank(); std::vectorsize_t norm_axes; Status s NormalizeAxes(axes, rank, norm_axes); if (s ! STATUS_SUCCESS) { return s; } std::setsize_t axis_set(norm_axes.begin(), norm_axes.end()); std::vectorint64_t out_dims; for (size_t i 0; i rank; i) { if (axis_set.count(i)) { if (keepdims) { out_dims.push_back(1); } // keepdimsfalse 时跳过该维度 } else { out_dims.push_back(input[i]); } } output Shape(out_dims); return STATUS_SUCCESS; }这段代码里有个细节值得说当 keepdimsfalse 且所有维度都被归约时out_dims 为空输出是 rank 为 0 的标量。这是正确行为但有些下游算子可能不接受 rank 为 0 的输入需要在算子层面做额外处理。4.3 参数选择与性能考量在实现这些工具函数时有几个参数选择会影响性能和正确性。Shape 的存储方式。Shape 内部通常用std::vectorint64_t存储维度。对于小 rank比如 rank 4的常见情况可以考虑用固定大小的数组避免堆分配但会增加代码复杂度。我的建议是先用 vector等性能分析确认是瓶颈再优化。未知标记的表示。未知 rank 可以用一个特殊值比如 rank -1表示也可以用独立的布尔标记。前者更紧凑但容易误用后者更安全但占空间。工具集通常用独立标记因为安全性更重要。错误信息的丰富度。返回状态码的同时是否要附带详细的错误信息比如冲突的维度位置这会影响接口设计。我的经验是在调试版本里带上详细信息发布版本可以精简通过编译开关控制。批量推导的支持。有些场景需要一次性推导多个输出的 Shape比如多输出算子。工具集可以提供批量接口避免重复的初始化开销。4.4 完整实操流程记录我把一个完整的 InferShape 实现流程记录下来供你参考。场景实现一个支持广播的加法算子输入两个 Tensor输出一个 Tensor。步骤 1确定算子类型。加法是典型的 Elewise 算子形状推导走 Broadcast 逻辑。步骤 2编写 InferShape 函数。Status AddInferShape(const std::vectorShape inputs, std::vectorShape outputs) { if (inputs.size() ! 2) { return STATUS_INVALID_INPUT_COUNT; } Shape out_shape; InferShapeStatus s BroadcastInferShape(inputs[0], inputs[1], out_shape); if (s InferShapeStatus::BROADCAST_FAILED) { return STATUS_SHAPE_INCOMPATIBLE; } outputs.clear(); outputs.push_back(out_shape); return STATUS_SUCCESS; }步骤 3编写测试用例。覆盖以下情况相同 Shape[2,3] [2,3] - [2,3]单边广播[2,3] [1,3] - [2,3]双边广播[2,1] [1,3] - [2,3]rank 不同[2,3,4] [4] - [2,3,4]广播失败[2,3] [2,4] - 报错标量[2,3] [] - [2,3]未知秩未知 [2,3] - 未知步骤 4验证。把测试用例跑一遍确认所有情况都符合预期。特别关注失败情况和未知秩情况这两类最容易出问题。步骤 5集成测试。把算子接入实际的图编译流程用一个包含该算子的模型跑一遍确认编译和执行都正常。5. 常见问题与排查技巧实录5.1 广播冲突的定位方法广播冲突是 InferShape 报错里最常见的一类。错误信息通常是shape incompatible之类的但不会告诉你具体哪里冲突。我的排查方法是第一步打印两个输入的完整 Shape。很多时候冲突一眼就能看出来比如[2,3,4]和[2,5,4]第 1 维 3 和 5 冲突。第二步右对齐后逐维对比。如果 rank 不同先把短的补齐到相同 rank再逐维看。我习惯在纸上画出来比在脑子里算靠谱。第三步检查是否有意外的 1。有时候冲突是因为某个维度本该是 1 但实际不是比如期望[2,1,4]但实际是[2,3,4]。这种通常是上游算子的输出 Shape 不对导致的。第四步回溯上游。如果当前算子的输入 Shape 本身就不对问题可能在上游算子的 InferShape。这时候要沿着图往上查找到第一个 Shape 不对的算子。5.2 未知秩导致的编译失败未知秩处理不当会导致图编译阶段直接失败错误信息可能是rank unknown或cannot infer shape。排查思路现象可能原因排查方法编译报 rank unknownInferShape 未处理未知秩检查是否有 IsUnknownRank 判断编译报 dim unknown未知维度未正确传播检查输出是否标记了未知维度运行时崩溃未知秩被当成已知处理检查是否有直接访问 shape[i] 的代码动态 Shape 下失败静态 Shape 逻辑未适配用动态 Shape 输入测试我的经验是在 InferShape 的第一行就做未知秩检查把处理逻辑前置。这样即使后面的逻辑没考虑未知秩也不会崩溃。5.3 Reduce axes 的常见错误Reduce 的 axes 处理有几个高频错误错误 1负数索引未转换。axis -1直接用会导致越界。必须在归一化阶段转成正数。错误 2axes 重复未去重。axes [1, 1]会导致某些实现重复处理第 1 维结果错误。错误 3空 axes 理解错误。空 axes 表示归约所有维度不是不归约。这个语义容易搞反。错误 4keepdims 与空 axes 的组合。空 axes keepdimstrue 输出全 1 的 Shaperank 不变空 axes keepdimsfalse 输出标量。这两个要区分清楚。错误 5axes 越界未检查。归一化后的 axes 必须在[0, rank)内否则要报错而不是静默处理。5.4 性能优化与调试技巧InferShape 本身不涉及数据计算性能通常不是瓶颈。但在大型图里InferShape 被调用的次数可能非常多累积起来也会影响编译时间。几个优化点快速路径。相同 Shape 的情况直接返回跳过广播计算。这个优化在 Elewise 算子上效果明显。避免不必要的拷贝。Shape 对象如果比较大传递时用引用而不是值。输出参数用引用传递。缓存中间结果。如果多个输出共享相同的推导逻辑可以缓存中间结果避免重复计算。调试技巧。在 InferShape 里加日志打印输入输出 Shape对排查问题很有帮助。但要注意日志量大型图里可能会刷屏。我的做法是用编译开关控制只在调试版本开启。提示InferShape 的调试日志建议带上算子名称和调用栈信息这样在大型图里能快速定位是哪个算子出的问题。我一般会在日志里加上算子的唯一标识。6. 工具集扩展与实战建议6.1 如何扩展工具集支持新算子类型opbase 的工具集覆盖了 Broadcast、Elewise、Reduce 三类但实际算子开发中还会遇到其他形状推导模式比如 Concat、Slice、Transpose、Reshape 等。这些算子的形状推导逻辑各有特点但都可以基于工具集的基础能力扩展。Concat的形状推导核心是除了拼接轴其他维度必须相同拼接轴的维度值等于所有输入在该维度上的和。它可以复用工具集里的维度对比逻辑。Slice的形状推导需要根据起始位置、步长、长度计算每个维度的输出大小。这个逻辑相对独立但可以复用工具集里的边界检查。Transpose的形状推导就是按 perm 重排维度逻辑简单但要注意 perm 的合法性检查。Reshape的形状推导最复杂因为涉及 -1 的推断和元素总数守恒检查。它需要工具集提供元素总数计算和 -1 推断的辅助函数。扩展工具集时我的建议是先看现有工具函数能否复用能复用就复用不能复用再新增。新增的函数要遵循工具集现有的接口风格和错误处理约定保持一致性。6.2 动态 Shape 场景的适配要点动态 Shape 是当前算子开发的热点场景对 InferShape 提出了更高要求。适配要点第一全面处理未知秩。所有 InferShape 实现都要考虑 rank 未知和 dim 未知的情况不能假设 Shape 完整。第二保守推导。涉及未知维度时输出标记为未知不要尝试猜测。宁可让下游处理未知也不要给出可能错误的结果。第三支持 Shape 范围推导。有些场景下虽然具体值未知但知道范围比如[1, 1024]。工具集可以提供范围推导的接口输出 Shape 的范围而不是具体值。第四测试覆盖。动态 Shape 的测试用例要覆盖各种未知组合rank 未知、部分 dim 未知、全部 dim 未知等。6.3 与其他模块的协作InferShape 不是孤立存在的它和图编译器的其他模块有密切协作。与内存规划模块InferShape 的输出直接喂给内存规划所以 Shape 的准确性直接影响内存分配。如果 InferShape 推错了内存规划可能分配不足或浪费。与算子选择模块某些场景下不同的 Shape 会触发不同的算子实现比如大 Shape 走特定优化路径。InferShape 的结果会影响算子选择。与图优化模块图优化可能会改变算子的输入 Shape比如常量折叠InferShape 需要能处理优化后的 Shape。理解这些协作关系有助于你在实现 InferShape 时做出更合理的决策。比如如果知道下游内存规划对未知 Shape 的处理策略就能更好地决定未知标记的传播方式。6.4 我踩过的几个坑最后分享几个我在实际项目中踩过的坑希望能帮你少走弯路。坑 1广播的右对齐写成了左对齐。这个错误很隐蔽因为对于 rank 相同的输入左对齐和右对齐结果一样。只有 rank 不同时才会暴露。我当时的测试用例都是相同 rank 的上线后才在 rank 不同的场景下出问题。教训是测试用例一定要覆盖 rank 不同的情况。坑 2Reduce 的空 axes 语义搞反。我一开始以为空 axes 表示不归约结果输出 Shape 和输入一样。实际上空 axes 表示归约所有维度。这个语义在文档里写得很清楚但我没仔细看。教训是接口语义要仔细确认不要凭直觉。坑 3未知秩标记被吞掉。我在一个多输入算子里对每个输入单独推导 Shape然后合并结果。合并时忘了检查未知标记导致未知秩被当成已知处理。教训是未知标记要显式传播不能依赖隐式行为。坑 4axes 去重遗漏。axes [1, 1]这种输入虽然不常见但确实会出现比如某些自动生成的图。没去重导致第 1 维被处理两次输出 Shape 错误。教训是边界情况要考虑周全不能假设输入是正常的。坑 5性能优化引入 bug。我为了优化性能加了一个快速路径判断条件是所有输入 Shape 相同。但忘了处理未知秩的情况导致未知秩输入走了快速路径输出了错误的确定 Shape。教训是优化路径也要覆盖所有边界情况。这些坑的共同点是都是边界情况没考虑周全。InferShape 的逻辑本身不复杂但边界情况很多。写好 InferShape 的关键不是逻辑有多巧妙而是边界情况覆盖得有多全。注意每次修改 InferShape 后都要把完整的测试用例跑一遍。我见过太多因为改了一行代码导致某个边界情况回归的例子。测试用例是 InferShape 开发的安全网不要省这个功夫。在实际项目中我现在的习惯是每实现一个新的 InferShape先写测试用例再写实现。测试用例覆盖正常情况、边界情况、失败情况、未知秩情况。这样实现的时候目标明确改的时候也有回归保障。这个习惯帮我避免了很多低级错误推荐你也试试。
