| |
| """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 <start_of_turn>model ... <end_of_turn> turns.""" |
| spans = [] |
| for m in re.finditer(r"<start_of_turn>model\n", text): |
| start = m.end() |
| em = re.search(r"\n<end_of_turn>", 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: |
| |
| 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"] |
| |
| |
| 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) |
|
|
| |
| 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() |
|
|