"""Step 0 - assert the sliding mask band structure, and record which path mask creation actually takes under the run config (attn_implementation=sdpa). Correctness gate: transformers can SILENTLY skip mask creation and pass attention_mask=None, which would make the 25 sliding layers do full causal attention - a different model from the one vLLM serves. This converts "probably fine because we are slow" into a measurement. CPU only. No weights. No GPU. """ import torch from transformers import AutoConfig from transformers.masking_utils import ( create_causal_mask, create_sliding_window_causal_mask, ) MODEL = "/tank/aimodels/gemma4-26b-a4b-it-heretic-bf16" N = 16384 W = 1024 PAD = " " + " " * 20 cfg = AutoConfig.from_pretrained(MODEL) text = cfg.get_text_config() text._attn_implementation = "sdpa" print("sliding_window %s" % text.sliding_window) print("layers %d (%d sliding / %d full)" % ( len(text.layer_types), text.layer_types.count("sliding_attention"), text.layer_types.count("full_attention"))) print("_attn_implementation %s" % text._attn_implementation) print() def build(attn_2d, label): batch = attn_2d.shape[0] if attn_2d is not None else 1 embeds = torch.zeros(batch, N, 8, dtype=torch.bfloat16) pos = torch.arange(N).unsqueeze(0) kw = dict(config=text, inputs_embeds=embeds, attention_mask=attn_2d, past_key_values=None, position_ids=pos) full = create_causal_mask(**kw) slide = create_sliding_window_causal_mask(**kw) print("--- %s ---" % label) for name, m in (("full_attention", full), ("sliding_attention", slide)): if m is None: print(" %-20s None -> flash / is_causal path AVAILABLE" % name) continue print(" %-20s tensor shape=%s dtype=%s" % (name, tuple(m.shape), m.dtype)) allowed = m if m.dtype == torch.bool else (m == 0) per_row = allowed[0, 0].sum(-1) print("%sallowed/row min=%d max=%d mean=%.1f" % ( PAD, per_row.min().item(), per_row.max().item(), per_row.float().mean().item())) if name == "sliding_attention": ok = per_row.max().item() <= W print("%sBAND <= %d ? %s" % (PAD, W, "PASS" if ok else "FAIL")) sat = (per_row >= W).nonzero() if sat.numel(): print("%ssaturates at row %d" % (PAD, sat[0].item())) else: print("%slast row allows %d of %d (%s)" % ( PAD, per_row[-1].item(), N, "causal-full OK" if per_row[-1].item() == N else "UNEXPECTED")) print() # 1. no 2D mask at all - the "constraints silently dropped" scenario build(None, "attention_mask=None (no padding info)") # 2. all-ones 2D mask - equal-length batch, no padding build(torch.ones(2, N, dtype=torch.long), "all-ones 2D (no padding)") # 3. REAL right-padded batch - what collate_mixed actually produces real = torch.ones(2, N, dtype=torch.long) real[1, 6000:] = 0 build(real, "right-padded 2D (what collate_mixed emits)")