| |
| """ |
| Frox AI Morph 1.1 — Evaluation Harness |
| |
| Runs three checks any time you finish a training phase: |
| 1. Perplexity on WikiText-2 (language modeling quality) |
| 2. Sanity generations on a fixed prompt set (coherence spot-check) |
| 3. Throughput benchmark (tokens/sec at a few sequence lengths) |
| |
| Usage: |
| python scripts/evaluate.py --model ./frox-morph-1-1-output/sft_final |
| python scripts/evaluate.py --model ./frox-morph-1-1-output/sft_final --skip-generation |
| """ |
| from __future__ import annotations |
|
|
| import argparse |
| import json |
| import sys |
| import time |
| from pathlib import Path |
|
|
| sys.path.insert(0, str(Path(__file__).resolve().parent.parent)) |
|
|
| import torch |
|
|
| from inference.engine.morph_engine import MorphInferenceEngine |
| from training.pipeline.trainer import evaluate_perplexity |
| from utils.common import print_banner, timer |
|
|
|
|
| SANITY_PROMPTS = [ |
| "What is the capital of France?", |
| "Write a haiku about the ocean.", |
| "Explain what a neural network is in two sentences.", |
| "def fibonacci(n):\n # Complete this function", |
| "What's 17 times 23?", |
| "Give me three tips for staying focused while studying.", |
| ] |
|
|
|
|
| def run_perplexity(engine: MorphInferenceEngine) -> float: |
| print("\n📊 Perplexity (WikiText-2)") |
| ppl = evaluate_perplexity( |
| engine.model.language_model, engine.tokenizer, engine.device, |
| max_samples=500, seq_len=512, |
| use_amp=engine.device.type == "cuda", amp_dtype=engine.dtype, |
| ) |
| print(f" Perplexity: {ppl}") |
| return ppl |
|
|
|
|
| def run_sanity_generations(engine: MorphInferenceEngine) -> list: |
| print("\n🧪 Sanity Generations") |
| results = [] |
| for prompt in SANITY_PROMPTS: |
| response = engine.generate( |
| [{"role": "user", "content": prompt}], |
| max_new_tokens=150, temperature=0.7, |
| ) |
| has_content = len(response.strip()) > 5 |
| has_repetition = _check_repetition(response) |
| status = "✓" if has_content and not has_repetition else "⚠" |
| print(f"\n {status} Q: {prompt}") |
| print(f" A: {response[:200]}{'...' if len(response) > 200 else ''}") |
| results.append({ |
| "prompt": prompt, "response": response, |
| "has_content": has_content, "has_repetition": has_repetition, |
| }) |
| return results |
|
|
|
|
| def _check_repetition(text: str, min_repeat: int = 4) -> bool: |
| """Flag degenerate repetition (a common failure mode of undertrained models).""" |
| words = text.split() |
| if len(words) < min_repeat * 2: |
| return False |
| for i in range(len(words) - min_repeat): |
| window = tuple(words[i:i + min_repeat]) |
| rest = words[i + min_repeat:i + min_repeat * 2] |
| if tuple(rest[:min_repeat]) == window: |
| return True |
| return False |
|
|
|
|
| def run_throughput_benchmark(engine: MorphInferenceEngine) -> dict: |
| print("\n⚡ Throughput Benchmark") |
| results = {} |
| for max_tokens in (50, 200, 500): |
| t0 = time.perf_counter() |
| _ = engine.generate( |
| [{"role": "user", "content": "Tell me a short story about a robot."}], |
| max_new_tokens=max_tokens, temperature=0.7, |
| ) |
| elapsed = time.perf_counter() - t0 |
| tok_s = max_tokens / elapsed |
| results[f"{max_tokens}_tokens"] = {"elapsed_s": round(elapsed, 2), "tok_per_s": round(tok_s, 1)} |
| print(f" {max_tokens:>4} tokens: {elapsed:.2f}s ({tok_s:.1f} tok/s)") |
| return results |
|
|
|
|
| def main(): |
| parser = argparse.ArgumentParser(description="Evaluate Frox AI Morph 1.1") |
| parser.add_argument("--model", type=str, required=True) |
| parser.add_argument("--skip-perplexity", action="store_true") |
| parser.add_argument("--skip-generation", action="store_true") |
| parser.add_argument("--skip-throughput", action="store_true") |
| parser.add_argument("--output", type=str, default="./eval_results.json") |
| args = parser.parse_args() |
|
|
| print_banner() |
| engine = MorphInferenceEngine.from_pretrained(args.model) |
|
|
| report = {"model_path": args.model, "stats": engine.stats()} |
|
|
| if not args.skip_perplexity: |
| with timer("Perplexity eval"): |
| report["perplexity"] = run_perplexity(engine) |
|
|
| if not args.skip_generation: |
| with timer("Sanity generations"): |
| report["sanity_generations"] = run_sanity_generations(engine) |
| n_ok = sum(1 for r in report["sanity_generations"] |
| if r["has_content"] and not r["has_repetition"]) |
| report["sanity_pass_rate"] = f"{n_ok}/{len(SANITY_PROMPTS)}" |
|
|
| if not args.skip_throughput: |
| with timer("Throughput benchmark"): |
| report["throughput"] = run_throughput_benchmark(engine) |
|
|
| Path(args.output).write_text(json.dumps(report, indent=2, default=str)) |
| print(f"\n✅ Full report saved to {args.output}") |
|
|
| if "perplexity" in report: |
| print(f"\n{'='*50}") |
| print(f"SUMMARY: perplexity={report['perplexity']} | " |
| f"sanity={report.get('sanity_pass_rate', 'skipped')}") |
| print(f"{'='*50}") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|