{"record":{"id":"eaeb5ad0b0afb38c","repo":"sgl-project/sglang","slug":"pad-masks-and-att-masks-must-be-batch-seq","errorCode":null,"errorMessage":"pad_masks and att_masks must be [batch, seq]","messagePattern":"pad_masks and att_masks must be \\[batch, seq\\]","errorType":"validation","errorClass":"ValueError","httpStatus":null,"severity":"error","filePath":"python/sglang/multimodal_gen/runtime/models/vlas/pi05_core.py","lineNumber":854,"sourceCode":"    fraction = torch.linspace(\n        0.0,\n        1.0,\n        dimension // 2,\n        dtype=torch.float64,\n        device=time.device,\n    )\n    period = min_period * (max_period / min_period) ** fraction\n    scaling = 1.0 / period * 2 * math.pi\n    sin_input = scaling[None, :] * time[:, None].to(torch.float64)\n    return torch.cat([torch.sin(sin_input), torch.cos(sin_input)], dim=1)\n\n\ndef make_att_2d_masks(\n    pad_masks: torch.Tensor,\n    att_masks: torch.Tensor,\n) -> torch.Tensor:\n    if att_masks.ndim != 2 or pad_masks.ndim != 2:\n        raise ValueError(\"pad_masks and att_masks must be [batch, seq]\")\n    cumsum = torch.cumsum(att_masks, dim=1)\n    att_2d_masks = cumsum[:, None, :] <= cumsum[:, :, None]\n    pad_2d_masks = pad_masks[:, None, :] * pad_masks[:, :, None]\n    return att_2d_masks & pad_2d_masks\n\n\ndef trim_trailing_padding_tokens(\n    tokens: torch.Tensor,\n    token_masks: torch.Tensor,\n) -> tuple[torch.Tensor, torch.Tensor]:\n    token_len = int(token_masks.sum(dim=1).max().item())\n    if token_len <= 0 or token_len >= tokens.shape[1]:\n        return tokens, token_masks\n    return tokens[:, :token_len], token_masks[:, :token_len]\n\n\ndef prepare_optional_full_attention_mask(\n    att_2d_masks: torch.Tensor,","sourceCodeStart":836,"sourceCodeEnd":872,"githubUrl":"https://github.com/sgl-project/sglang/blob/0132848349585cfe6aae51c4941cbae872505f8a/python/sglang/multimodal_gen/runtime/models/vlas/pi05_core.py#L836-L872","documentation":"make_att_2d_masks builds a 2-D [batch, seq, seq] attention mask from a padding mask and an attention mask via cumulative-sum logic that only works when both inputs are 2-D [batch, seq]. The guard rejects any other rank. It is used by encode_prefix and denoise_step, so the offending tensors usually come from the caller's prefix/action preparation code.","triggerScenarios":"Calling make_att_2d_masks with pad_masks or att_masks of rank != 2 — e.g. a 3-D mask [batch, heads, seq], a 1-D mask [seq], or a mask that still has a trailing singleton dim like [batch, seq, 1].","commonSituations":"Passing masks straight from a tokenizer that returns [batch, 1, seq, seq]; per-head attention masks from another backend; forgetting .squeeze(-1) after an unsqueeze elsewhere.","solutions":["Squeeze extra dims: pad_masks = pad_masks.squeeze(-1), att_masks = att_masks.squeeze(1)","Ensure masks are boolean/integer of shape [batch, seq] before calling encode_prefix/denoise_step","Add asserts pad_masks.ndim == 2 and att_masks.ndim == 2 in your mask-prep helper"],"exampleFix":"# before\natt_2d = make_att_2d_masks(pad_masks[:, :, 0], att_masks)  # 1-D pad mask\n\n# after\natt_2d = make_att_2d_masks(pad_masks, att_masks)  # both [batch, seq]","handlingStrategy":"validation","validationCode":"for name, m in ((\"pad_masks\", pad_masks), (\"att_masks\", att_masks)):\n    assert m.ndim == 2, f\"{name} must be [batch, seq], got {tuple(m.shape)}\"","typeGuard":"def is_2d_mask(m: torch.Tensor) -> bool:\n    return isinstance(m, torch.Tensor) and m.ndim == 2","tryCatchPattern":"try:\n    att2d = make_att_2d_masks(pad_masks, att_masks)\nexcept ValueError:\n    pad_masks = pad_masks.reshape(pad_masks.shape[0], -1)\n    att_masks = att_masks.reshape(att_masks.shape[0], -1)\n    att2d = make_att_2d_masks(pad_masks, att_masks)","preventionTips":["Normalize masks to [batch, seq] at the boundary of your pipeline; never pass 4-D tokenizer masks through","Document the expected mask rank next to every call site of encode_prefix/denoise_step"],"tags":["pi05","attention-mask","tensor-rank","input-validation"],"backgroundTag":"shape-validation-failed","analyzedSha":"0132848349585cfe6aae51c4941cbae872505f8a","analyzedAt":"2026-08-28T05:10:05.995Z","schemaVersion":2},"datasetVersion":"2026-08-28T06:17:29.519Z"}