Files
OBLITERATUS/obliteratus/capability_check.py
T
2026-08-19 22:25:38 -04:00

215 lines
7.0 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
#!/usr/bin/env python3
"""Capability check: compare abliterated model against stock on MMLU via lm-eval-harness.
Quick verification that abliteration surgery didn't lobotomize the model.
Uses lm-evaluation-harness for proper log-probability scoring — the same
methodology used by OrcaRouter, Coletti, and other abliteration releases.
Requires: pip install lm-eval
Usage:
obliteratus capability-check \\
--abliterated outputs/my-abliterated-model \\
--stock Qwen/Qwen3.8-27B \\
--device mps
# Quick mode (5 subjects, ~1 min per model):
obliteratus capability-check --abliterated ... --stock ... --quick
# Custom subjects:
obliteratus capability-check --abliterated ... --stock ... \\
--subjects mmlu_abstract_algebra,mmlu_computer_security
Lessons learned:
- DO NOT use custom generate-and-extract for MMLU. Log-probability scoring
(what lm-eval does) gives results comparable to published numbers.
Custom generation + letter extraction underperforms by 20+ pp.
- DO NOT use repetition_penalty for benchmarking. It interferes with
reasoning chains and degrades scores. Only use it for long-form generation.
- Stock Qwen3.8-27B scores ~87% on MMLU via lm-eval (0-shot).
If your stock score is much lower, your test setup is broken.
"""
from __future__ import annotations
import json
import logging
import subprocess
import sys
import tempfile
from pathlib import Path
logger = logging.getLogger(__name__)
QUICK_SUBJECTS = [
"mmlu_abstract_algebra",
"mmlu_computer_security",
"mmlu_us_foreign_policy",
"mmlu_high_school_biology",
"mmlu_professional_medicine",
]
DEFAULT_LIMIT = 5 # per subject; 57 subjects × 5 = 285 questions (comparable to OrcaRouter n=300)
def _run_lm_eval(model_path: str, tasks: str, limit: int, device: str,
output_dir: str, dtype: str = "bfloat16") -> dict:
"""Run lm-eval-harness and return parsed results."""
cmd = [
sys.executable, "-m", "lm_eval",
"--model", "hf",
"--model_args", f"pretrained={model_path},dtype={dtype},trust_remote_code=True",
"--tasks", tasks,
"--num_fewshot", "0",
"--limit", str(limit),
"--batch_size", "1",
"--device", device,
"--output_path", output_dir,
]
logger.info("Running: %s", " ".join(cmd[-8:]))
result = subprocess.run(cmd, capture_output=True, text=True, timeout=3600)
if result.returncode != 0:
logger.error("lm-eval failed:\n%s", result.stderr[-1000:])
raise RuntimeError(f"lm-eval exited with code {result.returncode}")
# Parse results from output directory
results_files = list(Path(output_dir).rglob("results*.json"))
if not results_files:
raise FileNotFoundError(f"No results files in {output_dir}")
with open(results_files[0]) as f:
return json.load(f)
def capability_check(
abliterated_path: str,
stock_path: str,
device: str = "auto",
dtype: str = "bfloat16",
quick: bool = False,
subjects: list[str] | None = None,
limit: int = DEFAULT_LIMIT,
output_dir: str | None = None,
) -> dict:
"""Compare abliterated vs stock model on MMLU.
Args:
abliterated_path: Path or HF repo for abliterated model.
stock_path: Path or HF repo for stock model.
device: Device (auto, cuda, mps, cpu).
dtype: Model dtype.
quick: Use 5 subjects instead of full MMLU.
subjects: Custom subject list (overrides quick).
limit: Questions per subject.
output_dir: Where to save results.
Returns:
dict with abliterated_acc, stock_acc, delta_pp.
"""
if subjects:
tasks = ",".join(subjects)
elif quick:
tasks = ",".join(QUICK_SUBJECTS)
else:
tasks = "mmlu"
if output_dir is None:
output_dir = tempfile.mkdtemp(prefix="obliteratus_capcheck_")
out = Path(output_dir)
# Run abliterated
logger.info("=== ABLITERATED ===")
abl_results = _run_lm_eval(
abliterated_path, tasks, limit, device,
str(out / "abliterated"), dtype
)
# Run stock
logger.info("=== STOCK ===")
stock_results = _run_lm_eval(
stock_path, tasks, limit, device,
str(out / "stock"), dtype
)
# Extract aggregate MMLU accuracy
abl_acc = None
stock_acc = None
for key in ["mmlu", tasks.split(",")[0]]:
if key in abl_results.get("results", {}):
abl_acc = abl_results["results"][key].get("acc,none")
break
for key in ["mmlu", tasks.split(",")[0]]:
if key in stock_results.get("results", {}):
stock_acc = stock_results["results"][key].get("acc,none")
break
# If running individual subjects, compute mean
if abl_acc is None:
accs = [v["acc,none"] for k, v in abl_results["results"].items()
if "acc,none" in v and not k.startswith("mmlu -")]
abl_acc = sum(accs) / len(accs) if accs else 0
if stock_acc is None:
accs = [v["acc,none"] for k, v in stock_results["results"].items()
if "acc,none" in v and not k.startswith("mmlu -")]
stock_acc = sum(accs) / len(accs) if accs else 0
delta = (abl_acc - stock_acc) * 100
summary = {
"abliterated_acc": round(abl_acc, 4),
"stock_acc": round(stock_acc, 4),
"delta_pp": round(delta, 1),
"tasks": tasks,
"limit": limit,
"method": "lm-eval-harness 0-shot log-likelihood",
}
# Save summary
with open(out / "capability_summary.json", "w") as f:
json.dump(summary, f, indent=2)
return summary
def main():
import argparse
logging.basicConfig(level=logging.INFO, format="%(message)s")
p = argparse.ArgumentParser(
description="Compare abliterated vs stock model on MMLU via lm-eval-harness."
)
p.add_argument("--abliterated", required=True, help="Abliterated model path or HF repo")
p.add_argument("--stock", required=True, help="Stock model path or HF repo")
p.add_argument("--device", default="auto")
p.add_argument("--dtype", default="bfloat16")
p.add_argument("--quick", action="store_true", help="5 subjects only (~1 min per model)")
p.add_argument("--subjects", type=str, default=None, help="Comma-separated subject list")
p.add_argument("--limit", type=int, default=DEFAULT_LIMIT, help="Questions per subject")
p.add_argument("--output-dir", type=str, default=None)
args = p.parse_args()
subjects = args.subjects.split(",") if args.subjects else None
result = capability_check(
args.abliterated, args.stock,
device=args.device, dtype=args.dtype,
quick=args.quick, subjects=subjects,
limit=args.limit, output_dir=args.output_dir,
)
print(f"\n{'='*50}")
print(f"CAPABILITY CHECK ({result['tasks']})")
print(f"{'='*50}")
print(f"Stock: {result['stock_acc']*100:.1f}%")
print(f"Abliterated: {result['abliterated_acc']*100:.1f}%")
print(f"Delta: {result['delta_pp']:+.1f}pp")
print(f"Method: {result['method']}")
if __name__ == "__main__":
main()