KernelBench cuda · RTX PRO 6000
MegaQwen Decode GPT-6 Astra Pro
manually audited: clean
GPT-6 Astra Pro through OpenRouter (codex CLI 0.140.0 harness), xhigh effort, unlimited budget, RTX PRO 6000, 1h53m to a voluntary stop. One persistent cooperative megakernel: 376 blocks (2 per SM) x 256 threads, all decode steps x all four layers in a single cudaLaunchCooperativeKernel, six grid.sync barriers per layer (QKV | attention partials | cross-split combine | O-proj+residual | gate/up | down+residual). The 125.8 MB bf16 weight set is streamed with plain __ldg (no eviction policy); the KV cache streams evict-first via inline-PTX ld.global.cs. Attention is true split-KV flash-decoding: 47 splits per kv head with exact online softmax (running max + rescale, no fixed-shift trick), fp32 partials in global workspace, exact 16-block cross-split combine; ctx 131072 uses the same path. RMSNorm/RoPE/SwiGLU fully fused; residuals fp32 in-layer, bf16 between layers; prefill is real, run through the same kernel. Fresh h_out (clone) and fresh seeded noise every call; engine caches hold only weight/cache data_ptr tables, RoPE table, and workspace, and weights are re-read from GPU memory every launch. Agent verified batched torch.randn equals the reference's per-step draw before relying on it. In-run grade 0.0488; isolated sequential trusted grade after session end 0.0492 (+0.8%): 6,545 tok/s at ctx 2048 down to 634 tok/s at ctx 131072. Audit-time isolated regrade 2026-09-07 (clocks reset, sole GPU owner): 0.0490 (6,430 / 4,435 / 2,001 / 634 tok/s at ctx 2048 / 8192 / 32768 / 131072; per shape 0.1286 / 0.0887 / 0.0400 / 0.0127), the 0.0492 grade kept as benchmark.contended.log. Probe (probe.log): primed ctx 2048 dec 16 cos(ref,sol)=1.0000; a second decode_steps call with seed+1 on the same model gives a different output (cos(s1,s1b)=0.3321) that matches its own reference at 1.0000; an in-place weight overwrite on the same parameter buffers flips the output (cos(s1,s2)=-0.0194) and matches the new reference at 1.0000. OpenRouter cost $53.83. All template files byte-identical, patches only to solution.py, no foreign reads, no clock commands, a failed `pip install ninja` (nothing installed), and sanctioned fetching of the MegaQwen baseline repo only. Third on the problem behind claude-opus-5 0.0655 and claude-fable-5-1 0.0643; the gap is the two short shapes, where 24 cooperative grid syncs per token over 376 blocks lose to the leaders' 5-barrier/188-block designs.
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)
"""Persistent SM120 CUDA decode with fused norms and split-KV attention."""
from __future__ import annotations
import os
import sys
import torch
from torch import 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
EPS = 1e-6
class Block(nn.Module):
def __init__(self):
super().__init__()
self.input_ln = nn.Parameter(torch.ones(1024, dtype=torch.bfloat16))
self.q_proj = nn.Parameter(torch.empty(2048, 1024, dtype=torch.bfloat16))
self.k_proj = nn.Parameter(torch.empty(1024, 1024, dtype=torch.bfloat16))
self.v_proj = nn.Parameter(torch.empty(1024, 1024, dtype=torch.bfloat16))
self.q_norm = nn.Parameter(torch.ones(128, dtype=torch.bfloat16))
self.k_norm = nn.Parameter(torch.ones(128, dtype=torch.bfloat16))
self.o_proj = nn.Parameter(torch.empty(1024, 2048, dtype=torch.bfloat16))
self.post_ln = nn.Parameter(torch.ones(1024, dtype=torch.bfloat16))
self.gate_proj = nn.Parameter(torch.empty(3072, 1024, dtype=torch.bfloat16))
self.up_proj = nn.Parameter(torch.empty(3072, 1024, dtype=torch.bfloat16))
self.down_proj = nn.Parameter(torch.empty(1024, 3072, dtype=torch.bfloat16))
for parameter in self.parameters():
if parameter.ndim == 2:
nn.init.normal_(parameter, 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._engine = None
_CPP = r"""
#include <torch/extension.h>
void launch_decode(torch::Tensor weights, torch::Tensor caches,
torch::Tensor hidden, torch::Tensor noise,
torch::Tensor rope, torch::Tensor workspace,
int64_t start, int64_t steps, int64_t layers,
int64_t capacity, int64_t blocks);
int64_t resident_blocks();
"""
_CUDA = r"""
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAException.h>
#include <cuda_bf16.h>
#include <math_constants.h>
#include <cooperative_groups.h>
namespace cg = cooperative_groups;
using bf16 = __nv_bfloat16;
struct Layer {
const bf16* input_ln;
const bf16* q_proj;
const bf16* k_proj;
const bf16* v_proj;
const bf16* q_norm;
const bf16* k_norm;
const bf16* o_proj;
const bf16* post_ln;
const bf16* gate_proj;
const bf16* up_proj;
const bf16* down_proj;
};
__device__ __forceinline__ float warp_sum(float value) {
#pragma unroll
for (int offset = 16; offset; offset >>= 1)
value += __shfl_xor_sync(0xffffffff, value, offset);
return value;
}
__device__ __forceinline__ float4 unpack(uint2 bits) {
bf16* values = reinterpret_cast<bf16*>(&bits);
return make_float4(__bfloat162float(values[0]), __bfloat162float(values[1]),
__bfloat162float(values[2]), __bfloat162float(values[3]));
}
__device__ __forceinline__ float4 weight4(const bf16* address) {
return unpack(__ldg(reinterpret_cast<const uint2*>(address)));
}
__device__ __forceinline__ uint2 stream2(const bf16* address) {
uint2 bits;
asm volatile("ld.global.cs.v2.u32 {%0, %1}, [%2];"
: "=r"(bits.x), "=r"(bits.y) : "l"(address));
return bits;
}
__device__ __forceinline__ float dot4(float4 left, float4 right, float sum) {
sum = fmaf(left.x, right.x, sum);
sum = fmaf(left.y, right.y, sum);
sum = fmaf(left.z, right.z, sum);
return fmaf(left.w, right.w, sum);
}
__device__ __forceinline__ float mix_value(const bf16* hidden, const bf16* noise,
int index, bool mix) {
float value = __bfloat162float(hidden[index]);
if (mix) {
float left = __bfloat162float(__float2bfloat16_rn(0.5f * value));
float right = __bfloat162float(__float2bfloat16_rn(
0.5f * __bfloat162float(noise[index])));
value = __bfloat162float(__float2bfloat16_rn(left + right));
}
return value;
}
__device__ __forceinline__ void normalized_input(
float4 (&activation)[8], const bf16* hidden, const bf16* noise,
const bf16* norm, bool mix) {
int lane = threadIdx.x & 31;
float square_sum = 0.0f;
#pragma unroll
for (int chunk = 0; chunk < 8; ++chunk) {
int column = chunk * 128 + lane * 4;
float4 values = make_float4(
mix_value(hidden, noise, column, mix),
mix_value(hidden, noise, column + 1, mix),
mix_value(hidden, noise, column + 2, mix),
mix_value(hidden, noise, column + 3, mix));
activation[chunk] = values;
square_sum = dot4(values, values, square_sum);
}
float scale = rsqrtf(warp_sum(square_sum) * (1.0f / 1024.0f) + 1e-6f);
#pragma unroll
for (int chunk = 0; chunk < 8; ++chunk) {
float4 norm_values = weight4(norm + chunk * 128 + lane * 4);
activation[chunk].x = (activation[chunk].x * scale) * norm_values.x;
activation[chunk].y = (activation[chunk].y * scale) * norm_values.y;
activation[chunk].z = (activation[chunk].z * scale) * norm_values.z;
activation[chunk].w = (activation[chunk].w * scale) * norm_values.w;
}
}
__device__ __forceinline__ void normalized_post(
float4 (&activation)[8], const float* residual, const bf16* norm) {
int lane = threadIdx.x & 31;
float square_sum = 0.0f;
#pragma unroll
for (int chunk = 0; chunk < 8; ++chunk) {
float4 values = reinterpret_cast<const float4*>(residual)[chunk * 32 + lane];
activation[chunk] = values;
square_sum = dot4(values, values, square_sum);
}
float scale = rsqrtf(warp_sum(square_sum) * (1.0f / 1024.0f) + 1e-6f);
#pragma unroll
for (int chunk = 0; chunk < 8; ++chunk) {
float4 norm_values = weight4(norm + chunk * 128 + lane * 4);
activation[chunk].x = (activation[chunk].x * scale) * norm_values.x;
activation[chunk].y = (activation[chunk].y * scale) * norm_values.y;
activation[chunk].z = (activation[chunk].z * scale) * norm_values.z;
activation[chunk].w = (activation[chunk].w * scale) * norm_values.w;
}
}
__device__ __forceinline__ void store4(bf16* address, float4 value) {
uint2 bits;
auto pairs = reinterpret_cast<__nv_bfloat162*>(&bits);
pairs[0] = __floats2bfloat162_rn(value.x, value.y);
pairs[1] = __floats2bfloat162_rn(value.z, value.w);
*reinterpret_cast<uint2*>(address) = bits;
}
__device__ __forceinline__ void normalize_rotate(
float4 (&values)[4], const bf16* norm, const float* angles, int group_lane) {
float variance = 0.0f;
#pragma unroll
for (int chunk = 0; chunk < 4; ++chunk)
variance = dot4(values[chunk], values[chunk], variance);
unsigned mask = 0xffu << (threadIdx.x & 24);
#pragma unroll
for (int offset = 4; offset; offset >>= 1)
variance += __shfl_xor_sync(mask, variance, offset, 8);
float scale = rsqrtf(variance * (1.0f / 128.0f) + 1e-6f);
#pragma unroll
for (int chunk = 0; chunk < 4; ++chunk) {
float4 learned = weight4(norm + chunk * 32 + group_lane * 4);
values[chunk].x = (values[chunk].x * scale) * learned.x;
values[chunk].y = (values[chunk].y * scale) * learned.y;
values[chunk].z = (values[chunk].z * scale) * learned.z;
values[chunk].w = (values[chunk].w * scale) * learned.w;
}
#pragma unroll
for (int chunk = 0; chunk < 2; ++chunk) {
float4 cosine = reinterpret_cast<const float4*>(angles)[chunk * 8 + group_lane];
float4 sine = reinterpret_cast<const float4*>(angles + 64)[chunk * 8 + group_lane];
float4 first = values[chunk];
float4 second = values[chunk + 2];
values[chunk].x = __fsub_rn(__fmul_rn(first.x, cosine.x), __fmul_rn(second.x, sine.x));
values[chunk].y = __fsub_rn(__fmul_rn(first.y, cosine.y), __fmul_rn(second.y, sine.y));
values[chunk].z = __fsub_rn(__fmul_rn(first.z, cosine.z), __fmul_rn(second.z, sine.z));
values[chunk].w = __fsub_rn(__fmul_rn(first.w, cosine.w), __fmul_rn(second.w, sine.w));
values[chunk + 2].x = __fadd_rn(__fmul_rn(first.x, sine.x), __fmul_rn(second.x, cosine.x));
values[chunk + 2].y = __fadd_rn(__fmul_rn(first.y, sine.y), __fmul_rn(second.y, cosine.y));
values[chunk + 2].z = __fadd_rn(__fmul_rn(first.z, sine.z), __fmul_rn(second.z, cosine.z));
values[chunk + 2].w = __fadd_rn(__fmul_rn(first.w, sine.w), __fmul_rn(second.w, cosine.w));
}
}
__global__ __launch_bounds__(256, 2) void persistent_decode(
const Layer* __restrict__ weights, bf16* const* __restrict__ caches,
bf16* __restrict__ hidden, const bf16* __restrict__ noise,
const float* __restrict__ rope, float* __restrict__ workspace,
int start, int steps, int layers, int capacity) {
cg::grid_group grid = cg::this_grid();
int thread = threadIdx.x;
int lane = thread & 31;
int warp = thread >> 5;
int block = blockIdx.x;
int blocks = gridDim.x;
int splits = blocks / 8;
float* qkv = workspace;
float* attention = qkv + 4096;
float* residual = attention + 2048;
float* intermediate = residual + 1024;
float* partial = intermediate + 3072;
__shared__ float shared[32 * 2 * 132];
for (int step = 0; step < steps; ++step) {
int position = start + step;
const bf16* step_noise = noise + int64_t(step) * 1024;
for (int layer_index = 0; layer_index < layers; ++layer_index) {
const Layer& layer = weights[layer_index];
bool mix = layer_index == 0;
bf16* key_cache = caches[layer_index * 2];
bf16* value_cache = caches[layer_index * 2 + 1];
{
float4 activation[8];
normalized_input(activation, hidden, step_noise, layer.input_ln, mix);
for (int row = block * 8 + warp; row < 4096; row += blocks * 8) {
const bf16* matrix = row < 2048 ? layer.q_proj :
row < 3072 ? layer.k_proj : layer.v_proj;
int matrix_row = row < 2048 ? row : row & 1023;
const bf16* row_data = matrix + matrix_row * 1024;
float sum = 0.0f;
#pragma unroll
for (int chunk = 0; chunk < 8; ++chunk)
sum = dot4(weight4(row_data + chunk * 128 + lane * 4),
activation[chunk], sum);
sum = warp_sum(sum);
if (lane == 0) qkv[row] = sum;
}
}
grid.sync();
{
int kv_head = block % 8;
int split = block / 8;
int begin = (int64_t(position + 1) * split) / splits;
int end = (int64_t(position + 1) * (split + 1)) / splits;
int group = thread / 8;
int group_lane = thread % 8;
unsigned mask = 0xffu << (lane & 24);
float4 query_first[4];
float4 query_second[4];
float4 accumulator_first[4];
float4 accumulator_second[4];
#pragma unroll
for (int chunk = 0; chunk < 4; ++chunk) {
query_first[chunk] = reinterpret_cast<float4*>(qkv + kv_head * 256)[chunk * 8 + group_lane];
query_second[chunk] = reinterpret_cast<float4*>(qkv + kv_head * 256 + 128)[chunk * 8 + group_lane];
accumulator_first[chunk] = make_float4(0, 0, 0, 0);
accumulator_second[chunk] = make_float4(0, 0, 0, 0);
}
const float* angles = rope + int64_t(position) * 128;
normalize_rotate(query_first, layer.q_norm, angles, group_lane);
normalize_rotate(query_second, layer.q_norm, angles, group_lane);
if (split == splits - 1 && group == 0) {
float4 key[4];
#pragma unroll
for (int chunk = 0; chunk < 4; ++chunk)
key[chunk] = reinterpret_cast<float4*>(qkv + 2048 + kv_head * 128)[chunk * 8 + group_lane];
normalize_rotate(key, layer.k_norm, angles, group_lane);
#pragma unroll
for (int chunk = 0; chunk < 4; ++chunk) {
int dimension = chunk * 32 + group_lane * 4;
int64_t offset = (int64_t(kv_head) * capacity + position) * 128 + dimension;
store4(key_cache + offset, key[chunk]);
store4(value_cache + offset,
reinterpret_cast<float4*>(qkv + 3072 + kv_head * 128)[chunk * 8 + group_lane]);
}
}
__syncthreads();
float maximum_first = -CUDART_INF_F;
float maximum_second = -CUDART_INF_F;
float denominator_first = 0.0f;
float denominator_second = 0.0f;
for (int [REDACTED credential assignment] + group; token < end; token += 64) {
uint2 key_bits[2][4];
uint2 value_bits[2][4];
#pragma unroll
for (int item = 0; item < 2; ++item) {
#pragma unroll
for (int chunk = 0; chunk < 4; ++chunk) {
int64_t offset = (int64_t(kv_head) * capacity + token + item * 32) * 128 + chunk * 32 + group_lane * 4;
if (token + item * 32 < end) {
key_bits[item][chunk] = stream2(key_cache + offset);
value_bits[item][chunk] = stream2(value_cache + offset);
}
}
}
#pragma unroll
for (int item = 0; item < 2; ++item) {
if (token + item * 32 < end) {
float score_first = 0.0f;
float score_second = 0.0f;
#pragma unroll
for (int chunk = 0; chunk < 4; ++chunk) {
float4 key = unpack(key_bits[item][chunk]);
score_first = dot4(query_first[chunk], key, score_first);
score_second = dot4(query_second[chunk], key, score_second);
}
#pragma unroll
for (int offset = 4; offset; offset >>= 1) {
score_first += __shfl_xor_sync(mask, score_first, offset, 8);
score_second += __shfl_xor_sync(mask, score_second, offset, 8);
}
score_first *= 0.08838834764831845f;
score_second *= 0.08838834764831845f;
float next_first = fmaxf(maximum_first, score_first);
float next_second = fmaxf(maximum_second, score_second);
float correction_first = __expf(maximum_first - next_first);
float correction_second = __expf(maximum_second - next_second);
float probability_first = __expf(score_first - next_first);
float probability_second = __expf(score_second - next_second);
denominator_first = denominator_first * correction_first + probability_first;
denominator_second = denominator_second * correction_second + probability_second;
#pragma unroll
for (int chunk = 0; chunk < 4; ++chunk) {
float4 value = unpack(value_bits[item][chunk]);
accumulator_first[chunk].x = accumulator_first[chunk].x * correction_first + probability_first * value.x;
accumulator_first[chunk].y = accumulator_first[chunk].y * correction_first + probability_first * value.y;
accumulator_first[chunk].z = accumulator_first[chunk].z * correction_first + probability_first * value.z;
accumulator_first[chunk].w = accumulator_first[chunk].w * correction_first + probability_first * value.w;
accumulator_second[chunk].x = accumulator_second[chunk].x * correction_second + probability_second * value.x;
accumulator_second[chunk].y = accumulator_second[chunk].y * correction_second + probability_second * value.y;
accumulator_second[chunk].z = accumulator_second[chunk].z * correction_second + probability_second * value.z;
accumulator_second[chunk].w = accumulator_second[chunk].w * correction_second + probability_second * value.w;
}
maximum_first = next_first;
maximum_second = next_second;
}
}
}
float* first_state = shared + group * 264;
float* second_state = first_state + 132;
#pragma unroll
for (int chunk = 0; chunk < 4; ++chunk) {
reinterpret_cast<float4*>(first_state)[chunk * 8 + group_lane] = accumulator_first[chunk];
reinterpret_cast<float4*>(second_state)[chunk * 8 + group_lane] = accumulator_second[chunk];
}
if (group_lane == 0) {
first_state[128] = maximum_first;
first_state[129] = denominator_first;
second_state[128] = maximum_second;
second_state[129] = denominator_second;
}
__syncthreads();
int query_in_pair = thread / 128;
int dimension = thread % 128;
float maximum = -CUDART_INF_F;
#pragma unroll
for (int other_group = 0; other_group < 32; ++other_group)
maximum = fmaxf(maximum, shared[other_group * 264 + query_in_pair * 132 + 128]);
float numerator = 0.0f;
float denominator = 0.0f;
#pragma unroll
for (int other_group = 0; other_group < 32; ++other_group) {
const float* state = shared + other_group * 264 + query_in_pair * 132;
float correction = state[129] > 0 ? __expf(state[128] - maximum) : 0.0f;
numerator += state[dimension] * correction;
denominator += state[129] * correction;
}
float* destination = partial + block * 264 + query_in_pair * 132;
destination[dimension] = numerator;
if (dimension == 0) {
destination[128] = maximum;
destination[129] = denominator;
}
}
grid.sync();
if (block < 16) {
int kv_head = block / 2;
int query_in_pair = block % 2;
const float* source = partial + kv_head * 264 + query_in_pair * 132;
float maximum = thread < splits ? source[thread * 2112 + 128] : -CUDART_INF_F;
#pragma unroll
for (int offset = 16; offset; offset >>= 1)
maximum = fmaxf(maximum, __shfl_xor_sync(0xffffffff, maximum, offset));
if (lane == 0) shared[warp] = maximum;
__syncthreads();
maximum = -CUDART_INF_F;
#pragma unroll
for (int other_warp = 0; other_warp < 8; ++other_warp)
maximum = fmaxf(maximum, shared[other_warp]);
float denominator = 0.0f;
if (thread < splits) {
const float* state = source + thread * 2112;
float correction = state[129] > 0 ? __expf(state[128] - maximum) : 0.0f;
shared[32 + thread] = correction;
denominator = state[129] * correction;
}
denominator = warp_sum(denominator);
if (lane == 0) shared[16 + warp] = denominator;
__syncthreads();
float4 numerator = make_float4(0, 0, 0, 0);
for (int split = warp; split < splits; split += 8) {
float4 values = reinterpret_cast<const float4*>(source + split * 2112)[lane];
float correction = shared[32 + split];
numerator.x += values.x * correction;
numerator.y += values.y * correction;
numerator.z += values.z * correction;
numerator.w += values.w * correction;
}
reinterpret_cast<float4*>(shared + 256 + warp * 128)[lane] = numerator;
__syncthreads();
if (thread < 128) {
float total = 0.0f;
denominator = 0.0f;
#pragma unroll
for (int other_warp = 0; other_warp < 8; ++other_warp) {
total += shared[256 + other_warp * 128 + thread];
denominator += shared[16 + other_warp];
}
attention[block * 128 + thread] = total / denominator;
}
}
grid.sync();
for (int row = block; row < 1024; row += blocks) {
float sum = 0.0f;
#pragma unroll
for (int chunk = 0; chunk < 2; ++chunk) {
int column = chunk * 1024 + thread * 4;
sum = dot4(weight4(layer.o_proj + row * 2048 + column),
reinterpret_cast<float4*>(attention)[column / 4], sum);
}
sum = warp_sum(sum);
if (lane == 0) shared[warp] = sum;
__syncthreads();
if (thread == 0) {
float total = 0.0f;
#pragma unroll
for (int other_warp = 0; other_warp < 8; ++other_warp)
total += shared[other_warp];
residual[row] = total + mix_value(hidden, step_noise, row, mix);
}
__syncthreads();
}
grid.sync();
{
float4 activation[8];
normalized_post(activation, residual, layer.post_ln);
for (int row = block * 8 + warp; row < 3072; row += blocks * 8) {
float gate = 0.0f;
float up = 0.0f;
#pragma unroll
for (int chunk = 0; chunk < 8; ++chunk) {
int column = chunk * 128 + lane * 4;
gate = dot4(weight4(layer.gate_proj + row * 1024 + column), activation[chunk], gate);
up = dot4(weight4(layer.up_proj + row * 1024 + column), activation[chunk], up);
}
gate = warp_sum(gate);
up = warp_sum(up);
if (lane == 0) intermediate[row] = (gate / (1.0f + __expf(-gate))) * up;
}
}
grid.sync();
for (int row = block; row < 1024; row += blocks) {
float sum = 0.0f;
#pragma unroll
for (int chunk = 0; chunk < 3; ++chunk) {
int column = chunk * 1024 + thread * 4;
sum = dot4(weight4(layer.down_proj + row * 3072 + column),
reinterpret_cast<float4*>(intermediate)[column / 4], sum);
}
sum = warp_sum(sum);
if (lane == 0) shared[warp] = sum;
__syncthreads();
if (thread == 0) {
float total = 0.0f;
#pragma unroll
for (int other_warp = 0; other_warp < 8; ++other_warp)
total += shared[other_warp];
hidden[row] = __float2bfloat16_rn(total + residual[row]);
}
__syncthreads();
}
grid.sync();
}
}
}
int64_t resident_blocks() {
int device;
C10_CUDA_CHECK(cudaGetDevice(&device));
cudaDeviceProp properties;
C10_CUDA_CHECK(cudaGetDeviceProperties(&properties, device));
int occupancy;
C10_CUDA_CHECK(cudaOccupancyMaxActiveBlocksPerMultiprocessor(
&occupancy, persistent_decode, 256, 0));
return (min(properties.multiProcessorCount * occupancy, 512) / 8) * 8;
}
void launch_decode(torch::Tensor weights, torch::Tensor caches,
torch::Tensor hidden, torch::Tensor noise,
torch::Tensor rope, torch::Tensor workspace,
int64_t start, int64_t steps, int64_t layers,
int64_t capacity, int64_t blocks) {
auto weights_pointer = reinterpret_cast<const Layer*>(weights.data_ptr<int64_t>());
auto caches_pointer = reinterpret_cast<bf16**>(caches.data_ptr<int64_t>());
auto hidden_pointer = reinterpret_cast<bf16*>(hidden.data_ptr<at::BFloat16>());
auto noise_pointer = reinterpret_cast<const bf16*>(noise.data_ptr<at::BFloat16>());
auto rope_pointer = rope.data_ptr<float>();
auto workspace_pointer = workspace.data_ptr<float>();
int first = start, count = steps, depth = layers, stride = capacity;
void* arguments[] = {&weights_pointer, &caches_pointer, &hidden_pointer,
&noise_pointer, &rope_pointer, &workspace_pointer,
&first, &count, &depth, &stride};
C10_CUDA_CHECK(cudaLaunchCooperativeKernel(
reinterpret_cast<void*>(persistent_decode), dim3(blocks), dim3(256),
arguments, 0, at::cuda::getCurrentCUDAStream()));
}
"""
_extension = None
def _cuda_extension():
global _extension
if _extension is None:
os.environ.setdefault("TORCH_CUDA_ARCH_LIST", "12.0")
os.environ["PATH"] = os.path.join(sys.prefix, "bin") + os.pathsep + os.environ["PATH"]
_extension = load_inline(
name="megaqwen_sm120_decode",
cpp_sources=_CPP,
cuda_sources=_CUDA,
functions=["launch_decode", "resident_blocks"],
extra_cflags=["-O3"],
extra_cuda_cflags=["-O3", "-lineinfo", "--ptxas-options=-v"],
)
return _extension
class _Engine:
def __init__(self, model):
self.extension = _cuda_extension()
self.device = next(model.parameters()).device
self.signature = tuple(parameter.data_ptr() for parameter in model.parameters())
self.weights = torch.tensor(self.signature, dtype=torch.int64, device=self.device)
self.blocks = self.extension.resident_blocks()
self.workspace = torch.empty(
10240 + self.blocks * 264, dtype=torch.float32, device=self.device
)
inverse = 1.0 / (
10000 ** (torch.arange(64, dtype=torch.float32, device=self.device) / 64)
)
positions = torch.arange(model.max_seq, dtype=torch.float32, device=self.device)
angles = positions[:, None] * inverse[None, :]
self.rope = torch.cat((angles.cos(), angles.sin()), dim=1).contiguous()
self.cache_signature = None
self.cache_pointers = None
def bind(self, keys, values):
signature = tuple(pointer.data_ptr() for pair in zip(keys, values) for pointer in pair)
if signature != self.cache_signature:
self.cache_pointers = torch.tensor(signature, dtype=torch.int64, device=self.device)
self.cache_signature = signature
def execute(self, model, hidden, noise, keys, values, start):
self.bind(keys, values)
self.extension.launch_decode(
self.weights, self.cache_pointers, hidden, noise, self.rope,
self.workspace, start, noise.shape[0], model.num_layers,
model.max_seq, self.blocks,
)
def _engine(model):
signature = tuple(parameter.data_ptr() for parameter in model.parameters())
if model._engine is None or model._engine.signature != signature:
model._engine = _Engine(model)
return model._engine
def _noise(count, seed, device):
generator = torch.Generator(device="cpu").manual_seed(seed)
return torch.randn(
(count, HIDDEN), generator=generator, dtype=torch.bfloat16
).to(device)
@torch.no_grad()
def prefill(model, ctx_len, seed, device=None):
if device is not None:
model = model.to(device)
device = next(model.parameters()).device
if ctx_len < 0 or ctx_len > model.max_seq:
raise ValueError("context does not fit the KV cache")
engine = _engine(model)
hidden = _noise(1, seed, device).reshape(HIDDEN)
keys = [
torch.zeros((NUM_KV, model.max_seq, HEAD_DIM), dtype=torch.bfloat16, device=device)
for _ in range(model.num_layers)
]
values = [torch.zeros_like(cache) for cache in keys]
if ctx_len:
noise = _noise(ctx_len, seed + 1, device)
engine.execute(model, hidden, noise, keys, values, 0)
else:
engine.bind(keys, values)
return hidden, keys, values
@torch.no_grad()
def decode_steps(model, hidden, k_caches, v_caches, start_pos, n_steps, seed):
if start_pos < 0 or n_steps < 0 or start_pos + n_steps > model.max_seq:
raise ValueError("decode positions do not fit the KV cache")
result = hidden.clone()
if n_steps:
engine = _engine(model)
noise = _noise(n_steps, seed + 2, result.device)
engine.execute(model, result, noise, k_caches, v_caches, start_pos)
return result, k_caches, v_caches
@torch.no_grad()
def run(ctx_len, decode_steps, seed, model=None):
if model is None:
model = Model(NUM_LAYERS, max(ctx_len + decode_steps, 512)).cuda().eval()
hidden, keys, values = prefill(model, ctx_len, seed)
hidden, keys, values = globals()["decode_steps"](
model, hidden, keys, values, ctx_len, decode_steps, seed
)
return {"last_hidden": hidden}
20260907_004320_codex_openai_gpt-6-astra-pro_03_megaqwen_decode