File size: 4,789 Bytes
7ac42fb
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3632d3d
7ac42fb
 
 
 
 
3632d3d
7ac42fb
 
 
 
 
 
 
 
3632d3d
7ac42fb
 
 
3632d3d
7ac42fb
 
 
 
 
 
 
 
 
 
 
3632d3d
 
 
 
 
 
 
7ac42fb
 
 
 
 
 
 
3632d3d
 
7ac42fb
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/usr/bin/env python
"""Evaluate the FunctionGemma LoRA adapter vs the base model on the held-out split.

Reports masked eval loss and next-token accuracy over model (assistant) turns
only, on the same 90/10 split used for training.
"""
import argparse
import re
import torch
from datasets import load_dataset
from peft import PeftModel
from transformers import AutoModelForCausalLM, AutoTokenizer

BASE = "unsloth/functiongemma-270m-it"
ADAPTER = "victor/functiongemma-270m-agent-sft-lora"


def model_spans(text):
    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 collate(rows, pad_id):
    ids = [r["input_ids"] for r in rows]
    att = [r["attention_mask"] for r in rows]
    lab = [r["labels"] for r in rows]
    ml = max(len(x) for x in ids)
    ids_t = torch.full((len(rows), ml), pad_id, dtype=torch.long)
    att_t = torch.zeros((len(rows), ml), dtype=torch.long)
    lab_t = torch.full((len(rows), ml), -100, dtype=torch.long)
    for i, (a, m, l) in enumerate(zip(ids, att, lab)):
        ids_t[i, : len(a)] = torch.tensor(a)
        att_t[i, : len(m)] = torch.tensor(m)
        lab_t[i, : len(l)] = torch.tensor(l)
    return ids_t, att_t, lab_t


@torch.no_grad()
def evaluate(model, ds, tokenizer, batch, device):
    model.eval()
    nll_sum, tok_sum, corr_sum = 0.0, 0, 0
    for i in range(0, len(ds), batch):
        rows = [ds[j] for j in range(i, min(i + batch, len(ds)))]
        ids, att, labs = collate(rows, tokenizer.pad_token_id)
        ids, att, labs = ids.to(device), att.to(device), labs.to(device)
        out = model(input_ids=ids, attention_mask=att, labels=labs)
        # accuracy over masked (assistant) tokens with causal shift
        logits = out.logits[:, :-1]  # (B, T-1, V)
        targets = labs[:, 1:]
        mask = targets != -100
        preds = logits.argmax(-1)
        corr = (preds == targets) & mask
        n = mask.sum().item()
        tok_sum += n
        corr_sum += corr.sum().item()
        if n > 0:
            nll_sum += out.loss.item() * n
        del out, logits
    loss = nll_sum / max(tok_sum, 1)
    acc = corr_sum / max(tok_sum, 1)
    return loss, acc


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--max_length", type=int, default=8192)
    ap.add_argument("--batch", type=int, default=1)
    args = ap.parse_args()

    device = "cuda" if torch.cuda.is_available() else "cpu"
    print("device:", device, "| batch:", args.batch)

    tokenizer = AutoTokenizer.from_pretrained(BASE, trust_remote_code=True)
    if tokenizer.pad_token is None:
        tokenizer.pad_token = tokenizer.eos_token

    ds = load_dataset("victor/functiongemma-agent-sft", split="train")
    split = ds.train_test_split(test_size=0.1, seed=42)
    eval_ds = split["test"]
    eval_ds = eval_ds.map(
        lambda r: tokenize_row(r, tokenizer, args.max_length), remove_columns=["text"]
    )
    lens = sorted(len(r["input_ids"]) for r in eval_ds)
    print(
        "held-out rows:", len(eval_ds),
        "| len p50=%d p95=%d max=%d" % (
            lens[len(lens) // 2], lens[int(len(lens) * 0.95)], lens[-1],
        ),
    )

    dtype = torch.bfloat16 if torch.cuda.is_bf16_supported() else torch.float16

    print("\n=== Base model ===")
    base = AutoModelForCausalLM.from_pretrained(BASE, trust_remote_code=True, torch_dtype=dtype).to(device)
    base_loss, base_acc = evaluate(base, eval_ds, tokenizer, args.batch, device)
    print(f"base  eval_loss={base_loss:.4f}  eval_token_acc={base_acc:.4f}")
    del base
    torch.cuda.empty_cache()

    print("\n=== Finetuned (base + LoRA adapter) ===")
    ft = AutoModelForCausalLM.from_pretrained(BASE, trust_remote_code=True, torch_dtype=dtype).to(device)
    ft = PeftModel.from_pretrained(ft, ADAPTER).to(device)
    ft_loss, ft_acc = evaluate(ft, eval_ds, tokenizer, args.batch, device)
    print(f"ft    eval_loss={ft_loss:.4f}  eval_token_acc={ft_acc:.4f}")

    print("\n=== Summary ===")
    print(f"eval_loss      : base {base_loss:.4f} -> ft {ft_loss:.4f}")
    print(f"eval_tok_acc   : base {base_acc:.4f} -> ft {ft_acc:.4f}")


if __name__ == "__main__":
    main()