148 lines
6.0 KiB
Python
148 lines
6.0 KiB
Python
"""CompassionBench response collection: base, ft, ft_rag (local mlx) + kimi_k3 (frontier).
|
|
|
|
Deviates from the task-7 brief's collect_local() on purpose (controller ruling):
|
|
the brief reloads the model via rag.answer() for every one of the 150 ft_rag
|
|
prompts, which means 150 full 7B model loads. Instead we load the model +
|
|
tokenizer ONCE per system pass, and for ft_rag inline retrieval (buddhagpt.index.search)
|
|
+ prompt construction (buddhagpt.rag.build_prompt) around the single preloaded
|
|
model/tokenizer.
|
|
|
|
Resumable: skips any (prompt_id, system) row already present in the output
|
|
file, so a restart after an interruption does not redo work. Writes are
|
|
flushed after every row so a kill -9 mid-run loses at most the in-flight
|
|
generation.
|
|
|
|
Run order: kimi_k3 first (fast, validates the OpenRouter plumbing end to end
|
|
before committing to the multi-hour local passes), then base, ft, ft_rag.
|
|
|
|
kimi_k3 uses a ThreadPoolExecutor(max_workers=8) (same pattern as
|
|
scripts/gen_data.py's run(): a lock around the shared output-file append) since
|
|
it's a remote API call and safely parallelizable. The three local mlx passes
|
|
stay strictly sequential/single-stream — one GPU, no benefit to threading, and
|
|
mlx generation is not thread-safe across concurrent calls on one model.
|
|
"""
|
|
import json
|
|
import threading
|
|
from concurrent.futures import ThreadPoolExecutor
|
|
from pathlib import Path
|
|
|
|
from buddhagpt.collect import load_bench, load_done
|
|
|
|
BENCH = Path("eval/compassionbench.yaml")
|
|
OUT = Path("data/responses.jsonl")
|
|
BASE_MODEL = "mlx-community/Qwen2.5-7B-Instruct-4bit"
|
|
FT_MODEL = "models/buddhagpt-7b-v1"
|
|
RAG_DB = Path("data/lancedb")
|
|
KIMI_MODEL = "moonshotai/kimi-k3"
|
|
|
|
|
|
def run_frontier_pass(items: list[dict], out: Path, model: str = KIMI_MODEL, max_workers: int = 8):
|
|
from buddhagpt.llm import openrouter_client, chat
|
|
|
|
client = openrouter_client()
|
|
system = model.split("/")[-1].replace("-", "_")
|
|
done = load_done(out)
|
|
remaining = [it for it in items if (it["id"], system) not in done]
|
|
print(f"[{system}] {len(remaining)}/{len(items)} remaining", flush=True)
|
|
if not remaining:
|
|
return
|
|
|
|
lock = threading.Lock()
|
|
totals = {"input": 0, "output": 0}
|
|
fail_count = 0
|
|
done_count = 0
|
|
out_f = out.open("a") # append: incremental persistence, survives interruption
|
|
|
|
def gen_one(it):
|
|
nonlocal fail_count, done_count
|
|
try:
|
|
# kimi-k3 emits hidden reasoning tokens even with reasoning.enabled=False,
|
|
# and they draw from the same max_tokens budget as the visible answer —
|
|
# at 1000 tokens, harder prompts exhausted the budget mid-reasoning and
|
|
# returned empty content (finish_reason="length", content=None). 3000
|
|
# leaves enough room for reasoning + a full visible answer.
|
|
text, usage = chat(client, model, [{"role": "user", "content": it["prompt"]}], max_tokens=8000)
|
|
if not text or not text.strip():
|
|
raise ValueError("empty response content (likely reasoning-token budget exhaustion)")
|
|
except Exception as e:
|
|
with lock:
|
|
fail_count += 1
|
|
print(f"[{system}] FAILED {it['id']}: {e}", flush=True)
|
|
return
|
|
with lock:
|
|
totals["input"] += usage["input"]; totals["output"] += usage["output"]
|
|
done_count += 1
|
|
out_f.write(json.dumps({"prompt_id": it["id"], "system": system, "response": text}) + "\n")
|
|
out_f.flush()
|
|
print(f"[{system}] {done_count}/{len(remaining)} done ({it['id']})", flush=True)
|
|
|
|
with ThreadPoolExecutor(max_workers=max_workers) as pool:
|
|
list(pool.map(gen_one, remaining))
|
|
|
|
out_f.close()
|
|
# Report raw token totals; actual $ cost is read from the OpenRouter activity
|
|
# dashboard (per-token pricing for kimi-k3 not hardcoded here to avoid drift).
|
|
print(
|
|
f"[{system}] calls: {len(remaining)} | failed: {fail_count} | tokens: {totals}",
|
|
flush=True,
|
|
)
|
|
|
|
|
|
def run_local_pass(items: list[dict], system: str, model_path: str, out: Path, use_rag: bool):
|
|
from mlx_lm import load, generate
|
|
|
|
done = load_done(out)
|
|
remaining = [it for it in items if (it["id"], system) not in done]
|
|
print(f"[{system}] {len(remaining)}/{len(items)} remaining", flush=True)
|
|
if not remaining:
|
|
return
|
|
|
|
print(f"[{system}] loading model from {model_path} ...", flush=True)
|
|
model, tokenizer = load(model_path)
|
|
|
|
if use_rag:
|
|
from buddhagpt.index import search
|
|
from buddhagpt.rag import build_prompt
|
|
|
|
with out.open("a") as f:
|
|
for i, it in enumerate(remaining, 1):
|
|
if use_rag:
|
|
passages = search(RAG_DB, it["prompt"], k=4)
|
|
messages = build_prompt(it["prompt"], passages)
|
|
else:
|
|
messages = [{"role": "user", "content": it["prompt"]}]
|
|
prompt = tokenizer.apply_chat_template(
|
|
messages, tokenize=False, add_generation_prompt=True
|
|
)
|
|
response = generate(model, tokenizer, prompt=prompt, max_tokens=500)
|
|
f.write(json.dumps({"prompt_id": it["id"], "system": system, "response": response}) + "\n")
|
|
f.flush()
|
|
print(f"[{system}] {i}/{len(remaining)} done ({it['id']})", flush=True)
|
|
|
|
|
|
def main():
|
|
items = load_bench(BENCH)
|
|
OUT.parent.mkdir(parents=True, exist_ok=True)
|
|
|
|
# 1. Frontier reference first — fast, validates plumbing before the long local runs.
|
|
# Parallel (8 workers): remote API calls, safely concurrent.
|
|
print("=== kimi_k3 (frontier reference) ===", flush=True)
|
|
run_frontier_pass(items, OUT)
|
|
|
|
# 2. Local passes — each loads the model exactly once.
|
|
print("=== base ===", flush=True)
|
|
run_local_pass(items, "base", BASE_MODEL, OUT, use_rag=False)
|
|
|
|
print("=== ft ===", flush=True)
|
|
run_local_pass(items, "ft", FT_MODEL, OUT, use_rag=False)
|
|
|
|
print("=== ft_rag ===", flush=True)
|
|
run_local_pass(items, "ft_rag", FT_MODEL, OUT, use_rag=True)
|
|
|
|
total = sum(1 for _ in OUT.open())
|
|
print(f"done. total rows: {total}", flush=True)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|