Bug Report: MAJOR issues with parallel generation and left-padding (different response from normal generatiln)

#50
by atbolsh - opened

Title: Gemma 4 Unified: left-padding a multimodal prompt by a specific
amount (pad=7 here) corrupts prefill logits — argmax flips to <audio|>;
neighboring pad lengths are fine

System Info

  • transformers version: 5.14.1 (output below is from 5.13.1; the two
    versions produce bit-identical results for this script, corruption
    included)
  • Platform: Linux-6.17.0-35-generic-x86_64-with-glibc2.39
  • Python version: 3.12.13
  • Huggingface_hub version: 1.18.0
  • Safetensors version: 0.8.0
  • Accelerate version: 1.14.0
  • Accelerate config: not found
  • DeepSpeed version: not installed
  • PyTorch version (accelerator?): 2.12.0+cu130 (CUDA)
  • Using distributed or parallel set-up in script?: No
  • Using GPU in script?: Yes (single GPU, device_map="auto")
  • GPU type: NVIDIA RTX PRO 6000 Blackwell Workstation Edition

Who can help?

@zucchini-nlp

Reproduction

With Gemma4UnifiedForConditionalGeneration (google/gemma-4-12B-it, bf16,
SDPA), left-padding a multimodal prompt (282 tokens, one image) by exactly
7 positions corrupts the prefill logits: max next-token logit delta vs the
unpadded forward jumps from ~0.5 (bf16 wobble) to 40, and the argmax
flips from 'The'/'I' to the audio special token <audio|> — on an
image prompt. Every other pad length in the sweep (1–6, 8, 9, 15, 16, 63,
64) stays within wobble, so this is not gradual numerical noise; it is a
discrete failure at a specific offset:

B. padded, batch=1, pad=6    max|dLogit|=  0.5000 mean=0.13692 argmax_flipped=False
B. padded, batch=1, pad=7    max|dLogit|= 40.1875 mean=11.57746 argmax_flipped=True
    solo top5: 'I':27.25  'The':27.25  'You':22.62  'There':21.12  'In':21.00
    this top5: '<audio|>':25.25  'a':11.50  '-':9.69  'I':8.94  'T':8.81
B. padded, batch=1, pad=8    max|dLogit|=  0.5312 mean=0.12108 argmax_flipped=False

Notes:

  • Batch size is irrelevant: batch=1 with pad=7 corrupts identically to a
    batch of 2 where the shorter row is padded by 7 (the standard left-pad
    collation for variable-length batched generate). The unpadded row of
    that batch stays clean (max delta 0.55).
  • Not a collation artifact: every sequence-aligned tensor (input_ids,
    attention_mask, token type ids) is padded in lockstep, the padded row's
    suffix is asserted byte-identical to its solo encoding, and
    position_ids follow the attention mask.
  • Text-only prompts are unaffected at every pad length tried (max delta
    ~0.5–0.66); the image input is required.
  • Reproduces bit-identically on transformers 5.13.1 and 5.14.1.
  • Ambiguity we did not resolve: for this prompt, pad=7 is also the pad that
    makes the total padded length 289 — we have not varied the prompt to
    distinguish "pad length 7" from "total length 289" as the trigger. The
    script below makes that a two-minute experiment.
  • User-visible symptom: decoding strips the special token, so greedy
    generate on a padded-by-7 row silently returns degraded caption-style
    text instead of an instruction-following reply — it looks like sampling
    variance, not a broken forward:
greedy solo   : 'The image contains red, yellow, green, and blue squares on a light gray background.'
greedy batched: 'a yellow rectangle, a red rectangle, a green rectangle, and a blue rectangle.'

Standalone script (no external assets; draws its own image, sweeps pad
lengths at batch=1, then runs the batch-of-2 and greedy comparisons; exits
1 when the bug fires):

from __future__ import annotations

import argparse
import sys
import tempfile
from pathlib import Path

import torch
from PIL import Image, ImageDraw
from transformers import AutoModelForMultimodalLM, AutoProcessor

SHORT_Q = "In one short sentence, what colors do you see?"
LONG_Q = ("Answer briefly: is the grid mostly empty? Explain in one "
          "sentence why you think so.")


def make_image(path: Path, size: int = 96) -> Path:
    """Deterministic grid of colored squares on light gray."""
    img = Image.new("RGB", (size, size), (235, 235, 235))
    draw = ImageDraw.Draw(img)
    cell = size // 6
    colors = [(200, 40, 40), (40, 160, 40), (220, 180, 30), (60, 60, 200)]
    for i, (gx, gy) in enumerate([(0, 1), (2, 3), (3, 0), (4, 4), (1, 5),
                                  (5, 2), (2, 0), (0, 4)]):
        draw.rectangle(
            [gx * cell, gy * cell, (gx + 1) * cell - 1, (gy + 1) * cell - 1],
            fill=colors[i % 4],
        )
    img.save(path)
    return path


def encode(processor, image_path: Path, text: str) -> dict:
    messages = [{"role": "user", "content": [
        {"type": "image", "url": str(image_path)},
        {"type": "text", "text": text},
    ]}]
    return processor.apply_chat_template(
        messages, tokenize=True, add_generation_prompt=True,
        return_dict=True, return_tensors="pt",
    )


def left_pad(enc: dict, pad_len: int, pad_token_id: int) -> dict:
    """Left-pad EVERY sequence-aligned tensor in lockstep with input_ids."""
    seq_len = enc["input_ids"].shape[1]
    out = {}
    for k, v in enc.items():
        if not isinstance(v, torch.Tensor):
            continue
        if (not v.dtype.is_floating_point and v.dim() == 2
                and v.shape[1] == seq_len):
            fill = pad_token_id if k == "input_ids" else 0
            out[k] = torch.cat([v.new_full((1, pad_len), fill), v], dim=1)
        else:
            out[k] = v
    # The padded row's suffix must be byte-identical to the solo encoding.
    assert torch.equal(out["input_ids"][0, pad_len:], enc["input_ids"][0])
    assert out["attention_mask"][0, :pad_len].sum() == 0
    return out


def last_logits(model, enc: dict) -> torch.Tensor:
    device = next(model.parameters()).device
    dtype = next(model.parameters()).dtype
    inputs = {}
    for k, v in enc.items():
        if isinstance(v, torch.Tensor):
            v = v.to(device)
            inputs[k] = v.to(dtype) if v.dtype.is_floating_point else v
    # Same position_ids convention generate() uses for left-padded prefill.
    mask = inputs["attention_mask"]
    inputs["position_ids"] = (mask.long().cumsum(-1) - 1).clamp(min=0)
    with torch.inference_mode():
        return model(**inputs).logits[:, -1].float().cpu()


def top5(tok, logits: torch.Tensor) -> str:
    vals, ids = logits.topk(5)
    return "  ".join(f"{tok.decode([int(i)])!r}:{float(v):.2f}"
                     for v, i in zip(vals, ids))


def report(tok, label: str, solo: torch.Tensor, other: torch.Tensor) -> bool:
    delta = (other - solo).abs()
    flipped = int(other.argmax()) != int(solo.argmax())
    print(f"{label:<28s} max|dLogit|={float(delta.max()):8.4f} "
          f"mean={float(delta.mean()):.5f} argmax_flipped={flipped}")
    if flipped:
        print(f"    solo top5: {top5(tok, solo)}")
        print(f"    this top5: {top5(tok, other)}")
    return flipped


def main() -> int:
    ap = argparse.ArgumentParser()
    ap.add_argument("--model", default="google/gemma-4-12B-it")
    args = ap.parse_args()

    import transformers
    print(f"transformers=={transformers.__version__}  "
          f"torch=={torch.__version__}")

    processor = AutoProcessor.from_pretrained(args.model, padding_side="left")
    model = AutoModelForMultimodalLM.from_pretrained(
        args.model, dtype="auto", attn_implementation="sdpa",
        device_map="auto",
    ).eval()
    tok = getattr(processor, "tokenizer", processor)
    pad_id = tok.pad_token_id

    with tempfile.TemporaryDirectory() as tmp:
        img = make_image(Path(tmp) / "grid.png")
        enc_s = encode(processor, img, SHORT_Q)
        enc_l = encode(processor, img, LONG_Q)
    len_s, len_l = enc_s["input_ids"].shape[1], enc_l["input_ids"].shape[1]
    assert len_l > len_s, (len_s, len_l)
    print(f"prompt lengths: short={len_s} long={len_l} tokens")

    solo_s = last_logits(model, enc_s)[0]
    solo_l = last_logits(model, enc_l)[0]
    padded_s = left_pad(enc_s, len_l - len_s, pad_id)

    print(f"\nsolo short top5: {top5(tok, solo_s)}")
    # B: same row, alone in the batch, over a sweep of pad lengths --
    # the corruption tracks the pad length (alignment), not batch size.
    corrupted = False
    for pad_len in (1, 2, 3, 4, 5, 6, 7, 8, 9, 15, 16, 63, 64):
        corrupted |= report(
            tok, f"B. padded, batch=1, pad={pad_len}",
            solo_s, last_logits(model, left_pad(enc_s, pad_len, pad_id))[0],
        )

    # C: same padded row, batch of 2 with the (unpadded) long prompt.
    stacked = {k: torch.cat([padded_s[k], enc_l[k]], dim=0)
               for k in padded_s if isinstance(padded_s[k], torch.Tensor)}
    both = last_logits(model, stacked)
    corrupted |= report(tok, "C. padded, batch=2", solo_s, both[0])
    report(tok, "   unpadded control row", solo_l, both[1])

    # User-visible symptom: greedy generate, solo vs the same batch of 2.
    gen_kwargs = dict(max_new_tokens=40, do_sample=False)
    device = next(model.parameters()).device

    def to_dev(d):
        return {k: (v.to(device) if isinstance(v, torch.Tensor) else v)
                for k, v in d.items()}

    with torch.inference_mode():
        g_solo = model.generate(**to_dev(enc_s), **gen_kwargs)
        g_batch = model.generate(**to_dev(stacked), **gen_kwargs)
    solo_text = processor.decode(g_solo[0][len_s:], skip_special_tokens=True)
    batch_text = processor.decode(g_batch[0][len_l:], skip_special_tokens=True)
    print(f"\ngreedy solo   : {solo_text.strip()!r}")
    print(f"greedy batched: {batch_text.strip()!r}")

    print("\nExpected: every condition within max|dLogit| <~ 1 of solo "
          "(bf16 wobble). BUG: a padded condition shows max|dLogit| in "
          "the tens with the argmax flipped to a modality special token, "
          "and the batched greedy text diverges from solo.")
    return 1 if corrupted else 0


if __name__ == "__main__":
    sys.exit(main())

Full output (transformers 5.13.1; 5.14.1 output is bit-identical):

transformers==5.13.1  torch==2.12.0+cu130
prompt lengths: short=282 long=289 tokens

solo short top5: 'I':27.25  'The':27.25  'You':22.62  'There':21.12  'In':21.00
B. padded, batch=1, pad=1    max|dLogit|=  0.5156 mean=0.13077 argmax_flipped=False
B. padded, batch=1, pad=2    max|dLogit|=  0.5625 mean=0.12874 argmax_flipped=False
B. padded, batch=1, pad=3    max|dLogit|=  0.5820 mean=0.14670 argmax_flipped=False
B. padded, batch=1, pad=4    max|dLogit|=  0.5469 mean=0.12501 argmax_flipped=False
B. padded, batch=1, pad=5    max|dLogit|=  0.4453 mean=0.10723 argmax_flipped=False
B. padded, batch=1, pad=6    max|dLogit|=  0.5000 mean=0.13692 argmax_flipped=False
B. padded, batch=1, pad=7    max|dLogit|= 40.1875 mean=11.57746 argmax_flipped=True
    solo top5: 'I':27.25  'The':27.25  'You':22.62  'There':21.12  'In':21.00
    this top5: '<audio|>':25.25  'a':11.50  '-':9.69  'I':8.94  'T':8.81
B. padded, batch=1, pad=8    max|dLogit|=  0.5312 mean=0.12108 argmax_flipped=False
B. padded, batch=1, pad=9    max|dLogit|=  0.6875 mean=0.07407 argmax_flipped=False
B. padded, batch=1, pad=15   max|dLogit|=  0.7031 mean=0.12159 argmax_flipped=False
B. padded, batch=1, pad=16   max|dLogit|=  0.4648 mean=0.08500 argmax_flipped=False
B. padded, batch=1, pad=63   max|dLogit|=  0.5000 mean=0.08344 argmax_flipped=False
B. padded, batch=1, pad=64   max|dLogit|=  0.4062 mean=0.07283 argmax_flipped=False
C. padded, batch=2           max|dLogit|= 40.1875 mean=11.60843 argmax_flipped=True
    solo top5: 'I':27.25  'The':27.25  'You':22.62  'There':21.12  'In':21.00
    this top5: '<audio|>':25.25  'a':11.25  '-':9.38  'I':9.00  'T':8.88
   unpadded control row      max|dLogit|=  0.5469 mean=0.16452 argmax_flipped=False

greedy solo   : 'The image contains red, yellow, green, and blue squares on a light gray background.'
greedy batched: 'a yellow rectangle, a red rectangle, a green rectangle, and a blue rectangle.'

Expected behavior

A left-padded multimodal row should produce (numerically close to) the same
logits as the identical row unpadded, at every pad length — as it already
does at 12 of the 13 pad lengths tried and for text-only inputs at all of
them. A discrete corruption at one specific offset, flipping the argmax to
an audio token on an image prompt, suggests an alignment/stride
assumption in the padded multimodal path (vision-block attention mask or
image-feature placement). Practical impact: standard left-padded batched
generate over variable-length multimodal prompts silently degrades
whichever rows happen to land on a bad pad offset.

UPDATE: I've narrowed down the problem. It only occurs when the total padded length == 1 mod 32. So "7" is not the magic number, "1 mod 32" is the magic condition.

Google org

Hi @atbolsh ,

Thanks for reporting the issue and providing the details. We are able to reproduce the issue and have escalated this issue to our internal team for further investigation.

@thnamratha it turned out to be an upstream, pytorch problem. They have patched it, but the bug fix will not land until the next version release: https://github.com/pytorch/pytorch/commit/4ebacbf99bec9418e05fcf7d915dadcc367c43e5

Until then, I would recommend adding adaptive left-padding as a work-around, much like what I did in this file, lines 287-347: https://github.com/atbolsh/game_saddle/blob/master/agent/model.py

I am told that transformers does not plan on patching this, since it's already handled upstream on pytorch. If you are running into this issue as a user, I recommend either pulling the pytorch nightly (as of this writing the fix is not in any major / minor / patch versions), or using adaptive padding as I am.

Sign up or log in to comment