Download train.py from victor/functiongemma-sft-train: direct link, hf CLI and curl.
- Browser
- Download file 5.9 kB
-
https://huggingface.co/victor/functiongemma-sft-train/resolve/main/train.py
- Command line
-
hf download hf://victor/functiongemma-sft-train/train.py
-
curl -L -o train.py https://huggingface.co/victor/functiongemma-sft-train/resolve/main/train.py
5.9 kB
| #!/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 <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: | |
| # 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() | |