#!/usr/bin/env python3
"""
Find the exact context length at which a model stops fitting, and test whether
that number is predictable rather than empirical.

    python3 threshold.py --model qwen3.5:9b-q4_K_M

`ctxsweep.py` showed that Ollama's default context can push layers onto the CPU
while VRAM sits unused, costing 2.6x on one model tested. That is a useful
warning but a useless tool: "try smaller numbers until it goes fast" is not an
answer anyone can act on without a benchmark rig.

If Ollama's planner is doing what it appears to -- estimating
`weights + KV cache(num_ctx) + buffers` and offloading layers until the estimate
fits -- then the tipping point is a *computable* property of the model and the
card, and the KV cache should grow linearly in num_ctx with a slope set by the
architecture.

This bisects for the largest fully-resident context, then measures resident VRAM
across several context sizes and fits a line to it. Two things fall out: the
number you should actually set, and a check on whether the mechanism is what we
think it is. If the fit is linear with a slope matching
`2 * n_layers * n_kv_heads * head_dim * bytes_per_element`, the model is right.
"""

from __future__ import annotations

import argparse
import json
import re
import subprocess
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 ps_split


def probe(model: str, ctx: int) -> dict:
    """Load at this context, record where the layers went and what it cost."""
    evict(model)
    wait_until(lambda: ollama_vram_mb() < 200, timeout=60)
    time.sleep(1.5)
    try:
        api("/api/generate",
            {"model": model, "prompt": "hi", "stream": False, "keep_alive": "60s",
             "options": {"num_ctx": ctx, "num_predict": 1}}, timeout=600)
    except Exception as e:
        # A load that errors is emphatically not resident. Returning True here
        # made the bisection read every failure as a success and walk straight
        # to the advertised ceiling.
        return {"num_ctx": ctx, "error": f"{type(e).__name__}: {e}",
                "processor_split": "load failed", "resident": False}
    split = ps_split(model)
    vram = ollama_vram_mb()
    evict(model)
    return {"num_ctx": ctx, "processor_split": split, "ollama_vram_mb": vram,
            "resident": split == "100% GPU"}


def arch(model: str) -> dict:
    """Architecture numbers straight from `ollama show`, for the prediction."""
    try:
        out = subprocess.run(["ollama", "show", model], capture_output=True,
                             text=True, timeout=30).stdout
    except (subprocess.SubprocessError, OSError):
        return {}
    want = {"context length": "ctx_max", "embedding length": "n_embd",
            "parameters": "params", "quantization": "quant",
            "architecture": "arch"}
    got: dict = {}
    for line in out.splitlines():
        for k, name in want.items():
            if line.strip().lower().startswith(k):
                got[name] = line.split()[-1]
    return got


def bisect_threshold(model: str, lo: int, hi: int, log: list) -> int:
    """Largest num_ctx that still loads 100% on GPU. -1 if even `lo` spills."""
    r = probe(model, lo)
    log.append(r)
    print(f"    ctx={lo:>7} {r.get('processor_split','?'):>16} "
          f"{r.get('ollama_vram_mb','?'):>6} MiB")
    if not r["resident"]:
        return -1

    best = lo
    while lo < hi:
        mid = (lo + hi + 1) // 2
        # Snap to a round step; Ollama pads context internally and single-token
        # granularity would spend a dozen loads resolving noise.
        mid = max(lo + 256, (mid // 256) * 256)
        if mid >= hi and mid != hi:
            mid = hi
        r = probe(model, mid)
        log.append(r)
        print(f"    ctx={mid:>7} {r.get('processor_split','?'):>16} "
              f"{str(r.get('ollama_vram_mb','-')):>6} MiB"
              + (f"   {r['error'][:60]}" if r.get("error") else ""))
        if r["resident"]:
            best, lo = mid, mid
        else:
            hi = mid - 256
        if hi - lo < 256:
            break
    return best


def fit_line(points: list[tuple[int, int]]) -> tuple[float, float]:
    """Least squares on (num_ctx, resident_MiB). Slope is KV MiB per token."""
    n = len(points)
    sx = sum(p[0] for p in points); sy = sum(p[1] for p in points)
    sxx = sum(p[0] * p[0] for p in points); sxy = sum(p[0] * p[1] for p in points)
    denom = n * sxx - sx * sx
    if not denom:
        return 0.0, sy / n
    slope = (n * sxy - sx * sy) / denom
    return slope, (sy - slope * sx) / n


def main() -> None:
    ap = argparse.ArgumentParser()
    ap.add_argument("--model", required=True)
    ap.add_argument("--min-ctx", type=int, default=512)
    ap.add_argument("--out", default=None)
    a = ap.parse_args()

    host = host_info()
    a_info = arch(a.model)
    ctx_max = int(a_info.get("ctx_max", 32768) or 32768)
    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   model={a.model}")
    print(f"architecture: {a_info}")
    if others:
        print(f"other processes on the card: {others}")

    print("\n  bisecting for the largest fully-resident context")
    log: list[dict] = []
    threshold = bisect_threshold(a.model, a.min_ctx, ctx_max, log)

    # Measure resident VRAM across the fitting range, to test linearity.
    print("\n  measuring KV growth below the threshold")
    fit_pts = []
    if threshold > 0:
        for c in sorted({a.min_ctx, threshold // 4, threshold // 2,
                         (threshold * 3) // 4, threshold}):
            if c < a.min_ctx:
                continue
            r = next((x for x in log if x["num_ctx"] == c), None) or probe(a.model, c)
            if r.get("resident") and r.get("ollama_vram_mb"):
                fit_pts.append((c, r["ollama_vram_mb"]))
                print(f"    ctx={c:>7}  {r['ollama_vram_mb']:>6} MiB")

    slope, intercept = fit_line(fit_pts) if len(fit_pts) >= 2 else (0.0, 0.0)
    kv_kib_per_tok = slope * 1024

    print(f"\n  max fully-resident context : {threshold}")
    print(f"  model advertises           : {ctx_max}")
    if len(fit_pts) >= 2:
        print(f"  weights + buffers (fit)    : {intercept:.0f} MiB")
        print(f"  KV cache per token (fit)   : {kv_kib_per_tok:.2f} KiB")
        free = int(host["vram_total_mb"]) - sum(others.values())
        pred = (free - intercept) * 1024 / kv_kib_per_tok if kv_kib_per_tok else 0
        print(f"  predicted threshold        : {pred:.0f}  "
              f"(measured {threshold}, {100*threshold/pred:.0f}% of prediction)")

    data = {"host": host, "model": a.model, "architecture": a_info,
            "other_processes_at_start": others,
            "measured_at": datetime.now(timezone.utc).isoformat(),
            "advertised_ctx": ctx_max,
            "max_resident_ctx": threshold,
            "fit": {"weights_plus_buffers_mb": round(intercept, 1),
                    "kv_kib_per_token": round(kv_kib_per_tok, 3),
                    "points": fit_pts},
            "probes": log}
    slug = (host["gpu"] or "gpu").lower().replace(" ", "-").replace("/", "-")
    out = Path(a.out or
               f"results/threshold--{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()
