# /// script # dependencies = [ # "trl>=0.24.0", # "peft>=0.17.0", # "transformers>=5.14.0", # "datasets>=4.0", # "trackio", # "Pillow>=10.0", # "accelerate>=1.0", # "torch>=2.5", # "torchvision>=0.20", # "num2words", # ] # /// """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): action(...) 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 + / ) 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" # so the prompt is a clean prefix def to_messages(example): # IMPORTANT: match inference (gemma4/server.py via mlx_vlm) EXACTLY — # one user turn holding [image] + "SYSTEM\n\nuser", then the assistant. # Gemma's chat template has no separate system role, and mlx_vlm wraps a # raw prompt as a single user turn, so we mirror that here to keep the # train and inference token streams identical. 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 = everything up to (and including) the assistant header, so # its token length marks where the completion begins. 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) # Second pass with the SAME images so the measured prompt length includes # the expanded soft-image tokens — the completion offset is then exact. 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 # padding for tid in image_token_ids: labels[labels == tid] = -100 # image structural tokens for i, plen in enumerate(prompt_lens): labels[i, : int(plen)] = -100 # prompt (assistant-only loss) 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", ) # Adapt attention + MLP projections across the language tower; leave the # vision encoder frozen (grounding signal comes from the LM reading tokens). 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}, # we collate raw rows remove_unused_columns=False, max_length=None, # never truncate image tokens 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, # A private model repo forces Trackio's per-push "static" space private # too, which it rejects (browser-only snapshots must be public) and that # crashes the Hub push. Disable the static snapshot; the live Trackio # dashboard during training still works. 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()