"""MegaQwen-style Qwen3-0.6B 4-layer decode, custom CUDA on RTX PRO 6000 (SM120). Memory-bound autoregressive decode. Everything is a hand-written CUDA kernel: mix -> [ RMSNorm+QKV gemv -> Q/K RMSNorm+RoPE+KV store -> split-KV flash attn -> O proj + residual -> RMSNorm+gate/up gemv -> SwiGLU -> down + residual ] x4 Numerics mirror reference.py: fp32 math with bf16 weights upcast, K/V rounded to bf16 in cache, block output rounded to bf16. A host loop launches the per-step kernels with `pos` as a plain argument (no device counters). """ from __future__ import annotations import glob import os import site import torch import torch.nn as nn from torch.utils.cpp_extension import load_inline def _setup_cuda_toolkit() -> str | None: """Point CUDA_HOME at the venv's cu13 toolkit. torch here is a cu130 build, but the box only ships a system CUDA 12.8 toolkit (major mismatch -> cpp_extension hard-fails). The matching cu13 nvcc/cccl/nvvm/crt are pip-installed into nvidia/cu13; wire CUDA_HOME to it. """ cu = None try: import nvidia for base in list(getattr(nvidia, "__path__", [])): cand = os.path.join(base, "cu13") if os.path.exists(os.path.join(cand, "bin", "nvcc")): cu = cand break except Exception: pass if cu is None: roots = list(site.getsitepackages()) if hasattr(site, "getsitepackages") else [] roots += [p for p in os.sys.path if p.endswith("site-packages")] for sp in roots: cand = os.path.join(sp, "nvidia", "cu13") if os.path.exists(os.path.join(cand, "bin", "nvcc")): cu = cand break if cu is None: return None lib, lib64 = os.path.join(cu, "lib"), os.path.join(cu, "lib64") if os.path.isdir(lib) and not os.path.exists(lib64): try: os.symlink("lib", lib64) except OSError: pass # The pip runtime ships only versioned libs (libcudart.so.13); the linker's # `-lcudart` needs an unversioned .so symlink. for versioned in glob.glob(os.path.join(lib, "lib*.so.*")): stem = versioned while "." in os.path.basename(stem) and not stem.endswith(".so"): stem = stem.rsplit(".", 1)[0] if not stem.endswith(".so"): continue if not os.path.exists(stem): try: os.symlink(os.path.basename(versioned), stem) except OSError: pass os.environ["CUDA_HOME"] = cu os.environ["CUDA_PATH"] = cu os.environ["PATH"] = os.path.join(cu, "bin") + os.pathsep + os.environ.get("PATH", "") # torch.cpp_extension caches CUDA_HOME at import (resolved from `which nvcc`, # which is the GPU-lock wrapper on PATH -> wrong/12.8). Force it to cu13 so # load_inline uses cu13/bin/nvcc directly (matches torch's cu130 build). import torch.utils.cpp_extension as _cpp _cpp.CUDA_HOME = cu return cu OP_TYPE = "megaqwen_decode" HIDDEN = 1024 INTERMEDIATE = 3072 NUM_Q = 16 NUM_KV = 8 HEAD_DIM = 128 NUM_LAYERS = 4 EPS = 1e-6 MAXSPLIT = 2048 # ---------------------------------------------------------------------------- # CUDA # ---------------------------------------------------------------------------- CUDA_SRC = r""" #include #include #include #include typedef __nv_bfloat16 bf16; #define H 1024 #define II 3072 #define NQ 16 #define NKV 8 #define DH 128 #define HALF 64 #define NL 4 #define EPSILON 1e-6f #define SCALE 0.08838834764831845f // 1/sqrt(128) #define MAXSPLIT 2048 #define PW 132 // padded partial stride: [m, l, _, _, acc[128]] #define PACC 4 // acc offset within a partial slot (16-byte aligned) __device__ __forceinline__ float b2f(bf16 x){ return __bfloat162float(x); } __device__ __forceinline__ bf16 f2b(float x){ return __float2bfloat16(x); } __device__ __forceinline__ float to_f(float x){ return x; } __device__ __forceinline__ float to_f(bf16 x){ return __bfloat162float(x); } __device__ __forceinline__ void store_out(float* p, int i, float v){ p[i]=v; } __device__ __forceinline__ void store_out(bf16* p, int i, float v){ p[i]=f2b(v); } // ---- mix: blk0 = round_bf16(0.5*randn + 0.5*hidden) ---- __global__ void mix_kernel(bf16* __restrict__ blk0, const bf16* __restrict__ rnd, const bf16* __restrict__ hid){ int i = blockIdx.x*blockDim.x + threadIdx.x; if(i fp32 h_normed. 1 block, 256 threads. ---- template __global__ void rmsnorm_kernel(float* __restrict__ out, const T* __restrict__ x_in, const bf16* __restrict__ nw){ __shared__ float red[8]; int tid=threadIdx.x, lane=tid&31, warp=tid>>5; float v0=to_f(x_in[tid]), v1=to_f(x_in[tid+256]), v2=to_f(x_in[tid+512]), v3=to_f(x_in[tid+768]); float p=v0*v0+v1*v1+v2*v2+v3*v3; for(int o=16;o>0;o>>=1) p+=__shfl_down_sync(0xffffffffu,p,o); if(lane==0) red[warp]=p; __syncthreads(); if(warp==0){ float t=(lane<8)?red[lane]:0.f; for(int o=4;o>0;o>>=1) t+=__shfl_down_sync(0xffffffffu,t,o); if(lane==0) red[0]=t; } __syncthreads(); float inv=rsqrtf(red[0]/(float)H + EPSILON); out[tid] = v0*inv*b2f(nw[tid]); out[tid+256] = v1*inv*b2f(nw[tid+256]); out[tid+512] = v2*inv*b2f(nw[tid+512]); out[tid+768] = v3*inv*b2f(nw[tid+768]); } // ---- pure QKV gemv over pre-normed h [H]. grid 512, block 256 (1 warp/row). ---- __global__ void qkv_kernel( float* __restrict__ q_raw, float* __restrict__ k_raw, float* __restrict__ v_raw, const bf16* __restrict__ Wq, const bf16* __restrict__ Wk, const bf16* __restrict__ Wv, const float* __restrict__ h){ __shared__ float sh[H]; int tid=threadIdx.x, lane=tid&31, warp=tid>>5; for(int i=tid;i=4096) return; const bf16* wr; float* outp; int oi; if(grow<2048){ wr=Wq+(size_t)grow*H; outp=q_raw; oi=grow; } else if(grow<3072){ int r=grow-2048; wr=Wk+(size_t)r*H; outp=k_raw; oi=r; } else { int r=grow-3072; wr=Wv+(size_t)r*H; outp=v_raw; oi=r; } float acc=0.f; for(int i=lane*8;i(wr+i); const bf16* wb=reinterpret_cast(&wv); #pragma unroll for(int t=0;t<8;t++) acc+=b2f(wb[t])*sh[i+t]; } for(int o=16;o>0;o>>=1) acc+=__shfl_down_sync(0xffffffffu,acc,o); if(lane==0) outp[oi]=acc; } // ---- pure gate/up gemv over pre-normed h [H]. grid 768 (6144 rows / 8). ---- // gate_up layout: [0,3072)=gate, [3072,6144)=up __global__ void gateup_kernel( float* __restrict__ gate_up, const bf16* __restrict__ Wg, const bf16* __restrict__ Wu, const float* __restrict__ h){ __shared__ float sh[H]; int tid=threadIdx.x, lane=tid&31, warp=tid>>5; for(int i=tid;i=6144) return; const bf16* wr; int oi; if(grow<3072){ wr=Wg+(size_t)grow*H; oi=grow; } else { int r=grow-3072; wr=Wu+(size_t)r*H; oi=3072+r; } float acc=0.f; for(int i=lane*8;i(wr+i); const bf16* wb=reinterpret_cast(&wv); #pragma unroll for(int t=0;t<8;t++) acc+=b2f(wb[t])*sh[i+t]; } for(int o=16;o>0;o>>=1) acc+=__shfl_down_sync(0xffffffffu,acc,o); if(lane==0) gate_up[oi]=acc; } // ---- plain gemv with residual add. out[row] = resid[row] + sum_j W[row,j]*x[j] ---- // K-split gemv with residual add. block=256 (8 warps); SK warps share a row so // M small (=H) still yields 8/SK rows/block * M/(8/SK) blocks = M*SK warps. template __global__ void gemv_resid_kernel(OUT* __restrict__ out, const bf16* __restrict__ W, const float* __restrict__ x_in, const RES* __restrict__ resid, int M, int K){ extern __shared__ float sh[]; // sh[0..K-1]=x; sh[K..]=partials[ROWS*SK] const int ROWS = 8/SK; int tid=threadIdx.x, lane=tid&31, warp=tid>>5; for(int i=tid;i=M) return; const bf16* wr=W+(size_t)grow*K; int kslice=K/SK, kstart=ks*kslice, kend=kstart+kslice; float acc=0.f; for(int i=kstart+lane*8;i(wr+i); const bf16* wb=reinterpret_cast(&wv); #pragma unroll for(int t=0;t<8;t++) acc+=b2f(wb[t])*sh[i+t]; } for(int o=16;o>0;o>>=1) acc+=__shfl_down_sync(0xffffffffu,acc,o); float* part=sh+K; if(lane==0) part[r*SK+ks]=acc; __syncthreads(); if(ks==0 && lane==0){ float s=0.f; #pragma unroll for(int j=0;j>5; for(int o=16;o>0;o>>=1) ss+=__shfl_down_sync(0xffffffffu,ss,o); if(lane==0) red[warp]=ss; __syncthreads(); if(warp==0){ float t=(lane<4)?red[lane]:0.f; for(int o=2;o>0;o>>=1) t+=__shfl_down_sync(0xffffffffu,t,o); if(lane==0) red[0]=t; } __syncthreads(); float inv=rsqrtf(red[0]/(float)DH + EPSILON); sh[d]=val*inv*b2f(nw[d]); __syncthreads(); float outv; if(d high occupancy, small nsplit. __global__ void attn_partial_kernel( float* __restrict__ partials, const float* __restrict__ q_out, const bf16* __restrict__ kc_layer, const bf16* __restrict__ vc_layer, int seqlen, int nsplit, int max_seq, int W){ int blk = blockIdx.x; int kvhead = blk % NKV; int split = blk / NKV; int warp = threadIdx.x >> 5, lane = threadIdx.x & 31, dd0 = lane*4; int qh0=2*kvhead, qh1=qh0+1; int psplit = (seqlen + nsplit - 1)/nsplit; int p0 = split*psplit, p1 = min(p0+psplit, seqlen); // this warp's contiguous slice of the split int np = p1 - p0; if(np<0) np=0; int wc = (np + W - 1)/W; int wp0 = p0 + warp*wc, wp1 = min(wp0+wc, p1); float4 q0 = *reinterpret_cast(q_out + qh0*DH + dd0); float4 q1 = *reinterpret_cast(q_out + qh1*DH + dd0); float m0=-1e30f,l0=0.f,a0x=0,a0y=0,a0z=0,a0w=0; float m1=-1e30f,l1=0.f,a1x=0,a1y=0,a1z=0,a1w=0; const bf16* kh=kc_layer + (size_t)kvhead*max_seq*DH + dd0; const bf16* vh=vc_layer + (size_t)kvhead*max_seq*DH + dd0; #pragma unroll 4 for(int p=wp0;p(kh + (size_t)p*DH); const bf16* kb=reinterpret_cast(&kf); float k0=b2f(kb[0]),k1=b2f(kb[1]),k2=b2f(kb[2]),k3=b2f(kb[3]); float d0d=q0.x*k0+q0.y*k1+q0.z*k2+q0.w*k3; float d1d=q1.x*k0+q1.y*k1+q1.z*k2+q1.w*k3; for(int o=16;o>0;o>>=1){ d0d+=__shfl_xor_sync(0xffffffffu,d0d,o); d1d+=__shfl_xor_sync(0xffffffffu,d1d,o); } float s0=d0d*SCALE, s1=d1d*SCALE; float2 vf = *reinterpret_cast(vh + (size_t)p*DH); const bf16* vb=reinterpret_cast(&vf); float v0=b2f(vb[0]),v1=b2f(vb[1]),v2=b2f(vb[2]),v3=b2f(vb[3]); float m0n=fmaxf(m0,s0), c0=__expf(m0-m0n), pe0=__expf(s0-m0n); l0=l0*c0+pe0; a0x=a0x*c0+pe0*v0; a0y=a0y*c0+pe0*v1; a0z=a0z*c0+pe0*v2; a0w=a0w*c0+pe0*v3; m0=m0n; float m1n=fmaxf(m1,s1), c1=__expf(m1-m1n), pe1=__expf(s1-m1n); l1=l1*c1+pe1; a1x=a1x*c1+pe1*v0; a1y=a1y*c1+pe1*v1; a1z=a1z*c1+pe1*v2; a1w=a1w*c1+pe1*v3; m1=m1n; } // shared: m0s[W], l0s[W], m1s[W], l1s[W], acc0[W*DH], acc1[W*DH] extern __shared__ float sh[]; float* m0s=sh; float* l0s=m0s+W; float* m1s=l0s+W; float* l1s=m1s+W; float* acc0=l1s+W; float* acc1=acc0+W*DH; if(lane==0){ m0s[warp]=m0; l0s[warp]=l0; m1s[warp]=m1; l1s[warp]=l1; } *reinterpret_cast(acc0 + warp*DH + dd0) = make_float4(a0x,a0y,a0z,a0w); *reinterpret_cast(acc1 + warp*DH + dd0) = make_float4(a1x,a1y,a1z,a1w); __syncthreads(); // merge W warps -> one partial (threads 0..127 handle dims) if(threadIdx.x < DH){ int d=threadIdx.x; float M0=-1e30f, M1=-1e30f; for(int w=0;w attn_out (fully parallel) ---- // grid = 16 (one per q head), block 128 (one per dim). __global__ void attn_combine_kernel(float* __restrict__ attn_out, const float* __restrict__ partials, int nsplit){ int qh=blockIdx.x, tid=threadIdx.x; // 128 threads extern __shared__ float sh[]; // sm[0..nsplit-1], sl[nsplit..2nsplit-1] float* sm = sh; float* sl = sh + nsplit; __shared__ float red[128]; const float* base = partials + (size_t)qh*nsplit*PW; // cooperative coalesced load of m,l for(int s=tid;s0;o>>=1){ if(tid0;o>>=1){ if(tidMAXSPLIT) nsplit=MAXSPLIT; mix_kernel<<<(H+255)/256,256,0,st>>>(BLK, RND+(size_t)s*H, HID); for(int L=0;L<<<1,256,0,st>>>(HN, blkL, IN_LN+L*LNL); qkv_kernel<<<512,256,0,st>>>(QRAW,KRAW,VRAW, WQ+L*WQL,WK+L*WKVL,WV+L*WKVL, HN); qknr_kernel<<<24,128,0,st>>>(QOUT, kcL, vcL, QRAW,KRAW,VRAW, QN+L*QKNL, KN+L*QKNL, pos, (int)max_seq); attn_partial_kernel<<>>(PART, QOUT, kcL, vcL, seqlen, nsplit, (int)max_seq, W); attn_combine_kernel<<>>(AOUT, PART, nsplit); gemv_resid_kernel<<>>(HATT, WO+L*WOL, AOUT, blkL, H, NQ*DH); rmsnorm_kernel<<<1,256,0,st>>>(HN, HATT, PO_LN+L*LNL); gateup_kernel<<<768,256,0,st>>>(GU, WG+L*WGL, WU+L*WGL, HN); silu_kernel<<<(II+255)/256,256,0,st>>>(ACT, GU); gemv_resid_kernel<<>>(blkOut, WD+L*WDL, ACT, HATT, H, II); } } } """ CPP_SRC = r""" #include void run_steps( torch::Tensor input_ln, torch::Tensor post_ln, torch::Tensor q_norm, torch::Tensor k_norm, torch::Tensor q_proj, torch::Tensor k_proj, torch::Tensor v_proj, torch::Tensor o_proj, torch::Tensor gate_proj, torch::Tensor up_proj, torch::Tensor down_proj, torch::Tensor k_cache_all, torch::Tensor v_cache_all, torch::Tensor hidden, torch::Tensor randn, torch::Tensor blk, torch::Tensor q_raw, torch::Tensor k_raw, torch::Tensor v_raw, torch::Tensor q_out, torch::Tensor partials, torch::Tensor attn_out, torch::Tensor h_attn, torch::Tensor gate_up, torch::Tensor act, torch::Tensor h_normed, int64_t start_pos, int64_t n_steps, int64_t max_seq, int64_t pos_per_split, int64_t warps_per_split); """ _EXT = None def _ext(): global _EXT if _EXT is None: cu = _setup_cuda_toolkit() inc = [f"-I{os.path.join(cu, 'include')}"] if cu else [] cap = torch.cuda.get_device_capability() arch = f"{cap[0]}{cap[1]}" # e.g. "100" (B200/sm_100), "120" (RTX PRO 6000) gencode = [f"-gencode=arch=compute_{arch},code=sm_{arch}", f"-gencode=arch=compute_{arch},code=compute_{arch}"] _EXT = load_inline( name=f"megaqwen_decode_ext_sm{arch}", cpp_sources=CPP_SRC, cuda_sources=CUDA_SRC, functions=["run_steps"], extra_cflags=["-O3"] + inc, extra_cuda_cflags=["-O3", "--expt-relaxed-constexpr"] + gencode + inc, verbose=False, ) return _EXT # ---------------------------------------------------------------------------- # Model (same state_dict as reference.Model) # ---------------------------------------------------------------------------- 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.dim() >= 2: 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._prepped = False # ---- one-time: stack weights + allocate scratch on device ---- def _prep(self): if self._prepped: return dev = next(self.parameters()).device bl = self.blocks def stack(name): return torch.stack([getattr(b, name) for b in bl]).contiguous() self.w = {n: stack(n) for n in ["input_ln", "post_ln", "q_norm", "k_norm", "q_proj", "k_proj", "v_proj", "o_proj", "gate_proj", "up_proj", "down_proj"]} f32 = dict(dtype=torch.float32, device=dev) bf = dict(dtype=torch.bfloat16, device=dev) self._hidden = torch.zeros(HIDDEN, **bf) self._blk = torch.zeros(NUM_LAYERS, HIDDEN, **bf) self._q_raw = torch.zeros(NUM_Q * HEAD_DIM, **f32) self._k_raw = torch.zeros(NUM_KV * HEAD_DIM, **f32) self._v_raw = torch.zeros(NUM_KV * HEAD_DIM, **f32) self._q_out = torch.zeros(NUM_Q * HEAD_DIM, **f32) self._partials = torch.zeros(NUM_Q * MAXSPLIT * 132, **f32) self._attn_out = torch.zeros(NUM_Q * HEAD_DIM, **f32) self._h_attn = torch.zeros(HIDDEN, **f32) self._gate_up = torch.zeros(2 * INTERMEDIATE, **f32) self._act = torch.zeros(INTERMEDIATE, **f32) self._hn = torch.zeros(HIDDEN, **f32) # Persistent KV cache + randn buffer (stable pointers for CUDA graphs). self._maxdec = 64 self._randn_buf = torch.zeros(self._maxdec, HIDDEN, **bf) self._kc = torch.zeros(NUM_LAYERS, NUM_KV, self.max_seq, HEAD_DIM, **bf) self._vc = torch.zeros(NUM_LAYERS, NUM_KV, self.max_seq, HEAD_DIM, **bf) self._graphs = {} self._prepped = True def _run_eager(self, k_cache_all, v_cache_all, randn, start_pos, n_steps, pos_per_split, warps): w = self.w _ext().run_steps( w["input_ln"], w["post_ln"], w["q_norm"], w["k_norm"], w["q_proj"], w["k_proj"], w["v_proj"], w["o_proj"], w["gate_proj"], w["up_proj"], w["down_proj"], k_cache_all, v_cache_all, self._hidden, randn, self._blk, self._q_raw, self._k_raw, self._v_raw, self._q_out, self._partials, self._attn_out, self._h_attn, self._gate_up, self._act, self._hn, int(start_pos), int(n_steps), int(self.max_seq), int(pos_per_split), int(warps)) def _decode(self, hidden, randn, start_pos, n_steps, pos_per_split, warps): """Graph-replayed decode over persistent _kc/_vc. randn is [n_steps, H].""" n = int(n_steps) # `hidden` may alias self._hidden (prefill returns it); snapshot it so the # eager warmup below cannot clobber the input we reset to before replay. hidden_in = hidden.clone() if not _USE_GRAPH or n > self._maxdec: self._hidden.copy_(hidden_in) self._run_eager(self._kc, self._vc, randn, start_pos, n, pos_per_split, warps) return key = (int(start_pos), n, int(pos_per_split), int(warps)) g = self._graphs.get(key) if g is None: self._randn_buf[:n].copy_(randn) self._hidden.copy_(hidden_in) # Warmup (eager) loads kernels & does lazy init before capture. self._run_eager(self._kc, self._vc, self._randn_buf, start_pos, n, pos_per_split, warps) torch.cuda.synchronize() g = torch.cuda.CUDAGraph() with torch.cuda.graph(g): self._run_eager(self._kc, self._vc, self._randn_buf, start_pos, n, pos_per_split, warps) self._graphs[key] = g self._randn_buf[:n].copy_(randn) self._hidden.copy_(hidden_in) g.replay() _USE_GRAPH = True def _split_cfg(ctx_len: int) -> tuple[int, int]: """(pos_per_split, warps_per_split). W warps/block cooperate on each split, so nsplit stays small (cheap combine) while occupancy = 8*nsplit*W stays high.""" if ctx_len <= 2048: return max(1, ctx_len // 128), 4 # nsplit~128, W=4 if ctx_len <= 8192: return max(1, ctx_len // 128), 8 # nsplit~64, W=8 if ctx_len <= 32768: return max(1, ctx_len // 256), 8 # nsplit~128, W=8 return max(1, ctx_len // 512), 4 # 131072: nsplit~256, W=4 def _pos_per_split(ctx_len: int) -> int: return _split_cfg(ctx_len)[0] def _warps_per_split(ctx_len: int) -> int: return _split_cfg(ctx_len)[1] def _randn(seed: int, n: int, device) -> torch.Tensor: g = torch.Generator(device="cpu") g.manual_seed(seed) return torch.randn((n, HIDDEN), generator=g, dtype=torch.bfloat16).to(device) @torch.no_grad() def prefill(model, ctx_len: int, seed: int): model = model.cuda().eval() model._prep() dev = model._hidden.device model._kc.zero_() model._vc.zero_() g = torch.Generator(device="cpu") g.manual_seed(seed) h0 = torch.randn(HIDDEN, generator=g, dtype=torch.bfloat16).to(dev) model._hidden.copy_(h0) rnd = _randn(seed + 1, ctx_len, dev) model._run_eager(model._kc, model._vc, rnd, 0, ctx_len, _pos_per_split(ctx_len), _warps_per_split(ctx_len)) k_list = [model._kc[i] for i in range(NUM_LAYERS)] v_list = [model._vc[i] for i in range(NUM_LAYERS)] return model._hidden, k_list, v_list @torch.no_grad() def decode_steps(model, hidden, k_caches, v_caches, start_pos, n_steps, seed): model._prep() dev = model._hidden.device rnd = _randn(seed + 2, n_steps, dev) model._decode(hidden, rnd, start_pos, n_steps, _pos_per_split(start_pos), _warps_per_split(start_pos)) return model._hidden, k_caches, v_caches _decode_impl = decode_steps def run(ctx_len, decode_steps, seed, model=None, max_seq=None): n_decode = int(decode_steps) if model is None: ms = max_seq or max(ctx_len + n_decode, 512) model = Model(NUM_LAYERS, ms) model = model.cuda().eval() h, k, v = prefill(model, ctx_len, seed) h, k, v = _decode_impl(model, h, k, v, ctx_len, n_decode, seed) return {"last_hidden": h.detach(), "ctx_len": ctx_len, "decode_steps": n_decode} def get_init_inputs(): return [NUM_LAYERS, 131072] def get_inputs(): return []