已合并
fix: support custom torch extension package #193
RuiWang_创建于 8月27日
fix: support custom torch extension package #193
已合并
RuiWang_创建于 8月27日
共 15 个文件变更+886-706
@@ -153,7 +153,7 @@ python3 -m ttk e2e --backend npusim -i cases.csv -t add_f32_01
153python3 -m ttk kernel -i examples/case_store/kernel/add.csv153python3 -m ttk kernel -i examples/case_store/kernel/add.csv
154python3 -m ttk geir -i examples/case_store/kernel/add.csv154python3 -m ttk geir -i examples/case_store/kernel/add.csv
155python3 -m ttk aclnn -i examples/case_store/aclnn/aclnn_cat.csv155python3 -m ttk aclnn -i examples/case_store/aclnn/aclnn_cat.csv
156-python3 -m ttk e2e -i examples/case_store/e2e/torch_add.csv156+python3 -m ttk e2e -i examples/case_store/e2e/torch_ops.csv
157 157 
158# 调试单个失败用例158# 调试单个失败用例
159python3 -m ttk kernel -i cases.csv -t case_name --dump-on-fail --single-log159python3 -m ttk kernel -i cases.csv -t case_name --dump-on-fail --single-log
@@ -38,7 +38,7 @@ python3 -m ttk kernel --backend npusim -i cases.csv -t add_01 --sim-cores 0 --si
38python3 -m ttk aclnn --backend npusim -i examples/case_store/aclnn/aclnn_add.csv --sim-cores 038python3 -m ttk aclnn --backend npusim -i examples/case_store/aclnn/aclnn_add.csv --sim-cores 0
39 39 
40# E2E(eager)额外出流水图40# E2E(eager)额外出流水图
41-python3 -m ttk e2e --backend npusim -i examples/case_store/e2e/torch_add.csv -t add_f32_01 --sim-report41+python3 -m ttk e2e --backend npusim -i examples/case_store/e2e/torch_ops.csv -t add_f32_01 --sim-report
42```42```
43 43 
44## 产物44## 产物
@@ -19,6 +19,7 @@ E2E 用例的参数(张量数量/顺序/dtype、关键字参数等)来自框
19| `torch.add` | 模块函数 | 直接调用 |19| `torch.add` | 模块函数 | 直接调用 |
20| `torch.nn.functional.relu` | 子模块函数 | 调用子模块中的函数 |20| `torch.nn.functional.relu` | 子模块函数 | 调用子模块中的函数 |
21| `torch.Tensor.relu_` | Tensor 方法 | 通过 Tensor 实例调用(原地操作) |21| `torch.Tensor.relu_` | Tensor 方法 | 通过 Tensor 实例调用(原地操作) |
22+| `torch.ops.cann_ops_transformer.causal_conv1d_fn` | torch.ops 自定义算子 | torch extension 包注册的算子 |
22 23 
23## 必填字段24## 必填字段
24 25 
@@ -29,6 +29,8 @@ description: 编写/定义算子测试规范(TestSpec,尚未跑或要新增
29 29 
30> 注册名即 `get_spec_attr` 的查找 key:Kernel/GEIR 用 `op_name`,ACLNN/E2E 用 `api_name`。GEIR 复用 Kernel 的 spec,无需额外编写。30> 注册名即 `get_spec_attr` 的查找 key:Kernel/GEIR 用 `op_name`,ACLNN/E2E 用 `api_name`。GEIR 复用 Kernel 的 spec,无需额外编写。
31 31 
32+> **torch.ops 自定义算子(torch extension 包,如 `cann_ops_transformer`)**:E2E 的 `api_name` 写 4 段 `torch.ops.<ns>.<op>`,参数来自算子自带 schema(非 torch API 签名),**无需手写签名**。详见 `ttk-how-write-case` 的 E2E 引用「torch.ops 自定义算子」。
33+ 
32```python34```python
33__spec__ = {"abs": "AbsTestSpec"}35__spec__ = {"abs": "AbsTestSpec"}
34 36 
@@ -199,6 +201,7 @@ python3 -m ttk aclnn -i cases.csv --plugin /path/to/assets_a/,/path/to/assets_b/
1992. **参数顺序**:函数参数名和顺序与算子定义文件(def.cpp / aclnn*.h)中的输入参数一致2012. **参数顺序**:函数参数名和顺序与算子定义文件(def.cpp / aclnn*.h)中的输入参数一致
2003. **kwargs 始终接收**:通过 `**kwargs` 接收元信息(dtypes、shapes、formats、soc_version 等,完整字段见 `references/kernel-plugin.md` / `aclnn-plugin.md`)2023. **kwargs 始终接收**:通过 `**kwargs` 接收元信息(dtypes、shapes、formats、soc_version 等,完整字段见 `references/kernel-plugin.md` / `aclnn-plugin.md`)
2014. **返回类型**:`golden` 返回列表(每个输出一个元素);`customize_inputs` 返回与输入结构一致的 tuple2034. **返回类型**:`golden` 返回列表(每个输出一个元素);`customize_inputs` 返回与输入结构一致的 tuple
204+5. **禁止 `register()` 手写签名**:不要在 golden 文件里 `from ttk...import ParamInfo/APIParamInfo/FrameworkApiInfoKeeper` 再调 `register()` 声明算子签名——签名应来自算子自身(E2E 的 `torch.ops` 算子从 `_schemas` 自动解析;Kernel/GEIR 从 `def.cpp`;ACLNN 从 `aclnn*.h`)。手写 `register()` 会把 golden 耦合到 ttk 内部类(ttk 改内部类时所有此类 golden 失效),且与算子真实签名漂移。`register()` 仅在算子**无 schema**(裸 `impl`、未 `define`)时作兜底。
202 205 
203## 调试206## 调试
204 207 
@@ -57,10 +57,10 @@ python3 -m ttk kernel --backend npusim -i examples/case_store/kernel/add.csv --s
57python3 -m ttk aclnn --backend npusim -i examples/case_store/aclnn/aclnn_add.csv --sim-cores 057python3 -m ttk aclnn --backend npusim -i examples/case_store/aclnn/aclnn_add.csv --sim-cores 0
58 58 
59# E2E 模式(torch_npu,eager 执行)59# E2E 模式(torch_npu,eager 执行)
60-python3 -m ttk e2e --backend npusim -i examples/case_store/e2e/torch_add.csv -t add_f32_0160+python3 -m ttk e2e --backend npusim -i examples/case_store/e2e/torch_ops.csv -t add_f32_01
61 61 
62# E2E 模式 + 额外生成仿真流水图62# E2E 模式 + 额外生成仿真流水图
63-python3 -m ttk e2e --backend npusim -i examples/case_store/e2e/torch_add.csv -t add_f32_01 --sim-report63+python3 -m ttk e2e --backend npusim -i examples/case_store/e2e/torch_ops.csv -t add_f32_01 --sim-report
64 64 
65# 额外生成仿真流水图(Kernel / ACLNN)65# 额外生成仿真流水图(Kernel / ACLNN)
66python3 -m ttk kernel --backend npusim -i examples/case_store/kernel/add.csv --sim-cores 0 --sim-report66python3 -m ttk kernel --backend npusim -i examples/case_store/kernel/add.csv --sim-cores 0 --sim-report
@@ -11,6 +11,7 @@
11| `torch.add` | 模块函数 | 直接调用 |11| `torch.add` | 模块函数 | 直接调用 |
12| `torch.nn.functional.relu` | 子模块函数 | 调用子模块中的函数 |12| `torch.nn.functional.relu` | 子模块函数 | 调用子模块中的函数 |
13| `torch.Tensor.relu_` | Tensor方法 | 通过Tensor实例调用(原地操作) |13| `torch.Tensor.relu_` | Tensor方法 | 通过Tensor实例调用(原地操作) |
14+| `torch.ops.cann_ops_transformer.causal_conv1d_fn` | torch.ops 自定义算子 | torch extension 包注册的算子 |
14 15 
15## 用例标识16## 用例标识
16 17 
@@ -37,16 +38,6 @@
37| `attributes` | DICT | 否 | `{}` | 框架API关键字参数。如 `{'alpha': 1.0}`。API签名中的必选非张量参数**必须**提供 |38| `attributes` | DICT | 否 | `{}` | 框架API关键字参数。如 `{'alpha': 1.0}`。API签名中的必选非张量参数**必须**提供 |
38| `golden_api` | STRING | 否 | `""` | 替代Golden计算的API。如 `torch.nn.functional.conv2d`。设为 `"disable"` 可禁用Golden生成 |39| `golden_api` | STRING | 否 | `""` | 替代Golden计算的API。如 `torch.nn.functional.conv2d`。设为 `"disable"` 可禁用Golden生成 |
39 40 
40-## 参考用例
41- 
42-`examples/case_store/e2e/` 目录下提供了各种场景的示例:
43- 
44-| 文件 | 涵盖场景 |
45-|------|----------|
46-| `torch_add.csv` | 基本用例、`torch.add` |
47-| `torch_npu_conv2d.csv` | `torch_npu.npu_conv2d`、`golden_api` |
48-| `tf_ops.csv` | TensorFlow API |
49- 
50## 批一致性字段41## 批一致性字段
51 42 
52配合 `--deterministic-level 3` 使用,用于跨用例输出切片比对。详见 [确定性计算与批一致性](../Deterministic_Compute.md)。43配合 `--deterministic-level 3` 使用,用于跨用例输出切片比对。详见 [确定性计算与批一致性](../Deterministic_Compute.md)。
@@ -82,11 +73,16 @@ full,torch.add,"((10,8),)",(100,),,,...
82 73 
83## 参考用例74## 参考用例
84 75 
85-`examples/case_store/e2e/` 目录下的示例:76+`examples/case_store/e2e/` 目录下的示例(torch 用例均在 `torch_ops.csv`,按场景分行):
86 77 
87-| 文件 | 框架 | 验证特性 | 关键列 |78+| 用例 | 框架 | 验证特性 | 关键列 |
88|------|------|---------|--------|79|------|------|---------|--------|
89-| `tf_ops.csv` | TensorFlow | 多 API(tf.raw_ops/nn/math/linalg)+ 多算子 | `tensor_view_shapes`、`attributes` |80+| `tf_ops.csv`(多行) | TensorFlow | 多 API(tf.raw_ops/nn/math/linalg)+ 多算子 | `tensor_view_shapes`、`attributes` |
90-| `torch_add.csv` | torch | add/abs/relu/mm + inplace(`relu_`)/out 变体 + alpha 属性 | `attributes`、`output_tensor_indexes` |81+| `torch_ops.csv` — `add_f32_01`/`add_f16_01`/`add_broadcast` | torch | `torch.add` 基本变体(f32/f16/broadcast) | `attributes`(alpha) |
91-| `torch_npu_conv2d.csv` | torch_npu | NPU 专属 API + 自定义 golden | `golden_api`、`output_tensor_indexes` |82+| `torch_ops.csv` — `abs_01` | torch | `torch.abs` 单输入 | — |
83+| `torch_ops.csv` — `relu_01` | torch | `torch.nn.functional.relu` | `attributes`(inplace) |
84+| `torch_ops.csv` — `relu_inp` | torch | `torch.Tensor.relu_` 原地(自动 `output_tensor_indexes=(0,)`) | `output_tensor_indexes` |
85+| `torch_ops.csv` — `add_out`/`mm_out` | torch | `torch.add`/`torch.mm` 带 `out` 输出 | `output_tensor_indexes` |
86+| `torch_ops.csv` — `npu_conv2d_f16`/`npu_conv2d_f32` | torch_npu | `torch_npu.npu_conv2d` + `golden_api` 替代 golden(需设备支持) | `golden_api` |
87+| `torch_ops.csv` — `causal_conv1d_fn_basic` | torch.ops | `torch.ops` 自定义算子,torchops-自定义算子torch-extension-包),需 `--plugin` golden | `api_name` |
92| `torch_add.xlsx` | torch | xlsx 多 sheet(T1/T2)输入验证 | — |88| `torch_add.xlsx` | torch | xlsx 多 sheet(T1/T2)输入验证 | — |
@@ -93,10 +93,10 @@ E2E 模式由统一的 Backend 抽象层(`ttk.core_modules.framework_api.backe
93 93 
94```shell94```shell
95# 自动按配置 hardware segment 选择可用后端95# 自动按配置 hardware segment 选择可用后端
96-python3 -m ttk e2e -i torch_add.csv96+python3 -m ttk e2e -i torch_ops.csv
97 97 
98-# 强制 CPU 后端(常用于Golden生成)98+# 强制 CPU 后端(常用于Golden生成;选 CPU 可跑的用例)
99-python3 -m ttk e2e -i torch_add.csv --cpu99+python3 -m ttk e2e -i torch_ops.csv -t add_f32_01 --cpu
100```100```
101 101 
102## 执行流程102## 执行流程
@@ -115,13 +115,13 @@ E2E 模式默认会采集Profiling性能数据:
115 115 
116```shell116```shell
117# 默认执行(含Profiling)117# 默认执行(含Profiling)
118-python3 -m ttk e2e -i torch_add.csv118+python3 -m ttk e2e -i torch_ops.csv
119 119 
120# 禁用Profiling120# 禁用Profiling
121-python3 -m ttk e2e -i torch_add.csv --no-prof121+python3 -m ttk e2e -i torch_ops.csv --no-prof
122 122 
123# 设置执行次数(默认板端3次)123# 设置执行次数(默认板端3次)
124-python3 -m ttk e2e -i torch_add.csv --run=5124+python3 -m ttk e2e -i torch_ops.csv --run=5
125```125```
126 126 
127## 性能相关参数127## 性能相关参数
@@ -136,29 +136,29 @@ python3 -m ttk e2e -i torch_add.csv --run=5
136 136 
137```shell137```shell
138# 使用全部可用NPU卡138# 使用全部可用NPU卡
139-python3 -m ttk e2e -i torch_add.csv139+python3 -m ttk e2e -i torch_ops.csv
140 140 
141# 使用2张卡141# 使用2张卡
142-python3 -m ttk e2e -i torch_add.csv --dev=2142+python3 -m ttk e2e -i torch_ops.csv --dev=2
143 143 
144# 每张卡2个进程144# 每张卡2个进程
145-python3 -m ttk e2e -i torch_add.csv --pc=2145+python3 -m ttk e2e -i torch_ops.csv --pc=2
146 146 
147# 指定使用卡0147# 指定使用卡0
148-python3 -m ttk e2e -i torch_add.csv --device-whitelist=0148+python3 -m ttk e2e -i torch_ops.csv --device-whitelist=0
149```149```
150 150 
151# 调试151# 调试
152 152 
153```shell153```shell
154# 调试单个用例154# 调试单个用例
155-python3 -m ttk e2e -i torch_add.csv -t add_f32_01 --single-log155+python3 -m ttk e2e -i torch_ops.csv -t add_f32_01 --single-log
156 156 
157# 固定随机种子(可复现)157# 固定随机种子(可复现)
158-python3 -m ttk e2e -i torch_add.csv --seed 42158+python3 -m ttk e2e -i torch_ops.csv --seed 42
159 159 
160# 仅校验CSV用例格式(不下设备)160# 仅校验CSV用例格式(不下设备)
161-python3 -m ttk e2e -i torch_add.csv --validate161+python3 -m ttk e2e -i torch_ops.csv --validate
162```162```
163 163 
164Dump 调试详见[Dump 数据调试](../Dump_Debug.md)。164Dump 调试详见[Dump 数据调试](../Dump_Debug.md)。
@@ -167,16 +167,16 @@ Dump 调试详见[Dump 数据调试](../Dump_Debug.md)。
167 167 
168```shell168```shell
169# torch.add 基础测试(自动选择可用后端)169# torch.add 基础测试(自动选择可用后端)
170-python3 -m ttk e2e -i examples/case_store/e2e/torch_add.csv170+python3 -m ttk e2e -i examples/case_store/e2e/torch_ops.csv -t add_f32_01
171 171 
172-# torch_npu.npu_conv2d(使用golden_api)172+# torch_npu.npu_conv2d(使用golden_api,需设备支持该算子)
173-python3 -m ttk e2e -i examples/case_store/e2e/torch_npu_conv2d.csv173+python3 -m ttk e2e -i examples/case_store/e2e/torch_ops.csv -t npu_conv2d_f16
174 174 
175-# 强制 CPU 后端175+# 强制 CPU 后端(选 CPU 可跑的用例)
176-python3 -m ttk e2e -i examples/case_store/e2e/torch_add.csv --cpu176+python3 -m ttk e2e -i examples/case_store/e2e/torch_ops.csv -t add_f32_01 --cpu
177 177 
178# 输出结果178# 输出结果
179-python3 -m ttk e2e -i torch_add.csv -o results.csv179+python3 -m ttk e2e -i torch_ops.csv -o results.csv
180 180 
181# Excel 多 sheet 用例(默认首个工作表;--sheet 指定工作表)181# Excel 多 sheet 用例(默认首个工作表;--sheet 指定工作表)
182python3 -m ttk e2e -i examples/case_store/e2e/torch_add.xlsx182python3 -m ttk e2e -i examples/case_store/e2e/torch_add.xlsx
@@ -1,3 +0,0 @@
1-testcase_name,api_name,tensor_view_shapes,tensor_dtypes,tensor_formats,attributes,output_tensor_indexes,golden_api
2-npu_conv2d_f16,torch_npu.npu_conv2d,"((1,3,224,224),(64,3,7,7),(64,))","('float16','float16','float16')",,"{'stride':[2,2],'padding':[3,3],'dilation':[1,1],'groups':1}",,torch.nn.functional.conv2d
3-npu_conv2d_f32,torch_npu.npu_conv2d,"((1,3,32,32),(16,3,3,3),(16,))","('float32','float32','float32')",,"{'stride':[1,1],'padding':[1,1],'dilation':[1,1],'groups':1}",,torch.nn.functional.conv2d
Rexamples/case_store/e2e/torch_add.csv→examples/case_store/e2e/torch_ops.csv+4-1
@@ -1,4 +1,4 @@
1-testcase_name,api_name,tensor_view_shapes,tensor_dtypes,tensor_formats,attributes,output_tensor_indexes1+testcase_name,api_name,tensor_view_shapes,tensor_dtypes,tensor_formats,attributes,output_tensor_indexes,golden_api
2add_f32_01,torch.add,"((2,3,4),(2,3,4))","('float32','float32')",,{'alpha':1.0},2add_f32_01,torch.add,"((2,3,4),(2,3,4))","('float32','float32')",,{'alpha':1.0},
3add_f16_01,torch.add,"((128,256),(128,256))","('float16','float16')",,{'alpha':1.0},3add_f16_01,torch.add,"((128,256),(128,256))","('float16','float16')",,{'alpha':1.0},
4add_broadcast,torch.add,"((4,1,3),(1,5,3))","('float32','float32')",,{'alpha':2.0},4add_broadcast,torch.add,"((4,1,3),(1,5,3))","('float32','float32')",,{'alpha':2.0},
@@ -7,3 +7,6 @@ relu_01,torch.nn.functional.relu,"((64,128),)","('float16',)",,{'inplace':False}
7relu_inp,torch.Tensor.relu_,"((2,3),)","('float32',)",,,7relu_inp,torch.Tensor.relu_,"((2,3),)","('float32',)",,,
8add_out,torch.add,"((2,3),(2,3),(2,3))","('float32','float32','float32')",,{'alpha':1.0},"(2,)"8add_out,torch.add,"((2,3),(2,3),(2,3))","('float32','float32','float32')",,{'alpha':1.0},"(2,)"
9mm_out,torch.mm,"((2,3),(3,4),(2,4))","('float32','float32','float32')",,{},"(2,)"9mm_out,torch.mm,"((2,3),(3,4),(2,4))","('float32','float32','float32')",,{},"(2,)"
10+npu_conv2d_f16,torch_npu.npu_conv2d,"((1,3,224,224),(64,3,7,7),(64,))","('float16','float16','float16')",,"{'stride':[2,2],'padding':[3,3],'dilation':[1,1],'groups':1}",,torch.nn.functional.conv2d
11+npu_conv2d_f32,torch_npu.npu_conv2d,"((1,3,32,32),(16,3,3,3),(16,))","('float32','float32','float32')",,"{'stride':[1,1],'padding':[1,1],'dilation':[1,1],'groups':1}",,torch.nn.functional.conv2d
12+causal_conv1d_fn_basic,torch.ops.cann_ops_transformer.causal_conv1d_fn,"((2,4,16),(2,16),(16,),(2,1,16),(3,))","('float16','float16','float16','float16','int32')",,{},,
@@ -376,34 +376,6 @@ def _detect_framework_from_csv(input_files, sheet=None):
376 return first_framework376 return first_framework
377 377 
378 378 
379-def _preload_plugin_modules(sw):
380- """Import all .py modules under --plugin dirs before testcase validation.
381- 
382- golden.py modules call FrameworkApiInfoKeeper().register() at import time
383- to declare API param info for OpOverloadPacket APIs that can't be parsed
384- via __annotations__. Without this preload, validation runs before the
385- plugin is lazily loaded, causing INPUT_COUNT_EXCEEDED false failures.
386- """
387- if not sw.plugin_path:
388- return
389- import importlib.util
390- 
391- for plugin_dir in sw.plugin_path:
392- plugin_path = pathlib.Path(plugin_dir)
393- if not plugin_path.is_dir():
394- continue
395- for py_file in plugin_path.glob("*.py"):
396- if py_file.name.startswith("_"):
397- continue
398- mod_name = py_file.stem
399- try:
400- spec = importlib.util.spec_from_file_location(mod_name, str(py_file))
401- mod = importlib.util.module_from_spec(spec)
402- spec.loader.exec_module(mod)
403- except Exception as e:
404- logging.debug(f"Preload plugin module {py_file} skipped: {e}")
405- 
406- 
407def run_with_switches(sw):379def run_with_switches(sw):
408 from ttk.core_modules.tbe_logging import default_logging_config380 from ttk.core_modules.tbe_logging import default_logging_config
409 from ttk.utilities import set_global_storage381 from ttk.utilities import set_global_storage
@@ -419,8 +391,6 @@ def run_with_switches(sw):
419 391 
420 logging.info(f"Command: ttk {sw.test_mode} -i {sw.input_files[0] if sw.input_files else ''}")392 logging.info(f"Command: ttk {sw.test_mode} -i {sw.input_files[0] if sw.input_files else ''}")
421 393 
422- _preload_plugin_modules(sw)
423- 
424 if sw.test_mode == "framework-api":394 if sw.test_mode == "framework-api":
425 from ttk.core_modules.framework_api.instance import FrameworkApiInstance395 from ttk.core_modules.framework_api.instance import FrameworkApiInstance
426 396 
@@ -15,8 +15,6 @@ Handles both module functions (torch.add) and Tensor methods (torch.Tensor.relu_
15Also supports TF APIs (tf.raw_ops.Add, tf.nn.relu, etc.).15Also supports TF APIs (tf.raw_ops.Add, tf.nn.relu, etc.).
16"""16"""
17 17 
18-from ttk.utilities.torch_ops_package_loader import TorchOpsPackageLoader
19- 
20_MODULE_ALIAS = {}18_MODULE_ALIAS = {}
21 19 
22 20 
@@ -44,7 +42,8 @@ def resolve_api(api_name: str):
44 42 
45 return resolve_callable_str(api_name), False43 return resolve_callable_str(api_name), False
46 44 
47- TorchOpsPackageLoader.ensure_registered(api_name)45+ # The extension package was imported by FrameworkApiInfoKeeper.get during
46+ # validation in the parent process; forked workers inherit sys.modules.
48 47 
49 # torch.Tensor.xxx -> Tensor method48 # torch.Tensor.xxx -> Tensor method
50 if len(parts) >= 3 and parts[1] == "Tensor":49 if len(parts) >= 3 and parts[1] == "Tensor":
@@ -10,10 +10,10 @@ Validates testcase parameters against API signatures.
10"""10"""
11 11 
12import logging12import logging
13-from typing import Optional, Dict13+from typing import Dict, Optional
14 14 
15from ttk.utilities import Singleton15from ttk.utilities import Singleton
16-from ttk.utilities.simple_param_extractor import APIParamInfo, get_api_params, register_api_params, ParamInfo16+from ttk.utilities.simple_param_extractor import APIParamInfo, get_api_params, register_api_params
17from ttk.utilities.torch_ops_package_loader import TorchOpsPackageLoader17from ttk.utilities.torch_ops_package_loader import TorchOpsPackageLoader
18 18 
19 19 
@@ -64,6 +64,12 @@ class FrameworkApiInfoKeeper(metaclass=Singleton):
64 f"but testcase configured {tensor_count}. "64 f"but testcase configured {tensor_count}. "
65 f"(source: {info.source})"65 f"(source: {info.source})"
66 )66 )
67+ if scalar_count != api_scalar_count:
68+ return (
69+ f"API [{api_name}] has {api_scalar_count} scalar parameters, "
70+ f"but testcase configured {scalar_count}. "
71+ f"(source: {info.source})"
72+ )
67 return None73 return None
68 74 
69 def get_tensor_distribution(self, api_name: str) -> tuple:75 def get_tensor_distribution(self, api_name: str) -> tuple:
@@ -22,7 +22,6 @@ from ttk.core_modules.plugin_loader import get_plugin_function
22from ttk.utilities import get22from ttk.utilities import get
23from ttk.utilities.container_utils import apply_as_list23from ttk.utilities.container_utils import apply_as_list
24from ttk.utilities.data import RandomData, resolve_custom_numpy_dtypes24from ttk.utilities.data import RandomData, resolve_custom_numpy_dtypes
25-from ttk.utilities.torch_ops_package_loader import TorchOpsPackageLoader
26 25 
27 26 
28def generate_inputs(testcase, switches, backend, plan, stored_inputs=None):27def generate_inputs(testcase, switches, backend, plan, stored_inputs=None):
@@ -36,9 +35,6 @@ def generate_inputs(testcase, switches, backend, plan, stored_inputs=None):
36 Plugin has no return value, modifies tensor arrays in-place via x[:] = value.35 Plugin has no return value, modifies tensor arrays in-place via x[:] = value.
37 Called with same ParamPlan arg order as profiling execution and golden generation.36 Called with same ParamPlan arg order as profiling execution and golden generation.
38 """37 """
39- # Custom input runs before API resolution in worker processes. Register the
40- # installed torch.ops package before a callback calls a companion metadata op.
41- TorchOpsPackageLoader.ensure_registered(testcase.api_name)
42 if stored_inputs is not None:38 if stored_inputs is not None:
43 testcase.np_storages = list(stored_inputs)39 testcase.np_storages = list(stored_inputs)
44 raw_inputs = build_views_from_storages(testcase)40 raw_inputs = build_views_from_storages(testcase)
@@ -108,6 +104,7 @@ def np_to_torch_inputs(testcase, raw_inputs):
108 same approach as aclnn input_generation.104 same approach as aclnn input_generation.
109 """105 """
110 import torch106 import torch
107+ 
111 from ttk.utilities.dtypes import numpy_to_torch_tensor108 from ttk.utilities.dtypes import numpy_to_torch_tensor
112 109 
113 np_storages = getattr(testcase, "np_storages", None)110 np_storages = getattr(testcase, "np_storages", None)
@@ -141,6 +138,7 @@ def np_to_tf_inputs(testcase, raw_inputs):
141 Tensors sourced from attributes use tf.constant for graph-mode const folding.138 Tensors sourced from attributes use tf.constant for graph-mode const folding.
142 """139 """
143 import tensorflow as tf140 import tensorflow as tf
141+ 
144 from ttk.utilities.dtypes import normalize_to_tf_dtype142 from ttk.utilities.dtypes import normalize_to_tf_dtype
145 143 
146 const_indexes = getattr(testcase, "const_input_indexes", None) or set()144 const_indexes = getattr(testcase, "const_input_indexes", None) or set()
@@ -169,7 +167,7 @@ def assign_tensor_value(arr, val, label):
169 else:167 else:
170 arr[:] = spec168 arr[:] = spec
171 except ValueError as e:169 except ValueError as e:
172- raise ValueError(f"Specify tensor [{label}] from `attributes` fail: {e}")170+ raise ValueError(f"Specify tensor [{label}] from `attributes` fail: {e}") from e
173 171 
174 172 
175def override_tensors_from_attributes(testcase, raw_inputs):173def override_tensors_from_attributes(testcase, raw_inputs):
@@ -8,9 +8,18 @@
8# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.8# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
9# See LICENSE in the root of the software repository for the full text of the License.9# See LICENSE in the root of the software repository for the full text of the License.
10 10 
11-"""Load installed packages that register custom ``torch.ops`` namespaces."""11+"""Auto-discover and import ``torch.ops`` extension packages on demand.
12+ 
13+A ``torch.ops.<namespace>.<op>`` API lives in an installed extension package
14+whose importable name equals the namespace (the torch extension convention).
15+This loader derives the package name from the api_name itself -- there is no
16+hardcoded namespace registry -- and imports it lazily the first time a matching
17+api_name is consulted at the registry match point (``FrameworkApiInfoKeeper.get``).
18+"""
12 19 
13import importlib20import importlib
21+import importlib.machinery
22+import importlib.util
14import os23import os
15import pathlib24import pathlib
16import sys25import sys
@@ -22,76 +31,90 @@ class TorchOpsPackageRegistrationError(RuntimeError):
22 31 
23 32 
24class TorchOpsPackageLoader:33class TorchOpsPackageLoader:
25- """Register known custom ``torch.ops`` namespaces before parsing or execution."""34+ """Import ``torch.ops`` extension packages on demand, by namespace name."""
26 35 
27- NAMESPACE_PACKAGES = {
28- "cann_ops_transformer": "cann_ops_transformer",
29- "cann_ops_nn": "cann_ops_nn",
30- }
31 CANN_ROOT_ENV_VARS = (36 CANN_ROOT_ENV_VARS = (
32 "ASCEND_HOME_PATH",37 "ASCEND_HOME_PATH",
33 "ASCEND_TOOLKIT_HOME",38 "ASCEND_TOOLKIT_HOME",
34 "ASCEND_AICPU_PATH",39 "ASCEND_AICPU_PATH",
35 )40 )
36 LOCK = threading.RLock()41 LOCK = threading.RLock()
42+ _cann_dirs = None
43+ _paths_inserted = False
37 44 
38 @classmethod45 @classmethod
39- def package_for_api(cls, api_name):46+ def _extension_namespace(cls, api_name):
47+ """Return the namespace for ``torch.ops.<ns>.<op>``, else None.
48+ 
49+ Built-in torch namespaces (aten, prim, ...) have no importable package
50+ and are left for torch itself to register.
51+ """
40 parts = api_name.split(".")52 parts = api_name.split(".")
41- if len(parts) < 4 or parts[0:2] != ["torch", "ops"]:53+ if len(parts) < 4 or parts[0] != "torch" or parts[1] != "ops":
42 return None54 return None
43- return cls.NAMESPACE_PACKAGES.get(parts[2])55+ return parts[2]
44 56 
45 @classmethod57 @classmethod
46- def site_package_candidates(cls):58+ def cann_site_packages(cls):
47- candidates = []59+ """CANN ``<root>/python/site-packages`` dirs from the environment."""
48- seen = set()60+ if cls._cann_dirs is None:
49- for env_name in cls.CANN_ROOT_ENV_VARS:61+ dirs, seen = [], set()
50- root = os.environ.get(env_name)62+ for env_name in cls.CANN_ROOT_ENV_VARS:
51- if not root:63+ root = os.environ.get(env_name)
52- continue64+ if not root:
53- candidate = pathlib.Path(root).expanduser() / "python" / "site-packages"65+ continue
54- candidate = candidate.resolve()66+ candidate = pathlib.Path(root).expanduser().joinpath("python", "site-packages").resolve()
55- if candidate not in seen:67+ if candidate.is_dir() and candidate not in seen:
56- candidates.append(candidate)68+ dirs.append(str(candidate))
57- seen.add(candidate)69+ seen.add(candidate)
58- return candidates70+ cls._cann_dirs = dirs
71+ return cls._cann_dirs
59 72 
60 @classmethod73 @classmethod
61 def environment_summary(cls):74 def environment_summary(cls):
62- return ", ".join(75+ return ", ".join(f"{name}={os.environ.get(name) or '<unset>'}" for name in cls.CANN_ROOT_ENV_VARS)
63- f"{name}={os.environ.get(name) or '<unset>'}"
64- for name in cls.CANN_ROOT_ENV_VARS
65- )
66 76 
67 @classmethod77 @classmethod
68 def ensure_registered(cls, api_name):78 def ensure_registered(cls, api_name):
69- package_name = cls.package_for_api(api_name)79+ """Import the extension package for ``torch.ops.<ns>.<op>`` if needed.
70- if package_name is None:80+ 
81+ Idempotent and thread-safe. The package name is the namespace itself
82+ (torch extension convention). CANN site-packages are probed first so a
83+ namespace name never collides with an unrelated global package; a
84+ global ``find_spec`` fallback covers pip-installed extensions. Built-in
85+ namespaces (aten/prim/...) resolve to nothing and are skipped.
86+ """
87+ ns = cls._extension_namespace(api_name)
88+ if ns is None or ns in sys.modules:
71 return89 return
72 90 
73 with cls.LOCK:91 with cls.LOCK:
74- candidates = cls.site_package_candidates()92+ cann_dirs = cls.cann_site_packages()
75- inserted = []93+ spec = importlib.machinery.PathFinder.find_spec(ns, path=cann_dirs) if cann_dirs else None
76- for candidate in reversed(candidates):94+ if spec is not None:
77- candidate_str = str(candidate)95+ cls._insert_paths(cann_dirs)
78- if candidate.is_dir() and candidate_str not in sys.path:96+ else:
79- sys.path.insert(0, candidate_str)97+ spec = importlib.util.find_spec(ns)
80- inserted.append(candidate_str)98+ if spec is None:
81- importlib.invalidate_caches()99+ return # built-in namespace, or genuinely missing
82- 
83 try:100 try:
84- if package_name not in sys.modules:101+ importlib.import_module(ns)
85- importlib.import_module(package_name)
86 except Exception as error:102 except Exception as error:
87- for candidate_str in inserted:103+ checked = ", ".join(cann_dirs) or "<none>"
88- if candidate_str in sys.path:
89- sys.path.remove(candidate_str)
90- checked = ", ".join(str(path) for path in candidates) or "<none>"
91 raise TorchOpsPackageRegistrationError(104 raise TorchOpsPackageRegistrationError(
92- f"Cannot register {api_name!r}: importing installed package "105+ f"Cannot register {api_name!r}: importing torch.ops namespace "
93- f"{package_name!r} failed with {type(error).__name__}: {error}. "106+ f"package {ns!r} failed with {type(error).__name__}: {error}. "
94 f"CANN environment: {cls.environment_summary()}. "107 f"CANN environment: {cls.environment_summary()}. "
95 f"Checked site-packages: {checked}. Source the target CANN "108 f"Checked site-packages: {checked}. Source the target CANN "
96- "environment and install the matching transformer package."109+ "environment and install the matching extension package."
97 ) from error110 ) from error
111+ 
112+ @classmethod
113+ def _insert_paths(cls, cann_dirs):
114+ if cls._paths_inserted:
115+ return
116+ for candidate in reversed(cann_dirs):
117+ if candidate not in sys.path:
118+ sys.path.insert(0, candidate)
119+ importlib.invalidate_caches()
120+ cls._paths_inserted = True