已关闭
[Bug-Report|缺陷反馈]: extract_api_params 按 api_name 前缀判断有无 schema,指向同一 OpOverloadPacket 的其他写法静默只取到 1 个参数 #122
qianzehong创建于  15 天前关闭于  12 天前
qianzehong
15 天前 创建

Describe the current behavior / 问题描述 (Mandatory / 必填)

extract_api_params 判断"这个 api 有没有权威 FunctionSchema"用的是 api_name 字符串前缀,而不是解析出来的对象本身。结果是:指向同一个 OpOverloadPacket 的另一种合法写法拿不到参数,且失败得很安静——不报错,只是返回 1 个假参数,最终表现为 INPUT_COUNT_EXCEEDED

走查 HEAD a862aaettk/utilities/simple_param_extractor.py

# 1877 行 _extract_api_params_impl
if api_name.startswith('torch.ops.'):
    result = _extract_params_from_aten_schemas(api_name)
# 1136 行 _extract_params_from_aten_schemas —— 它自己已经在按对象判定
obj = _resolve_function(api_name)
if obj is None or not hasattr(obj, '_schemas'):
    return None

1136 行的 hasattr(obj, '_schemas') 才是正确判据;1877 行的前缀门控是冗余的,而且过严——名字不以 torch.ops. 开头就根本不进这个分支,即使 _resolve_function 明明能解析出带 _schemas 的对象。

此时会退到 _resolve_function(2239 行)的通用路径:

mod_name = '.'.join(parts[:-1])
func_name = parts[-1]
return getattr(importlib.import_module(mod_name), func_name, None)

对于自定义算子包,这条路会拿到 OpOverloadPacket,随后 docstring / TypeError / __annotations__ 三种猜法对它全部失效,最后从 __annotations__ 里捞出一个 ['__file__']

即便把 1877 行的门控去掉,1152 行仍会崩:

namespace = api_name.split('.')[2]      # 写死按 torch.ops.<ns>.<op> 四段拆名

两段式的名字直接 IndexError: list index out of range

Steps to reproduce the issue / 复现步骤 (Mandatory / 必填)

环境:Ascend950PR_9579,CANN 9.2.0,ops-transformer master 构建的 torch extension wheel。

import torch, torch_npu, cann_ops_transformer
from ttk.utilities.simple_param_extractor import (
    extract_api_params, _resolve_function, _extract_params_from_aten_schemas)

top  = getattr(cann_ops_transformer, "und_gen_qkv_rms_norm_rope_cache")
tops = torch.ops.cann_ops_transformer.und_gen_qkv_rms_norm_rope_cache
print(top is tops)                       # True —— 同一个对象
print(list(top._schemas.keys()))         # ['']  —— schema 就挂在上面

for n in ("torch.ops.cann_ops_transformer.und_gen_qkv_rms_norm_rope_cache",
          "cann_ops_transformer.und_gen_qkv_rms_norm_rope_cache"):
    o = _resolve_function(n)
    print(n, type(o).__name__, hasattr(o, "_schemas"), len(extract_api_params(n).params))

_extract_params_from_aten_schemas("cann_ops_transformer.und_gen_qkv_rms_norm_rope_cache")

实测输出:

True
['']
torch.ops.cann_ops_transformer.und_gen_qkv_rms_norm_rope_cache  OpOverloadPacket True  17
cann_ops_transformer.und_gen_qkv_rms_norm_rope_cache            OpOverloadPacket True   1
IndexError: list index out of range

即:两个名字解析出的是同一个对象、hasattr(_schemas) 都为 True,取参结果却是 17 vs 1。

取到的那 1 个参数是 ['__file__']source__annotations__

Describe the expected behavior / 期望结果 (Mandatory / 必填)

判据改成看对象而不是看名字:凡是 _resolve_function 能解析出带 _schemas 的对象,就走权威 schema 那条路,与 api_name 怎么写无关。

具体两处:

  1. 1877 行:去掉 api_name.startswith('torch.ops.') 门控,直接调 _extract_params_from_aten_schemas,由它内部 1136 行的 hasattr(obj, '_schemas') 自行短路返回 None
  2. 1152 行namespace 不再从 api_name 拆,改从 schema 对象自身取(schema.name 形如 cann_ops_transformer::<op>),或退化成 api_name.rsplit('.', 1)[0],避免写死四段结构。

影响面

不只是"换个写法更方便"的问题——cann/ops-transformer 仓内已有算子踩到

posembedding/qkv_rms_norm_rope_cache_with_k_scale/tests/assets/qkv_rms_norm_rope_cache_with_k_scale_e2e_golden.py 注册了两个 api_name:

"cann_ops_transformer.qkv_rms_norm_rope_cache_with_k_scale":  "...FunctionalTestSpec",
"cann_ops_transformer.qkv_rms_norm_rope_cache_with_k_scale_": "...InplaceTestSpec",

该算子的 wrapper 用 get_as_library() 注册了 schema,所以顶层名字同样被 OpOverloadPacket 顶掉。实测这两个 api_name 取参都只得到 1 个 ['__file__'],e2e 跑起来必然 INPUT_COUNT_EXCEEDED

另外这个写法在 ops-transformer 侧是有依据的:torch_extension/README.md 的 Quick Start 给的就是 cann_ops_transformer.ops.<op>,而 30+ 个算子文档与 docs/zh/torch_api_list.md 给的是顶层 cann_ops_transformer.<op>。用户按文档写 api_name 是很自然的事,而现在这么写会静默失败。

附:一个小建议

1 个参数且名为 __file__ 这种结果,基本可以确定是解析失败而非真实签名。若修复成本较高,至少可以在 extract_api_params 返回前加一条告警(例如参数名里出现 dunder、或参数数为 1 且对象是 OpOverloadPacket),避免用户拿着 INPUT_COUNT_EXCEEDED 往自己的用例上找原因。

likedislike
RuiWang_成员
13 天前 评论:

您好,感谢您的反馈,当前问题在处理中

likedislike
CANN-robotCANN-robot成员
12 天前 关闭了 issue
CANN-robotCANN-robot成员
12 天前 添加了label:resolved