A CPU 0-D float64 tensor used together with an NPU tensor can be incorrectly lowered into a fused NPU kernel after device-based lowering dispatch.
The CPU scalar is passed to the generated kernel through .item():
Here, in_ptr1 is the value returned by arg2_1.item().
Expected Behavior
The CPU float64 scalar should be converted through the fallback path before the NPU kernel is launched. The NPU kernel should receive a supported float32 scalar and should not perform float64 computation internally.
Actual Behavior
Because the input tensor is on CPU, the community lowering is selected for prims.convert_element_type. The conversion is then fused with the NPU pointwise operation, leaving a float64 scalar in the generated kernel.
This causes NPU compilation to fail because the kernel contains an unsupported float64 input/conversion path.
Root Cause
The device-dispatch logic selects lowering based on the input tensor device. It does not account for the special case where a CPU 0-D float64 tensor is used in an NPU-containing graph and later becomes a scalar kernel argument through .item().
在提交新问题之前,请确保您已经在社区中搜索过相关问题,并使用了社区中提供的资源/工具后,仍未找到满意的解决方式。
⚠️ 安全信息提醒:请仔细检查提供的文本内容,确保其不包含敏感数据信息,包括但不限于:
在分享配置信息或代码示例时,请将敏感信息脱敏处理,或使用
<TOKEN>等占位符替代原有内容。环境信息
A5
torch_npu: master
🐛 问题描述
Description
A CPU 0-D float64 tensor used together with an NPU tensor can be incorrectly lowered into a fused NPU kernel after device-based lowering dispatch.
The CPU scalar is passed to the generated kernel through .item():
arg1_1 = rand_strided( (6528,), (1,), device="npu:0", dtype=torch.bfloat16 ) arg2_1 = rand_strided( (), (), device="cpu", dtype=torch.float64 ) triton_poi_fused_0.run(arg1_1, arg2_1.item(), ...)The generated kernel contains the conversion:
Here, in_ptr1 is the value returned by arg2_1.item().
Expected Behavior
The CPU float64 scalar should be converted through the fallback path before the NPU kernel is launched. The NPU kernel should receive a supported float32 scalar and should not perform float64 computation internally.
Actual Behavior
Because the input tensor is on CPU, the community lowering is selected for prims.convert_element_type. The conversion is then fused with the NPU pointwise operation, leaving a float64 scalar in the generated kernel.
This causes NPU compilation to fail because the kernel contains an unsupported float64 input/conversion path.
Root Cause
The device-dispatch logic selects lowering based on the input tensor device. It does not account for the special case where a CPU 0-D float64 tensor is used in an NPU-containing graph and later becomes a scalar kernel argument through .item().
欢迎加入社区,感谢您对社区的贡献 🎉!