victor's picture
victor HF Staff
Put eval.py
3632d3d verified
Raw
History Blame Contribute Delete
4.79 kB
#!/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()