#!/usr/bin/env python3
"""
Measure what Ollama's default context sizing costs you on a free card.

    python3 ctxsweep.py --model dolphin3:8b

Found while running `contention.py`: a 4.9 GB model on an otherwise-idle 10 GiB
card was running 12% of its layers on the CPU. It had room. Ollama's allocation
planner sizes the KV cache for the model's advertised context length, and when
that plan does not fit it offloads layers until it does -- leaving part of the
card unused and paying the CPU-spill penalty for a context you are not using.

`num_ctx` is the knob. This sweeps it and records residency, split and
throughput at each value, so the cost of the default is a number rather than a
suspicion.
"""

from __future__ import annotations

import argparse
import json
import sys
import time
from datetime import datetime, timezone
from pathlib import Path

sys.path.insert(0, str(Path(__file__).resolve().parent))

from bench import api, evict, host_info, ollama_vram_mb, processes, wait_until
from contention import PROMPT, ps_split


def one(model: str, ctx: int | None) -> dict:
    evict(model)
    wait_until(lambda: ollama_vram_mb() < 200, timeout=60)
    time.sleep(2)

    opts: dict = {"num_predict": 96}
    if ctx:
        opts["num_ctx"] = ctx
    t = time.time()
    r = api("/api/generate", {"model": model, "prompt": PROMPT, "stream": False,
                              "keep_alive": "5m", "options": opts}, timeout=600)
    ev, ns = r.get("eval_count", 0), r.get("eval_duration", 0)
    row = {
        "num_ctx": ctx,
        "total_seconds": round(time.time() - t, 2),
        "load_seconds": round(r.get("load_duration", 0) / 1e9, 2),
        "tokens_per_second": round(ev / (ns / 1e9), 2) if ns else None,
        "ollama_vram_mb": ollama_vram_mb(),
        "processor_split": ps_split(model),
    }
    evict(model)
    return row


def main() -> None:
    ap = argparse.ArgumentParser()
    ap.add_argument("--model", required=True)
    ap.add_argument("--ctx", default="0,131072,32768,16384,8192,4096",
                    help="context sizes to try; 0 means leave it to Ollama")
    ap.add_argument("--out", default=None)
    a = ap.parse_args()

    host = host_info()
    others = {k: v for k, v in processes().items() if "ollama" not in k.lower()}
    print(f"\n{host['gpu']}  {host['vram_total_mb']} MiB total   model={a.model}")
    if others:
        print(f"other processes on the card: {others}")
    print("\n  the card is otherwise idle for all rows below\n")

    rows = []
    for c in [int(x) for x in a.ctx.split(",")]:
        row = one(a.model, c or None)
        rows.append(row)
        print(f"  num_ctx={str(c or 'default'):>8}  {row['processor_split']:>16}  "
              f"{str(row['tokens_per_second']):>7} tok/s   "
              f"{row['ollama_vram_mb']:>5} MiB resident")

    data = {"host": host, "model": a.model, "prompt": PROMPT,
            "other_processes_at_start": others,
            "measured_at": datetime.now(timezone.utc).isoformat(),
            "steps": rows}
    slug = (host["gpu"] or "gpu").lower().replace(" ", "-").replace("/", "-")
    out = Path(a.out or
               f"results/ctxsweep--{slug}--{a.model.replace(':','-').replace('/','-')}.json")
    out.parent.mkdir(parents=True, exist_ok=True)
    out.write_text(json.dumps(data, indent=2))
    print(f"\nwrote {out}")


if __name__ == "__main__":
    main()
