# Fused W4A16 hybrid decode: the whole per-token forward (KDA + MLA + MoE, # all int4 dequant-GEMVs, norms, residuals, state updates) runs as ONE custom # CUDA kernel launch (a cooperative-grid megakernel with grid.sync between # stages). step() below invokes that single kernel exactly once and performs # no other GPU work. import torch import torch.nn as nn from torch.utils.cpp_extension import load_inline GROUP_SIZE = 128 SB_N = 36416 SF_BASE = 62824 def _pack_int4(w_q): lo = w_q[[REDACTED: IP]] & 0xF hi = w_q[[REDACTED: IP]] & 0xF return (lo | (hi << 4)).contiguous() def quantize(w_io, group=GROUP_SIZE): K, N = w_io.shape ng = K // group wg = w_io.view(ng, group, N).float() wmin = wg.min(dim=1, keepdim=True).values wmax = wg.max(dim=1, keepdim=True).values scales = (wmax - wmin).clamp_min(1e-8) / 15.0 zeros = (-wmin / scales).round().clamp(0, 15) w_q = ((wg / scales) + zeros).round().clamp(0, 15).to(torch.uint8).view(K, N) return _pack_int4(w_q), scales.squeeze(1).to(torch.bfloat16), zeros.squeeze(1).to(torch.bfloat16) class QuantLinear(nn.Module): def __init__(self, in_f, out_f, group=GROUP_SIZE): super().__init__() assert in_f % group == 0 and in_f % 2 == 0 self.in_f, self.out_f, self.group = in_f, out_f, group ng = in_f // group self.register_buffer("w_q", torch.zeros(in_f // 2, out_f, dtype=torch.uint8)) self.register_buffer("scales", torch.zeros(ng, out_f, dtype=torch.bfloat16)) self.register_buffer("zeros", torch.zeros(ng, out_f, dtype=torch.bfloat16)) def init_random(self, gen, std=0.02): w = torch.randn(self.in_f, self.out_f, generator=gen) * std wq, s, z = quantize(w, self.group) self.w_q.copy_(wq) self.scales.copy_(s) self.zeros.copy_(z) class QuantExperts(nn.Module): def __init__(self, n, in_f, out_f, group=GROUP_SIZE): super().__init__() self.n, self.in_f, self.out_f, self.group = n, in_f, out_f, group ng = in_f // group self.register_buffer("w_q", torch.zeros(n, in_f // 2, out_f, dtype=torch.uint8)) self.register_buffer("scales", torch.zeros(n, ng, out_f, dtype=torch.bfloat16)) self.register_buffer("zeros", torch.zeros(n, ng, out_f, dtype=torch.bfloat16)) def init_random(self, gen, std=0.02): for e in range(self.n): w = torch.randn(self.in_f, self.out_f, generator=gen) * std wq, s, z = quantize(w, self.group) self.w_q[e].copy_(wq) self.scales[e].copy_(s) self.zeros[e].copy_(z) class KDA(nn.Module): def __init__(self, cfg): super().__init__() self.cfg = cfg H, Dk, d = cfg.kda_heads, cfg.kda_head_dim, cfg.hidden self.q_proj = QuantLinear(d, H * Dk, cfg.group) self.k_proj = QuantLinear(d, H * Dk, cfg.group) self.v_proj = QuantLinear(d, H * Dk, cfg.group) self.g_proj = QuantLinear(d, H * Dk, cfg.group) self.beta_proj = nn.Linear(d, H, bias=False, dtype=cfg.dtype) self.conv_w = nn.Parameter(torch.empty(3, H * Dk, cfg.short_conv, dtype=cfg.dtype)) self.o_proj = QuantLinear(H * Dk, d, cfg.group) self.scale = Dk ** -0.5 class MLA(nn.Module): def __init__(self, cfg): super().__init__() self.cfg = cfg H, d = cfg.mla_heads, cfg.hidden self.q_proj = QuantLinear(d, H * (cfg.qk_nope + cfg.qk_rope), cfg.group) self.kv_a = QuantLinear(d, cfg.kv_lora + cfg.qk_rope, cfg.group) self.kv_b = QuantLinear(cfg.kv_lora, H * (cfg.qk_nope + cfg.v_head), cfg.group) self.o_proj = QuantLinear(H * cfg.v_head, d, cfg.group) self.scale = (cfg.qk_nope + cfg.qk_rope) ** -0.5 class MoE(nn.Module): def __init__(self, cfg): super().__init__() self.cfg = cfg d, m, E = cfg.hidden, cfg.moe_inter, cfg.n_experts self.router = nn.Linear(d, E, bias=False, dtype=cfg.dtype) self.gate = QuantExperts(E, d, m, cfg.group) self.up = QuantExperts(E, d, m, cfg.group) self.down = QuantExperts(E, m, d, cfg.group) self.s_gate = QuantExperts(cfg.n_shared, d, m, cfg.group) self.s_up = QuantExperts(cfg.n_shared, d, m, cfg.group) self.s_down = QuantExperts(cfg.n_shared, m, d, cfg.group) class Block(nn.Module): def __init__(self, cfg, kind): super().__init__() self.kind = kind self.attn_norm = nn.Parameter(torch.ones(cfg.hidden, dtype=cfg.dtype)) self.moe_norm = nn.Parameter(torch.ones(cfg.hidden, dtype=cfg.dtype)) self.attn = KDA(cfg) if kind == "K" else MLA(cfg) self.moe = MoE(cfg) _mod = None def _get_mod(): global _mod if _mod is None: _mod = load_inline(name='kimi_megak', cpp_sources=_CPP_DECL, cuda_sources=_CUDA_SRC, functions=["mega_launch"], extra_cuda_cflags=["-O3"]) return _mod class Model(nn.Module): def __init__(self, cfg): super().__init__() assert cfg.hidden == 2304 and cfg.kda_heads == 32 and cfg.kda_head_dim == 128 assert cfg.short_conv == 4 and cfg.mla_heads == 32 and cfg.kv_lora == 512 assert cfg.qk_nope == 128 and cfg.qk_rope == 64 and cfg.v_head == 128 assert cfg.n_experts == 64 and cfg.n_active == 8 and cfg.n_shared == 1 assert cfg.moe_inter == 1024 and cfg.group == 128 assert tuple(cfg.pattern) == ("K", "K", "K", "M") self.cfg = cfg self.blocks = nn.ModuleList(Block(cfg, k) for k in cfg.pattern) self.reset_parameters() self._sb = None self._si = None self._sf = None self._bc = None self._bk = None self._cap = 0 def reset_parameters(self): g = torch.Generator(device="cpu").manual_seed(1234) for mod in self.modules(): if isinstance(mod, (QuantLinear, QuantExperts)): mod.init_random(g) elif isinstance(mod, nn.Linear): nn.init.normal_(mod.weight, 0.0, 0.02, generator=g) elif isinstance(mod, KDA): nn.init.normal_(mod.conv_w, 0.0, 0.1, generator=g) def _ensure(self, dev, need): if self._sb is None or self._sb.device != dev: self._sb = torch.empty(SB_N, dtype=torch.bfloat16, device=dev) self._si = torch.empty(8, dtype=torch.int32, device=dev) self._bc = None if self._bc is None or self._bc.device != dev or self._cap < need: self._cap = need + 64 self._bc = torch.empty((self._cap, 512), dtype=torch.bfloat16, device=dev) self._bk = torch.empty((self._cap, 64), dtype=torch.bfloat16, device=dev) self._sf = torch.empty(SF_BASE + 32 * self._cap, dtype=torch.float32, device=dev) def step(self, hidden, state): cur = int(state[3]["c_kv"].shape[0]) need = cur + 1 self._ensure(hidden.device, need) args = [] args.append(hidden) args.append(self._sb) args.append(self._sf) args.append(self._si) args.append(self.blocks[0].attn.q_proj.w_q) args.append(self.blocks[0].attn.q_proj.scales) args.append(self.blocks[0].attn.q_proj.zeros) args.append(self.blocks[0].attn.k_proj.w_q) args.append(self.blocks[0].attn.k_proj.scales) args.append(self.blocks[0].attn.k_proj.zeros) args.append(self.blocks[0].attn.v_proj.w_q) args.append(self.blocks[0].attn.v_proj.scales) args.append(self.blocks[0].attn.v_proj.zeros) args.append(self.blocks[0].attn.g_proj.w_q) args.append(self.blocks[0].attn.g_proj.scales) args.append(self.blocks[0].attn.g_proj.zeros) args.append(self.blocks[0].attn.o_proj.w_q) args.append(self.blocks[0].attn.o_proj.scales) args.append(self.blocks[0].attn.o_proj.zeros) args.append(self.blocks[0].attn.beta_proj.weight) args.append(self.blocks[0].attn.conv_w) args.append(state[0]["S"]) args.append(state[0]["cq"]) args.append(state[0]["ck"]) args.append(state[0]["cv"]) args.append(self.blocks[0].attn_norm) args.append(self.blocks[0].moe_norm) args.append(self.blocks[0].moe.router.weight) args.append(self.blocks[0].moe.gate.w_q) args.append(self.blocks[0].moe.gate.scales) args.append(self.blocks[0].moe.gate.zeros) args.append(self.blocks[0].moe.up.w_q) args.append(self.blocks[0].moe.up.scales) args.append(self.blocks[0].moe.up.zeros) args.append(self.blocks[0].moe.down.w_q) args.append(self.blocks[0].moe.down.scales) args.append(self.blocks[0].moe.down.zeros) args.append(self.blocks[0].moe.s_gate.w_q) args.append(self.blocks[0].moe.s_gate.scales) args.append(self.blocks[0].moe.s_gate.zeros) args.append(self.blocks[0].moe.s_up.w_q) args.append(self.blocks[0].moe.s_up.scales) args.append(self.blocks[0].moe.s_up.zeros) args.append(self.blocks[0].moe.s_down.w_q) args.append(self.blocks[0].moe.s_down.scales) args.append(self.blocks[0].moe.s_down.zeros) args.append(self.blocks[1].attn.q_proj.w_q) args.append(self.blocks[1].attn.q_proj.scales) args.append(self.blocks[1].attn.q_proj.zeros) args.append(self.blocks[1].attn.k_proj.w_q) args.append(self.blocks[1].attn.k_proj.scales) args.append(self.blocks[1].attn.k_proj.zeros) args.append(self.blocks[1].attn.v_proj.w_q) args.append(self.blocks[1].attn.v_proj.scales) args.append(self.blocks[1].attn.v_proj.zeros) args.append(self.blocks[1].attn.g_proj.w_q) args.append(self.blocks[1].attn.g_proj.scales) args.append(self.blocks[1].attn.g_proj.zeros) args.append(self.blocks[1].attn.o_proj.w_q) args.append(self.blocks[1].attn.o_proj.scales) args.append(self.blocks[1].attn.o_proj.zeros) args.append(self.blocks[1].attn.beta_proj.weight) args.append(self.blocks[1].attn.conv_w) args.append(state[1]["S"]) args.append(state[1]["cq"]) args.append(state[1]["ck"]) args.append(state[1]["cv"]) args.append(self.blocks[1].attn_norm) args.append(self.blocks[1].moe_norm) args.append(self.blocks[1].moe.router.weight) args.append(self.blocks[1].moe.gate.w_q) args.append(self.blocks[1].moe.gate.scales) args.append(self.blocks[1].moe.gate.zeros) args.append(self.blocks[1].moe.up.w_q) args.append(self.blocks[1].moe.up.scales) args.append(self.blocks[1].moe.up.zeros) args.append(self.blocks[1].moe.down.w_q) args.append(self.blocks[1].moe.down.scales) args.append(self.blocks[1].moe.down.zeros) args.append(self.blocks[1].moe.s_gate.w_q) args.append(self.blocks[1].moe.s_gate.scales) args.append(self.blocks[1].moe.s_gate.zeros) args.append(self.blocks[1].moe.s_up.w_q) args.append(self.blocks[1].moe.s_up.scales) args.append(self.blocks[1].moe.s_up.zeros) args.append(self.blocks[1].moe.s_down.w_q) args.append(self.blocks[1].moe.s_down.scales) args.append(self.blocks[1].moe.s_down.zeros) args.append(self.blocks[2].attn.q_proj.w_q) args.append(self.blocks[2].attn.q_proj.scales) args.append(self.blocks[2].attn.q_proj.zeros) args.append(self.blocks[2].attn.k_proj.w_q) args.append(self.blocks[2].attn.k_proj.scales) args.append(self.blocks[2].attn.k_proj.zeros) args.append(self.blocks[2].attn.v_proj.w_q) args.append(self.blocks[2].attn.v_proj.scales) args.append(self.blocks[2].attn.v_proj.zeros) args.append(self.blocks[2].attn.g_proj.w_q) args.append(self.blocks[2].attn.g_proj.scales) args.append(self.blocks[2].attn.g_proj.zeros) args.append(self.blocks[2].attn.o_proj.w_q) args.append(self.blocks[2].attn.o_proj.scales) args.append(self.blocks[2].attn.o_proj.zeros) args.append(self.blocks[2].attn.beta_proj.weight) args.append(self.blocks[2].attn.conv_w) args.append(state[2]["S"]) args.append(state[2]["cq"]) args.append(state[2]["ck"]) args.append(state[2]["cv"]) args.append(self.blocks[2].attn_norm) args.append(self.blocks[2].moe_norm) args.append(self.blocks[2].moe.router.weight) args.append(self.blocks[2].moe.gate.w_q) args.append(self.blocks[2].moe.gate.scales) args.append(self.blocks[2].moe.gate.zeros) args.append(self.blocks[2].moe.up.w_q) args.append(self.blocks[2].moe.up.scales) args.append(self.blocks[2].moe.up.zeros) args.append(self.blocks[2].moe.down.w_q) args.append(self.blocks[2].moe.down.scales) args.append(self.blocks[2].moe.down.zeros) args.append(self.blocks[2].moe.s_gate.w_q) args.append(self.blocks[2].moe.s_gate.scales) args.append(self.blocks[2].moe.s_gate.zeros) args.append(self.blocks[2].moe.s_up.w_q) args.append(self.blocks[2].moe.s_up.scales) args.append(self.blocks[2].moe.s_up.zeros) args.append(self.blocks[2].moe.s_down.w_q) args.append(self.blocks[2].moe.s_down.scales) args.append(self.blocks[2].moe.s_down.zeros) args.append(self.blocks[3].attn.q_proj.w_q) args.append(self.blocks[3].attn.q_proj.scales) args.append(self.blocks[3].attn.q_proj.zeros) args.append(self.blocks[3].attn.kv_a.w_q) args.append(self.blocks[3].attn.kv_a.scales) args.append(self.blocks[3].attn.kv_a.zeros) args.append(self.blocks[3].attn.kv_b.w_q) args.append(self.blocks[3].attn.kv_b.scales) args.append(self.blocks[3].attn.kv_b.zeros) args.append(self.blocks[3].attn.o_proj.w_q) args.append(self.blocks[3].attn.o_proj.scales) args.append(self.blocks[3].attn.o_proj.zeros) args.append(self.blocks[3].attn_norm) args.append(self.blocks[3].moe_norm) args.append(self.blocks[3].moe.router.weight) args.append(self.blocks[3].moe.gate.w_q) args.append(self.blocks[3].moe.gate.scales) args.append(self.blocks[3].moe.gate.zeros) args.append(self.blocks[3].moe.up.w_q) args.append(self.blocks[3].moe.up.scales) args.append(self.blocks[3].moe.up.zeros) args.append(self.blocks[3].moe.down.w_q) args.append(self.blocks[3].moe.down.scales) args.append(self.blocks[3].moe.down.zeros) args.append(self.blocks[3].moe.s_gate.w_q) args.append(self.blocks[3].moe.s_gate.scales) args.append(self.blocks[3].moe.s_gate.zeros) args.append(self.blocks[3].moe.s_up.w_q) args.append(self.blocks[3].moe.s_up.scales) args.append(self.blocks[3].moe.s_up.zeros) args.append(self.blocks[3].moe.s_down.w_q) args.append(self.blocks[3].moe.s_down.scales) args.append(self.blocks[3].moe.s_down.zeros) args.append(state[3]["c_kv"]) args.append(state[3]["k_rope"]) args.append(self._bc) args.append(self._bk) _get_mod().mega_launch(*args, cur, self._cap) state[3]["c_kv"] = self._bc[:need] state[3]["k_rope"] = self._bk[:need] return hidden, state _CPP_DECL = 'void mega_launch(torch::Tensor h, torch::Tensor SBp, torch::Tensor SFp, torch::Tensor SIp, torch::Tensor L0qw, torch::Tensor L0qs, torch::Tensor L0qz, torch::Tensor L0kw, torch::Tensor L0ks, torch::Tensor L0kz, torch::Tensor L0vw, torch::Tensor L0vs, torch::Tensor L0vz, torch::Tensor L0gw, torch::Tensor L0gs, torch::Tensor L0gz, torch::Tensor L0ow, torch::Tensor L0os, torch::Tensor L0oz, torch::Tensor L0bw, torch::Tensor L0cw, torch::Tensor L0S, torch::Tensor L0cq, torch::Tensor L0ck, torch::Tensor L0cv, torch::Tensor L0an, torch::Tensor L0mn, torch::Tensor L0rt, torch::Tensor L0gqw, torch::Tensor L0gqs, torch::Tensor L0gqz, torch::Tensor L0uqw, torch::Tensor L0uqs, torch::Tensor L0uqz, torch::Tensor L0dqw, torch::Tensor L0dqs, torch::Tensor L0dqz, torch::Tensor L0sgqw, torch::Tensor L0sgqs, torch::Tensor L0sgqz, torch::Tensor L0suqw, torch::Tensor L0suqs, torch::Tensor L0suqz, torch::Tensor L0sdqw, torch::Tensor L0sdqs, torch::Tensor L0sdqz, torch::Tensor L1qw, torch::Tensor L1qs, torch::Tensor L1qz, torch::Tensor L1kw, torch::Tensor L1ks, torch::Tensor L1kz, torch::Tensor L1vw, torch::Tensor L1vs, torch::Tensor L1vz, torch::Tensor L1gw, torch::Tensor L1gs, torch::Tensor L1gz, torch::Tensor L1ow, torch::Tensor L1os, torch::Tensor L1oz, torch::Tensor L1bw, torch::Tensor L1cw, torch::Tensor L1S, torch::Tensor L1cq, torch::Tensor L1ck, torch::Tensor L1cv, torch::Tensor L1an, torch::Tensor L1mn, torch::Tensor L1rt, torch::Tensor L1gqw, torch::Tensor L1gqs, torch::Tensor L1gqz, torch::Tensor L1uqw, torch::Tensor L1uqs, torch::Tensor L1uqz, torch::Tensor L1dqw, torch::Tensor L1dqs, torch::Tensor L1dqz, torch::Tensor L1sgqw, torch::Tensor L1sgqs, torch::Tensor L1sgqz, torch::Tensor L1suqw, torch::Tensor L1suqs, torch::Tensor L1suqz, torch::Tensor L1sdqw, torch::Tensor L1sdqs, torch::Tensor L1sdqz, torch::Tensor L2qw, torch::Tensor L2qs, torch::Tensor L2qz, torch::Tensor L2kw, torch::Tensor L2ks, torch::Tensor L2kz, torch::Tensor L2vw, torch::Tensor L2vs, torch::Tensor L2vz, torch::Tensor L2gw, torch::Tensor L2gs, torch::Tensor L2gz, torch::Tensor L2ow, torch::Tensor L2os, torch::Tensor L2oz, torch::Tensor L2bw, torch::Tensor L2cw, torch::Tensor L2S, torch::Tensor L2cq, torch::Tensor L2ck, torch::Tensor L2cv, torch::Tensor L2an, torch::Tensor L2mn, torch::Tensor L2rt, torch::Tensor L2gqw, torch::Tensor L2gqs, torch::Tensor L2gqz, torch::Tensor L2uqw, torch::Tensor L2uqs, torch::Tensor L2uqz, torch::Tensor L2dqw, torch::Tensor L2dqs, torch::Tensor L2dqz, torch::Tensor L2sgqw, torch::Tensor L2sgqs, torch::Tensor L2sgqz, torch::Tensor L2suqw, torch::Tensor L2suqs, torch::Tensor L2suqz, torch::Tensor L2sdqw, torch::Tensor L2sdqs, torch::Tensor L2sdqz, torch::Tensor Mqw, torch::Tensor Mqs, torch::Tensor Mqz, torch::Tensor Maw, torch::Tensor Mas, torch::Tensor Maz, torch::Tensor Mbw, torch::Tensor Mbs, torch::Tensor Mbz, torch::Tensor Mow, torch::Tensor Mos, torch::Tensor Moz, torch::Tensor Man, torch::Tensor Mmn, torch::Tensor Mrt, torch::Tensor Mgqw, torch::Tensor Mgqs, torch::Tensor Mgqz, torch::Tensor Muqw, torch::Tensor Muqs, torch::Tensor Muqz, torch::Tensor Mdqw, torch::Tensor Mdqs, torch::Tensor Mdqz, torch::Tensor Msgqw, torch::Tensor Msgqs, torch::Tensor Msgqz, torch::Tensor Msuqw, torch::Tensor Msuqs, torch::Tensor Msuqz, torch::Tensor Msdqw, torch::Tensor Msdqs, torch::Tensor Msdqz, torch::Tensor oc, torch::Tensor ok, torch::Tensor nc, torch::Tensor nk, long long curlen, long long capx);\n' _CUDA_SRC = '\n#include \n#include \n#include \n#include \nnamespace cg = cooperative_groups;\n\n__device__ __forceinline__ float bf2f(__nv_bfloat16 x) { return __bfloat162float(x); }\n__device__ __forceinline__ __nv_bfloat16 f2bf(float x) { return __float2bfloat16_rn(x); }\n\n// Warp-group int4 GEMV: 32 lanes collaborate on one output (shuffle reduce).\n// Short per-lane chains + 32x MLP instead of one 2304-deep latency chain.\n__device__ void h_w4(cg::grid_group &gg, const __nv_bfloat16* x,\n const unsigned char* wq, const __nv_bfloat16* sc, const __nv_bfloat16* ze,\n __nv_bfloat16* y, int K, int N, float* shm, int tid, int TT) {\n int lane = tid & 31;\n int grp = tid >> 5;\n int NG = TT >> 5;\n int ng = K >> 7;\n for (int n = grp; n < N; n += NG) {\n float acc = 0.0f;\n for (int ggi = 0; ggi < ng; ++ggi) {\n float s = bf2f(sc[ggi * N + n]);\n float z = bf2f(ze[ggi * N + n]);\n int k0 = ggi << 7;\n for (int k = k0 + lane; k < k0 + 128; k += 32) {\n unsigned char b = wq[(k >> 1) * N + n];\n int qv = (k & 1) ? (b >> 4) : (b & 0xF);\n acc += bf2f(x[k]) * (((float)qv - z) * s);\n }\n }\n for (int o = 16; o > 0; o >>= 1) acc += __shfl_down_sync(0xFFFFFFFFu, acc, o);\n if (lane == 0) y[n] = f2bf(acc);\n }\n gg.sync();\n}\n\n__device__ void h_w4f(cg::grid_group &gg, const __nv_bfloat16* x,\n const unsigned char* wq, const __nv_bfloat16* sc, const __nv_bfloat16* ze,\n float* y, int K, int N, float* shm, int tid, int TT) {\n int lane = tid & 31;\n int grp = tid >> 5;\n int NG = TT >> 5;\n int ng = K >> 7;\n for (int n = grp; n < N; n += NG) {\n float acc = 0.0f;\n for (int ggi = 0; ggi < ng; ++ggi) {\n float s = bf2f(sc[ggi * N + n]);\n float z = bf2f(ze[ggi * N + n]);\n int k0 = ggi << 7;\n for (int k = k0 + lane; k < k0 + 128; k += 32) {\n unsigned char b = wq[(k >> 1) * N + n];\n int qv = (k & 1) ? (b >> 4) : (b & 0xF);\n acc += bf2f(x[k]) * (((float)qv - z) * s);\n }\n }\n for (int o = 16; o > 0; o >>= 1) acc += __shfl_down_sync(0xFFFFFFFFu, acc, o);\n if (lane == 0) y[n] = acc;\n }\n gg.sync();\n}\n\n// gate+up: one warp-group per (gate_n == up_n index), two accumulators\n__device__ void h_w4dual(cg::grid_group &gg, const __nv_bfloat16* x,\n const unsigned char* wqg, const __nv_bfloat16* scg, const __nv_bfloat16* zeg,\n const unsigned char* wqu, const __nv_bfloat16* scu, const __nv_bfloat16* zeu,\n float* yg, float* yu, int K, int N, float* shm, int tid, int TT) {\n int lane = tid & 31;\n int grp = tid >> 5;\n int NG = TT >> 5;\n int ng = K >> 7;\n for (int n = grp; n < N; n += NG) {\n float ag = 0.0f, au = 0.0f;\n for (int ggi = 0; ggi < ng; ++ggi) {\n float sg = bf2f(scg[ggi * N + n]);\n float zg = bf2f(zeg[ggi * N + n]);\n float su = bf2f(scu[ggi * N + n]);\n float zu = bf2f(zeu[ggi * N + n]);\n int k0 = ggi << 7;\n for (int k = k0 + lane; k < k0 + 128; k += 32) {\n float xv = bf2f(x[k]);\n unsigned char bg = wqg[(k >> 1) * N + n];\n unsigned char bu = wqu[(k >> 1) * N + n];\n int qg = (k & 1) ? (bg >> 4) : (bg & 0xF);\n int qu = (k & 1) ? (bu >> 4) : (bu & 0xF);\n ag += xv * (((float)qg - zg) * sg);\n au += xv * (((float)qu - zu) * su);\n }\n }\n for (int o = 16; o > 0; o >>= 1) {\n ag += __shfl_down_sync(0xFFFFFFFFu, ag, o);\n au += __shfl_down_sync(0xFFFFFFFFu, au, o);\n }\n if (lane == 0) {\n yg[n] = ag;\n yu[n] = au;\n }\n }\n gg.sync();\n}\n\n__device__ void h_w4down(cg::grid_group &gg, const float* xf,\n const unsigned char* wq, const __nv_bfloat16* sc, const __nv_bfloat16* ze,\n float* acc, float w, __nv_bfloat16* hres, int K, int N, float* shm, int tid, int TT) {\n int lane = tid & 31;\n int grp = tid >> 5;\n int NG = TT >> 5;\n int ng = K >> 7;\n for (int n = grp; n < N; n += NG) {\n float d = 0.0f;\n for (int ggi = 0; ggi < ng; ++ggi) {\n float s = bf2f(sc[ggi * N + n]);\n float z = bf2f(ze[ggi * N + n]);\n int k0 = ggi << 7;\n for (int k = k0 + lane; k < k0 + 128; k += 32) {\n unsigned char b = wq[(k >> 1) * N + n];\n int qv = (k & 1) ? (b >> 4) : (b & 0xF);\n d += xf[k] * (((float)qv - z) * s);\n }\n }\n for (int o = 16; o > 0; o >>= 1) d += __shfl_down_sync(0xFFFFFFFFu, d, o);\n if (lane == 0) {\n float an = acc[n] + w * d;\n acc[n] = an;\n hres[n] = f2bf(bf2f(hres[n]) + an);\n }\n }\n gg.sync();\n}\n\n__device__ void h_moe_gateup(cg::grid_group &gg, const __nv_bfloat16* x,\n const unsigned char* gqw, const __nv_bfloat16* gqs, const __nv_bfloat16* gqz,\n const unsigned char* uqw, const __nv_bfloat16* uqs, const __nv_bfloat16* uqz,\n const int* topi, float* gf8, float* uf8, float* hh8, int K, int N,\n float* shm, int tid, int TT) {\n int lane = tid & 31;\n int grp = tid >> 5;\n int NG = TT >> 5;\n int ng = K >> 7;\n for (int t = grp; t < 8 * N; t += NG) {\n int e = t / N;\n int n = t % N;\n int ej = topi[e];\n const unsigned char* eqg = gqw + (int64_t)ej * 1179648;\n const __nv_bfloat16* esg = gqs + (int64_t)ej * 18432;\n const __nv_bfloat16* ezg = gqz + (int64_t)ej * 18432;\n const unsigned char* equ = uqw + (int64_t)ej * 1179648;\n const __nv_bfloat16* esu = uqs + (int64_t)ej * 18432;\n const __nv_bfloat16* ezu = uqz + (int64_t)ej * 18432;\n float ag = 0.0f, au = 0.0f;\n for (int ggi = 0; ggi < ng; ++ggi) {\n float sg = bf2f(esg[ggi * N + n]);\n float zg = bf2f(ezg[ggi * N + n]);\n float su = bf2f(esu[ggi * N + n]);\n float zu = bf2f(ezu[ggi * N + n]);\n int k0 = ggi << 7;\n for (int k = k0 + lane; k < k0 + 128; k += 32) {\n float xv = bf2f(x[k]);\n unsigned char bg = eqg[(k >> 1) * N + n];\n unsigned char bu = equ[(k >> 1) * N + n];\n int qg = (k & 1) ? (bg >> 4) : (bg & 0xF);\n int qu = (k & 1) ? (bu >> 4) : (bu & 0xF);\n ag += xv * (((float)qg - zg) * sg);\n au += xv * (((float)qu - zu) * su);\n }\n }\n for (int o = 16; o > 0; o >>= 1) {\n ag += __shfl_down_sync(0xFFFFFFFFu, ag, o);\n au += __shfl_down_sync(0xFFFFFFFFu, au, o);\n }\n if (lane == 0) {\n gf8[e * 1024 + n] = ag;\n uf8[e * 1024 + n] = au;\n hh8[e * 1024 + n] = (ag / (1.0f + expf(-ag))) * au;\n }\n }\n gg.sync();\n}\n\n__device__ void h_moe_down(cg::grid_group &gg, const float* hh8,\n const unsigned char* dqw, const __nv_bfloat16* dqs, const __nv_bfloat16* dqz,\n const int* topi, const float* topw, float* acc, int K, int N,\n float* shm, int tid, int TT) {\n int lane = tid & 31;\n int grp = tid >> 5;\n int NG = TT >> 5;\n int ng = K >> 7;\n for (int n = grp; n < N; n += NG) {\n float tot = 0.0f;\n for (int e = 0; e < 8; ++e) {\n float wt = topw[e];\n int ej = topi[e];\n const unsigned char* eq = dqw + (int64_t)ej * 1179648;\n const __nv_bfloat16* es = dqs + (int64_t)ej * 18432;\n const __nv_bfloat16* ez = dqz + (int64_t)ej * 18432;\n const float* hx = hh8 + e * 1024;\n float d = 0.0f;\n for (int ggi = 0; ggi < ng; ++ggi) {\n float s = bf2f(es[ggi * N + n]);\n float z = bf2f(ez[ggi * N + n]);\n int k0 = ggi << 7;\n for (int k = k0 + lane; k < k0 + 128; k += 32) {\n unsigned char b = eq[(k >> 1) * N + n];\n int qv = (k & 1) ? (b >> 4) : (b & 0xF);\n d += hx[k] * (((float)qv - z) * s);\n }\n }\n for (int o = 16; o > 0; o >>= 1) d += __shfl_down_sync(0xFFFFFFFFu, d, o);\n if (lane == 0) tot += wt * d;\n }\n if (lane == 0) acc[n] += tot;\n }\n gg.sync();\n}\n\n__device__ void h_w4x4(cg::grid_group &gg, const __nv_bfloat16* x,\n const unsigned char* wq0, const __nv_bfloat16* sc0, const __nv_bfloat16* ze0,\n const unsigned char* wq1, const __nv_bfloat16* sc1, const __nv_bfloat16* ze1,\n const unsigned char* wq2, const __nv_bfloat16* sc2, const __nv_bfloat16* ze2,\n const unsigned char* wq3, const __nv_bfloat16* sc3, const __nv_bfloat16* ze3,\n __nv_bfloat16* y0, __nv_bfloat16* y1, __nv_bfloat16* y2, __nv_bfloat16* y3,\n int K, int N, float* shm, int tid, int TT) {\n int lane = tid & 31;\n int grp = tid >> 5;\n int NG = TT >> 5;\n int ng = K >> 7;\n for (int t = grp; t < 4 * N; t += NG) {\n int p = t / N;\n int n = t % N;\n const unsigned char* wq = (p == 0) ? wq0 : ((p == 1) ? wq1 : ((p == 2) ? wq2 : wq3));\n const __nv_bfloat16* sc = (p == 0) ? sc0 : ((p == 1) ? sc1 : ((p == 2) ? sc2 : sc3));\n const __nv_bfloat16* ze = (p == 0) ? ze0 : ((p == 1) ? ze1 : ((p == 2) ? ze2 : ze3));\n float acc = 0.0f;\n for (int ggi = 0; ggi < ng; ++ggi) {\n float s = bf2f(sc[ggi * N + n]);\n float z = bf2f(ze[ggi * N + n]);\n int k0 = ggi << 7;\n for (int k = k0 + lane; k < k0 + 128; k += 32) {\n unsigned char b = wq[(k >> 1) * N + n];\n int qv = (k & 1) ? (b >> 4) : (b & 0xF);\n acc += bf2f(x[k]) * (((float)qv - z) * s);\n }\n }\n for (int o = 16; o > 0; o >>= 1) acc += __shfl_down_sync(0xFFFFFFFFu, acc, o);\n if (lane == 0) {\n __nv_bfloat16* y = (p == 0) ? y0 : ((p == 1) ? y1 : ((p == 2) ? y2 : y3));\n y[n] = f2bf(acc);\n }\n }\n gg.sync();\n}\n\n__device__ void h_w4_resadd(cg::grid_group &gg, const __nv_bfloat16* x,\n const unsigned char* wq, const __nv_bfloat16* sc, const __nv_bfloat16* ze,\n __nv_bfloat16* h, int K, int N, float* shm, int tid, int TT) {\n int lane = tid & 31;\n int grp = tid >> 5;\n int NG = TT >> 5;\n int ng = K >> 7;\n for (int n = grp; n < N; n += NG) {\n float acc = 0.0f;\n for (int ggi = 0; ggi < ng; ++ggi) {\n float s = bf2f(sc[ggi * N + n]);\n float z = bf2f(ze[ggi * N + n]);\n int k0 = ggi << 7;\n for (int k = k0 + lane; k < k0 + 128; k += 32) {\n unsigned char b = wq[(k >> 1) * N + n];\n int qv = (k & 1) ? (b >> 4) : (b & 0xF);\n acc += bf2f(x[k]) * (((float)qv - z) * s);\n }\n }\n for (int o = 16; o > 0; o >>= 1) acc += __shfl_down_sync(0xFFFFFFFFu, acc, o);\n if (lane == 0) h[n] = f2bf(bf2f(h[n]) + acc);\n }\n gg.sync();\n}\n\n__device__ void h_bfgemv(cg::grid_group &gg, const __nv_bfloat16* x,\n const __nv_bfloat16* W, float* y, int K, int N, float* shm, int tid, int TT) {\n int lane = tid & 31;\n int grp = tid >> 5;\n int NG = TT >> 5;\n for (int n = grp; n < N; n += NG) {\n float acc = 0.0f;\n const __nv_bfloat16* row = W + (int64_t)n * K;\n for (int k = lane; k < K; k += 32) acc += bf2f(x[k]) * bf2f(row[k]);\n for (int o = 16; o > 0; o >>= 1) acc += __shfl_down_sync(0xFFFFFFFFu, acc, o);\n if (lane == 0) y[n] = acc;\n }\n gg.sync();\n}\n\n__device__ void h_rms(cg::grid_group &gg, const __nv_bfloat16* x,\n const __nv_bfloat16* w, __nv_bfloat16* y, int N, float* shm) {\n if (blockIdx.x == 0) {\n float s = 0.0f;\n for (int i = threadIdx.x; i < N; i += blockDim.x) {\n float v = bf2f(x[i]);\n s += v * v;\n }\n shm[threadIdx.x] = s;\n __syncthreads();\n for (int st = blockDim.x >> 1; st > 0; st >>= 1) {\n if (threadIdx.x < st) shm[threadIdx.x] += shm[threadIdx.x + st];\n __syncthreads();\n }\n float inv = rsqrtf(shm[0] / (float)N + 1e-6f);\n for (int i = threadIdx.x; i < N; i += blockDim.x)\n y[i] = f2bf(bf2f(x[i]) * inv * bf2f(w[i]));\n }\n gg.sync();\n}\n\n\n__device__ void h_topk(cg::grid_group &gg, const float* lg, int* topi, float* topw) {\n if (blockIdx.x == 0 && threadIdx.x == 0) {\n float pr[64];\n float m = -1e30f;\n for (int e = 0; e < 64; ++e) m = fmaxf(m, lg[e]);\n float s = 0.0f;\n for (int e = 0; e < 64; ++e) { pr[e] = expf(lg[e] - m); s += pr[e]; }\n float ws = 0.0f;\n for (int j = 0; j < 8; ++j) {\n int bi = -1;\n float bv = -1.0f;\n for (int e = 0; e < 64; ++e) {\n float v = pr[e] / s;\n if (v > bv) { bv = v; bi = e; }\n }\n if (bi < 0) { bi = 0; bv = 0.0f; }\n topi[j] = bi;\n topw[j] = bv;\n ws += bv;\n pr[bi] = -1.0f;\n }\n for (int j = 0; j < 8; ++j) topw[j] = topw[j] / ws * 2.446f;\n }\n gg.sync();\n}\n\n// short causal depthwise conv + SiLU over q/k/v (each 4096), update windows\n\n\n__device__ void h_conv(cg::grid_group &gg, __nv_bfloat16* q, __nv_bfloat16* k,\n __nv_bfloat16* v, __nv_bfloat16* cq, __nv_bfloat16* ck, __nv_bfloat16* cv,\n const __nv_bfloat16* cw, int tid, int TT) {\n for (int t = tid; t < 12288; t += TT) {\n int a = t >> 12;\n int c = t & 4095;\n __nv_bfloat16* buf = (a == 0) ? q : ((a == 1) ? k : v);\n __nv_bfloat16* win = (a == 0) ? cq : ((a == 1) ? ck : cv);\n float nv = bf2f(buf[c]);\n float w0 = bf2f(win[c]);\n float w1 = bf2f(win[4096 + c]);\n float w2 = bf2f(win[8192 + c]);\n const __nv_bfloat16* cc = cw + ((int64_t)a * 4096 + c) * 4;\n float s = w0 * bf2f(cc[0]) + w1 * bf2f(cc[1]) + w2 * bf2f(cc[2]) + nv * bf2f(cc[3]);\n float o = s / (1.0f + expf(-s));\n buf[c] = f2bf(o);\n __nv_bfloat16 orig = __float2bfloat16_rn(nv);\n win[c] = win[4096 + c];\n win[4096 + c] = win[8192 + c];\n win[8192 + c] = orig;\n }\n gg.sync();\n}\n\n// KDA gated-delta recurrence: S in/out fp32 (32,128,128), i-loop 2-way split\n\n\n__device__ void h_kdarec(cg::grid_group &gg, const __nv_bfloat16* q,\n const __nv_bfloat16* k, const __nv_bfloat16* v, const __nv_bfloat16* g,\n const float* beta, float* S, __nv_bfloat16* obuf, float QS, int tid, int TT) {\n for (int t = tid; t < 4096; t += TT) {\n int hh = t >> 7;\n int j = t & 127;\n float* Sh = S + (int64_t)hh * 16384;\n const __nv_bfloat16* kh = k + hh * 128;\n const __nv_bfloat16* vh = v + hh * 128;\n const __nv_bfloat16* gh = g + hh * 128;\n const __nv_bfloat16* qh = q + hh * 128;\n float b = beta[hh];\n float p0 = 0.0f, p1 = 0.0f;\n for (int i = 0; i < 64; ++i) {\n float s0 = Sh[i * 128 + j];\n float s1 = Sh[(i + 64) * 128 + j];\n s0 *= 1.0f / (1.0f + expf(bf2f(gh[i])));\n s1 *= 1.0f / (1.0f + expf(bf2f(gh[i + 64])));\n Sh[i * 128 + j] = s0;\n Sh[(i + 64) * 128 + j] = s1;\n p0 += s0 * bf2f(kh[i]);\n p1 += s1 * bf2f(kh[i + 64]);\n }\n float diff = bf2f(vh[j]) - (p0 + p1);\n float o0 = 0.0f, o1 = 0.0f;\n for (int i = 0; i < 64; ++i) {\n float k0 = bf2f(kh[i]), k1 = bf2f(kh[i + 64]);\n float s0 = Sh[i * 128 + j] + b * k0 * diff;\n float s1 = Sh[(i + 64) * 128 + j] + b * k1 * diff;\n Sh[i * 128 + j] = s0;\n Sh[(i + 64) * 128 + j] = s1;\n o0 += s0 * bf2f(qh[i]);\n o1 += s1 * bf2f(qh[i + 64]);\n }\n obuf[hh * 128 + j] = f2bf((o0 + o1) * QS);\n }\n gg.sync();\n}\n\n\n__device__ void h_res1(cg::grid_group &gg, __nv_bfloat16* h, const __nv_bfloat16* a,\n int N, int tid, int TT) {\n for (int i = tid; i < N; i += TT) h[i] = f2bf(bf2f(h[i]) + bf2f(a[i]));\n gg.sync();\n}\n\n\n__device__ void h_resacc(cg::grid_group &gg, __nv_bfloat16* h, const float* acc,\n int N, int tid, int TT) {\n for (int i = tid; i < N; i += TT) h[i] = f2bf(bf2f(h[i]) + acc[i]);\n gg.sync();\n}\n\n\n__device__ void h_hh(cg::grid_group &gg, const float* gf, const float* uf, float* hh,\n int N, int tid, int TT) {\n for (int i = tid; i < N; i += TT) {\n float x = gf[i];\n hh[i] = (x / (1.0f + expf(-x))) * uf[i];\n }\n gg.sync();\n}\n\n\n__device__ void h_zero(cg::grid_group &gg, float* p, int N, int tid, int TT) {\n for (int i = tid; i < N; i += TT) p[i] = 0.0f;\n gg.sync();\n}\n\n// MLA: rope q/k, copy old cache rows, append new row\n\n\n__device__ void h_rope(cg::grid_group &gg, __nv_bfloat16* mlaq, __nv_bfloat16* mlakv,\n const __nv_bfloat16* oc, const __nv_bfloat16* okk, __nv_bfloat16* nc,\n __nv_bfloat16* nk, int pos, int tid, int TT) {\n for (int t = tid; t < pos * 512; t += TT) nc[t] = oc[t];\n for (int t = tid; t < pos * 64; t += TT) nk[t] = okk[t];\n if (blockIdx.x == 0) {\n for (int t2 = threadIdx.x; t2 < 2048; t2 += blockDim.x) {\n int hh = t2 >> 6;\n int dd = t2 & 63;\n int p = dd >> 1;\n float inv = 1.0f / powf(10000.0f, (2.0f * (float)p) / 64.0f);\n float ang = (float)pos * inv;\n float co = cosf(ang), si = sinf(ang);\n __nv_bfloat16* rp = mlaq + hh * 192 + 128;\n float x0 = bf2f(rp[2 * p]), x1 = bf2f(rp[2 * p + 1]);\n rp[2 * p] = f2bf(x0 * co - x1 * si);\n rp[2 * p + 1] = f2bf(x1 * co + x0 * si);\n }\n for (int t2 = threadIdx.x; t2 < 64; t2 += blockDim.x) {\n int p = t2 >> 1;\n float inv = 1.0f / powf(10000.0f, (2.0f * (float)p) / 64.0f);\n float ang = (float)pos * inv;\n float co = cosf(ang), si = sinf(ang);\n float x0 = bf2f(mlakv[512 + 2 * p]), x1 = bf2f(mlakv[512 + 2 * p + 1]);\n float y0 = x0 * co - x1 * si, y1 = x1 * co + x0 * si;\n mlakv[512 + 2 * p] = f2bf(y0);\n mlakv[512 + 2 * p + 1] = f2bf(y1);\n nk[pos * 64 + 2 * p] = f2bf(y0);\n nk[pos * 64 + 2 * p + 1] = f2bf(y1);\n }\n for (int t2 = threadIdx.x; t2 < 512; t2 += blockDim.x) nc[pos * 512 + t2] = mlakv[t2];\n }\n gg.sync();\n}\n\n// absorbed query: 4 r per thread\n\n\n__device__ void h_qabs(cg::grid_group &gg, const __nv_bfloat16* mlaq,\n const unsigned char* Wbq, const __nv_bfloat16* Wbs, const __nv_bfloat16* Wbz,\n float* qabs, int tid, int TT) {\n for (int u = tid; u < 4096; u += TT) {\n int hh = u >> 7;\n int r0 = (u & 127) << 2;\n const __nv_bfloat16* qn = mlaq + hh * 192;\n float a0 = 0.0f, a1 = 0.0f, a2 = 0.0f, a3 = 0.0f;\n for (int d = 0; d < 128; ++d) {\n int n = hh * 256 + d;\n float qv = bf2f(qn[d]);\n int gb0 = (r0 >> 7) * 8192 + n;\n int gb1 = ((r0 + 3) >> 7) * 8192 + n;\n unsigned char bA = Wbq[(r0 >> 1) * 8192 + n];\n unsigned char bB = Wbq[((r0 + 2) >> 1) * 8192 + n];\n a0 += (((float)(bA & 0xF) - bf2f(Wbz[gb0])) * bf2f(Wbs[gb0])) * qv;\n a1 += (((float)(bA >> 4) - bf2f(Wbz[gb0])) * bf2f(Wbs[gb0])) * qv;\n a2 += (((float)(bB & 0xF) - bf2f(Wbz[gb1])) * bf2f(Wbs[gb1])) * qv;\n a3 += (((float)(bB >> 4) - bf2f(Wbz[gb1])) * bf2f(Wbs[gb1])) * qv;\n }\n qabs[hh * 512 + r0] = a0;\n qabs[hh * 512 + r0 + 1] = a1;\n qabs[hh * 512 + r0 + 2] = a2;\n qabs[hh * 512 + r0 + 3] = a3;\n }\n gg.sync();\n}\n\n// scores with 4-way split accumulation (no tail code: quarters cover L exactly)\n\n\n__device__ void h_softmax(cg::grid_group &gg, float* sc, int L, int cap, float* ms, float* shm) {\n if (blockIdx.x < 32) {\n int hh = blockIdx.x;\n float* row = sc + (int64_t)hh * cap;\n float m = -1e30f;\n for (int l = threadIdx.x; l < L; l += blockDim.x) m = fmaxf(m, row[l]);\n shm[threadIdx.x] = m;\n __syncthreads();\n for (int st = blockDim.x >> 1; st > 0; st >>= 1) {\n if (threadIdx.x < st) shm[threadIdx.x] = fmaxf(shm[threadIdx.x], shm[threadIdx.x + st]);\n __syncthreads();\n }\n if (threadIdx.x == 0) ms[hh] = shm[0];\n }\n gg.sync();\n int HL = 32 * L;\n for (int t = blockIdx.x * blockDim.x + threadIdx.x; t < HL; t += gridDim.x * blockDim.x) {\n int hh = t / L;\n int l = t % L;\n sc[(int64_t)hh * cap + l] = expf(sc[(int64_t)hh * cap + l] - ms[hh]);\n }\n gg.sync();\n if (blockIdx.x < 32) {\n int hh = blockIdx.x;\n float* row = sc + (int64_t)hh * cap;\n float s = 0.0f;\n for (int l = threadIdx.x; l < L; l += blockDim.x) s += row[l];\n shm[threadIdx.x] = s;\n __syncthreads();\n for (int st = blockDim.x >> 1; st > 0; st >>= 1) {\n if (threadIdx.x < st) shm[threadIdx.x] += shm[threadIdx.x + st];\n __syncthreads();\n }\n if (threadIdx.x == 0) ms[32 + hh] = shm[0];\n }\n gg.sync();\n for (int t = blockIdx.x * blockDim.x + threadIdx.x; t < HL; t += gridDim.x * blockDim.x) {\n int hh = t / L;\n int l = t % L;\n sc[(int64_t)hh * cap + l] /= ms[32 + hh];\n }\n gg.sync();\n}\n\n// y[hh*512+r] = sum_l p[hh,l] * nc[l,r]; block/head, l-loop 4 quarters\n\n\n// scores: warp-group per (head, token); r and rope dims lane-split\n// scores: warp-group per (head, token); r and rope dims lane-split\n__device__ void h_scores(cg::grid_group &gg, const float* qabs, const __nv_bfloat16* nc,\n const __nv_bfloat16* mlaq, const __nv_bfloat16* nk, int L, int cap,\n float* scores, float MS, int tid, int TT) {\n int lane = tid & 31;\n int grp = tid >> 5;\n int NG = TT >> 5;\n int HL = 32 * L;\n for (int t = grp; t < HL; t += NG) {\n int hh = t / L;\n int l = t % L;\n const float* qa = qabs + hh * 512;\n const __nv_bfloat16* cr = nc + (int64_t)l * 512;\n float acc = 0.0f;\n for (int r = lane; r < 512; r += 32) acc += qa[r] * bf2f(cr[r]);\n const __nv_bfloat16* qr = mlaq + hh * 192 + 128;\n const __nv_bfloat16* kr = nk + (int64_t)l * 64;\n for (int dd = lane; dd < 64; dd += 32) acc += bf2f(qr[dd]) * bf2f(kr[dd]);\n for (int o = 16; o > 0; o >>= 1) acc += __shfl_down_sync(0xFFFFFFFFu, acc, o);\n if (lane == 0) scores[(int64_t)hh * cap + l] = acc * MS;\n }\n gg.sync();\n}\n\n// y[hh*512+r] = sum_l p[hh,l] * nc[l,r]; warp per (head, r-block-of-32),\n// lane owns one r: row reads coalesced; l-range quartered for MLP\n__device__ void h_yml(cg::grid_group &gg, const float* sc, const __nv_bfloat16* nc,\n int L, int cap, float* yml, int tid, int TT) {\n int lane = tid & 31;\n int grp = tid >> 5;\n int NG = TT >> 5;\n int q1 = L >> 2, q2 = L >> 1, q3 = q1 + q2;\n for (int t = grp; t < 512; t += NG) {\n int hh = t >> 4;\n int r = ((t & 15) << 5) + lane;\n const float* pr = sc + (int64_t)hh * cap;\n const __nv_bfloat16* base = nc + r;\n float a0 = 0.0f, a1 = 0.0f, a2 = 0.0f, a3 = 0.0f;\n for (int l = 0; l < q1; ++l) a0 += pr[l] * bf2f(base[(int64_t)l * 512]);\n for (int l = q1; l < q2; ++l) a1 += pr[l] * bf2f(base[(int64_t)l * 512]);\n for (int l = q2; l < q3; ++l) a2 += pr[l] * bf2f(base[(int64_t)l * 512]);\n for (int l = q3; l < L; ++l) a3 += pr[l] * bf2f(base[(int64_t)l * 512]);\n yml[hh * 512 + r] = a0 + a1 + a2 + a3;\n }\n gg.sync();\n}\n\n__device__ void h_mlao(cg::grid_group &gg, const float* yml,\n const unsigned char* Wbq, const __nv_bfloat16* Wbs, const __nv_bfloat16* Wbz,\n __nv_bfloat16* o4) {\n if (blockIdx.x < 32) {\n int hh = blockIdx.x;\n const float* yh = yml + hh * 512;\n for (int u = threadIdx.x; u < 32; u += blockDim.x) {\n int dd0 = u << 2;\n float a0 = 0.0f, a1 = 0.0f, a2 = 0.0f, a3 = 0.0f;\n for (int r = 0; r < 512; ++r) {\n float yv = yh[r];\n int gb = (r >> 7) * 8192 + hh * 256 + 128;\n unsigned char b0 = Wbq[(r >> 1) * 8192 + hh * 256 + 128 + dd0];\n int q0 = (r & 1) ? (b0 >> 4) : (b0 & 0xF);\n a0 += yv * (((float)q0 - bf2f(Wbz[gb + dd0])) * bf2f(Wbs[gb + dd0]));\n unsigned char b1 = Wbq[(r >> 1) * 8192 + hh * 256 + 128 + dd0 + 1];\n int q1 = (r & 1) ? (b1 >> 4) : (b1 & 0xF);\n a1 += yv * (((float)q1 - bf2f(Wbz[gb + dd0 + 1])) * bf2f(Wbs[gb + dd0 + 1]));\n unsigned char b2 = Wbq[(r >> 1) * 8192 + hh * 256 + 128 + dd0 + 2];\n int q2 = (r & 1) ? (b2 >> 4) : (b2 & 0xF);\n a2 += yv * (((float)q2 - bf2f(Wbz[gb + dd0 + 2])) * bf2f(Wbs[gb + dd0 + 2]));\n unsigned char b3 = Wbq[(r >> 1) * 8192 + hh * 256 + 128 + dd0 + 3];\n int q3 = (r & 1) ? (b3 >> 4) : (b3 & 0xF);\n a3 += yv * (((float)q3 - bf2f(Wbz[gb + dd0 + 3])) * bf2f(Wbs[gb + dd0 + 3]));\n }\n o4[hh * 128 + dd0] = f2bf(a0);\n o4[hh * 128 + dd0 + 1] = f2bf(a1);\n o4[hh * 128 + dd0 + 2] = f2bf(a2);\n o4[hh * 128 + dd0 + 3] = f2bf(a3);\n }\n }\n gg.sync();\n}\n\n\n__global__ void megakernel(\n __nv_bfloat16* h,\n __nv_bfloat16* SBp,\n float* SFp,\n int* SIp,\n const unsigned char* L0qw,\n const __nv_bfloat16* L0qs,\n const __nv_bfloat16* L0qz,\n const unsigned char* L0kw,\n const __nv_bfloat16* L0ks,\n const __nv_bfloat16* L0kz,\n const unsigned char* L0vw,\n const __nv_bfloat16* L0vs,\n const __nv_bfloat16* L0vz,\n const unsigned char* L0gw,\n const __nv_bfloat16* L0gs,\n const __nv_bfloat16* L0gz,\n const unsigned char* L0ow,\n const __nv_bfloat16* L0os,\n const __nv_bfloat16* L0oz,\n const __nv_bfloat16* L0bw,\n const __nv_bfloat16* L0cw,\n float* L0S,\n __nv_bfloat16* L0cq,\n __nv_bfloat16* L0ck,\n __nv_bfloat16* L0cv,\n const __nv_bfloat16* L0an,\n const __nv_bfloat16* L0mn,\n const __nv_bfloat16* L0rt,\n const unsigned char* L0gqw,\n const __nv_bfloat16* L0gqs,\n const __nv_bfloat16* L0gqz,\n const unsigned char* L0uqw,\n const __nv_bfloat16* L0uqs,\n const __nv_bfloat16* L0uqz,\n const unsigned char* L0dqw,\n const __nv_bfloat16* L0dqs,\n const __nv_bfloat16* L0dqz,\n const unsigned char* L0sgqw,\n const __nv_bfloat16* L0sgqs,\n const __nv_bfloat16* L0sgqz,\n const unsigned char* L0suqw,\n const __nv_bfloat16* L0suqs,\n const __nv_bfloat16* L0suqz,\n const unsigned char* L0sdqw,\n const __nv_bfloat16* L0sdqs,\n const __nv_bfloat16* L0sdqz,\n const unsigned char* L1qw,\n const __nv_bfloat16* L1qs,\n const __nv_bfloat16* L1qz,\n const unsigned char* L1kw,\n const __nv_bfloat16* L1ks,\n const __nv_bfloat16* L1kz,\n const unsigned char* L1vw,\n const __nv_bfloat16* L1vs,\n const __nv_bfloat16* L1vz,\n const unsigned char* L1gw,\n const __nv_bfloat16* L1gs,\n const __nv_bfloat16* L1gz,\n const unsigned char* L1ow,\n const __nv_bfloat16* L1os,\n const __nv_bfloat16* L1oz,\n const __nv_bfloat16* L1bw,\n const __nv_bfloat16* L1cw,\n float* L1S,\n __nv_bfloat16* L1cq,\n __nv_bfloat16* L1ck,\n __nv_bfloat16* L1cv,\n const __nv_bfloat16* L1an,\n const __nv_bfloat16* L1mn,\n const __nv_bfloat16* L1rt,\n const unsigned char* L1gqw,\n const __nv_bfloat16* L1gqs,\n const __nv_bfloat16* L1gqz,\n const unsigned char* L1uqw,\n const __nv_bfloat16* L1uqs,\n const __nv_bfloat16* L1uqz,\n const unsigned char* L1dqw,\n const __nv_bfloat16* L1dqs,\n const __nv_bfloat16* L1dqz,\n const unsigned char* L1sgqw,\n const __nv_bfloat16* L1sgqs,\n const __nv_bfloat16* L1sgqz,\n const unsigned char* L1suqw,\n const __nv_bfloat16* L1suqs,\n const __nv_bfloat16* L1suqz,\n const unsigned char* L1sdqw,\n const __nv_bfloat16* L1sdqs,\n const __nv_bfloat16* L1sdqz,\n const unsigned char* L2qw,\n const __nv_bfloat16* L2qs,\n const __nv_bfloat16* L2qz,\n const unsigned char* L2kw,\n const __nv_bfloat16* L2ks,\n const __nv_bfloat16* L2kz,\n const unsigned char* L2vw,\n const __nv_bfloat16* L2vs,\n const __nv_bfloat16* L2vz,\n const unsigned char* L2gw,\n const __nv_bfloat16* L2gs,\n const __nv_bfloat16* L2gz,\n const unsigned char* L2ow,\n const __nv_bfloat16* L2os,\n const __nv_bfloat16* L2oz,\n const __nv_bfloat16* L2bw,\n const __nv_bfloat16* L2cw,\n float* L2S,\n __nv_bfloat16* L2cq,\n __nv_bfloat16* L2ck,\n __nv_bfloat16* L2cv,\n const __nv_bfloat16* L2an,\n const __nv_bfloat16* L2mn,\n const __nv_bfloat16* L2rt,\n const unsigned char* L2gqw,\n const __nv_bfloat16* L2gqs,\n const __nv_bfloat16* L2gqz,\n const unsigned char* L2uqw,\n const __nv_bfloat16* L2uqs,\n const __nv_bfloat16* L2uqz,\n const unsigned char* L2dqw,\n const __nv_bfloat16* L2dqs,\n const __nv_bfloat16* L2dqz,\n const unsigned char* L2sgqw,\n const __nv_bfloat16* L2sgqs,\n const __nv_bfloat16* L2sgqz,\n const unsigned char* L2suqw,\n const __nv_bfloat16* L2suqs,\n const __nv_bfloat16* L2suqz,\n const unsigned char* L2sdqw,\n const __nv_bfloat16* L2sdqs,\n const __nv_bfloat16* L2sdqz,\n const unsigned char* Mqw,\n const __nv_bfloat16* Mqs,\n const __nv_bfloat16* Mqz,\n const unsigned char* Maw,\n const __nv_bfloat16* Mas,\n const __nv_bfloat16* Maz,\n const unsigned char* Mbw,\n const __nv_bfloat16* Mbs,\n const __nv_bfloat16* Mbz,\n const unsigned char* Mow,\n const __nv_bfloat16* Mos,\n const __nv_bfloat16* Moz,\n const __nv_bfloat16* Man,\n const __nv_bfloat16* Mmn,\n const __nv_bfloat16* Mrt,\n const unsigned char* Mgqw,\n const __nv_bfloat16* Mgqs,\n const __nv_bfloat16* Mgqz,\n const unsigned char* Muqw,\n const __nv_bfloat16* Muqs,\n const __nv_bfloat16* Muqz,\n const unsigned char* Mdqw,\n const __nv_bfloat16* Mdqs,\n const __nv_bfloat16* Mdqz,\n const unsigned char* Msgqw,\n const __nv_bfloat16* Msgqs,\n const __nv_bfloat16* Msgqz,\n const unsigned char* Msuqw,\n const __nv_bfloat16* Msuqs,\n const __nv_bfloat16* Msuqz,\n const unsigned char* Msdqw,\n const __nv_bfloat16* Msdqs,\n const __nv_bfloat16* Msdqz,\n const __nv_bfloat16* oc,\n const __nv_bfloat16* ok,\n __nv_bfloat16* nc,\n __nv_bfloat16* nk,\n long long curlen, long long capx) {\n cg::grid_group gg = cg::this_grid();\n __shared__ float shm[4096];\n int tid = blockIdx.x * blockDim.x + threadIdx.x;\n int TT = gridDim.x * blockDim.x;\n const int64_t EGQ = 1179648, EGS = 18432, EDQ = 1179648, EDS = 18432;\n __nv_bfloat16 *xn = SBp + 0, *hn = SBp + 2304, *att = SBp + 4608;\n __nv_bfloat16 *q = SBp + 9216, *k = SBp + 13312, *v = SBp + 17408, *g = SBp + 21504;\n __nv_bfloat16 *o4 = SBp + 25600, *mlaq = SBp + 29696, *mlakv = SBp + 35840;\n float *qabs = SFp + 0, *yml = SFp + 16384, *logits = SFp + 32768, *topw = SFp + 32832;\n float *acc = SFp + 32840, *hh = SFp + 35144, *gf = SFp + 36168, *uf = SFp + 37192;\n float *beta = SFp + 38216, *gf8 = SFp + 38248, *uf8 = SFp + 46440;\n float *hh8 = SFp + 54632, *scores = SFp + 62824;\n int *topi = SIp;\n int Lcur = (int)curlen;\n int Ccap = (int)capx;\n int Ltot = Lcur + 1;\n h_rms(gg, h, L0an, xn, 2304, shm);\n h_w4x4(gg, xn, L0qw, L0qs, L0qz, L0kw, L0ks, L0kz, L0vw, L0vs, L0vz, L0gw, L0gs, L0gz, q, k, v, g, 2304, 4096, shm, tid, TT);\n h_conv(gg, q, k, v, L0cq, L0ck, L0cv, L0cw, tid, TT);\n h_bfgemv(gg, xn, L0bw, beta, 2304, 32, shm, tid, TT);\n h_kdarec(gg, q, k, v, g, beta, L0S, o4, 0.0883883476f, tid, TT);\n h_w4_resadd(gg, o4, L0ow, L0os, L0oz, h, 4096, 2304, shm, tid, TT);\n h_rms(gg, h, L0mn, hn, 2304, shm);\n h_bfgemv(gg, hn, L0rt, logits, 2304, 64, shm, tid, TT);\n h_topk(gg, logits, topi, topw);\n h_zero(gg, acc, 2304, tid, TT);\n h_moe_gateup(gg, hn, L0gqw, L0gqs, L0gqz, L0uqw, L0uqs, L0uqz, topi, gf8, uf8, hh8, 2304, 1024, shm, tid, TT);\n h_moe_down(gg, hh8, L0dqw, L0dqs, L0dqz, topi, topw, acc, 1024, 2304, shm, tid, TT);\n h_w4dual(gg, hn, L0sgqw, L0sgqs, L0sgqz, L0suqw, L0suqs, L0suqz, gf, uf, 2304, 1024, shm, tid, TT);\n h_hh(gg, gf, uf, hh, 1024, tid, TT);\n h_w4down(gg, hh, L0sdqw, L0sdqs, L0sdqz, acc, 1.0f, h, 1024, 2304, shm, tid, TT);\n h_rms(gg, h, L1an, xn, 2304, shm);\n h_w4x4(gg, xn, L1qw, L1qs, L1qz, L1kw, L1ks, L1kz, L1vw, L1vs, L1vz, L1gw, L1gs, L1gz, q, k, v, g, 2304, 4096, shm, tid, TT);\n h_conv(gg, q, k, v, L1cq, L1ck, L1cv, L1cw, tid, TT);\n h_bfgemv(gg, xn, L1bw, beta, 2304, 32, shm, tid, TT);\n h_kdarec(gg, q, k, v, g, beta, L1S, o4, 0.0883883476f, tid, TT);\n h_w4_resadd(gg, o4, L1ow, L1os, L1oz, h, 4096, 2304, shm, tid, TT);\n h_rms(gg, h, L1mn, hn, 2304, shm);\n h_bfgemv(gg, hn, L1rt, logits, 2304, 64, shm, tid, TT);\n h_topk(gg, logits, topi, topw);\n h_zero(gg, acc, 2304, tid, TT);\n h_moe_gateup(gg, hn, L1gqw, L1gqs, L1gqz, L1uqw, L1uqs, L1uqz, topi, gf8, uf8, hh8, 2304, 1024, shm, tid, TT);\n h_moe_down(gg, hh8, L1dqw, L1dqs, L1dqz, topi, topw, acc, 1024, 2304, shm, tid, TT);\n h_w4dual(gg, hn, L1sgqw, L1sgqs, L1sgqz, L1suqw, L1suqs, L1suqz, gf, uf, 2304, 1024, shm, tid, TT);\n h_hh(gg, gf, uf, hh, 1024, tid, TT);\n h_w4down(gg, hh, L1sdqw, L1sdqs, L1sdqz, acc, 1.0f, h, 1024, 2304, shm, tid, TT);\n h_rms(gg, h, L2an, xn, 2304, shm);\n h_w4x4(gg, xn, L2qw, L2qs, L2qz, L2kw, L2ks, L2kz, L2vw, L2vs, L2vz, L2gw, L2gs, L2gz, q, k, v, g, 2304, 4096, shm, tid, TT);\n h_conv(gg, q, k, v, L2cq, L2ck, L2cv, L2cw, tid, TT);\n h_bfgemv(gg, xn, L2bw, beta, 2304, 32, shm, tid, TT);\n h_kdarec(gg, q, k, v, g, beta, L2S, o4, 0.0883883476f, tid, TT);\n h_w4_resadd(gg, o4, L2ow, L2os, L2oz, h, 4096, 2304, shm, tid, TT);\n h_rms(gg, h, L2mn, hn, 2304, shm);\n h_bfgemv(gg, hn, L2rt, logits, 2304, 64, shm, tid, TT);\n h_topk(gg, logits, topi, topw);\n h_zero(gg, acc, 2304, tid, TT);\n h_moe_gateup(gg, hn, L2gqw, L2gqs, L2gqz, L2uqw, L2uqs, L2uqz, topi, gf8, uf8, hh8, 2304, 1024, shm, tid, TT);\n h_moe_down(gg, hh8, L2dqw, L2dqs, L2dqz, topi, topw, acc, 1024, 2304, shm, tid, TT);\n h_w4dual(gg, hn, L2sgqw, L2sgqs, L2sgqz, L2suqw, L2suqs, L2suqz, gf, uf, 2304, 1024, shm, tid, TT);\n h_hh(gg, gf, uf, hh, 1024, tid, TT);\n h_w4down(gg, hh, L2sdqw, L2sdqs, L2sdqz, acc, 1.0f, h, 1024, 2304, shm, tid, TT);\n h_rms(gg, h, Man, xn, 2304, shm);\n h_w4(gg, xn, Mqw, Mqs, Mqz, mlaq, 2304, 6144, shm, tid, TT);\n h_w4(gg, xn, Maw, Mas, Maz, mlakv, 2304, 576, shm, tid, TT);\n h_rope(gg, mlaq, mlakv, oc, ok, nc, nk, Lcur, tid, TT);\n h_qabs(gg, mlaq, Mbw, Mbs, Mbz, qabs, tid, TT);\n h_scores(gg, qabs, nc, mlaq, nk, Ltot, Ccap, scores, 0.0721687836f, tid, TT);\n h_softmax(gg, scores, Ltot, Ccap, logits, shm);\n h_yml(gg, scores, nc, Ltot, Ccap, yml, tid, TT);\n h_mlao(gg, yml, Mbw, Mbs, Mbz, o4);\n h_w4_resadd(gg, o4, Mow, Mos, Moz, h, 4096, 2304, shm, tid, TT);\n h_rms(gg, h, Mmn, hn, 2304, shm);\n h_bfgemv(gg, hn, Mrt, logits, 2304, 64, shm, tid, TT);\n h_topk(gg, logits, topi, topw);\n h_zero(gg, acc, 2304, tid, TT);\n h_moe_gateup(gg, hn, Mgqw, Mgqs, Mgqz, Muqw, Muqs, Muqz, topi, gf8, uf8, hh8, 2304, 1024, shm, tid, TT);\n h_moe_down(gg, hh8, Mdqw, Mdqs, Mdqz, topi, topw, acc, 1024, 2304, shm, tid, TT);\n h_w4dual(gg, hn, Msgqw, Msgqs, Msgqz, Msuqw, Msuqs, Msuqz, gf, uf, 2304, 1024, shm, tid, TT);\n h_hh(gg, gf, uf, hh, 1024, tid, TT);\n h_w4down(gg, hh, Msdqw, Msdqs, Msdqz, acc, 1.0f, h, 1024, 2304, shm, tid, TT);\n}\n\n\nvoid mega_launch(torch::Tensor h, torch::Tensor SBp, torch::Tensor SFp, torch::Tensor SIp, torch::Tensor L0qw, torch::Tensor L0qs, torch::Tensor L0qz, torch::Tensor L0kw, torch::Tensor L0ks, torch::Tensor L0kz, torch::Tensor L0vw, torch::Tensor L0vs, torch::Tensor L0vz, torch::Tensor L0gw, torch::Tensor L0gs, torch::Tensor L0gz, torch::Tensor L0ow, torch::Tensor L0os, torch::Tensor L0oz, torch::Tensor L0bw, torch::Tensor L0cw, torch::Tensor L0S, torch::Tensor L0cq, torch::Tensor L0ck, torch::Tensor L0cv, torch::Tensor L0an, torch::Tensor L0mn, torch::Tensor L0rt, torch::Tensor L0gqw, torch::Tensor L0gqs, torch::Tensor L0gqz, torch::Tensor L0uqw, torch::Tensor L0uqs, torch::Tensor L0uqz, torch::Tensor L0dqw, torch::Tensor L0dqs, torch::Tensor L0dqz, torch::Tensor L0sgqw, torch::Tensor L0sgqs, torch::Tensor L0sgqz, torch::Tensor L0suqw, torch::Tensor L0suqs, torch::Tensor L0suqz, torch::Tensor L0sdqw, torch::Tensor L0sdqs, torch::Tensor L0sdqz, torch::Tensor L1qw, torch::Tensor L1qs, torch::Tensor L1qz, torch::Tensor L1kw, torch::Tensor L1ks, torch::Tensor L1kz, torch::Tensor L1vw, torch::Tensor L1vs, torch::Tensor L1vz, torch::Tensor L1gw, torch::Tensor L1gs, torch::Tensor L1gz, torch::Tensor L1ow, torch::Tensor L1os, torch::Tensor L1oz, torch::Tensor L1bw, torch::Tensor L1cw, torch::Tensor L1S, torch::Tensor L1cq, torch::Tensor L1ck, torch::Tensor L1cv, torch::Tensor L1an, torch::Tensor L1mn, torch::Tensor L1rt, torch::Tensor L1gqw, torch::Tensor L1gqs, torch::Tensor L1gqz, torch::Tensor L1uqw, torch::Tensor L1uqs, torch::Tensor L1uqz, torch::Tensor L1dqw, torch::Tensor L1dqs, torch::Tensor L1dqz, torch::Tensor L1sgqw, torch::Tensor L1sgqs, torch::Tensor L1sgqz, torch::Tensor L1suqw, torch::Tensor L1suqs, torch::Tensor L1suqz, torch::Tensor L1sdqw, torch::Tensor L1sdqs, torch::Tensor L1sdqz, torch::Tensor L2qw, torch::Tensor L2qs, torch::Tensor L2qz, torch::Tensor L2kw, torch::Tensor L2ks, torch::Tensor L2kz, torch::Tensor L2vw, torch::Tensor L2vs, torch::Tensor L2vz, torch::Tensor L2gw, torch::Tensor L2gs, torch::Tensor L2gz, torch::Tensor L2ow, torch::Tensor L2os, torch::Tensor L2oz, torch::Tensor L2bw, torch::Tensor L2cw, torch::Tensor L2S, torch::Tensor L2cq, torch::Tensor L2ck, torch::Tensor L2cv, torch::Tensor L2an, torch::Tensor L2mn, torch::Tensor L2rt, torch::Tensor L2gqw, torch::Tensor L2gqs, torch::Tensor L2gqz, torch::Tensor L2uqw, torch::Tensor L2uqs, torch::Tensor L2uqz, torch::Tensor L2dqw, torch::Tensor L2dqs, torch::Tensor L2dqz, torch::Tensor L2sgqw, torch::Tensor L2sgqs, torch::Tensor L2sgqz, torch::Tensor L2suqw, torch::Tensor L2suqs, torch::Tensor L2suqz, torch::Tensor L2sdqw, torch::Tensor L2sdqs, torch::Tensor L2sdqz, torch::Tensor Mqw, torch::Tensor Mqs, torch::Tensor Mqz, torch::Tensor Maw, torch::Tensor Mas, torch::Tensor Maz, torch::Tensor Mbw, torch::Tensor Mbs, torch::Tensor Mbz, torch::Tensor Mow, torch::Tensor Mos, torch::Tensor Moz, torch::Tensor Man, torch::Tensor Mmn, torch::Tensor Mrt, torch::Tensor Mgqw, torch::Tensor Mgqs, torch::Tensor Mgqz, torch::Tensor Muqw, torch::Tensor Muqs, torch::Tensor Muqz, torch::Tensor Mdqw, torch::Tensor Mdqs, torch::Tensor Mdqz, torch::Tensor Msgqw, torch::Tensor Msgqs, torch::Tensor Msgqz, torch::Tensor Msuqw, torch::Tensor Msuqs, torch::Tensor Msuqz, torch::Tensor Msdqw, torch::Tensor Msdqs, torch::Tensor Msdqz, torch::Tensor oc, torch::Tensor ok, torch::Tensor nc, torch::Tensor nk, long long curlen, long long capx) {\n __nv_bfloat16* p_h = (__nv_bfloat16*)h.data_ptr();\n __nv_bfloat16* p_SBp = (__nv_bfloat16*)SBp.data_ptr();\n float* p_SFp = (float*)SFp.data_ptr();\n int* p_SIp = (int*)SIp.data_ptr();\n const unsigned char* p_L0qw = (const unsigned char*)L0qw.data_ptr();\n const __nv_bfloat16* p_L0qs = (const __nv_bfloat16*)L0qs.data_ptr();\n const __nv_bfloat16* p_L0qz = (const __nv_bfloat16*)L0qz.data_ptr();\n const unsigned char* p_L0kw = (const unsigned char*)L0kw.data_ptr();\n const __nv_bfloat16* p_L0ks = (const __nv_bfloat16*)L0ks.data_ptr();\n const __nv_bfloat16* p_L0kz = (const __nv_bfloat16*)L0kz.data_ptr();\n const unsigned char* p_L0vw = (const unsigned char*)L0vw.data_ptr();\n const __nv_bfloat16* p_L0vs = (const __nv_bfloat16*)L0vs.data_ptr();\n const __nv_bfloat16* p_L0vz = (const __nv_bfloat16*)L0vz.data_ptr();\n const unsigned char* p_L0gw = (const unsigned char*)L0gw.data_ptr();\n const __nv_bfloat16* p_L0gs = (const __nv_bfloat16*)L0gs.data_ptr();\n const __nv_bfloat16* p_L0gz = (const __nv_bfloat16*)L0gz.data_ptr();\n const unsigned char* p_L0ow = (const unsigned char*)L0ow.data_ptr();\n const __nv_bfloat16* p_L0os = (const __nv_bfloat16*)L0os.data_ptr();\n const __nv_bfloat16* p_L0oz = (const __nv_bfloat16*)L0oz.data_ptr();\n const __nv_bfloat16* p_L0bw = (const __nv_bfloat16*)L0bw.data_ptr();\n const __nv_bfloat16* p_L0cw = (const __nv_bfloat16*)L0cw.data_ptr();\n float* p_L0S = (float*)L0S.data_ptr();\n __nv_bfloat16* p_L0cq = (__nv_bfloat16*)L0cq.data_ptr();\n __nv_bfloat16* p_L0ck = (__nv_bfloat16*)L0ck.data_ptr();\n __nv_bfloat16* p_L0cv = (__nv_bfloat16*)L0cv.data_ptr();\n const __nv_bfloat16* p_L0an = (const __nv_bfloat16*)L0an.data_ptr();\n const __nv_bfloat16* p_L0mn = (const __nv_bfloat16*)L0mn.data_ptr();\n const __nv_bfloat16* p_L0rt = (const __nv_bfloat16*)L0rt.data_ptr();\n const unsigned char* p_L0gqw = (const unsigned char*)L0gqw.data_ptr();\n const __nv_bfloat16* p_L0gqs = (const __nv_bfloat16*)L0gqs.data_ptr();\n const __nv_bfloat16* p_L0gqz = (const __nv_bfloat16*)L0gqz.data_ptr();\n const unsigned char* p_L0uqw = (const unsigned char*)L0uqw.data_ptr();\n const __nv_bfloat16* p_L0uqs = (const __nv_bfloat16*)L0uqs.data_ptr();\n const __nv_bfloat16* p_L0uqz = (const __nv_bfloat16*)L0uqz.data_ptr();\n const unsigned char* p_L0dqw = (const unsigned char*)L0dqw.data_ptr();\n const __nv_bfloat16* p_L0dqs = (const __nv_bfloat16*)L0dqs.data_ptr();\n const __nv_bfloat16* p_L0dqz = (const __nv_bfloat16*)L0dqz.data_ptr();\n const unsigned char* p_L0sgqw = (const unsigned char*)L0sgqw.data_ptr();\n const __nv_bfloat16* p_L0sgqs = (const __nv_bfloat16*)L0sgqs.data_ptr();\n const __nv_bfloat16* p_L0sgqz = (const __nv_bfloat16*)L0sgqz.data_ptr();\n const unsigned char* p_L0suqw = (const unsigned char*)L0suqw.data_ptr();\n const __nv_bfloat16* p_L0suqs = (const __nv_bfloat16*)L0suqs.data_ptr();\n const __nv_bfloat16* p_L0suqz = (const __nv_bfloat16*)L0suqz.data_ptr();\n const unsigned char* p_L0sdqw = (const unsigned char*)L0sdqw.data_ptr();\n const __nv_bfloat16* p_L0sdqs = (const __nv_bfloat16*)L0sdqs.data_ptr();\n const __nv_bfloat16* p_L0sdqz = (const __nv_bfloat16*)L0sdqz.data_ptr();\n const unsigned char* p_L1qw = (const unsigned char*)L1qw.data_ptr();\n const __nv_bfloat16* p_L1qs = (const __nv_bfloat16*)L1qs.data_ptr();\n const __nv_bfloat16* p_L1qz = (const __nv_bfloat16*)L1qz.data_ptr();\n const unsigned char* p_L1kw = (const unsigned char*)L1kw.data_ptr();\n const __nv_bfloat16* p_L1ks = (const __nv_bfloat16*)L1ks.data_ptr();\n const __nv_bfloat16* p_L1kz = (const __nv_bfloat16*)L1kz.data_ptr();\n const unsigned char* p_L1vw = (const unsigned char*)L1vw.data_ptr();\n const __nv_bfloat16* p_L1vs = (const __nv_bfloat16*)L1vs.data_ptr();\n const __nv_bfloat16* p_L1vz = (const __nv_bfloat16*)L1vz.data_ptr();\n const unsigned char* p_L1gw = (const unsigned char*)L1gw.data_ptr();\n const __nv_bfloat16* p_L1gs = (const __nv_bfloat16*)L1gs.data_ptr();\n const __nv_bfloat16* p_L1gz = (const __nv_bfloat16*)L1gz.data_ptr();\n const unsigned char* p_L1ow = (const unsigned char*)L1ow.data_ptr();\n const __nv_bfloat16* p_L1os = (const __nv_bfloat16*)L1os.data_ptr();\n const __nv_bfloat16* p_L1oz = (const __nv_bfloat16*)L1oz.data_ptr();\n const __nv_bfloat16* p_L1bw = (const __nv_bfloat16*)L1bw.data_ptr();\n const __nv_bfloat16* p_L1cw = (const __nv_bfloat16*)L1cw.data_ptr();\n float* p_L1S = (float*)L1S.data_ptr();\n __nv_bfloat16* p_L1cq = (__nv_bfloat16*)L1cq.data_ptr();\n __nv_bfloat16* p_L1ck = (__nv_bfloat16*)L1ck.data_ptr();\n __nv_bfloat16* p_L1cv = (__nv_bfloat16*)L1cv.data_ptr();\n const __nv_bfloat16* p_L1an = (const __nv_bfloat16*)L1an.data_ptr();\n const __nv_bfloat16* p_L1mn = (const __nv_bfloat16*)L1mn.data_ptr();\n const __nv_bfloat16* p_L1rt = (const __nv_bfloat16*)L1rt.data_ptr();\n const unsigned char* p_L1gqw = (const unsigned char*)L1gqw.data_ptr();\n const __nv_bfloat16* p_L1gqs = (const __nv_bfloat16*)L1gqs.data_ptr();\n const __nv_bfloat16* p_L1gqz = (const __nv_bfloat16*)L1gqz.data_ptr();\n const unsigned char* p_L1uqw = (const unsigned char*)L1uqw.data_ptr();\n const __nv_bfloat16* p_L1uqs = (const __nv_bfloat16*)L1uqs.data_ptr();\n const __nv_bfloat16* p_L1uqz = (const __nv_bfloat16*)L1uqz.data_ptr();\n const unsigned char* p_L1dqw = (const unsigned char*)L1dqw.data_ptr();\n const __nv_bfloat16* p_L1dqs = (const __nv_bfloat16*)L1dqs.data_ptr();\n const __nv_bfloat16* p_L1dqz = (const __nv_bfloat16*)L1dqz.data_ptr();\n const unsigned char* p_L1sgqw = (const unsigned char*)L1sgqw.data_ptr();\n const __nv_bfloat16* p_L1sgqs = (const __nv_bfloat16*)L1sgqs.data_ptr();\n const __nv_bfloat16* p_L1sgqz = (const __nv_bfloat16*)L1sgqz.data_ptr();\n const unsigned char* p_L1suqw = (const unsigned char*)L1suqw.data_ptr();\n const __nv_bfloat16* p_L1suqs = (const __nv_bfloat16*)L1suqs.data_ptr();\n const __nv_bfloat16* p_L1suqz = (const __nv_bfloat16*)L1suqz.data_ptr();\n const unsigned char* p_L1sdqw = (const unsigned char*)L1sdqw.data_ptr();\n const __nv_bfloat16* p_L1sdqs = (const __nv_bfloat16*)L1sdqs.data_ptr();\n const __nv_bfloat16* p_L1sdqz = (const __nv_bfloat16*)L1sdqz.data_ptr();\n const unsigned char* p_L2qw = (const unsigned char*)L2qw.data_ptr();\n const __nv_bfloat16* p_L2qs = (const __nv_bfloat16*)L2qs.data_ptr();\n const __nv_bfloat16* p_L2qz = (const __nv_bfloat16*)L2qz.data_ptr();\n const unsigned char* p_L2kw = (const unsigned char*)L2kw.data_ptr();\n const __nv_bfloat16* p_L2ks = (const __nv_bfloat16*)L2ks.data_ptr();\n const __nv_bfloat16* p_L2kz = (const __nv_bfloat16*)L2kz.data_ptr();\n const unsigned char* p_L2vw = (const unsigned char*)L2vw.data_ptr();\n const __nv_bfloat16* p_L2vs = (const __nv_bfloat16*)L2vs.data_ptr();\n const __nv_bfloat16* p_L2vz = (const __nv_bfloat16*)L2vz.data_ptr();\n const unsigned char* p_L2gw = (const unsigned char*)L2gw.data_ptr();\n const __nv_bfloat16* p_L2gs = (const __nv_bfloat16*)L2gs.data_ptr();\n const __nv_bfloat16* p_L2gz = (const __nv_bfloat16*)L2gz.data_ptr();\n const unsigned char* p_L2ow = (const unsigned char*)L2ow.data_ptr();\n const __nv_bfloat16* p_L2os = (const __nv_bfloat16*)L2os.data_ptr();\n const __nv_bfloat16* p_L2oz = (const __nv_bfloat16*)L2oz.data_ptr();\n const __nv_bfloat16* p_L2bw = (const __nv_bfloat16*)L2bw.data_ptr();\n const __nv_bfloat16* p_L2cw = (const __nv_bfloat16*)L2cw.data_ptr();\n float* p_L2S = (float*)L2S.data_ptr();\n __nv_bfloat16* p_L2cq = (__nv_bfloat16*)L2cq.data_ptr();\n __nv_bfloat16* p_L2ck = (__nv_bfloat16*)L2ck.data_ptr();\n __nv_bfloat16* p_L2cv = (__nv_bfloat16*)L2cv.data_ptr();\n const __nv_bfloat16* p_L2an = (const __nv_bfloat16*)L2an.data_ptr();\n const __nv_bfloat16* p_L2mn = (const __nv_bfloat16*)L2mn.data_ptr();\n const __nv_bfloat16* p_L2rt = (const __nv_bfloat16*)L2rt.data_ptr();\n const unsigned char* p_L2gqw = (const unsigned char*)L2gqw.data_ptr();\n const __nv_bfloat16* p_L2gqs = (const __nv_bfloat16*)L2gqs.data_ptr();\n const __nv_bfloat16* p_L2gqz = (const __nv_bfloat16*)L2gqz.data_ptr();\n const unsigned char* p_L2uqw = (const unsigned char*)L2uqw.data_ptr();\n const __nv_bfloat16* p_L2uqs = (const __nv_bfloat16*)L2uqs.data_ptr();\n const __nv_bfloat16* p_L2uqz = (const __nv_bfloat16*)L2uqz.data_ptr();\n const unsigned char* p_L2dqw = (const unsigned char*)L2dqw.data_ptr();\n const __nv_bfloat16* p_L2dqs = (const __nv_bfloat16*)L2dqs.data_ptr();\n const __nv_bfloat16* p_L2dqz = (const __nv_bfloat16*)L2dqz.data_ptr();\n const unsigned char* p_L2sgqw = (const unsigned char*)L2sgqw.data_ptr();\n const __nv_bfloat16* p_L2sgqs = (const __nv_bfloat16*)L2sgqs.data_ptr();\n const __nv_bfloat16* p_L2sgqz = (const __nv_bfloat16*)L2sgqz.data_ptr();\n const unsigned char* p_L2suqw = (const unsigned char*)L2suqw.data_ptr();\n const __nv_bfloat16* p_L2suqs = (const __nv_bfloat16*)L2suqs.data_ptr();\n const __nv_bfloat16* p_L2suqz = (const __nv_bfloat16*)L2suqz.data_ptr();\n const unsigned char* p_L2sdqw = (const unsigned char*)L2sdqw.data_ptr();\n const __nv_bfloat16* p_L2sdqs = (const __nv_bfloat16*)L2sdqs.data_ptr();\n const __nv_bfloat16* p_L2sdqz = (const __nv_bfloat16*)L2sdqz.data_ptr();\n const unsigned char* p_Mqw = (const unsigned char*)Mqw.data_ptr();\n const __nv_bfloat16* p_Mqs = (const __nv_bfloat16*)Mqs.data_ptr();\n const __nv_bfloat16* p_Mqz = (const __nv_bfloat16*)Mqz.data_ptr();\n const unsigned char* p_Maw = (const unsigned char*)Maw.data_ptr();\n const __nv_bfloat16* p_Mas = (const __nv_bfloat16*)Mas.data_ptr();\n const __nv_bfloat16* p_Maz = (const __nv_bfloat16*)Maz.data_ptr();\n const unsigned char* p_Mbw = (const unsigned char*)Mbw.data_ptr();\n const __nv_bfloat16* p_Mbs = (const __nv_bfloat16*)Mbs.data_ptr();\n const __nv_bfloat16* p_Mbz = (const __nv_bfloat16*)Mbz.data_ptr();\n const unsigned char* p_Mow = (const unsigned char*)Mow.data_ptr();\n const __nv_bfloat16* p_Mos = (const __nv_bfloat16*)Mos.data_ptr();\n const __nv_bfloat16* p_Moz = (const __nv_bfloat16*)Moz.data_ptr();\n const __nv_bfloat16* p_Man = (const __nv_bfloat16*)Man.data_ptr();\n const __nv_bfloat16* p_Mmn = (const __nv_bfloat16*)Mmn.data_ptr();\n const __nv_bfloat16* p_Mrt = (const __nv_bfloat16*)Mrt.data_ptr();\n const unsigned char* p_Mgqw = (const unsigned char*)Mgqw.data_ptr();\n const __nv_bfloat16* p_Mgqs = (const __nv_bfloat16*)Mgqs.data_ptr();\n const __nv_bfloat16* p_Mgqz = (const __nv_bfloat16*)Mgqz.data_ptr();\n const unsigned char* p_Muqw = (const unsigned char*)Muqw.data_ptr();\n const __nv_bfloat16* p_Muqs = (const __nv_bfloat16*)Muqs.data_ptr();\n const __nv_bfloat16* p_Muqz = (const __nv_bfloat16*)Muqz.data_ptr();\n const unsigned char* p_Mdqw = (const unsigned char*)Mdqw.data_ptr();\n const __nv_bfloat16* p_Mdqs = (const __nv_bfloat16*)Mdqs.data_ptr();\n const __nv_bfloat16* p_Mdqz = (const __nv_bfloat16*)Mdqz.data_ptr();\n const unsigned char* p_Msgqw = (const unsigned char*)Msgqw.data_ptr();\n const __nv_bfloat16* p_Msgqs = (const __nv_bfloat16*)Msgqs.data_ptr();\n const __nv_bfloat16* p_Msgqz = (const __nv_bfloat16*)Msgqz.data_ptr();\n const unsigned char* p_Msuqw = (const unsigned char*)Msuqw.data_ptr();\n const __nv_bfloat16* p_Msuqs = (const __nv_bfloat16*)Msuqs.data_ptr();\n const __nv_bfloat16* p_Msuqz = (const __nv_bfloat16*)Msuqz.data_ptr();\n const unsigned char* p_Msdqw = (const unsigned char*)Msdqw.data_ptr();\n const __nv_bfloat16* p_Msdqs = (const __nv_bfloat16*)Msdqs.data_ptr();\n const __nv_bfloat16* p_Msdqz = (const __nv_bfloat16*)Msdqz.data_ptr();\n const __nv_bfloat16* p_oc = (const __nv_bfloat16*)oc.data_ptr();\n const __nv_bfloat16* p_ok = (const __nv_bfloat16*)ok.data_ptr();\n __nv_bfloat16* p_nc = (__nv_bfloat16*)nc.data_ptr();\n __nv_bfloat16* p_nk = (__nv_bfloat16*)nk.data_ptr();\n static bool stack_ok = false;\n if (!stack_ok) {\n cudaDeviceSetLimit(cudaLimitStackSize, 16384);\n stack_ok = true;\n }\n void* KA[] = {&p_h, &p_SBp, &p_SFp, &p_SIp, &p_L0qw, &p_L0qs, &p_L0qz, &p_L0kw, &p_L0ks, &p_L0kz, &p_L0vw, &p_L0vs, &p_L0vz, &p_L0gw, &p_L0gs, &p_L0gz, &p_L0ow, &p_L0os, &p_L0oz, &p_L0bw, &p_L0cw, &p_L0S, &p_L0cq, &p_L0ck, &p_L0cv, &p_L0an, &p_L0mn, &p_L0rt, &p_L0gqw, &p_L0gqs, &p_L0gqz, &p_L0uqw, &p_L0uqs, &p_L0uqz, &p_L0dqw, &p_L0dqs, &p_L0dqz, &p_L0sgqw, &p_L0sgqs, &p_L0sgqz, &p_L0suqw, &p_L0suqs, &p_L0suqz, &p_L0sdqw, &p_L0sdqs, &p_L0sdqz, &p_L1qw, &p_L1qs, &p_L1qz, &p_L1kw, &p_L1ks, &p_L1kz, &p_L1vw, &p_L1vs, &p_L1vz, &p_L1gw, &p_L1gs, &p_L1gz, &p_L1ow, &p_L1os, &p_L1oz, &p_L1bw, &p_L1cw, &p_L1S, &p_L1cq, &p_L1ck, &p_L1cv, &p_L1an, &p_L1mn, &p_L1rt, &p_L1gqw, &p_L1gqs, &p_L1gqz, &p_L1uqw, &p_L1uqs, &p_L1uqz, &p_L1dqw, &p_L1dqs, &p_L1dqz, &p_L1sgqw, &p_L1sgqs, &p_L1sgqz, &p_L1suqw, &p_L1suqs, &p_L1suqz, &p_L1sdqw, &p_L1sdqs, &p_L1sdqz, &p_L2qw, &p_L2qs, &p_L2qz, &p_L2kw, &p_L2ks, &p_L2kz, &p_L2vw, &p_L2vs, &p_L2vz, &p_L2gw, &p_L2gs, &p_L2gz, &p_L2ow, &p_L2os, &p_L2oz, &p_L2bw, &p_L2cw, &p_L2S, &p_L2cq, &p_L2ck, &p_L2cv, &p_L2an, &p_L2mn, &p_L2rt, &p_L2gqw, &p_L2gqs, &p_L2gqz, &p_L2uqw, &p_L2uqs, &p_L2uqz, &p_L2dqw, &p_L2dqs, &p_L2dqz, &p_L2sgqw, &p_L2sgqs, &p_L2sgqz, &p_L2suqw, &p_L2suqs, &p_L2suqz, &p_L2sdqw, &p_L2sdqs, &p_L2sdqz, &p_Mqw, &p_Mqs, &p_Mqz, &p_Maw, &p_Mas, &p_Maz, &p_Mbw, &p_Mbs, &p_Mbz, &p_Mow, &p_Mos, &p_Moz, &p_Man, &p_Mmn, &p_Mrt, &p_Mgqw, &p_Mgqs, &p_Mgqz, &p_Muqw, &p_Muqs, &p_Muqz, &p_Mdqw, &p_Mdqs, &p_Mdqz, &p_Msgqw, &p_Msgqs, &p_Msgqz, &p_Msuqw, &p_Msuqs, &p_Msuqz, &p_Msdqw, &p_Msdqs, &p_Msdqz, &p_oc, &p_ok, &p_nc, &p_nk, &curlen, &capx};\n cudaError_t e = cudaLaunchCooperativeKernel((void*)megakernel, dim3(256), dim3(256),\n KA, 0, c10::cuda::getCurrentCUDAStream().stream());\n if (e != cudaSuccess) throw std::runtime_error(cudaGetErrorString(e));\n}\n'