
【Bug已解决】GroupQueryAttentionFusion produces an invalid graph when a GQA node has more than 9 inputs (e.g. attention_bias) — sum of input arg count is not equal to size of input defs 解决方案一、现象长什么样ONNX Runtime 的图优化器里有一个GroupQueryAttentionFusionpass它把若干 Q/K/V 处理节点融合成一个GroupQueryAttention算子。当被融合的 GQA 节点带有额外输入比如attention_bias使输入数 9时融合产出的图是非法的节点声明的输入定义数量和实际传入的输入参数数量对不上ORT 在图校验阶段直接报错。现象# 现象 A图校验失败 # FAIL: Node (GroupQueryAttention) has inputs count mismatch: # sum of input arg count is not equal to size of input defs # 节点 input defs 说有 9 个输入实际融进去了 10 个 # 现象 B只在带 attention_bias 的 GQA 触发 # 普通 GQAQ/K/V 必要的 mask/scale 共 ≤9 个输入融合正常 # 一旦加了 attention_bias 这个第 10 个输入融合逻辑炸 # 现象 C不报错但运行结果错 # 某些版本里校验被绕过融合节点“吞掉”了多出来的输入 # 导致 attention_bias 没生效注意力结果静默偏差最坑的是现象 C图校验有时被关掉或版本不同于是非法的融合节点带着缺失的 bias 继续跑注意力算错且不报错极难自查。二、背景GroupQueryAttentionGQA算子在 ONNX 里有一组固定的输入槽query、key、value以及可选的past_key、past_value、attention_bias、past_sequence_length等。不同实现支持的输入数不同常见是 9 个以内。ORT 的GroupQueryAttentionFusion在遍历计算图、识别可融合的子图时会把识别到的输入按固定顺序塞进融合节点的 input 列表。问题出在融合逻辑对“输入数量”做了硬编码上限如MAX_GQA_INPUTS 9或固定长度的 input defs 模板当 GQA 节点实际有 10 个输入多了attention_bias时融合代码要么① 只拷贝了前 9 个输入第 10 个丢了现象 C要么 ② 把 10 个输入都塞进节点但节点的input defsop schema 声明的输入数仍是 9导致图校验报“数量不匹配”现象 A。这是图优化器审查里典型的坑**融合 pass 对算子输入数做了固定假设没随 op schema 的实际输入数动态调整」。三、根因融合逻辑硬编码输入上限GroupQueryAttentionFusion假设 GQA 最多 9 个输入遇到带attention_bias的 10 输入节点时处理错位。input defs 与实际输入数不一致融合产物节点的 input 列表长度 ≠ op schema 声明的 input defs 数量图校验必然失败现象 A。缺少对“多输入 GQA”的图校验对拍CI 只测标准 9 输入 GQA带 bias 的 10 输入路径从未覆盖非法图长期存在。本质是GQA 融合 pass 对输入数做固定假设、未随 schema 动态适配导致多输入节点产出非法图且缺测试覆盖。四、最小可运行复现下面用 Python 模拟“融合把 10 个输入塞进只声明 9 个 input defs 的节点导致数量不匹配”GQA_INPUT_DEFS 9 # op schema 声明的输入数固定 def fuse_gqa_buggy(recognized_inputs: list): buggy: 直接把识别到的输入全塞进节点不管 defs 数量。 # 实际塞了 len(recognized_inputs) 个但节点 defs 仍是 9 node_inputs recognized_inputs if len(node_inputs) ! GQA_INPUT_DEFS: raise ValueError( fsum of input arg count ({len(node_inputs)}) is not equal to fsize of input defs ({GQA_INPUT_DEFS})) return node_inputs inputs_9 [q, k, v, pk, pv, mask, scale, b1, b2] inputs_10 inputs_9 [attention_bias] print(9-input ok:, len(fuse_gqa_buggy(inputs_9))) # 通过 try: fuse_gqa_buggy(inputs_10) # 10 输入 - 报错 except ValueError as e: print(REPRO A -, e) def fuse_gqa_fixed(recognized_inputs: list, schema_inputs: int): fixed: 按 schema 实际声明的输入数来构建节点。 if len(recognized_inputs) schema_inputs: raise ValueError(fGQA node has {len(recognized_inputs)} inputs, fschema supports {schema_inputs}) # 用 schema 实际输入数构建确保一致 return recognized_inputs[:schema_inputs] print(fixed 10-input handled:, len(fuse_gqa_fixed(inputs_10, 10))) # 按 schema10buggy(10)报“数量不匹配”fixed按 schema10 正确构建。五、解决方案第一层最小直接修复最小修复融合 pass 按GroupQueryAttentionop 的实际 schema 输入数动态构建节点不再假设固定 9// 修正从 op schema 取真实输入数而非硬编码 9 auto gqa_schema ctx-getSchema(GroupQueryAttention, opset); int expected_inputs static_castint(gqa_schema.inputs().size()); // 只融合那些输入数 schema 支持的节点超出则不做融合或扩展 schema if (recognized_inputs.size() expected_inputs) { LOGS(logger, WARNING) GQA node has recognized_inputs.size() inputs, exceeds schema expected_inputs ; skip fusion to avoid invalid graph; return false; // 不融合保留原图合法 } // 构建融合节点input 列表长度 expected_inputs这一层改动最小按 schema 真实输入数校验/构建非法图消失。但依赖“每处融合都查 schema”下看第二层。六、解决方案第二层结构性改进把“GQA 融合必须按 schema 实际输入数构建、且输入数超限时安全跳过”固化成单一事实来源。下面这个 dataclass 集中管理融合契约from dataclasses import dataclass, field from typing import List dataclass class GqaFusionInputPolicy: 单一事实来源GroupQueryAttention 融合的输入数契约。 # schema 实际支持的输入名含 attention_bias 时为 10 schema_inputs: tuple (query, key, value, past_key, past_value, attention_bias, past_sequence_length, cos, sin, qk_norm) def can_fuse(self, recognized: List[str]) - bool: # 所有识别到的输入都必须在 schema 内且数量不超 schema known set(self.schema_inputs) return all(r in known for r in recognized) and \ len(recognized) len(self.schema_inputs) def build_node_inputs(self, recognized: List[str]) - List[str]: if not self.can_fuse(recognized): raise ValueError( fcannot fuse GQA with inputs {recognized}: fexceeds schema {self.schema_inputs}) # 按 schema 顺序对齐保证节点 input 列表 schema 输入数 return [r for r in self.schema_inputs if r in set(recognized)]这一层的关键收益按 schema 动态schema_inputs是真实输入定义融合按它构建输入数永远一致超限安全跳过can_fuse在超出时返回 False保留原图合法而非产出非法图顺序对齐build_node_inputs按 schema 顺序排避免错槽单一事实来源所有 GQA 融合输入约定收口在GqaFusionInputPolicy。七、解决方案第三层断言 / CI 守护把第二层钉成 pytest挂进 CI覆盖多输入 GQAimport pytest from your_package.gqa_fusion import GqaFusionInputPolicy def test_9_input_fuses(): # 断言 1标准 9 输入可融合 p GqaFusionInputPolicy() ins [query, key, value, past_key, past_value, attention_bias, past_sequence_length, cos, sin] assert p.can_fuse(ins) assert len(p.build_node_inputs(ins)) 9 def test_10_input_with_bias_fuses(): # 断言 2带 attention_bias 的 10 输入也可融合schema 已含它 p GqaFusionInputPolicy() ins list(p.schema_inputs[:10]) assert p.can_fuse(ins) assert len(p.build_node_inputs(ins)) 10 def test_unknown_input_rejected(): # 断言 3含 schema 外输入必须拒绝融合保留原图 p GqaFusionInputPolicy() assert not p.can_fuse([query, key, value, mystery_input]) def test_node_input_count_matches_schema(): # 断言 4融合节点 input 数必 schema 输入数杜绝“数量不匹配” p GqaFusionInputPolicy() ins list(p.schema_inputs[:10]) built p.build_node_inputs(ins) assert len(built) len([s for s in p.schema_inputs if s in set(ins)])四条断言从“9 输入可融合”“10 输入带 bias 可融合”“未知输入拒绝”“数量匹配 schema”四面把非法图回归钉死在 CI。八、排查清单GQA 融合报“input arg count ! size of input defs”时被融合的 GQA 节点输入数是否 9如带了attention_bias是就确认融合逻辑是否硬编码了 9。融合产物的 input 列表长度是否等于 op schema 声明的输入数不等必校验失败现象 A。图校验被关时是否“吞掉”多出来的输入导致静默错开校验跑一遍确认现象 C。用第二层GqaFusionInputPolicy按 schema 动态构建 超限安全跳过。加第三层 pytest断言“9 输入可融合、10 输入带 bias 可融合、未知输入拒绝、数量匹配 schema”。任何图融合 pass 都必须按 op schema 真实输入数构建节点不能假设固定长度。九、小结GroupQueryAttentionFusion的非法图 bug 本质是融合 pass 对 GQA 输入数做了固定假设硬编码 9当节点带attention_bias达到 10 个输入时要么丢输入、要么产出“input 列表长度 ≠ schema 输入数”的非法节点图校验报“sum of input arg count is not equal to size of input defs”且校验被关时还会静默算错。修复分三层——第一层按 op schema 真实输入数动态校验/构建超限时安全跳过保留原图第二层用GqaFusionInputPolicy这个 dataclass 把融合输入契约收口成单一事实来源按 schema 顺序对齐第三层用四条 pytest 把“9 输入可融合、10 输入带 bias 可融合、未知输入拒绝、数量匹配 schema”钉死在 CI。核心心法图融合 pass 必须按 op schema 的真实输入数构建节点绝不能假设固定输入长度超出时必须安全跳过而非产出非法图。