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。
Describe the current behavior / 问题描述 (Mandatory / 必填)
extract_api_params判断"这个 api 有没有权威 FunctionSchema"用的是 api_name 字符串前缀,而不是解析出来的对象本身。结果是:指向同一个OpOverloadPacket的另一种合法写法拿不到参数,且失败得很安静——不报错,只是返回 1 个假参数,最终表现为INPUT_COUNT_EXCEEDED。走查 HEAD
a862aae,ttk/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 None1136 行的
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")实测输出:
即:两个名字解析出的是同一个对象、
hasattr(_schemas)都为 True,取参结果却是 17 vs 1。取到的那 1 个参数是
['__file__'],source为__annotations__。Describe the expected behavior / 期望结果 (Mandatory / 必填)
判据改成看对象而不是看名字:凡是
_resolve_function能解析出带_schemas的对象,就走权威 schema 那条路,与 api_name 怎么写无关。具体两处:
api_name.startswith('torch.ops.')门控,直接调_extract_params_from_aten_schemas,由它内部 1136 行的hasattr(obj, '_schemas')自行短路返回None。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往自己的用例上找原因。