Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
22 changes: 19 additions & 3 deletions agents/conductors/profiling/AGENTS.md
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,7 @@ prompt, PyAutoMind `issued/profiling_agent.md`.
|------|----------|-------|
| `campaign` | Which grid runs are done / CPU-unusable / missing on this tier, and how do I dispatch the rest? | dispatch plan (local sweep flags incl. the per-run timeout; A100 submit list) |
| `ingest` | Which probe JSONs aren't in the vram tables yet, and which results have no pin? | table-update rows, pin list, baseline + dashboard steps |
| `ingest --axis compile` | Which warm compile rows are unpinned, and which have drifted from their pin? | drifted rows (with pinned vs observed), unpinned keys, confirm/classify/re-pin steps |
| `triage` | What do the pinned-drift findings mean? | per-finding classification: stale pin → re-pin here; library regression → `bug/` via intake |

```
Expand All @@ -41,9 +42,24 @@ bucketed by **hardware**, with `mixed_precision` a separate field. The two
vocabularies do not interchange, so the compile axis maps tiers itself rather
than reusing `TIER_CONFIGS`.

`--axis compile` currently serves `campaign` (coverage); `ingest` and `triage`
reject it with exit 5 until the compile pins land, so a compile flag can never
silently return a runtime answer.
`--axis compile` serves `campaign` (coverage) and `ingest` (warm-pin drift);
`triage` rejects it with exit 5 until drift classification lands, so a compile
flag can never silently return a runtime answer.

**Drift is deliberately hard to trigger.** A row counts only if it is *newer*
than its pin, at least `2.0x` the pinned value, **and** at least `1.0 s` above it
in absolute terms. Rows predating the pin are the history the pin was chosen
over — flagging them would report the improvement that set the pin as a
regression. The ratio alone screams about sub-second cells where 100 ms of
jitter is 3x; the absolute floor alone misses a cheap cell degrading by an order
of magnitude. Both gates, generous, because host load alone has produced 7x
errors in this corpus and an alarm that cries wolf gets ignored.

Pins live in the workspace (`jax_compile/pins.json`) and are **sticky** — the
workspace's `update_pins.py` will not move an existing pin without `--repin`. If
pins auto-followed the newest measurement, re-deriving them after a cache
regression would bake the regression in and the surveillance would report
all-clear forever.

**Compile timings are host-load-sensitive** — the first measurements in
`jax_compile/README.md` were wrong by up to **7×** (851 s vs 117 s for the same
Expand Down
165 changes: 161 additions & 4 deletions agents/conductors/profiling/_profiling.py
Original file line number Diff line number Diff line change
Expand Up @@ -61,6 +61,20 @@
# the results tree rather than a corrupt probe record.
COMPILE_IDENTITY_FIELDS = ("hardware", "dataset_class", "instrument")

# Mirrors autolens_profiling/scripts/misc/jax_compile/pins.py. Duplicated rather
# than imported for the same reason the grid is read via ast: importing the
# workspace would drag the JAX stack into the Brain. Kept honest by a test that
# reads the workspace's own definition.
COMPARABILITY_FIELDS = ("hardware", "hostname", "jax_version", "mixed_precision", "cache_state")
CELL_FIELDS = ("dataset_class", "model_type", "instrument", "transform")
PIN_FIELDS = COMPARABILITY_FIELDS + CELL_FIELDS

# Drift thresholds. Generous on purpose: host load alone has produced 7x errors
# in this corpus, so a tight bound would flag a busy laptop as a regression and
# teach people to ignore the alarm.
COMPILE_DRIFT_RATIO = 2.0
COMPILE_DRIFT_FLOOR_S = 1.0


def workspace_root(explicit: str | None = None) -> Path:
if explicit:
Expand Down Expand Up @@ -140,6 +154,31 @@ def load_compile_corpus(ws: Path) -> "list[tuple[str, int, dict[str, Any]]]":
return out


def load_pins(ws: Path) -> list[dict[str, Any]]:
"""The workspace's warm-compile pins (`jax_compile/pins.json`)."""
path = compile_dir(ws) / "pins.json"
if not path.is_file():
return []
try:
data = json.loads(path.read_text())
except (OSError, ValueError):
return []
pins = data.get("pins") if isinstance(data, dict) else data
return [p for p in pins if isinstance(p, dict)] if isinstance(pins, list) else []


def pin_key_str(key: tuple) -> str:
parts = dict(zip(PIN_FIELDS, key))
cell = "/".join(
str(parts[f]) for f in ("dataset_class", "model_type", "instrument") if parts.get(f)
)
return (
f"{cell} [{parts.get('transform')}] "
f"@ {parts.get('hardware')}/{parts.get('hostname')} jax{parts.get('jax_version')}"
f"{' mp' if parts.get('mixed_precision') else ''} {parts.get('cache_state')}"
)


def compile_tier_of(hardware: str | None) -> str:
"""Which campaign tier a compile record belongs to.

Expand Down Expand Up @@ -341,6 +380,106 @@ def campaign_compile(ws: Path, tier: str) -> dict[str, Any]:
}


def ingest_compile(ws: Path) -> dict[str, Any]:
"""Which warm compile rows are unpinned, and which have drifted from a pin.

The surveillance the arc exists for: the persistent cache and
`--xla_gpu_autotune_level=0` are *settings*, so a config drift or an
`XLA_FLAGS` clobber puts the worst case back with nothing failing.

Every comparison here happens strictly inside one comparability key. Rows
from different hardware, hosts, jax versions, precisions or cache states are
never paired — that is not conservatism, it is the difference between a
signal and noise: compile timings are host-load-sensitive to a measured 7x,
and a `jax_version` bump recompiles once BY DESIGN rather than regressing.
"""
pins = load_pins(ws)
if not pins:
return {
"agent": "profiling",
"mode": "ingest",
"axis": "compile",
"pins": 0,
"unpinned": [],
"drifted": [],
"next_action": (
"no compile pins — run `python3 scripts/misc/jax_compile/update_pins.py --write` "
"in autolens_profiling first"
),
}

by_key = {tuple(p.get(f) for f in PIN_FIELDS): p for p in pins}
unpinned: list[dict[str, Any]] = []
drifted: list[dict[str, Any]] = []
seen: set[tuple] = set()

for rel, idx, rec in load_compile_corpus(ws):
if rec.get("cache_state") != "warm" or "compile_s" not in rec:
continue
key = tuple(rec.get(f) for f in PIN_FIELDS)
if any(k in (None, "") for k in key if k is not False):
continue
pin = by_key.get(key)
if pin is None:
if key not in seen:
seen.add(key)
unpinned.append({"record": f"{rel}[{idx}]", "pin": pin_key_str(key)})
continue
# Only rows NEWER than the pin can be drift. Every warm row predating
# the pin is the history the pin was chosen over — flagging those
# reports the improvement that set the pin as though it were a
# regression, which is how an alarm earns its way into being ignored.
if str(rec.get("timestamp") or "") <= str(pin.get("source_timestamp") or ""):
continue
expected, got = pin.get("compile_s"), rec.get("compile_s")
if not isinstance(expected, (int, float)) or not isinstance(got, (int, float)):
continue
if expected <= 0:
continue
ratio = got / expected
# Both gates, deliberately. The ratio alone screams about sub-second
# cells where a 100 ms jitter is 3x; the absolute delta alone misses a
# cheap cell degrading by an order of magnitude. Generous because host
# load alone has produced 7x errors in this corpus.
if ratio >= COMPILE_DRIFT_RATIO and abs(got - expected) >= COMPILE_DRIFT_FLOOR_S:
drifted.append(
{
"record": f"{rel}[{idx}]",
"pin": pin_key_str(key),
"pinned_s": expected,
"observed_s": got,
"ratio": round(ratio, 2),
"tag": rec.get("tag"),
}
)

return {
"agent": "profiling",
"mode": "ingest",
"axis": "compile",
"pins": len(pins),
"unpinned": unpinned,
"drifted": drifted,
"policy": (
f"Drift = a warm row NEWER than its pin, >= {COMPILE_DRIFT_RATIO}x the "
f"pinned value AND >= {COMPILE_DRIFT_FLOOR_S}s absolute, compared ONLY "
f"within {'/'.join(COMPARABILITY_FIELDS)}. Cross-key pairs and rows "
"predating the pin are never a regression."
),
"steps": [
"re-run the drifted cell warm to confirm it is not host load "
"(check the record's host_state against the pin's)",
"if confirmed, classify it — `pyauto-brain profiling triage --axis compile`",
"pin the unpinned rows: `python3 scripts/misc/jax_compile/update_pins.py --write`",
],
"next_action": (
"compile pins current — no warm drift"
if not drifted and not unpinned
else f"{len(drifted)} drifted, {len(unpinned)} unpinned warm key(s)"
),
}


# ---------------------------------------------------------------------------
# ingest
# ---------------------------------------------------------------------------
Expand Down Expand Up @@ -517,6 +656,21 @@ def emit_human(d: dict[str, Any]) -> None:
print("Dispatch plan:")
for s in d["dispatch_plan"]:
print(f" - {s}")
elif d["mode"] == "ingest" and d.get("axis") == "compile":
print(f"Compile pins: {d['pins']}")
print(f"Drifted: {len(d['drifted'])}")
for x in d["drifted"][:10]:
print(
f" {x['pin']}: pinned {x['pinned_s']}s -> observed "
f"{x['observed_s']}s ({x['ratio']}x, tag={x['tag']!r})"
)
print(f"Unpinned warm keys: {len(d['unpinned'])}")
for x in d["unpinned"][:10]:
print(f" {x['pin']}")
if d.get("policy"):
print(f"Policy: {d['policy']}")
for s in d.get("steps", []):
print(f" - {s}")
elif d["mode"] == "ingest":
print(f"Provenance: {d['provenance']}")
print(f"Probe updates: {len(d['probe_updates'])}")
Expand Down Expand Up @@ -557,10 +711,13 @@ def main(argv=None) -> int:
# ingest/triage own the compile axis in later phases of the arc (pins, then
# drift classification). Refusing now is deliberate: a mode that silently
# ignored --axis would report runtime findings under a compile flag.
if a.axis == "compile" and a.mode != "campaign":
# triage owns the compile axis in phase 3 (drift CLASSIFICATION). Refusing
# is deliberate: a mode that silently ignored --axis would report runtime
# findings under a compile flag.
if a.axis == "compile" and a.mode == "triage":
print(
f"profiling: --axis compile is not implemented for {a.mode!r} yet "
"(campaign only; ingest/triage land with the compile pins)",
"profiling: --axis compile is not implemented for 'triage' yet "
"(campaign + ingest only; classification lands next)",
file=sys.stderr,
)
return 5
Expand All @@ -573,7 +730,7 @@ def main(argv=None) -> int:
if a.mode == "campaign":
d = campaign_compile(ws, a.tier) if a.axis == "compile" else campaign(ws, a.tier)
elif a.mode == "ingest":
d = ingest(ws)
d = ingest_compile(ws) if a.axis == "compile" else ingest(ws)
else:
d = triage(ws)

Expand Down
Loading
Loading