修复 KunlunXIN XPU 上开启 FlagGems 后,非连续 source(例如 expand 出来的 SDPA mask)触发:
torch.AcceleratorError: CUDA error: invalid device function
一句话版本:原 KunlunXIN copy_ 把非连续 source 交给原生 aten.copy_ fallback;当前 XPU 兼容层的原生非连续 elementwise kernel 不可用。这里移除 source 必须连续的限制,让 source / destination 的真实 strides 直接交给已有 pointwise_dynamic copy kernel。
最小复现:
import torch
import flag_gems
flag_gems.enable()
a = torch.randn(128, device="cuda")
e = a.expand(4, 128)
c = torch.empty((4, 128), device="cuda")
c.copy_(e) # 失败
h = torch.empty((4, 128), dtype=torch.bfloat16, device="cuda")
h.copy_(e) # fp32 -> bf16,同样会走问题路径原实现里:
if not src.is_contiguous():
return False导致非连续 source 进入:
torch.ops.aten.copy_.default.redispatch(...)在当前 torch_xmlir / KunlunXIN 环境中,该原生非连续 copy kernel 最终报 invalid device function。
分流结构:
source / destination 是否可安全走 Triton?
├── 是 --> _copy_kernel.instantiate(rank)
└── 否 --> 原 aten.copy_ fallback
pointwise_dynamic 会按 rank 生成 kernel,并携带 source / destination 的真实
stride 展开 offset,因此不会把非连续 tensor 误当作连续平铺内存;同时沿用
KunlunXIN 后端的 12-CTA grid/tile 策略。
因此 expand、transpose、slice、permute、复合 view、非连续 destination,以及高 rank 输入都走同一条已验证路径。
与 gatherFIX 对齐:
├── README.md # 本文
├── copy.patch # FlagGems 最小算子 patch(只改 copy.py)
├── docs/
│ ├── COPY_REPORT.md # 根因、实现与测试报告
│ └── PERFORMANCE_ACCURACY.md # P800 性能 / 精度对比
├── results/
│ └── copy_perf_accuracy_p800.json
├── src/
│ ├── copy.py # 修复后的完整 KunlunXIN copy_ 文件
│ └── copy.py.orig # 修改前备份
└── tests/
├── repro_copy.py # 原问题最小复现
├── test_copy_standalone.py # 直连算子函数,16 用例,不经过 torch dispatch
├── test_rank_coverage.py # rank 1-6 与 destination stride 覆盖
├── gap_checks.py # empty / rank6 /非法 broadcast / overlap 写入
├── test_sdpa_e2e.py # expanded mask 进入真实 SDPA 计算
├── bench_copy.py # 性能 / 精度矩阵 benchmark
└── test_copy.py # FlagGems pytest 集成测试
cp src/copy.py \
/env/FlagGems/src/flag_gems/runtime/backend/_kunlunxin/ops/copy.py
find /env/FlagGems/src/flag_gems/runtime/backend/_kunlunxin \
-name __pycache__ -type d -exec rm -rf {} +cd /env/FlagGems
git apply /workspace/copyop/copy.patchcopy.patch 与 gatherFIX 的 gather.patch 一样,只包含算子实现的最小 diff;tests/test_copy.py 是可单独复制回 FlagGems 测试目录的集成测试。
cd tests
./repro_copy.py
./test_copy_standalone.py # 16 passed
./test_rank_coverage.py # 8 passed
./gap_checks.py
./test_sdpa_e2e.py这些 standalone 脚本都带 shebang 和可执行位,也可以继续用
python3 tests/xxx.py 执行。
FlagGems pytest 集成测试:
cd /env/FlagGems
pytest -q tests/test_copy_ops.py当前环境结果:
27 passed, 1 warning
性能观察:
CUDA_VISIBLE_DEVICES=0 \
python tests/bench_copy.py --warmup 30 --iters 200 \
--json results/copy_perf_accuracy_p800.json详细数据和精度矩阵见 docs/PERFORMANCE_ACCURACY.md。
- complex 转 real 仍保留原生 fallback,以维持 PyTorch warning 语义;
- 内部重叠 destination 会拒绝写入,与 PyTorch 行为一致;
- bench 只提供当前机器的趋势观察,不作为严格性能回归结论。
P800 (1024, 1024) fp32 最新观察:
contiguous aten dispatch: 0.042 ms/iter
expanded aten dispatch: 0.040 ms/iter
expanded direct pointwise: 0.026 ms/iter
legacy fixed-rank kernel: 2.401 ms/iter
实际 dispatch 与 direct kernel 的差距是 host/dispatch 开销;kernel 本身提速
91.58x,完整 copy_ 调用相对 legacy kernel 提速 59.29x。最终不新增手写
fixed-rank kernel,直接复用 pointwise_dynamic。