cortex.6.sol / evaluation /evaluate.py
openhands
openhands
feat(cortex): add CORTEX training pipeline and model audit
6fbe100
Raw History Blame Contribute Delete
4.6 kB
"""CORTEX evaluation.
Evaluation is deliberately separate from training. It answers two questions:
1. Is the model mechanically healthy? (loss on held-out data, perplexity)
2. Can it produce code at all? (exact-match / pass-rate on a held-out set)
The reported numbers are whatever they are. A small development model will score
badly on HumanEval; that is expected and is reported honestly rather than dressed up.
"""
from __future__ import annotations
import json
import math
from pathlib import Path
import torch
from model.cortex_model import CortexConfig
from training.train import build_token_stream, load_jsonl
__all__ = ["evaluate_loss", "evaluate_generation", "report"]
REPO_ROOT = Path(__file__).resolve().parent.parent
@torch.no_grad()
def evaluate_loss(model, token_data: torch.Tensor, batch_size: int = 2, device: str = "cpu") -> dict:
"""Average cross-entropy over the held-out token windows."""
model.eval().to(device)
n = token_data.shape[0]
total, count = 0.0, 0
for start in range(0, n, batch_size):
batch = token_data[start:start + batch_size].to(device)
out = model(batch[:, :-1], labels=batch[:, 1:])
loss = out["loss"].item()
total += loss * batch.shape[0]
count += batch.shape[0]
mean = total / max(count, 1)
return {
"loss": mean,
"perplexity": math.exp(min(mean, 50)),
"windows": count,
}
@torch.no_grad()
def evaluate_generation(model, tokenizer, records: list[dict], limit: int = 20,
max_new_tokens: int = 64, device: str = "cpu") -> dict:
"""Greedy-free sample generation, scored by whether the reference is reproduced.
This is a weak signal on a small model and is reported as such: it measures
whether the model can continue a prompt towards the reference solution, not
whether the code is correct.
"""
model.eval().to(device)
eos = tokenizer.token_to_id("<|place▁holder▁no▁0|>") or 1
exact, contains, samples = 0, 0, []
for record in records[:limit]:
text = record.get("text", "")
if "### Response" not in text:
continue
prompt, reference = text.split("### Response", 1)
prompt = prompt + "### Response"
ids = torch.tensor([tokenizer.encode(prompt, add_special_tokens=False).ids], device=device)
out = model.generate(ids, max_new_tokens=max_new_tokens, temperature=0.0,
eos_token_id=eos)
generated = tokenizer.decode(out[0][ids.shape[1]:].tolist())
reference = reference.strip()
exact += int(generated.strip() == reference)
contains += int(reference and reference in generated)
samples.append({
"prompt_tail": prompt[-80:].strip(),
"generated": generated[:200],
"reference": reference[:200],
})
n = max(len(samples), 1)
return {
"samples": samples,
"exact_match": exact / n,
"reference_contained": contains / n,
"n": len(samples),
}
def report(checkpoint_dir: str | Path, eval_dataset_id: str = "cortex-code-eval",
limit: int = 20, device: str = "cpu") -> dict:
"""Load a CORTEX checkpoint and produce a full evaluation report."""
from model.init import load_cortex_checkpoint
from tokenizer.cortex_tokenizer import load_tokenizer
from training.train import load_registry
model, config, provenance = load_cortex_checkpoint(checkpoint_dir)
tokenizer = load_tokenizer()
registry = load_registry()
entry = {d["id"]: d for d in registry["datasets"]}[eval_dataset_id]
records = load_jsonl(REPO_ROOT / entry["path"])
token_data = build_token_stream(records, tokenizer, config.max_position_embeddings // 4)
loss = evaluate_loss(model, token_data, device=device)
generation = evaluate_generation(model, tokenizer, records, limit=limit, device=device)
return {
"checkpoint": str(checkpoint_dir),
"model_name": config.model_name,
"parameters": provenance.get("parameter_count"),
"step": provenance.get("step"),
"eval_dataset": {
"id": eval_dataset_id,
"source": entry["source"],
"license": entry["license"],
},
"loss": loss,
"generation": {k: v for k, v in generation.items() if k != "samples"},
"samples": generation["samples"],
"caveat": (
"This is a development-scale model. Low scores are expected and are reported "
"as measured, not adjusted."
),
}