victor's picture
victor HF Staff
Put train.py
3a6b9fc verified
Raw
History Blame Contribute Delete
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()