Skip to content

feat: add optimized Kunlun rotary embedding backend - #1001

Draft
JoeZhang-0x000 wants to merge 1 commit into
InfiniTensor:masterfrom
JoeZhang-0x000:feat/kunlun-rope
Draft

JoeZhang-0x000 wants to merge 1 commit into
InfiniTensor:masterfrom
JoeZhang-0x000:feat/kunlun-rope

Conversation

@JoeZhang-0x000

Copy link
Copy Markdown

Summary

  • Add the Kunlun native RotaryEmbedding backend and only its required build/device/Python-dispatch plumbing.
  • Query available clusters (cap 12, fallback 8) and stripe adjacent workers across clusters; use native int64 positions with on-device bounds clamping.
  • Preserve the existing in-place Q/optional-K API, NeoX/GPT-J layouts, partial rotation, offset, inverse, strided batched views, and F16/BF16/F32 input/cache combinations.

Motivation

Depends on InfiniRT #48, revision e02d72ec75414d234cb1473994f6463856844acc. Build against that runtime until it merges. Supersedes the closed InfiniCore #1574: current operator implementations belong in InfiniOps. There is no existing Kunlun RoPE on upstream master, so this PR includes the minimal backend entry points as well as the optimization.

Type of Change

  • feat — new platform/operator backend

Platforms Affected

  • Kunlun (WITH_KUNLUN)
  • Build system / CMake
  • Python device mapping / wrapper generation

Smoke Test Result

P800 OAM (12 clusters), XTDK LLVM 15/xpu3, XRE runtime query 5.0 (bded6e7), GCC 9.4, PyTorch 2.5.1. Fresh native library and Python binding builds:

cmake -S . -B build -DWITH_KUNLUN=ON -DWITH_CPU=ON \
  -DXRE_ROOT=/path/to/xre -DXTDK_ROOT=/path/to/xtdk \
  -DINFINI_RT_ROOT=/path/to/pr-48-install \
  -DINFINI_OPS_OPS=rotary_embedding -DGENERATE_PYTHON_BINDINGS=ON \
  -DCMAKE_BUILD_TYPE=Release
cmake --build build -j4
python -m pytest tests/test_rotary_embedding.py --devices kunlun -q
# 42 passed; 10 Ascend-specific cases skipped.

For direct CMake installs, expose the extension as infini.ops (the wheel already uses that namespace). Tests include 18 mixed-dtype/NeoX/GPT-J graph cases with two changed-input/position replays each: negatives, 2**40, INT64_MAX, 512 rotary dimensions, offset/inverse, strided QKV and padded cache rows. All 24 BF16 benchmark shapes also match the CPU reference with zero eager error; baseline/candidate eager and changed-input graph output hashes match in every shape.

Test Results on Supported Platforms

Configuration Result
Kunlun native RoPE 42 passed, 10 platform-specific skips; 24 benchmark shapes passed eager + graph checks
CPU-only native build (add,mul,relu,cast) 369 passed, 3 skipped
Existing wrapper/public-header/torch/ninetoothed generator tests 65 passed
Other accelerators Hardware/SDK unavailable; native kernels unchanged

clang-format 21, Ruff and git diff --check passed. Full unrelated-operator suite not run; Kunlun currently implements RoPE only.

Benchmark / Performance Impact

Controlled comparison on one P800, BF16 packed QKV views, Qwen3-8B head dimension 128 and Q/KV heads 32/TP, 8/TP. TP labels describe per-rank shapes, not multi-card execution. Baseline is the same modern implementation with only the former 8-cluster/contiguous-worker scheme restored; candidate uses 12 clusters/striped workers. Same runtime/toolchain and native int64 API on both sides.

Median synchronized graph replay time per RoPE Q+K call (μs): 8 calls/graph, 20 replays/sample, 5 samples, 3 warm-up replays. Graph execution is checked after changing inputs and positions before timing.

8192-token prefill:

TP shape Baseline μs Optimized μs Speedup
1 3283.65 2223.96 1.476×
2 1708.09 1151.75 1.483×
4 874.31 587.60 1.488×
8 445.15 299.04 1.489×

Across 24 shapes (TP 1/2/4/8 × 1/4/16/64/512/8192 tokens), speedup is 0.957–1.489×. Six tiny 1–4-token cases regress by at most 0.56 μs (4.5%); larger shapes improve.

Notes for Reviewers

Kernel launch metadata uses fixed-width types because XTDK's device size_t/ptrdiff_t differ from the host ABI. The public operator API is unchanged; the legacy capability symbol is unnecessary here. No attention routing or other kernels are added, and no report documents, CSV files or benchmark scripts are committed.

Earlier legacy vLLM integration results, not rerun for this new API port: Qwen3-0.6B GSM8K native/plugin both 715/1319 (54.21%, gap 0 pp); Qwen3-8B static throughput minima versus the same vendor baseline at TP1/2/4/8: 111.12%/104.65%/98.47%/90.17%. Both sides used the same isolated vendor cache fix, with vendor attention. Original integration report.

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant