| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| """LoRA SFT of Gemma 4 for GUI grounding + web-agent action prediction. |
| |
| Trains on the flat dataset built by prep_data.py (image, system, user, |
| assistant, source). Each row becomes one multimodal chat example: |
| |
| system : the GUI-agent system prompt |
| user : [screenshot] + instruction + previous actions |
| assistant (label): <think>…</think> <code>action(...)</code> |
| |
| Only the assistant turn contributes to the loss (the prompt + image tokens are |
| masked to -100 in the collator). Adapters are pushed to the Hub. |
| |
| Launch via HF Jobs (see gemma4/launch_train.sh) — never run untethered; the |
| Jobs box is ephemeral so push_to_hub must be on. |
| """ |
|
|
| import argparse |
|
|
| import torch |
| import torch.nn as nn |
| from datasets import load_dataset |
| from peft import LoraConfig |
| from transformers import AutoModelForImageTextToText, AutoProcessor |
| from trl import SFTConfig, SFTTrainer |
|
|
| PROJ_SUFFIXES = ("q_proj", "k_proj", "v_proj", "o_proj", "gate_proj", "up_proj", "down_proj") |
|
|
|
|
| def find_lora_targets(model): |
| """FULL module names of the language-tower projection Linears. |
| |
| Gemma 4's vision tower wraps its projections in `Gemma4ClippableLinear` |
| (not `nn.Linear`), which PEFT can't adapt, and matching by bare suffix (e.g. |
| "q_proj") would also grab those and crash. So we walk the graph and return |
| the *fully-qualified* names of the real `nn.Linear` projections under |
| `language_model`. PEFT matches those exactly, leaving the vision encoder |
| frozen and skipping MoE expert Parameters (which aren't Modules). |
| """ |
| targets = [] |
| for name, module in model.named_modules(): |
| if not isinstance(module, nn.Linear): |
| continue |
| if "language_model" not in name: |
| continue |
| if name.split(".")[-1] in PROJ_SUFFIXES: |
| targets.append(name) |
| return targets |
|
|
|
|
| def build_collator(processor, model_config): |
| """Render our (system/user/assistant + image) rows into a batch. |
| |
| Gemma 4 uses an image-text-to-text processor: we pass the chat template |
| with an image placeholder in the user turn, tokenize, and mask padding plus |
| every image-structural token (the soft image token + <start_of_image> / |
| <end_of_image>) so loss is computed only over real text tokens. We read the |
| exact ids from the model config so this can't drift with the tokenizer. |
| """ |
| image_token_ids = { |
| tid for tid in [ |
| getattr(model_config, "image_token_id", None), |
| getattr(model_config, "boi_token_id", None), |
| getattr(model_config, "eoi_token_id", None), |
| getattr(processor, "image_token_id", None), |
| ] if isinstance(tid, int) |
| } |
| print(f"[train] masking image/pad token ids: {sorted(image_token_ids)}") |
| processor.tokenizer.padding_side = "right" |
|
|
| def to_messages(example): |
| |
| |
| |
| |
| |
| return [ |
| {"role": "user", "content": [ |
| {"type": "image"}, |
| {"type": "text", "text": f"{example['system']}\n\n{example['user']}"}, |
| ]}, |
| {"role": "assistant", "content": [{"type": "text", "text": example["assistant"]}]}, |
| ] |
|
|
| def collate(examples): |
| full_texts, prompt_texts, images = [], [], [] |
| for ex in examples: |
| 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 |
| batch["labels"] = labels |
| return batch |
|
|
| return collate |
|
|
|
|
| def main(): |
| ap = argparse.ArgumentParser() |
| ap.add_argument("--model", default="google/gemma-4-E4B-it") |
| ap.add_argument("--dataset", required=True) |
| ap.add_argument("--hub-model-id", required=True) |
| ap.add_argument("--epochs", type=float, default=1.0) |
| ap.add_argument("--batch-size", type=int, default=2) |
| ap.add_argument("--grad-accum", type=int, default=8) |
| ap.add_argument("--lr", type=float, default=2e-4) |
| ap.add_argument("--max-steps", type=int, default=-1) |
| ap.add_argument("--eval-frac", type=float, default=0.03) |
| ap.add_argument("--eval-steps", type=int, default=50) |
| ap.add_argument("--save-steps", type=int, default=100) |
| ap.add_argument("--project", default="gemma4-gui-agent") |
| ap.add_argument("--run-name", default="gemma4-e4b-lora") |
| ap.add_argument("--private", action="store_true", default=True) |
| args = ap.parse_args() |
|
|
| print(f"[train] loading dataset {args.dataset}") |
| ds = load_dataset(args.dataset, split="train") |
| split = ds.train_test_split(test_size=args.eval_frac, seed=42) |
| train_ds, eval_ds = split["train"], split["test"] |
| print(f"[train] {len(train_ds)} train / {len(eval_ds)} eval") |
|
|
| print(f"[train] loading {args.model}") |
| processor = AutoProcessor.from_pretrained(args.model) |
| model = AutoModelForImageTextToText.from_pretrained( |
| args.model, dtype=torch.bfloat16, attn_implementation="eager", |
| ) |
|
|
| |
| |
| targets = find_lora_targets(model) |
| print(f"[train] LoRA targets: {len(targets)} language-model Linear projections " |
| f"(e.g. {targets[:2]})") |
| if not targets: |
| raise RuntimeError("no language_model projection Linears found for LoRA") |
| peft_config = LoraConfig( |
| r=16, lora_alpha=32, lora_dropout=0.05, bias="none", |
| task_type="CAUSAL_LM", |
| target_modules=targets, |
| modules_to_save=None, |
| ) |
|
|
| sft_config = SFTConfig( |
| output_dir=args.hub_model_id.split("/")[-1], |
| per_device_train_batch_size=args.batch_size, |
| per_device_eval_batch_size=args.batch_size, |
| gradient_accumulation_steps=args.grad_accum, |
| num_train_epochs=args.epochs, |
| max_steps=args.max_steps, |
| learning_rate=args.lr, |
| warmup_ratio=0.03, |
| lr_scheduler_type="cosine", |
| logging_steps=5, |
| eval_strategy="steps", |
| eval_steps=args.eval_steps, |
| save_strategy="steps", |
| save_steps=args.save_steps, |
| save_total_limit=2, |
| bf16=True, |
| gradient_checkpointing=True, |
| gradient_checkpointing_kwargs={"use_reentrant": False}, |
| dataset_kwargs={"skip_prepare_dataset": True}, |
| remove_unused_columns=False, |
| max_length=None, |
| push_to_hub=True, |
| hub_model_id=args.hub_model_id, |
| hub_private_repo=args.private, |
| hub_strategy="every_save", |
| report_to="trackio", |
| run_name=args.run_name, |
| project=args.project, |
| |
| |
| |
| |
| trackio_static_space_id=False, |
| ) |
|
|
| trainer = SFTTrainer( |
| model=model, |
| args=sft_config, |
| train_dataset=train_ds, |
| eval_dataset=eval_ds, |
| data_collator=build_collator(processor, model.config), |
| peft_config=peft_config, |
| processing_class=processor, |
| ) |
|
|
| trainer.train() |
| trainer.save_model(sft_config.output_dir) |
| trainer.push_to_hub() |
| print(f"[train] pushed adapters to {args.hub_model_id}") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|