53 lines
2.5 KiB
Python
53 lines
2.5 KiB
Python
#!/usr/bin/env python3
|
|
"""Repeat the two consequential failures and a 49k-token recall probe."""
|
|
import argparse
|
|
import json
|
|
import pathlib
|
|
import time
|
|
|
|
from measure import request
|
|
|
|
|
|
def main():
|
|
p = argparse.ArgumentParser()
|
|
p.add_argument("--base", required=True)
|
|
p.add_argument("--model", required=True)
|
|
p.add_argument("--output", type=pathlib.Path, required=True)
|
|
p.add_argument("--temperature", type=float, default=0.2)
|
|
p.add_argument("--seeds", default="7,99")
|
|
p.add_argument("--skip-long", action="store_true")
|
|
p.add_argument("--go", action="store_true")
|
|
args = p.parse_args()
|
|
if not args.go:
|
|
p.error("Inference requires explicit --go")
|
|
tasks = {item["id"]: item for item in json.loads(pathlib.Path(__file__).with_name("tasks.json").read_text())}
|
|
seeds = [int(value) for value in args.seeds.split(",")]
|
|
report = {"model": args.model, "temperature": args.temperature,
|
|
"started": time.time(), "repeats": []}
|
|
for task_id in ("i3_code_debugging", "i6_state_vs_configuration"):
|
|
for seed in seeds:
|
|
payload = {"model": args.model, "stream": False, "temperature": args.temperature,
|
|
"seed": seed, "reasoning_effort": "medium", "max_tokens": 4096,
|
|
"messages": [{"role": "user", "content": tasks[task_id]["prompt"]}]}
|
|
answer = request(args.base, None, payload)
|
|
report["repeats"].append({"id": task_id, "seed": seed, **answer})
|
|
print(task_id, seed, answer["wall_seconds"], flush=True)
|
|
if not args.skip_long:
|
|
filler = ("In one sentence, explain why a 10 mm through-hole in a 40 mm cube "
|
|
"does not change its external dimensions. " * 1800)
|
|
prompt = "Merke dir diesen einmaligen Prüfcode: KIESEL-7319.\n" + filler + \
|
|
"\nNenne ausschließlich den Prüfcode aus dem Anfang der Nachricht."
|
|
payload = {"model": args.model, "stream": False, "temperature": 0,
|
|
"seed": 42, "reasoning_effort": "none", "max_tokens": 96,
|
|
"messages": [{"role": "user", "content": prompt}]}
|
|
report["long_recall"] = request(args.base, None, payload)
|
|
report["finished"] = time.time()
|
|
args.output.parent.mkdir(parents=True, exist_ok=True)
|
|
args.output.write_text(json.dumps(report, ensure_ascii=False, indent=2) + "\n")
|
|
if not args.skip_long:
|
|
print("long_recall", report["long_recall"]["wall_seconds"], flush=True)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|