File size: 3,943 Bytes
572ea6f
 
 
 
 
 
65b9081
572ea6f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
# /// script
# dependencies = [
#     "transformers>=5.14.0",
#     "datasets>=4.0",
#     "Pillow>=10.0",
#     "torch>=2.5",
#     "torchvision>=0.20",
#     "accelerate>=1.0",
#     "num2words",
# ]
# ///
"""CPU pre-flight for the Gemma 4 training collator.

Runs the EXACT collate path from train_sft.py on a few real dataset rows, using
the processor + model config (no 16GB weights), then decodes one example to
prove which tokens are supervised vs masked. Catches chat-template, token-id,
and shape bugs on cpu-basic (~$0.01) before we burn GPU minutes.

    python gemma4/validate_collator.py --dataset khalidFlex/gui-agent-smoke
"""

import argparse

import torch
from datasets import load_dataset
from transformers import AutoConfig, AutoProcessor


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--model", default="google/gemma-4-E4B-it")
    ap.add_argument("--dataset", default="khalidFlex/gui-agent-smoke")
    ap.add_argument("--n", type=int, default=4)
    args = ap.parse_args()

    print(f"[val] processor + config for {args.model}")
    processor = AutoProcessor.from_pretrained(args.model)
    config = AutoConfig.from_pretrained(args.model)

    image_token_ids = {
        tid for tid in [
            getattr(config, "image_token_id", None),
            getattr(config, "boi_token_id", None),
            getattr(config, "eoi_token_id", None),
        ] if isinstance(tid, int)
    }
    print(f"[val] image/structural token ids: {sorted(image_token_ids)}")
    processor.tokenizer.padding_side = "right"

    ds = load_dataset(args.dataset, split=f"train[:{args.n}]")
    print(f"[val] loaded {len(ds)} rows")

    def to_messages(ex):
        return [
            {"role": "user", "content": [
                {"type": "image"},
                {"type": "text", "text": f"{ex['system']}\n\n{ex['user']}"},
            ]},
            {"role": "assistant", "content": [{"type": "text", "text": ex["assistant"]}]},
        ]

    full_texts, prompt_texts, images = [], [], []
    for ex in ds:
        msgs = to_messages(ex)
        full_texts.append(processor.apply_chat_template(msgs, tokenize=False, add_generation_prompt=False))
        prompt_texts.append(processor.apply_chat_template(msgs[:-1], tokenize=False, add_generation_prompt=True))
        images.append([ex["image"].convert("RGB")])

    batch = processor(text=full_texts, images=images, return_tensors="pt", padding=True)
    prompt_batch = processor(text=prompt_texts, images=images, return_tensors="pt", padding=True)
    prompt_lens = prompt_batch["attention_mask"].sum(dim=1)

    labels = batch["input_ids"].clone()
    labels[batch["attention_mask"] == 0] = -100
    for tid in image_token_ids:
        labels[labels == tid] = -100
    for i, plen in enumerate(prompt_lens):
        labels[i, : int(plen)] = -100

    print("\n[val] === batch shapes ===")
    for k, v in batch.items():
        if isinstance(v, torch.Tensor):
            print(f"  {k}: {tuple(v.shape)} {v.dtype}")

    # Prove supervision: decode only the non-masked (label != -100) tokens of row 0.
    row0 = labels[0]
    supervised = batch["input_ids"][0][row0 != -100]
    n_sup = int((row0 != -100).sum())
    n_tot = int((batch["attention_mask"][0] == 1).sum())
    print(f"\n[val] row0 supervised tokens: {n_sup}/{n_tot} real tokens")
    print(f"[val] supervised text (should be ONLY the assistant answer):")
    print("   ", repr(processor.tokenizer.decode(supervised)))

    assert n_sup > 0, "no supervised tokens — masking is wrong!"
    assert n_sup < n_tot, "everything supervised — prompt not masked!"
    decoded = processor.tokenizer.decode(supervised)
    assert "<code>" in decoded or "click" in decoded or "final_answer" in decoded, \
        "supervised span doesn't look like an action — offset is off"
    print("\n[val] ✓ collator OK — prompt+image masked, assistant supervised.")


if __name__ == "__main__":
    main()