{"record":{"id":"d8ee28981e9f5952","repo":"sgl-project/sglang","slug":"mxfp8-fused-prologue-requires-interleaved-k-v-scal","errorCode":null,"errorMessage":"MXFP8 fused prologue requires interleaved K/V scale buffers with shape {sf_shape}, got {tuple(sfk.shape)} and {tuple(sfv.shape)}.","messagePattern":"MXFP8 fused prologue requires interleaved K/V scale buffers with shape (.+?), got (.+?) and (.+?)\\.","errorType":"exception","errorClass":"ValueError","httpStatus":null,"severity":"error","filePath":"python/sglang/kernels/ops/attention/inkling_attn_prologue.py","lineNumber":97,"sourceCode":"    do_store: bool = True,\n    mxfp8_quant: bool = False,\n    sfk: torch.Tensor | None = None,\n    sfv: torch.Tensor | None = None,\n    page_size: int = 128,\n    log_scaling_tau: torch.Tensor | None = None,\n) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor | None]:\n    \"\"\"Returns fresh contiguous (q_normed, k_normed, v_conv) [T, dq/dkv];\n    KV rows are also scattered into k_buf/v_buf at ``loc`` (the attention call\n    should pass save_kv_cache=False).\"\"\"\n    t = qkvr.shape[0]\n    if mxfp8_quant:\n        if dq % 128 != 0 or dkv % 128 != 0:\n            raise ValueError(\"MXFP8 fused prologue requires head_dim-aligned Q/K/V.\")\n        if sfk is None or sfv is None:\n            raise ValueError(\"MXFP8 fused prologue requires K/V scale buffers.\")\n        sf_shape = (k_buf.shape[0] // page_size, dkv // 128, 32, page_size // 32, 4)\n        if sfk.shape != sf_shape or sfv.shape != sf_shape:\n            raise ValueError(\n                \"MXFP8 fused prologue requires interleaved K/V scale buffers \"\n                f\"with shape {sf_shape}, got {tuple(sfk.shape)} and {tuple(sfv.shape)}.\"\n            )\n        if not sfk.is_contiguous() or not sfv.is_contiguous():\n            raise ValueError(\n                \"MXFP8 fused prologue requires contiguous interleaved SFK/SFV.\"\n            )\n        q_out = torch.empty(t, dq, dtype=torch.float8_e4m3fn, device=qkvr.device)\n        sfq_u8 = torch.empty(\n            (t, dq // 128, 128 // 32), dtype=torch.uint8, device=qkvr.device\n        )\n        sfk_u8 = sfk.view(torch.uint8)\n        sfv_u8 = sfv.view(torch.uint8)\n    else:\n        q_out = torch.empty(t, dq, dtype=qkvr.dtype, device=qkvr.device)\n        sfq_u8 = torch.empty(0, dtype=torch.uint8, device=qkvr.device)\n        sfk_u8 = torch.empty(0, dtype=torch.uint8, device=qkvr.device)\n        sfv_u8 = torch.empty(0, dtype=torch.uint8, device=qkvr.device)","sourceCodeStart":79,"sourceCodeEnd":115,"githubUrl":"https://github.com/sgl-project/sglang/blob/0132848349585cfe6aae51c4941cbae872505f8a/python/sglang/kernels/ops/attention/inkling_attn_prologue.py#L79-L115","documentation":"MXFP8 scale buffers must follow the interleaved layout the Triton quantization kernel writes: sf_shape = (k_buf.shape[0]//page_size, dkv//128, 32, page_size//32, 4) — i.e. [pages, scale-blocks-per-page, 32, rows-per-block-of-32, 4 uint8 sub-scales]. The check compares both sfk.shape and sfv.shape against this exact tuple; any deviation (wrong page_size assumption, wrong dkv//128 count, or a flat scale buffer) is rejected with the expected vs got shapes in the message.","triggerScenarios":"Calling inkling_attn_prologue_verify with mxfp8_quant=True where sfk/sfv were allocated with a different page_size, a different head grouping, or a legacy non-interleaved MXFP8 layout.","commonSituations":"Changing page_size in server args without reallocating scale buffers; version upgrades that changed the interleaved scale layout; reusing scale buffers allocated for dkv=128 on a dkv=256 layer.","solutions":["Reallocate sfk/sfv exactly as torch.empty(k_buf.shape[0]//page_size, dkv//128, 32, page_size//32, 4, dtype=torch.uint8, device=...)","Confirm the page_size used to allocate matches the one passed to the prologue call","After upgrading sglang, re-derive the shape from the current formula instead of hardcoding it"],"exampleFix":"# before\nsfk = torch.empty(num_pages, dkv//128, page_size, dtype=torch.uint8, device='cuda')\n# after\nsfk = torch.empty(k_buf.shape[0]//page_size, dkv//128, 32, page_size//32, 4, dtype=torch.uint8, device='cuda')","handlingStrategy":"validation","validationCode":"sf_shape = (k_buf.shape[0] // page_size, dkv // 128, 32, page_size // 32, 4)\nassert sfk.shape == sf_shape and sfv.shape == sf_shape, (sfk.shape, sfv.shape, sf_shape)","typeGuard":null,"tryCatchPattern":null,"preventionTips":["Compute sf_shape from k_buf/page_size at allocation time, never hardcode","Reallocate scale buffers whenever page_size or cache capacity changes"],"tags":["mxfp8","scale-buffers","tensor-shape","kv-cache","inkling"],"backgroundTag":"tensor-shape-mismatch","analyzedSha":"0132848349585cfe6aae51c4941cbae872505f8a","analyzedAt":"2026-08-28T05:10:05.995Z","schemaVersion":2},"datasetVersion":"2026-08-28T06:17:29.519Z"}