Skip to content

Add Strix Halo (gfx1151) Atlas Inference Hub-kernel path for Qwen3.5/3.6/3.8 Gated DeltaNet - #49127

Open
AzeezIsh wants to merge 1 commit into
huggingface:mainfrom
AzeezIsh:gdn-rocm-gfx1151
Open

AzeezIsh wants to merge 1 commit into
huggingface:mainfrom
AzeezIsh:gdn-rocm-gfx1151

Conversation

@AzeezIsh

@AzeezIsh AzeezIsh commented Sep 26, 2026 •

Copy link
Copy Markdown
Contributor

CPU CI GPU run-slow

On AMD Strix Halo (Ryzen AI Max, gfx1151, RDNA 3.5) the fla and causal_conv1d fast paths have no ROCm build, so Qwen3.5-family models (Qwen3.5, Qwen3.6-27B, Qwen3.8-27B) fall back to the slow pure-torch Gated DeltaNet (is_fast_path_available=False). This adds a rocm device entry, gated to that GPU, to the existing Qwen3_5GatedDeltaNet mapping from #46423, pointing at the ROCm build of the same Hub kernel (Atlas-Inference/gdn). Strix Halo users get the fast path automatically; every other GPU, including other ROCm GPUs, keeps its existing path.

  • One Device(type="rocm", properties=ROCMProperties(min_capability=115, max_capability=115)) entry next to the GB10 one in integrations/hub_kernels.py, same layer name and repo, revision pinned.
  • A usage note with measured numbers in docs/source/en/model_doc/qwen3_5.md.

The kernel is the same stateless kernels layer as the GB10 build. it reuses the host module's projections, gated RMSNorm and out_proj, and replaces only the conv1d, q/k L2-norm and delta-rule cores. On gfx1151 it runs the kernels the Atlas engine serves Qwen3.6/3.8-27B with on Strix Halo.

  • prefill: causal_conv1d_update_prefill → l2_norm_bf16 → FLA-chunked delta rule (recompute_wu → chunk_delta_h → chunk_fwd_o)
  • decode: causal_conv1d_update_l2norm_f32 → gated_delta_rule_decode_f32

What does this PR do?

The gated rocm entry in integrations/hub_kernels.py plus the docs note. @use_kernel_forward_from_hub("Qwen3_5GatedDeltaNet") decorators from #46423 already cover the dense and MoE classes.

Verified on Strix Halo (gfx1151, torch.cuda.get_device_capability() == (11, 5)), bf16, torch 2.14.0 ROCm 7.2 and ROCm 7.14 wheels, transformers main + kernels 0.17:

  • Parity against a float64 torch reference for every op, including ragged prefill lengths (1, 33, 65, 100, 257, 1024): 23/23 tests on both wheels.
  • kernelize() swaps the forward of all 24/24 GDN layers of Qwen/Qwen3.5-9B, and the profiler shows the kernel ops running.
  • End-user path: transformers main + this PR, from_pretrained(..., use_kernels=True) with an empty kernel cache pulls Atlas-Inference/gdn at the pinned revision. All 24/24 GDN layers of Qwen/Qwen3.5-9B run it, and greedy generate() produces identical token IDs

Measured through the end-user path (from_pretrained(..., use_kernels=True), kernel pulled from the Hub) on Qwen/Qwen3.5-9B (bf16, gfx1151, torch 2.14 + ROCm 7.2, 1024-token prompt, greedy decode of 256 tokens; mean of two runs):

use_kernels TTFT (prefill) Decode
False (PyTorch fallback) 1.33 s 5.77 tok/s
True (Atlas-Inference/gdn) 1.16 s (1.15x faster) 6.23 tok/s

About 74% of this prefill is the dense projection GEMMs, which the kernel doesn't touch. Per Gated DeltaNet layer of Qwen/Qwen3.8-27B (real weights), the kernel is 1.65x faster on a 1024-token prefill and 1.26x faster per decode step. The published build targets torch 2.14 with ROCm 7.2.

Follow-up to #46423 (the GB10/SM121 entry). The ROCm build needs kernel-builder with gfx1151 support, which landed in huggingface/kernels#839 (tracked in huggingface/kernels#834). No overlapping PR exists.

The same trust note applies as for the GB10 entry: the mapped kernel needs trust_remote_code=True until Atlas-Inference is on the kernels trusted-publisher allowlist, would love to enable this for continued future contributions without any blocks!

Code Agent Policy

  • I confirm that this is not a pure code agent PR.

Before submitting

Who can review?

@vasqu original PR merger, @danieldk first issue
Potentially @drbh (kernels integration) and @Cyrilvallez (text models / Qwen3.5) as well :)
If you could enable myself on the kernels trusted-publisher allowlist that would help for the future!

@github-actions

Copy link
Copy Markdown
Contributor

CI recap

Dashboard: View test results in Grafana
Latest run: 36276919751:2
Result: failure | Jobs: 2 | Tests: 11 | Failures: 1 | Duration: 59s

Code quality check failed: test jobs were skipped. Fix the code quality issues and push again to run tests.

@Rocketknight1

Copy link
Copy Markdown
Member

cc @Abdennacer-Badaoui

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.

2 participants