#!/usr/bin/env python """SFT of FunctionGemma-270M with LoRA on victor/functiongemma-agent-sft. The dataset is single pre-formatted FunctionGemma-native tool-calling text. We pre-tokenize into input_ids/attention_mask/labels where labels=-100 on every token EXCEPT the model (assistant) turns, so the loss only teaches the model to produce correct function calls / answers, not to memorize the tool definitions or user prompts. Usage: python train.py [--smoke] [--max_steps N] [--max_length N] [--epochs N] [--batch N] [--grad_accum N] [--lr F] [--gc] [--output REPO_ID] """ import argparse import os import re import torch from datasets import load_dataset from huggingface_hub import HfApi, create_repo from peft import LoraConfig from transformers import AutoModelForCausalLM, AutoTokenizer from trl import SFTConfig, SFTTrainer BASE = "unsloth/functiongemma-270m-it" def model_spans(text): """Character spans of the model ... turns.""" spans = [] for m in re.finditer(r"model\n", text): start = m.end() em = re.search(r"\n", text[start:]) if em: spans.append((start, start + em.end())) return spans def tokenize_row(row, tokenizer, max_length): text = row["text"] enc = tokenizer( text, return_offsets_mapping=True, truncation=True, max_length=max_length, ) spans = model_spans(text) labels = [] for (s, e), tid in zip(enc["offset_mapping"], enc["input_ids"]): keep = any(a <= e and s <= b for (a, b) in spans) labels.append(tid if keep else -100) return { "input_ids": enc["input_ids"], "attention_mask": enc["attention_mask"], "labels": labels, } def main(): ap = argparse.ArgumentParser() ap.add_argument("--smoke", action="store_true", help="tiny subset, no push") ap.add_argument("--max_steps", type=int, default=-1) ap.add_argument("--max_length", type=int, default=8192) ap.add_argument("--epochs", type=int, default=3) ap.add_argument("--batch", type=int, default=8) ap.add_argument("--grad_accum", type=int, default=4) ap.add_argument("--lr", type=float, default=5e-5) ap.add_argument("--gc", action="store_true", help="enable gradient checkpointing") ap.add_argument("--output", default="victor/functiongemma-270m-agent-sft-lora") args = ap.parse_args() tokenizer = AutoTokenizer.from_pretrained(BASE, trust_remote_code=True) if tokenizer.pad_token is None: tokenizer.pad_token = tokenizer.eos_token print("pad_token:", tokenizer.pad_token, "| pad_id:", tokenizer.pad_token_id) if not args.smoke: # Fail fast if the push token is missing/invalid. HfApi().whoami(token=os.environ["HF_TOKEN"]) create_repo(args.output, token=os.environ["HF_TOKEN"], exist_ok=True) print("auth OK; output repo ready:", args.output) ds = load_dataset("victor/functiongemma-agent-sft", split="train") print("total rows:", len(ds)) if args.smoke: ds = ds.select(range(min(160, len(ds)))) split = ds.train_test_split(test_size=0.1, seed=42) train_ds, eval_ds = split["train"], split["test"] print("train rows:", len(train_ds), "| eval rows:", len(eval_ds)) tmap = lambda r: tokenize_row(r, tokenizer, args.max_length) train_ds = train_ds.map(tmap, remove_columns=["text"]) eval_ds = eval_ds.map(tmap, remove_columns=["text"]) model = AutoModelForCausalLM.from_pretrained( BASE, trust_remote_code=True, torch_dtype=torch.bfloat16 if torch.cuda.is_bf16_supported() else torch.float16, ) model.config.pad_token_id = tokenizer.pad_token_id peft = LoraConfig( r=16, lora_alpha=32, lora_dropout=0.0, target_modules=[ "q_proj", "k_proj", "v_proj", "o_proj", "gate_proj", "up_proj", "down_proj", ], task_type="CAUSAL_LM", ) use_bf16 = torch.cuda.is_bf16_supported() cfg = SFTConfig( output_dir="./fg-lora", num_train_epochs=args.epochs, per_device_train_batch_size=args.batch, gradient_accumulation_steps=args.grad_accum, learning_rate=args.lr, lr_scheduler_type="cosine", warmup_steps=0.03, bf16=use_bf16, fp16=not use_bf16, max_length=args.max_length, packing=False, eval_strategy="steps", eval_steps=200, logging_steps=20, save_strategy="epoch", report_to="none", gradient_checkpointing=args.gc, push_to_hub=False, hub_model_id=args.output, ) if args.max_steps > 0: cfg.max_steps = args.max_steps trainer = SFTTrainer( model=model, args=cfg, train_dataset=train_ds, eval_dataset=eval_ds, processing_class=tokenizer, peft_config=peft, ) trainer.train() if args.smoke: print("SMOKE OK") return token = os.environ["HF_TOKEN"] # Remove files from the earlier buggy run that pushed the UNTRAINED base # model into the output repo, so the repo holds only the real adapter. for f in ["model.safetensors", "config.json", "generation_config.json", "README.md"]: try: HfApi().delete_file(path_in_repo=f, repo_id=args.output, token=token) print("deleted stale file:", f) except Exception as e: print("skip delete", f, "->", e) # Push the TRAINED LoRA adapter (trainer.model is the PEFT-wrapped model). trainer.model.save_pretrained("./fg-adapter", safe_serialization=True) trainer.model.push_to_hub(args.output, token=token) tokenizer.push_to_hub(args.output, token=token) print("Pushed trained LoRA adapter to", args.output) if __name__ == "__main__": main()