File size: 5,899 Bytes
fd5e892
 
 
 
 
 
 
 
 
 
 
c9a51f2
fd5e892
 
 
 
 
 
 
3a6b9fc
fd5e892
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
c9a51f2
 
fd5e892
c9a51f2
fd5e892
 
 
3a6b9fc
 
 
 
 
bd1a946
 
 
 
 
 
fd5e892
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
d7961df
fd5e892
 
 
 
 
c9a51f2
 
fd5e892
 
c9a51f2
fd5e892
 
 
b6f2e8e
 
fd5e892
 
 
 
 
 
 
 
 
b6f2e8e
fd5e892
 
 
 
 
3a6b9fc
 
 
 
 
 
 
 
 
 
 
 
 
 
 
fd5e892
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
#!/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()