#!/usr/bin/env python3 """Warm every length bucket of intern-decision's kernel autotune cache. The fast kernels (triton + fla) autotune per shape BUCKET: the first call that lands in a new bucket costs seconds of autotune. The cache lives in the container's `triton-cache` volume, so it survives ordinary recreates — this script is for AFTER AN IMAGE CHANGE (a new triton/fla version keys new entries). Idempotent: a second run is all fast. One noul /v1/systemone call per bucket, smallest up to the service's own MAX_TOKENS (read from /health; never hard-coded). Aims just under each bucket's top edge; the aim is calibrated from the service's own reported input_tokens after every call, so drift cannot leave a bucket cold. Prints each call's bucket, the model's own token count, and the wall time. Exits non-zero if any call fails or lands outside its bucket. scripts/intern-decision-warmup [--endpoint URL] [--bucket-width N] """ from __future__ import annotations import argparse import json import subprocess import sys import time import urllib.error import urllib.request UNIT = "the quick brown fox. " # ~5.13 tokens against this service (measured 2026-09-30) def post(endpoint: str, token: str, body: dict, timeout: float = 180.0): req = urllib.request.Request( endpoint.rstrip("/") + "/v1/systemone", data=json.dumps(body).encode(), method="POST", headers={"Authorization": f"Bearer {token}", "Content-Type": "application/json"}) t0 = time.perf_counter() try: with urllib.request.urlopen(req, timeout=timeout) as r: return r.status, json.loads(r.read()), time.perf_counter() - t0 except urllib.error.HTTPError as e: return e.code, json.loads(e.read() or b"{}"), time.perf_counter() - t0 def main() -> int: ap = argparse.ArgumentParser() ap.add_argument("--endpoint", default="http://intern-decision.fv.internal:8033") ap.add_argument("--bucket-width", type=int, default=2048, help="asserted bucket width; the acceptance run proves or corrects it") args = ap.parse_args() token = subprocess.run(["secret", "get", "intern-decision/api-token"], capture_output=True, text=True).stdout.strip() if not token: print("no token: secret get intern-decision/api-token", file=sys.stderr) return 2 health = json.load(urllib.request.urlopen(args.endpoint.rstrip("/") + "/health", timeout=15)) max_tokens = int(health["max_tokens"]) width = args.bucket_width buckets = list(range(width, max_tokens + 1, width)) print(f"endpoint={args.endpoint} max_tokens={max_tokens} bucket_width={width} " f"buckets={len(buckets)}") def ask(reps: int): return post(args.endpoint, token, {"state": UNIT * reps, "questions": {"warm": {"type": "noul", "instructions": "Is the system healthy?"}}}) # Two-point calibration of the tokenizer's linear model, input_tokens = a*reps + b. A single # probe cannot separate a from b and its "calibration" overcorrects on the next aim (measured: # the aim oscillates ±500 tokens around the bucket edge). Two probes 900 reps apart fit exactly. s1, o1, _ = ask(100) s2, o2, _ = ask(1000) if s1 != 200 or s2 != 200: print("calibration failed", s1, s2, file=sys.stderr) return 2 t1 = o1["usage"]["input_tokens"]; t2 = o2["usage"]["input_tokens"] a = (t2 - t1) / 900.0 b = t1 - a * 100 print(f"calibrated: input_tokens = {a:.4f}*reps + {b:.1f}") failed, total = False, 0.0 print(f"{'bucket<=':>9} {'tokens':>7} {'wall_s':>7} note") for hi in buckets: note = "" got = status = 0 wall = " n/a" for attempt in (1, 2): # a second aim fixes a rounding miss; never miss a bucket reps = max(1, round((hi - 60 - b) / a)) status, out, w = ask(reps) total += w wall = f"{w:.2f}" got = (out.get("usage") or {}).get("input_tokens", 0) if status != 200 or hi - width < got <= hi: break if status != 200: note = f"FAILED {status}: {(out.get('error') or {}).get('message', '')[:60]}" failed = True elif not (hi - width < got <= hi): note = f"OUT OF BUCKET ({hi-width} < {got} <= {hi})" failed = True print(f"{hi:>9} {got:>7} {wall:>7} {note}") print(f"total {total:.1f} s over {len(buckets)} buckets") return 1 if failed else 0 if __name__ == "__main__": sys.exit(main())