{"record":{"id":"c090e6c7c54ffde9","repo":"sgl-project/sglang","slug":"timestep-must-have-shape-b-s-9-d","errorCode":null,"errorMessage":"timestep must have shape [B, S, 9 * D]","messagePattern":"timestep must have shape \\[B, S, 9 \\* D\\]","errorType":"validation","errorClass":"ValueError","httpStatus":null,"severity":"error","filePath":"python/sglang/kernels/ops/diffusion/modulate/ltx2_ada_values_triton.py","lineNumber":143,"sourceCode":"    ).to(tl.bfloat16)\n\n    tl.store(out0_ptr + base, (table0 + temb0).to(tl.bfloat16), mask=mask)\n    tl.store(out1_ptr + base, (table1 + temb1).to(tl.bfloat16), mask=mask)\n    tl.store(out2_ptr + base, (table2 + temb2).to(tl.bfloat16), mask=mask)\n    tl.store(out3_ptr + base, (table3 + temb3).to(tl.bfloat16), mask=mask)\n    tl.store(out4_ptr + base, (table4 + temb4).to(tl.bfloat16), mask=mask)\n    tl.store(out5_ptr + base, (table5 + temb5).to(tl.bfloat16), mask=mask)\n    tl.store(out6_ptr + base, (table6 + temb6).to(tl.bfloat16), mask=mask)\n    tl.store(out7_ptr + base, (table7 + temb7).to(tl.bfloat16), mask=mask)\n    tl.store(out8_ptr + base, (table8 + temb8).to(tl.bfloat16), mask=mask)\n\n\ndef ltx2_ada_values9(\n    scale_shift_table: torch.Tensor,\n    timestep: torch.Tensor,\n) -> tuple[torch.Tensor, ...]:\n    if timestep.ndim != 3:\n        raise ValueError(\"timestep must have shape [B, S, 9 * D]\")\n    if not timestep.is_cuda or timestep.dtype != torch.bfloat16:\n        raise ValueError(\"timestep must be a CUDA bfloat16 tensor\")\n    if not timestep.is_contiguous():\n        raise ValueError(\"timestep must be contiguous\")\n    if scale_shift_table.ndim != 2 or scale_shift_table.shape[0] != 9:\n        raise ValueError(\"scale_shift_table must have shape [9, D]\")\n    if (\n        not scale_shift_table.is_cuda\n        or scale_shift_table.dtype not in (torch.bfloat16, torch.float32)\n        or scale_shift_table.stride(-1) != 1\n    ):\n        raise ValueError(\n            \"scale_shift_table must be CUDA, bf16/fp32, last-dim contiguous\"\n        )\n\n    total_params = int(scale_shift_table.shape[0])\n    hidden = int(scale_shift_table.shape[1])\n    if hidden <= 0 or timestep.shape[-1] != total_params * hidden:","sourceCodeStart":125,"sourceCodeEnd":161,"githubUrl":"https://github.com/sgl-project/sglang/blob/0132848349585cfe6aae51c4941cbae872505f8a/python/sglang/kernels/ops/diffusion/modulate/ltx2_ada_values_triton.py#L125-L161","documentation":"ltx2_ada_values9 computes the nine LTX2 AdaLN modulation values from a [B, S, 9*D] timestep embedding. The first validation requires a rank-3 tensor; anything else (e.g. 2D or 4D) is rejected with this ValueError.","triggerScenarios":"Passing timestep as 2D [B, 9*D] (single token, no sequence dim) or 4D [B, S, 1, 9*D]; shapes that were not expanded to include the sequence dimension before the call.","commonSituations":"Text-to-video pipelines where a global embedding is broadcast over S tokens — forgetting timestep[:, None, :] expansion; feeding raw scheduler timesteps instead of the projected embedding tensor.","solutions":["Reshape to 3D: timestep = timestep.view(B, S, 9*D) or timestep[:, None, :] when S=1","Verify you're passing the sinusoidal-timestep projection output, not raw scalar timesteps","Add an assert timestep.ndim == 3 before calling"],"exampleFix":"# before\nvals = ltx2_ada_values9(table, emb_2d)  # [B, 9*D]\n# after\nemb = emb_2d[:, None, :].expand(B, S, 9*D).contiguous()\nvals = ltx2_ada_values9(table, emb)","handlingStrategy":"validation","validationCode":"assert timestep.ndim == 3, timestep.shape","typeGuard":"def is_3d(t: torch.Tensor) -> bool:\n    return t.dim() == 3","tryCatchPattern":null,"preventionTips":["Expand [B, 9D] embeddings to [B, S, 9D] explicitly","Pass the projected timestep embedding, not raw timesteps"],"tags":["shape","ltx2","adaln","diffusion"],"backgroundTag":"tensor-shape-mismatch","analyzedSha":"0132848349585cfe6aae51c4941cbae872505f8a","analyzedAt":"2026-08-28T05:10:05.995Z","schemaVersion":2},"datasetVersion":"2026-08-28T06:17:29.519Z"}