#!/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"model\n", text): start = m.end() em = re.search(r"\n", 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()