KernelBench cuda · RTX PRO 6000

MegaQwen Decode Grok 4.6

5.42%geomean peak fraction across shapes

manually audited: clean

Flat multi-kernel Qwen3-0.6B-geometry decode (RMSNorm+QKV, GQA flash-decode, post-attn RMSNorm+SwiGLU). Optional CUDA-graph replay keyed on a pointer signature (hidden, caches, weights); recaptures when pointers change. template_mutated=false.

harnessgrokagent session48mtotal wall1h 9mcheck25sbenchmark14moutput tokensregimethroughput

Per-shape vs governing ceilingeach shape graded against whichever binds — bf16 compute or HBM bandwidth

No per-shape benchmark data archived for this run.

Kernel source (redacted)
"""Fast Qwen3-0.6B-geometry decode for RTX PRO 6000 (SM120).

Non-cooperative multi-kernel path: MegaQwen's cooperative grid.sync ceiling
does not scale to 188 SMs / 128k context. Fused RMSNorm+QKV GEMV, GQA
flash-decode with shared KV reads and streaming loads, fused post-attn
RMSNorm+SwiGLU, CUDA-graph replay of one token.
"""
from __future__ import annotations

import os
from typing import Optional

import torch
import torch.nn as nn
from torch.utils.cpp_extension import load_inline

HIDDEN = 1024
INTERMEDIATE = 3072
NUM_Q = 16
NUM_KV = 8
HEAD_DIM = 128
NUM_LAYERS = 4
Q_SIZE = NUM_Q * HEAD_DIM
KV_SIZE = NUM_KV * HEAD_DIM
EPS = 1e-6

_mod = None


def _cuda_src() -> str:
    return r'''
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <cuda_runtime.h>
#include <cuda_bf16.h>
#include <vector>
#include <algorithm>
#include <cstring>
#include <cmath>
#include <stdexcept>
#include <string>

constexpr int H = 1024;
constexpr int ISIZE = 3072;
constexpr int NQ = 16;
constexpr int NKV = 8;
constexpr int HD = 128;
constexpr int QS = 2048;
constexpr int KVS = 1024;
constexpr int WARP = 32;
constexpr int BLOCK = 256;
constexpr int NWARPS = BLOCK / WARP;
constexpr int GQA = 2;
constexpr int MAX_LAYERS = 8;
constexpr int MAX_CHUNKS = 32;
constexpr float RMS_EPS = 1e-6f;
constexpr float ATTN_SCALE = 0.08838834764831843f; // 1/sqrt(128)

struct LayerW {
    const __nv_bfloat16* input_ln;
    const __nv_bfloat16* q_proj;
    const __nv_bfloat16* k_proj;
    const __nv_bfloat16* v_proj;
    const __nv_bfloat16* q_norm;
    const __nv_bfloat16* k_norm;
    const __nv_bfloat16* o_proj;
    const __nv_bfloat16* post_ln;
    const __nv_bfloat16* gate;
    const __nv_bfloat16* up;
    const __nv_bfloat16* down;
};

__device__ __forceinline__ float warp_sum(float v) {
#pragma unroll
    for (int o = 16; o > 0; o >>= 1) v += __shfl_down_sync(0xffffffff, v, o);
    return v;
}
__device__ __forceinline__ float warp_max(float v) {
#pragma unroll
    for (int o = 16; o > 0; o >>= 1) v = fmaxf(v, __shfl_down_sync(0xffffffff, v, o));
    return v;
}
__device__ __forceinline__ float silu(float x) {
    return x * (1.f / (1.f + expf(-x)));
}

__device__ __forceinline__ uint4 ldg_u4(const void* p) {
    return __ldg(reinterpret_cast<const uint4*>(p));
}
__device__ __forceinline__ uint4 ldcs_u4(const void* p) {
    uint4 v;
    asm volatile("ld.global.cs.v4.u32 {%0,%1,%2,%3}, [%4];"
        : "=r"(v.x), "=r"(v.y), "=r"(v.z), "=r"(v.w)
        : "l"(p));
    return v;
}
__device__ __forceinline__ uint2 ldcs_u2(const void* p) {
    uint2 v;
    asm volatile("ld.global.cs.v2.u32 {%0,%1}, [%2];"
        : "=r"(v.x), "=r"(v.y)
        : "l"(p));
    return v;
}

__device__ __forceinline__ float dot8_act(const uint4 u, const float* act, int k) {
    auto* w = reinterpret_cast<const __nv_bfloat162*>(&u);
    float2 a = __bfloat1622float2(w[0]);
    float2 b = __bfloat1622float2(w[1]);
    float2 c = __bfloat1622float2(w[2]);
    float2 d = __bfloat1622float2(w[3]);
    return a.x * act[k]     + a.y * act[k + 1]
         + b.x * act[k + 2] + b.y * act[k + 3]
         + c.x * act[k + 4] + c.y * act[k + 5]
         + d.x * act[k + 6] + d.y * act[k + 7];
}

__device__ __forceinline__ float gemv_k1024(const __nv_bfloat16* row, const float* act, int lane) {
    float sum = 0.f;
#pragma unroll
    for (int i = 0; i < 4; ++i) {
        int k = lane * 8 + i * 256;
        sum += dot8_act(ldg_u4(row + k), act, k);
    }
    return warp_sum(sum);
}

__device__ __forceinline__ float gemv_k2048(const __nv_bfloat16* row, const float* act, int lane) {
    float sum = 0.f;
#pragma unroll
    for (int i = 0; i < 8; ++i) {
        int k = lane * 8 + i * 256;
        sum += dot8_act(ldg_u4(row + k), act, k);
    }
    return warp_sum(sum);
}

__device__ __forceinline__ float gemv_k3072(const __nv_bfloat16* row, const float* act, int lane) {
    float sum = 0.f;
#pragma unroll
    for (int i = 0; i < 12; ++i) {
        int k = lane * 8 + i * 256;
        sum += dot8_act(ldg_u4(row + k), act, k);
    }
    return warp_sum(sum);
}

// -------------------- kernels --------------------

__global__ void __launch_bounds__(256, 4)
rms_qkv_kernel(
    const __nv_bfloat16* __restrict__ x,
    const __nv_bfloat16* __restrict__ noise,   // nullable: if set, mix 0.5*noise+0.5*x
    const __nv_bfloat16* __restrict__ w_norm,
    const __nv_bfloat16* __restrict__ Wq,
    const __nv_bfloat16* __restrict__ Wk,
    const __nv_bfloat16* __restrict__ Wv,
    float* __restrict__ q, float* __restrict__ k, float* __restrict__ v,
    float* __restrict__ residual
) {
    extern __shared__ float smem[];
    float* s_act = smem;
    float* s_red = smem + H;
    const int tid = threadIdx.x;
    const int warp = tid / WARP;
    const int lane = tid % WARP;

    float sq = 0.f;
    for (int i = tid; i < H; i += BLOCK) {
        float val = __bfloat162float(x[i]);
        if (noise) val = 0.5f * __bfloat162float(noise[i]) + 0.5f * val;
        s_act[i] = val;
        residual[i] = val;
        sq += val * val;
    }
    sq = warp_sum(sq);
    if (lane == 0) s_red[warp] = sq;
    __syncthreads();
    if (warp == 0) {
        float s = (lane < NWARPS) ? s_red[lane] : 0.f;
        s = warp_sum(s);
        if (lane == 0) s_red[0] = rsqrtf(s / float(H) + RMS_EPS);
    }
    __syncthreads();
    float rstd = s_red[0];
    for (int i = tid; i < H; i += BLOCK)
        s_act[i] *= rstd * __bfloat162float(__ldg(w_norm + i));
    __syncthreads();

    constexpr int TOTAL = QS + KVS + KVS;
    for (int m = blockIdx.x * NWARPS + warp; m < TOTAL; m += gridDim.x * NWARPS) {
        const __nv_bfloat16* row;
        float* outp;
        if (m < QS) { row = Wq + (size_t)m * H; outp = q + m; }
        else if (m < QS + KVS) { row = Wk + (size_t)(m - QS) * H; outp = k + (m - QS); }
        else { row = Wv + (size_t)(m - QS - KVS) * H; outp = v + (m - QS - KVS); }
        float sum = gemv_k1024(row, s_act, lane);
        if (lane == 0) *outp = sum;
    }
}

__device__ __forceinline__ void rms_rope_head(
    const float* head_in, const __nv_bfloat16* nw,
    float* head_out, int pos, int tid, int bdx,
    float* s_red, float* s_n
) {
    const int half = HD / 2;
    const int lane = tid % WARP;
    const int warp = tid / WARP;
    float sq = 0.f;
    for (int i = tid; i < HD; i += bdx) sq += head_in[i] * head_in[i];
    sq = warp_sum(sq);
    if (lane == 0) s_red[warp] = sq;
    __syncthreads();
    if (warp == 0) {
        float s = (lane < (bdx / WARP)) ? s_red[lane] : 0.f;
        s = warp_sum(s);
        if (lane == 0) s_red[0] = rsqrtf(s / float(HD) + RMS_EPS);
    }
    __syncthreads();
    float rstd = s_red[0];
    for (int i = tid; i < HD; i += bdx)
        s_n[i] = head_in[i] * rstd * __bfloat162float(__ldg(nw + i));
    __syncthreads();
    for (int i = tid; i < half; i += bdx) {
        float inv = expf(-logf(10000.f) * (float)i / (float)half);
        float ang = (float)pos * inv;
        float c, s;
        __sincosf(ang, &s, &c);
        float x1 = s_n[i], x2 = s_n[i + half];
        head_out[i] = x1 * c - x2 * s;
        head_out[i + half] = x1 * s + x2 * c;
    }
    __syncthreads();
}

// GQA partial: one block = (kv_head, chunk). Fuses Q/K RMSNorm+RoPE+cache write
// (chunk 0) and online-softmax over cache[0:pos] plus the current token.
__global__ void __launch_bounds__(256, 3)
attn_gqa_partial_kernel(
    const float* __restrict__ q,
    const float* __restrict__ k,
    const float* __restrict__ v,
    const __nv_bfloat16* __restrict__ q_norm,
    const __nv_bfloat16* __restrict__ k_norm,
    const __nv_bfloat16* __restrict__ k_cache,
    const __nv_bfloat16* __restrict__ v_cache,
    __nv_bfloat16* __restrict__ k_cache_w,
    __nv_bfloat16* __restrict__ v_cache_w,
    float* __restrict__ partial_m,
    float* __restrict__ partial_l,
    float* __restrict__ partial_o,
    const int* __restrict__ d_pos,
    int max_seq, int num_chunks, int chunk_stride
) {
    const int kv_h = blockIdx.x;
    const int chunk = blockIdx.y;
    const int tid = threadIdx.x;
    const int warp = tid / WARP;
    const int lane = tid % WARP;
    const int pos = *d_pos;
    const int cache_old = pos; // positions [0, pos) already in cache
    const int qh0 = kv_h * GQA;
    const int qh1 = qh0 + 1;

    extern __shared__ float sm[];
    float* sq0 = sm;            // 128
    float* sq1 = sm + HD;       // 128
    float* sk  = sm + 2 * HD;   // 128
    float* sv  = sm + 3 * HD;   // 128
    float* sred = sm + 4 * HD;  // 8
    float* sn   = sm + 4 * HD + 8; // 128 scratch for rms_rope

    // All chunks redundantly rope Q/K for this KV head (no cross-block race).
    rms_rope_head(q + qh0 * HD, q_norm, sq0, pos, tid, BLOCK, sred, sn);
    rms_rope_head(q + qh1 * HD, q_norm, sq1, pos, tid, BLOCK, sred, sn);
    rms_rope_head(k + kv_h * HD, k_norm, sk, pos, tid, BLOCK, sred, sn);
    for (int i = tid; i < HD; i += BLOCK)
        sv[i] = v[kv_h * HD + i];
    __syncthreads();

    if (chunk == 0) {
        // packed [NKV, max_seq, 2, HD] — K then V at each position
        __nv_bfloat16* kvw = k_cache_w + ((size_t)kv_h * max_seq + pos) * (2 * HD);
        for (int i = tid; i < HD; i += BLOCK) {
            kvw[i] = __float2bfloat16(sk[i]);
            kvw[HD + i] = __float2bfloat16(sv[i]);
        }
    }

    int chunk_sz = (cache_old + num_chunks - 1) / num_chunks;
    if (chunk_sz < 1) chunk_sz = 1;
    int p0 = chunk * chunk_sz;
    int p1 = min(p0 + chunk_sz, cache_old);

    auto write_empty = [&](int qh) {
        if (tid == 0) {
            partial_m[qh * chunk_stride + chunk] = -INFINITY;
            partial_l[qh * chunk_stride + chunk] = 0.f;
        }
        for (int d = tid; d < HD; d += BLOCK)
            partial_o[(qh * chunk_stride + chunk) * HD + d] = 0.f;
    };

    // Current token is attached to chunk 0 so every position is covered once.
    const bool do_cur = (chunk == 0);
    if (p0 >= cache_old && !do_cur) {
        write_empty(qh0);
        write_empty(qh1);
        return;
    }

    float m0 = -INFINITY, l0 = 0.f, m1 = -INFINITY, l1 = 0.f;
    float o0[4] = {0, 0, 0, 0};
    float o1[4] = {0, 0, 0, 0};

    auto accumulate = [&](float s0, float s1, const float* vreg) {
        {
            float mnew = fmaxf(m0, s0);
            float ed = expf(m0 - mnew);
            float e = expf(s0 - mnew);
            l0 = l0 * ed + e;
#pragma unroll
            for (int j = 0; j < 4; ++j) o0[j] = o0[j] * ed + e * vreg[j];
            m0 = mnew;
        }
        {
            float mnew = fmaxf(m1, s1);
            float ed = expf(m1 - mnew);
            float e = expf(s1 - mnew);
            l1 = l1 * ed + e;
#pragma unroll
            for (int j = 0; j < 4; ++j) o1[j] = o1[j] * ed + e * vreg[j];
            m1 = mnew;
        }
    };

    auto score_pos = [&](const __nv_bfloat16* kp, float& s0, float& s1) {
        s0 = 0.f; s1 = 0.f;
#pragma unroll
        for (int d = lane * 4; d < HD; d += WARP * 4) {
            uint2 ku = ldcs_u2(kp + d);
            auto* kb = reinterpret_cast<const __nv_bfloat162*>(&ku);
            float2 kf0 = __bfloat1622float2(kb[0]);
            float2 kf1 = __bfloat1622float2(kb[1]);
            s0 += sq0[d] * kf0.x + sq0[d + 1] * kf0.y
                + sq0[d + 2] * kf1.x + sq0[d + 3] * kf1.y;
            s1 += sq1[d] * kf0.x + sq1[d + 1] * kf0.y
                + sq1[d + 2] * kf1.x + sq1[d + 3] * kf1.y;
        }
        s0 = warp_sum(s0) * ATTN_SCALE;
        s1 = warp_sum(s1) * ATTN_SCALE;
        s0 = __shfl_sync(0xffffffff, s0, 0);
        s1 = __shfl_sync(0xffffffff, s1, 0);
    };
    auto load_v = [&](const __nv_bfloat16* vp, float* vreg) {
#pragma unroll
        for (int j = 0, d = lane; d < HD; d += WARP, ++j) {
            unsigned short bits;
            asm volatile("ld.global.cs.u16 %0, [%1];"
                : "=h"(bits) : "l"(vp + d));
            vreg[j] = __bfloat162float(*reinterpret_cast<const __nv_bfloat16*>(&bits));
        }
    };

    if (p0 < cache_old) {
        int p = p0 + warp;
        // Dual-issue two positions when the chunk is long enough.
        for (; p + NWARPS < p1; p += 2 * NWARPS) {
            const __nv_bfloat16* kv0 = k_cache + ((size_t)kv_h * max_seq + p) * (2 * HD);
            const __nv_bfloat16* kv1 = k_cache + ((size_t)kv_h * max_seq + p + NWARPS) * (2 * HD);
            const __nv_bfloat16* kp0 = kv0;
            const __nv_bfloat16* vp0 = kv0 + HD;
            const __nv_bfloat16* kp1 = kv1;
            const __nv_bfloat16* vp1 = kv1 + HD;
            float s0a, s1a, s0b, s1b;
            score_pos(kp0, s0a, s1a);
            score_pos(kp1, s0b, s1b);
            float v0[4], v1[4];
            load_v(vp0, v0);
            load_v(vp1, v1);
            accumulate(s0a, s1a, v0);
            accumulate(s0b, s1b, v1);
        }
        for (; p < p1; p += NWARPS) {
            const __nv_bfloat16* kv = k_cache + ((size_t)kv_h * max_seq + p) * (2 * HD);
            const __nv_bfloat16* kp = kv;
            const __nv_bfloat16* vp = kv + HD;
            float s0, s1;
            score_pos(kp, s0, s1);
            float vreg[4];
            load_v(vp, vreg);
            accumulate(s0, s1, vreg);
        }
    }

    if (do_cur && warp == 0) {
        float s0 = 0.f, s1 = 0.f;
        for (int d = lane; d < HD; d += WARP) {
            s0 += sq0[d] * sk[d];
            s1 += sq1[d] * sk[d];
        }
        s0 = warp_sum(s0) * ATTN_SCALE;
        s1 = warp_sum(s1) * ATTN_SCALE;
        s0 = __shfl_sync(0xffffffff, s0, 0);
        s1 = __shfl_sync(0xffffffff, s1, 0);
        float vreg[4];
#pragma unroll
        for (int j = 0, d = lane; d < HD; d += WARP, ++j) vreg[j] = sv[d];
        accumulate(s0, s1, vreg);
    }

    __shared__ float sm0[8], sl0[8], sm1[8], sl1[8];
    __shared__ float so0[8][HD], so1[8][HD];
    if (lane == 0) { sm0[warp] = m0; sl0[warp] = l0; sm1[warp] = m1; sl1[warp] = l1; }
#pragma unroll
    for (int j = 0, d = lane; d < HD; d += WARP, ++j) {
        so0[warp][d] = o0[j];
        so1[warp][d] = o1[j];
    }
    __syncthreads();

    auto combine = [&](int qh, float* smx, float* slx, float so[][HD]) {
        if (warp != 0) return;
        float gm = -INFINITY;
        for (int w = 0; w < NWARPS; ++w)
            if (smx[w] > -1e30f) gm = fmaxf(gm, smx[w]);
        float gl = 0.f;
        float go[4] = {0, 0, 0, 0};
        for (int w = 0; w < NWARPS; ++w) {
            if (!(smx[w] > -1e30f)) continue;
            float sc = expf(smx[w] - gm);
            gl += slx[w] * sc;
#pragma unroll
            for (int j = 0, d = lane; d < HD; d += WARP, ++j)
                go[j] += so[w][d] * sc;
        }
        if (lane == 0) {
            partial_m[qh * chunk_stride + chunk] = gm;
            partial_l[qh * chunk_stride + chunk] = gl;
        }
#pragma unroll
        for (int j = 0, d = lane; d < HD; d += WARP, ++j)
            partial_o[(qh * chunk_stride + chunk) * HD + d] = go[j];
    };
    combine(qh0, sm0, sl0, so0);
    __syncthreads();
    combine(qh1, sm1, sl1, so1);
}

__global__ void attn_reduce_kernel(
    const float* __restrict__ partial_m,
    const float* __restrict__ partial_l,
    const float* __restrict__ partial_o,
    float* __restrict__ attn_out,
    int num_chunks, int chunk_stride
) {
    const int qh = blockIdx.x;
    const int tid = threadIdx.x;
    float gm = -INFINITY;
    for (int c = tid; c < num_chunks; c += blockDim.x)
        gm = fmaxf(gm, partial_m[qh * chunk_stride + c]);
    __shared__ float sred[8];
    int lane = tid % WARP, warp = tid / WARP;
    float wm = warp_max(gm);
    if (lane == 0) sred[warp] = wm;
    __syncthreads();
    if (warp == 0) {
        float v = (lane < blockDim.x / WARP) ? sred[lane] : -INFINITY;
        v = warp_max(v);
        if (lane == 0) sred[0] = v;
    }
    __syncthreads();
    gm = sred[0];
    for (int d = tid; d < HD; d += blockDim.x) {
        float gl = 0.f, go = 0.f;
        for (int c = 0; c < num_chunks; ++c) {
            float m = partial_m[qh * chunk_stride + c];
            if (!(m > -1e30f)) continue;
            float sc = expf(m - gm);
            gl += partial_l[qh * chunk_stride + c] * sc;
            go += partial_o[(qh * chunk_stride + c) * HD + d] * sc;
        }
        attn_out[qh * HD + d] = go / fmaxf(gl, 1e-20f);
    }
}

__global__ void __launch_bounds__(256, 4)
o_proj_kernel(
    const float* __restrict__ attn, const __nv_bfloat16* __restrict__ Wo,
    const float* __restrict__ residual, float* __restrict__ out
) {
    extern __shared__ float s_attn[];
    const int tid = threadIdx.x, warp = tid / WARP, lane = tid % WARP;
    for (int i = tid; i < QS; i += BLOCK) s_attn[i] = attn[i];
    __syncthreads();
    for (int m = blockIdx.x * NWARPS + warp; m < H; m += gridDim.x * NWARPS) {
        float sum = gemv_k2048(Wo + (size_t)m * QS, s_attn, lane);
        if (lane == 0) out[m] = sum + residual[m];
    }
}

__global__ void __launch_bounds__(256, 4)
rms_gate_up_kernel(
    const float* __restrict__ x,
    const __nv_bfloat16* __restrict__ w_norm,
    const __nv_bfloat16* __restrict__ Wg,
    const __nv_bfloat16* __restrict__ Wu,
    float* __restrict__ out
) {
    extern __shared__ float smem[];
    float* s_act = smem;
    float* s_red = smem + H;
    const int tid = threadIdx.x, warp = tid / WARP, lane = tid % WARP;

    float sq = 0.f;
    for (int i = tid; i < H; i += BLOCK) {
        float val = x[i];
        s_act[i] = val;
        sq += val * val;
    }
    sq = warp_sum(sq);
    if (lane == 0) s_red[warp] = sq;
    __syncthreads();
    if (warp == 0) {
        float s = (lane < NWARPS) ? s_red[lane] : 0.f;
        s = warp_sum(s);
        if (lane == 0) s_red[0] = rsqrtf(s / float(H) + RMS_EPS);
    }
    __syncthreads();
    float rstd = s_red[0];
    for (int i = tid; i < H; i += BLOCK)
        s_act[i] *= rstd * __bfloat162float(__ldg(w_norm + i));
    __syncthreads();

    for (int m = blockIdx.x * NWARPS + warp; m < ISIZE; m += gridDim.x * NWARPS) {
        const __nv_bfloat16* rg = Wg + (size_t)m * H;
        const __nv_bfloat16* ru = Wu + (size_t)m * H;
        float sg = 0.f, su = 0.f;
#pragma unroll
        for (int i = 0; i < 4; ++i) {
            int k = lane * 8 + i * 256;
            sg += dot8_act(ldg_u4(rg + k), s_act, k);
            su += dot8_act(ldg_u4(ru + k), s_act, k);
        }
        sg = warp_sum(sg);
        su = warp_sum(su);
        if (lane == 0) out[m] = silu(sg) * su;
    }
}

__global__ void __launch_bounds__(256, 4)
down_residual_kernel(
    const float* __restrict__ mid, const __nv_bfloat16* __restrict__ Wd,
    const float* __restrict__ residual, __nv_bfloat16* __restrict__ out
) {
    extern __shared__ float s_mid[];
    const int tid = threadIdx.x, warp = tid / WARP, lane = tid % WARP;
    for (int i = tid; i < ISIZE; i += BLOCK) s_mid[i] = mid[i];
    __syncthreads();
    for (int m = blockIdx.x * NWARPS + warp; m < H; m += gridDim.x * NWARPS) {
        float sum = gemv_k3072(Wd + (size_t)m * ISIZE, s_mid, lane);
        if (lane == 0) out[m] = __float2bfloat16(sum + residual[m]);
    }
}

__global__ void step_inc_kernel(int* __restrict__ d_pos, int* __restrict__ d_step) {
    if (threadIdx.x == 0 && blockIdx.x == 0) {
        ++(*d_pos);
        ++(*d_step);
    }
}

// -------------------- host --------------------

static int pick_blocks(int rows, int maxb) {
    int need = (rows + NWARPS - 1) / NWARPS;
    // Oversubscribe SMs — GEMV is latency-bound and non-cooperative.
    int cap = std::max(maxb * 2, 256);
    return std::max(1, std::min(need, cap));
}

static int choose_chunks(int cache_hint, int max_blocks) {
    int sm_chunks = std::max(1, max_blocks / NKV);
    sm_chunks = std::min(sm_chunks, MAX_CHUNKS);
    int by_len = std::max(1, (cache_hint + 63) / 64);
    int n = std::min(sm_chunks, by_len);
    if (cache_hint >= 128) n = std::max(n, std::min(sm_chunks, 4));
    if (cache_hint >= 512) n = std::max(n, std::min(sm_chunks, 8));
    if (cache_hint >= 2048) n = std::max(n, std::min(MAX_CHUNKS, 24));
    if (cache_hint >= 8192) n = MAX_CHUNKS;
    return std::max(1, std::min(n, MAX_CHUNKS));
}

static LayerW make_layer(const std::vector<torch::Tensor>& wf, int layer) {
    auto p = [&](int i) {
        return reinterpret_cast<const __nv_bfloat16*>(wf[layer * 11 + i].data_ptr());
    };
    return LayerW{p(0), p(1), p(2), p(3), p(4), p(5), p(6), p(7), p(8), p(9), p(10)};
}

static void ck(cudaError_t e, const char* what) {
    if (e != cudaSuccess) {
        throw std::runtime_error(std::string(what) + ": " + cudaGetErrorString(e));
    }
}

static void launch_layer(
    const __nv_bfloat16* x, const __nv_bfloat16* noise, const LayerW& w,
    __nv_bfloat16* k_cache, __nv_bfloat16* v_cache,
    __nv_bfloat16* y,
    float* residual, float* g_q, float* g_k, float* g_v, float* g_attn,
    float* g_mid, float* g_resid2,
    float* partial_m, float* partial_l, float* partial_o,
    const int* d_pos, int max_seq,
    int num_chunks, int chunk_stride, int max_blocks,
    cudaStream_t stream
) {
    size_t sh_h = (H + NWARPS) * sizeof(float);
    size_t sh_q = QS * sizeof(float);
    size_t sh_i = ISIZE * sizeof(float);
    size_t sh_attn = (4 * HD + 8 + HD) * sizeof(float);

    rms_qkv_kernel<<<pick_blocks(QS + KVS + KVS, max_blocks), BLOCK, sh_h, stream>>>(
        x, noise, w.input_ln, w.q_proj, w.k_proj, w.v_proj, g_q, g_k, g_v, residual);

    dim3 grid(NKV, num_chunks);
    attn_gqa_partial_kernel<<<grid, BLOCK, sh_attn, stream>>>(
        g_q, g_k, g_v, w.q_norm, w.k_norm,
        k_cache, v_cache, k_cache, v_cache,
        partial_m, partial_l, partial_o,
        d_pos, max_seq, num_chunks, chunk_stride);

    attn_reduce_kernel<<<NQ, 128, 0, stream>>>(
        partial_m, partial_l, partial_o, g_attn, num_chunks, chunk_stride);

    o_proj_kernel<<<pick_blocks(H, max_blocks), BLOCK, sh_q, stream>>>(
        g_attn, w.o_proj, residual, g_resid2);

    rms_gate_up_kernel<<<pick_blocks(ISIZE, max_blocks), BLOCK, sh_h, stream>>>(
        g_resid2, w.post_ln, w.gate, w.up, g_mid);

    down_residual_kernel<<<pick_blocks(H, max_blocks), BLOCK, sh_i, stream>>>(
        g_mid, w.down, g_resid2, y);
}

static void one_step(
    const __nv_bfloat16* noise_base,
    __nv_bfloat16* hidden,
    __nv_bfloat16* buf_a, __nv_bfloat16* buf_b,
    LayerW* layers, int num_layers,
    __nv_bfloat16** kptrs, __nv_bfloat16** vptrs,
    float* residual, float* g_q, float* g_k, float* g_v, float* g_attn,
    float* g_mid, float* g_resid2,
    float* partial_m, float* partial_l, float* partial_o,
    int* d_pos, int* d_step,
    int max_seq, int num_chunks, int chunk_stride, int max_blocks,
    cudaStream_t stream
) {
    for (int layer = 0; layer < num_layers; ++layer) {
        bool first = (layer == 0);
        bool last = (layer == num_layers - 1);
        const __nv_bfloat16* in_ptr = first ? hidden : ((layer % 2 == 1) ? buf_a : buf_b);
        __nv_bfloat16* out_ptr = last ? hidden : ((layer % 2 == 0) ? buf_a : buf_b);
        const __nv_bfloat16* noise = nullptr;
        if (first) {
            // noise row selected on device via d_step: we pass base+offset on host
            // but d_step changes across graph replays. Use a kernel-visible trick:
            // encode noise as base; mix kernel indexes by *d_step.
            noise = noise_base; // special: first-layer mix uses noise_base[*d_step]
        }
        // For graph-safe noise indexing, layer 0 uses noise_base + (*d_step)*H.
        // launch_layer takes a single noise pointer; pass a tagged pointer and
        // let rms_qkv read sequentially — we instead launch a tiny index copy
        // only for layer 0 via the d_step pointer baked into a wrapper below.
        launch_layer(
            in_ptr, first ? noise_base : nullptr, layers[layer],
            kptrs[layer], vptrs[layer], out_ptr,
            residual, g_q, g_k, g_v, g_attn, g_mid, g_resid2,
            partial_m, partial_l, partial_o,
            d_pos, max_seq, num_chunks, chunk_stride, max_blocks, stream
        );
    }
    step_inc_kernel<<<1, 1, 0, stream>>>(d_pos, d_step);
}

// Graph-safe mix: rms_qkv currently treats `noise` as a 1024-vector. For
// multi-step graphs we need noise[*d_step]. Use a dedicated mix that writes
// a staging vector, then layer 0 reads that (no noise ptr).
__global__ void mix_indexed_kernel(
    const __nv_bfloat16* __restrict__ noise_base,
    const __nv_bfloat16* __restrict__ h,
    __nv_bfloat16* __restrict__ out,
    const int* __restrict__ d_step
) {
    int i = blockIdx.x * blockDim.x + threadIdx.x;
    if (i < H) {
        const __nv_bfloat16* row = noise_base + (size_t)(*d_step) * H;
        float a = __bfloat162float(row[i]);
        float b = __bfloat162float(h[i]);
        out[i] = __float2bfloat16(0.5f * a + 0.5f * b);
    }
}

static void one_step_graph(
    const __nv_bfloat16* noise_base,
    __nv_bfloat16* hidden,
    __nv_bfloat16* mix_buf,
    __nv_bfloat16* buf_a, __nv_bfloat16* buf_b,
    LayerW* layers, int num_layers,
    __nv_bfloat16** kptrs, __nv_bfloat16** vptrs,
    float* residual, float* g_q, float* g_k, float* g_v, float* g_attn,
    float* g_mid, float* g_resid2,
    float* partial_m, float* partial_l, float* partial_o,
    int* d_pos, int* d_step,
    int max_seq, int num_chunks, int chunk_stride, int max_blocks,
    cudaStream_t stream
) {
    mix_indexed_kernel<<<(H + 255) / 256, 256, 0, stream>>>(
        noise_base, hidden, mix_buf, d_step);

    for (int layer = 0; layer < num_layers; ++layer) {
        bool first = (layer == 0);
        bool last = (layer == num_layers - 1);
        const __nv_bfloat16* in_ptr = first ? mix_buf : ((layer % 2 == 1) ? buf_a : buf_b);
        __nv_bfloat16* out_ptr = last ? hidden : ((layer % 2 == 0) ? buf_a : buf_b);
        launch_layer(
            in_ptr, nullptr, layers[layer],
            kptrs[layer], vptrs[layer], out_ptr,
            residual, g_q, g_k, g_v, g_attn, g_mid, g_resid2,
            partial_m, partial_l, partial_o,
            d_pos, max_seq, num_chunks, chunk_stride, max_blocks, stream
        );
    }
    step_inc_kernel<<<1, 1, 0, stream>>>(d_pos, d_step);
}

struct GraphCache {
    cudaGraphExec_t exec = nullptr;
    void* sig[48]{};
    int n_sig = 0;
    int num_chunks = 0;
};

static GraphCache g_cache;

static bool sig_equal(void** a, void** b, int n) {
    return std::memcmp(a, b, sizeof(void*) * n) == 0;
}

void run_steps_cuda(
    torch::Tensor hidden,
    torch::Tensor noise,
    const std::vector<torch::Tensor>& weights_flat,
    std::vector<torch::Tensor> k_caches,
    std::vector<torch::Tensor> v_caches,
    torch::Tensor buf_a,
    torch::Tensor buf_b,
    torch::Tensor mix_buf,
    torch::Tensor residual,
    torch::Tensor g_q, torch::Tensor g_k, torch::Tensor g_v,
    torch::Tensor g_attn, torch::Tensor g_mid,
    torch::Tensor g_resid2,
    torch::Tensor partial_m, torch::Tensor partial_l, torch::Tensor partial_o,
    torch::Tensor d_pos, torch::Tensor d_step,
    torch::Tensor packed_w,
    int start_pos, int n_steps, int max_seq, int num_layers, int max_blocks
) {
    if (n_steps <= 0) return;
    cudaStream_t stream = at::cuda::getCurrentCUDAStream().stream();
    int chunk_stride = (int)partial_m.size(1);
    int cache_hint = start_pos + n_steps;
    int num_chunks = choose_chunks(cache_hint, max_blocks);

    LayerW layers[MAX_LAYERS];
    for (int i = 0; i < num_layers; ++i) layers[i] = make_layer(weights_flat, i);

    __nv_bfloat16* kptrs[MAX_LAYERS];
    __nv_bfloat16* vptrs[MAX_LAYERS];
    for (int i = 0; i < num_layers; ++i) {
        kptrs[i] = reinterpret_cast<__nv_bfloat16*>(k_caches[i].data_ptr());
        vptrs[i] = reinterpret_cast<__nv_bfloat16*>(v_caches[i].data_ptr());
    }

    auto* h_ptr = reinterpret_cast<__nv_bfloat16*>(hidden.data_ptr());
    auto* noise_ptr = reinterpret_cast<__nv_bfloat16*>(noise.data_ptr());
    auto* ba = reinterpret_cast<__nv_bfloat16*>(buf_a.data_ptr());
    auto* bb = reinterpret_cast<__nv_bfloat16*>(buf_b.data_ptr());
    auto* mix = reinterpret_cast<__nv_bfloat16*>(mix_buf.data_ptr());
    float* p_res = residual.data_ptr<float>();
    float* p_q = g_q.data_ptr<float>();
    float* p_k = g_k.data_ptr<float>();
    float* p_v = g_v.data_ptr<float>();
    float* p_attn = g_attn.data_ptr<float>();
    float* p_mid = g_mid.data_ptr<float>();
    float* p_r2 = g_resid2.data_ptr<float>();
    float* p_pm = partial_m.data_ptr<float>();
    float* p_pl = partial_l.data_ptr<float>();
    float* p_po = partial_o.data_ptr<float>();
    int* p_pos = d_pos.data_ptr<int>();
    int* p_step = d_step.data_ptr<int>();

    // Reset device counters (async, captured separately from the graph).
    ck(cudaMemcpyAsync(p_pos, &start_pos, sizeof(int), cudaMemcpyHostToDevice, stream), "memcpy d_pos");
    int zero = 0;
    ck(cudaMemcpyAsync(p_step, &zero, sizeof(int), cudaMemcpyHostToDevice, stream), "memcpy d_step");

    void* sig[48] = {};
    int ns = 0;
    sig[ns++] = (void*)h_ptr;
    sig[ns++] = (void*)noise_ptr;
    sig[ns++] = (void*)ba;
    sig[ns++] = (void*)bb;
    sig[ns++] = (void*)mix;
    sig[ns++] = (void*)p_res;
    sig[ns++] = (void*)p_q;
    sig[ns++] = (void*)p_k;
    sig[ns++] = (void*)p_v;
    sig[ns++] = (void*)p_attn;
    sig[ns++] = (void*)p_mid;
    sig[ns++] = (void*)p_r2;
    sig[ns++] = (void*)p_pm;
    sig[ns++] = (void*)p_pl;
    sig[ns++] = (void*)p_po;
    sig[ns++] = (void*)p_pos;
    sig[ns++] = (void*)p_step;
    for (int i = 0; i < num_layers; ++i) {
        sig[ns++] = (void*)kptrs[i];
        sig[ns++] = (void*)vptrs[i];
        sig[ns++] = (void*)layers[i].q_proj;
    }
    sig[ns++] = (void*)(intptr_t)num_chunks;
    sig[ns++] = (void*)(intptr_t)num_layers;
    sig[ns++] = (void*)(intptr_t)max_seq;

    // Eager loop first (correctness). CUDA-graph replay is enabled when
    // capture succeeds on a private stream; otherwise we stay eager.
    bool reuse = g_cache.exec && g_cache.n_sig == ns && g_cache.num_chunks == num_chunks
                 && sig_equal(sig, g_cache.sig, ns);

    if (!reuse && n_steps >= 8) {
        if (g_cache.exec) {
            cudaGraphExecDestroy(g_cache.exec);
            g_cache.exec = nullptr;
        }
        // Isolate capture on a private stream after a full device sync so we
        // never fight a PyTorch stream that is already capturing.
        cudaDeviceSynchronize();
        cudaStream_t cap = nullptr;
        if (cudaStreamCreateWithFlags(&cap, cudaStreamNonBlocking) == cudaSuccess) {
            cudaGraph_t graph = nullptr;
            cudaError_t e0 = cudaStreamBeginCapture(cap, cudaStreamCaptureModeThreadLocal);
            if (e0 == cudaSuccess) {
                one_step_graph(
                    noise_ptr, h_ptr, mix, ba, bb,
                    layers, num_layers, kptrs, vptrs,
                    p_res, p_q, p_k, p_v, p_attn, p_mid, p_r2,
                    p_pm, p_pl, p_po,
                    p_pos, p_step,
                    max_seq, num_chunks, chunk_stride, max_blocks, cap
                );
                cudaError_t e1 = cudaStreamEndCapture(cap, &graph);
                if (e1 == cudaSuccess && graph) {
                    cudaGraphExec_t exec = nullptr;
                    if (cudaGraphInstantiate(&exec, graph, nullptr, nullptr, 0) == cudaSuccess) {
                        g_cache.exec = exec;
                        std::memcpy(g_cache.sig, sig, sizeof(void*) * ns);
                        g_cache.n_sig = ns;
                        g_cache.num_chunks = num_chunks;
                    }
                    cudaGraphDestroy(graph);
                }
            }
            cudaStreamDestroy(cap);
            cudaGetLastError();
        }
        reuse = g_cache.exec != nullptr;
    }

    if (reuse) {
        for (int i = 0; i < n_steps; ++i) {
            ck(cudaGraphLaunch(g_cache.exec, stream), "graph launch");
        }
    } else {
        for (int i = 0; i < n_steps; ++i) {
            one_step_graph(
                noise_ptr, h_ptr, mix, ba, bb,
                layers, num_layers, kptrs, vptrs,
                p_res, p_q, p_k, p_v, p_attn, p_mid, p_r2,
                p_pm, p_pl, p_po,
                p_pos, p_step,
                max_seq, num_chunks, chunk_stride, max_blocks, stream
            );
        }
    }
    ck(cudaGetLastError(), "after steps");
}

PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
    m.def("run_steps_cuda", &run_steps_cuda);
}
'''


def _get_mod():
    global _mod
    if _mod is not None:
        return _mod
    os.environ.setdefault("TORCH_CUDA_ARCH_LIST", "12.0")
    extra = [
        "-O3",
        "--use_fast_math",
        "-U__CUDA_NO_HALF_OPERATORS__",
        "-U__CUDA_NO_HALF_CONVERSIONS__",
        "-U__CUDA_NO_BFLOAT16_CONVERSIONS__",
        "-U__CUDA_NO_BFLOAT16_OPERATORS__",
        "--expt-relaxed-constexpr",
        "-std=c++17",
    ]
    try:
        major, minor = torch.cuda.get_device_capability(0)
        extra.append(f"-gencode=arch=compute_{major}{minor},code=sm_{major}{minor}")
    except Exception:
        pass
    _mod = load_inline(
        name="megaqwen_decode_sm120_v9",
        cpp_sources=[],
        cuda_sources=[_cuda_src()],
        extra_cuda_cflags=extra,
        verbose=False,
    )
    return _mod


class Block(nn.Module):
    def __init__(self):
        super().__init__()
        H, I, D = HIDDEN, INTERMEDIATE, HEAD_DIM
        self.input_ln = nn.Parameter(torch.ones(H, dtype=torch.bfloat16))
        self.q_proj = nn.Parameter(torch.empty(NUM_Q * D, H, dtype=torch.bfloat16))
        self.k_proj = nn.Parameter(torch.empty(NUM_KV * D, H, dtype=torch.bfloat16))
        self.v_proj = nn.Parameter(torch.empty(NUM_KV * D, H, dtype=torch.bfloat16))
        self.q_norm = nn.Parameter(torch.ones(D, dtype=torch.bfloat16))
        self.k_norm = nn.Parameter(torch.ones(D, dtype=torch.bfloat16))
        self.o_proj = nn.Parameter(torch.empty(H, NUM_Q * D, dtype=torch.bfloat16))
        self.post_ln = nn.Parameter(torch.ones(H, dtype=torch.bfloat16))
        self.gate_proj = nn.Parameter(torch.empty(I, H, dtype=torch.bfloat16))
        self.up_proj = nn.Parameter(torch.empty(I, H, dtype=torch.bfloat16))
        self.down_proj = nn.Parameter(torch.empty(H, I, dtype=torch.bfloat16))
        for p in self.parameters():
            if p is self.input_ln or p is self.post_ln or p is self.q_norm or p is self.k_norm:
                continue
            nn.init.normal_(p, std=0.02)


class Model(nn.Module):
    def __init__(self, num_layers: int = NUM_LAYERS, max_seq: int = 131072):
        super().__init__()
        self.num_layers = num_layers
        self.max_seq = max_seq
        self.blocks = nn.ModuleList([Block() for _ in range(num_layers)])
        self._ws = None
        self._max_blocks = None
        self._packed = None
        self._k = None
        self._v = None

    def _ensure_ws(self, device):
        if self._ws is not None and self._ws["residual"].device == device:
            return self._ws
        props = torch.cuda.get_device_properties(device)
        self._max_blocks = int(props.multi_processor_count)
        max_chunks = 32
        self._ws = {
            "residual": torch.empty(HIDDEN, device=device, dtype=torch.float32),
            "g_q": torch.empty(Q_SIZE, device=device, dtype=torch.float32),
            "g_k": torch.empty(KV_SIZE, device=device, dtype=torch.float32),
            "g_v": torch.empty(KV_SIZE, device=device, dtype=torch.float32),
            "g_attn": torch.empty(Q_SIZE, device=device, dtype=torch.float32),
            "g_mid": torch.empty(INTERMEDIATE, device=device, dtype=torch.float32),
            "g_resid2": torch.empty(HIDDEN, device=device, dtype=torch.float32),
            "partial_m": torch.empty(NUM_Q, max_chunks, device=device, dtype=torch.float32),
            "partial_l": torch.empty(NUM_Q, max_chunks, device=device, dtype=torch.float32),
            "partial_o": torch.empty(NUM_Q * max_chunks * HEAD_DIM, device=device, dtype=torch.float32),
            "buf_a": torch.empty(HIDDEN, device=device, dtype=torch.bfloat16),
            "buf_b": torch.empty(HIDDEN, device=device, dtype=torch.bfloat16),
            "mix_buf": torch.empty(HIDDEN, device=device, dtype=torch.bfloat16),
            "d_pos": torch.zeros(1, device=device, dtype=torch.int32),
            "d_step": torch.zeros(1, device=device, dtype=torch.int32),
            "h_work": torch.empty(HIDDEN, device=device, dtype=torch.bfloat16),
        }
        return self._ws

    def _ensure_caches(self, device):
        if self._k is None or self._k[0].device != device or self._k[0].shape[1] != self.max_seq:
            self._k, self._v = empty_caches(self.num_layers, self.max_seq, device)
        return self._k, self._v

    def pack_weights(self):
        chunks = []
        for b in self.blocks:
            for p in (
                b.input_ln, b.q_proj, b.k_proj, b.v_proj, b.q_norm, b.k_norm,
                b.o_proj, b.post_ln, b.gate_proj, b.up_proj, b.down_proj,
            ):
                chunks.append(p.detach().reshape(-1).contiguous())
        self._packed = torch.cat(chunks).contiguous()
        return self._packed

    def weight_list(self):
        out = []
        for b in self.blocks:
            out.extend(
                [
                    b.input_ln,
                    b.q_proj,
                    b.k_proj,
                    b.v_proj,
                    b.q_norm,
                    b.k_norm,
                    b.o_proj,
                    b.post_ln,
                    b.gate_proj,
                    b.up_proj,
                    b.down_proj,
                ]
            )
        return out


def empty_caches(num_layers: int, max_seq: int, device, dtype=torch.bfloat16):
    # Packed [kv, seq, 2, hd] so K and V at a position are 512B sequential.
    # Both lists alias the same storage; the kernel reads k_cache as packed.
    packed = [
        torch.zeros(NUM_KV, max_seq, 2, HEAD_DIM, device=device, dtype=dtype)
        for _ in range(num_layers)
    ]
    return packed, packed


def _seeded_hidden(seed: int, device) -> torch.Tensor:
    g = torch.Generator(device="cpu")
    g.manual_seed(seed)
    return torch.randn(HIDDEN, generator=g, dtype=torch.bfloat16).to(device)


def _run_steps(model: Model, hidden: torch.Tensor, k_caches, v_caches,
               start_pos: int, n_steps: int, noise: torch.Tensor) -> torch.Tensor:
    mod = _get_mod()
    device = hidden.device
    ws = model._ensure_ws(device)
    h_work = ws["h_work"]
    h_work.copy_(hidden)
    # Keep a stable noise pointer so CUDA graphs can be reused across trials.
    nbuf = ws.get("noise")
    if nbuf is None or nbuf.shape[0] < noise.shape[0] or nbuf.device != device:
        ws["noise"] = torch.empty(max(noise.shape[0], 128), HIDDEN, device=device, dtype=torch.bfloat16)
        nbuf = ws["noise"]
    nbuf[: noise.shape[0]].copy_(noise)
    empty_packed = ws.get("empty_packed")
    if empty_packed is None:
        ws["empty_packed"] = torch.empty(0, device=device, dtype=torch.bfloat16)
        empty_packed = ws["empty_packed"]
    mod.run_steps_cuda(
        h_work,
        nbuf,
        model.weight_list(),
        k_caches,
        v_caches,
        ws["buf_a"],
        ws["buf_b"],
        ws["mix_buf"],
        ws["residual"],
        ws["g_q"],
        ws["g_k"],
        ws["g_v"],
        ws["g_attn"],
        ws["g_mid"],
        ws["g_resid2"],
        ws["partial_m"],
        ws["partial_l"],
        ws["partial_o"],
        ws["d_pos"],
        ws["d_step"],
        empty_packed,
        int(start_pos),
        int(n_steps),
        int(model.max_seq),
        int(model.num_layers),
        int(model._max_blocks),
    )
    return h_work


@torch.no_grad()
def prefill(model: Model, ctx_len: int, seed: int, device=None):
    device = device or next(model.parameters()).device
    model = model.to(device).eval()
    assert ctx_len <= model.max_seq
    _get_mod()
    h = _seeded_hidden(seed, device)
    k_caches, v_caches = model._ensure_caches(device)
    for t in k_caches:
        t.zero_()
    for t in v_caches:
        t.zero_()
    g = torch.Generator(device="cpu")
    g.manual_seed(seed + 1)
    noise = torch.randn(ctx_len, HIDDEN, generator=g, dtype=torch.bfloat16, device="cpu").to(device)
    h = _run_steps(model, h, k_caches, v_caches, 0, ctx_len, noise)
    return h, k_caches, v_caches


@torch.no_grad()
def decode_steps(
    model: Model,
    hidden: torch.Tensor,
    k_caches,
    v_caches,
    start_pos: int,
    n_steps: int,
    seed: int,
):
    g = torch.Generator(device="cpu")
    g.manual_seed(seed + 2)
    noise = torch.randn(n_steps, HIDDEN, generator=g, dtype=torch.bfloat16, device="cpu").to(
        hidden.device
    )
    h = _run_steps(model, hidden, k_caches, v_caches, start_pos, n_steps, noise)
    return h, k_caches, v_caches


def run(
    ctx_len: int,
    n_decode: int,
    seed: int,
    model: Optional[Model] = None,
    max_seq: Optional[int] = None,
) -> dict:
    device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
    max_seq = max_seq or max(ctx_len + n_decode, 512)
    if model is None:
        model = Model(NUM_LAYERS, max_seq)
    else:
        if getattr(model, "max_seq", 0) < ctx_len + n_decode:
            raise ValueError(
                f"model.max_seq={getattr(model, 'max_seq', None)} too small for "
                f"ctx_len={ctx_len}+n_decode={n_decode}"
            )
    model = model.to(device).eval()
    h, k_caches, v_caches = prefill(model, ctx_len, seed, device=device)
    h, k_caches, v_caches = decode_steps(
        model, h, k_caches, v_caches, start_pos=ctx_len, n_steps=n_decode, seed=seed
    )
    return {
        "last_hidden": h.detach(),
        "ctx_len": ctx_len,
        "decode_steps": n_decode,
    }

20260814_000332_grok_grok-4.6_03_megaqwen_decode