changh95's picture
tt-model push diffusion-planner-p150 (container)
4d9b003 verified
Raw History Blame Contribute Delete
7.47 kB
# SPDX-License-Identifier: Apache-2.0
"""Optimization knobs: declared once, read from the environment once at model build, each an A/B switch.
The rule every port follows (PLAN.md section 1.1; RP section 1.2): a knob's default is the measured best, the
server pins the full set in ``tt-model.yaml serve.env`` (:meth:`Knobs.serve_env` renders it), and a host test can
check the pins against the Python defaults. Setting an ``experiment`` knob logs a warning, so an A/B run is never
mistaken for the published configuration.
Example::
KNOBS = Knobs("CENTERPOINT", [
Knob("FUSED_HEAD", True, "merged 64->384->15 head convs (False = one conv per head)"),
Knob("BFP8_WEIGHTS", False, "bfp8 backbone weights", experiment=True),
Knob("NUM_CQS", 1, "command queues", choices=(1, 2)),
])
knobs = KNOBS.read() # at build; KnobValues is immutable
if knobs.FUSED_HEAD: ...
"""
from __future__ import annotations
import logging
import os
from dataclasses import dataclass
from typing import Any, Dict, Iterable, Iterator, Mapping, Optional, Tuple
__all__ = ["Knob", "Knobs", "KnobValues", "parse_bool"]
log = logging.getLogger(__name__)
_TRUE = {"1", "true", "yes", "on", "y"}
_FALSE = {"0", "false", "no", "off", "n", ""}
def parse_bool(text: str) -> bool:
value = str(text).strip().lower()
if value in _TRUE:
return True
if value in _FALSE:
return False
raise ValueError(f"{text!r} is not a boolean (1/0, true/false, on/off, yes/no)")
@dataclass(frozen=True)
class Knob:
"""One switch. ``name`` is the suffix of the env variable ``<PREFIX>_<NAME>``; the type is the default's."""
name: str
default: Any
doc: str = ""
choices: Optional[Tuple[Any, ...]] = None
experiment: bool = False
def __post_init__(self) -> None:
if not self.name or not self.name.replace("_", "").isalnum() or self.name.upper() != self.name:
raise ValueError(f"knob name {self.name!r} must be UPPER_SNAKE_CASE")
if self.choices is not None and self.default not in self.choices:
raise ValueError(f"knob {self.name}: default {self.default!r} not in {self.choices}")
def parse(self, text: str) -> Any:
kind = type(self.default)
try:
if kind is bool:
value: Any = parse_bool(text)
elif kind is int:
value = int(text.strip())
elif kind is float:
value = float(text.strip())
else:
value = text.strip()
except ValueError as exc:
raise ValueError(f"{self.name}={text!r}: {exc}") from None
if self.choices is not None and value not in self.choices:
raise ValueError(f"{self.name}={value!r}: expected one of {self.choices}")
return value
def render(self, value: Any) -> str:
if isinstance(value, bool):
return "1" if value else "0"
return str(value)
class KnobValues(Mapping):
"""Immutable knob values with their source (``"default"`` / ``"env"`` / ``"override"``); attribute or item
access."""
def __init__(self, prefix: str, values: Dict[str, Any], sources: Dict[str, str]):
object.__setattr__(self, "_prefix", prefix)
object.__setattr__(self, "_values", dict(values))
object.__setattr__(self, "_sources", dict(sources))
def __getattr__(self, name: str) -> Any:
if name.startswith("_"):
raise AttributeError(name)
try:
return self._values[name]
except KeyError:
raise AttributeError(f"no knob {name!r} (prefix {self._prefix})") from None
def __setattr__(self, name: str, value: Any) -> None:
raise AttributeError("KnobValues is read-only (knobs are read once at build)")
def __getitem__(self, name: str) -> Any:
return self._values[name]
def __iter__(self) -> Iterator[str]:
return iter(self._values)
def __len__(self) -> int:
return len(self._values)
def source(self, name: str) -> str:
return self._sources[name]
def overridden(self) -> Dict[str, Any]:
"""Knobs not at their default source (set from the environment or by an explicit override)."""
return {k: v for k, v in self._values.items() if self._sources[k] != "default"}
def as_dict(self) -> Dict[str, Any]:
return dict(self._values)
def __repr__(self) -> str:
return f"KnobValues({self._prefix}: {self._values})"
class Knobs:
"""The declared knob set of one model (env prefix = the bundle's ``<ENV>``)."""
def __init__(self, prefix: str, knobs: Iterable[Knob]):
self.prefix = prefix.strip().upper()
self.knobs: Dict[str, Knob] = {}
for knob in knobs:
if knob.name in self.knobs:
raise ValueError(f"duplicate knob {knob.name}")
self.knobs[knob.name] = knob
def env_name(self, name: str) -> str:
return f"{self.prefix}_{name}"
def defaults(self) -> KnobValues:
return KnobValues(self.prefix, {n: k.default for n, k in self.knobs.items()},
{n: "default" for n in self.knobs})
def read(self, env: Optional[Mapping[str, str]] = None, **overrides: Any) -> KnobValues:
"""Defaults, then ``<PREFIX>_<NAME>`` environment values, then explicit ``overrides`` (e.g. from
``from_pretrained`` keyword arguments). Raises on unparsable values; warns on experiment knobs."""
env = os.environ if env is None else env
values, sources = {}, {}
for name, knob in self.knobs.items():
raw = env.get(self.env_name(name))
if raw is not None and raw.strip() != "":
values[name], sources[name] = knob.parse(raw), "env"
else:
values[name], sources[name] = knob.default, "default"
for name, value in overrides.items():
if name not in self.knobs:
raise KeyError(f"unknown knob {name!r}; have {sorted(self.knobs)}")
knob = self.knobs[name]
values[name], sources[name] = knob.parse(knob.render(value)), "override"
experiments = [n for n, k in self.knobs.items() if k.experiment and values[n] != k.default]
if experiments:
log.warning("%s: experiment knobs set (%s); results are not the published configuration", self.prefix,
", ".join(f"{self.env_name(n)}={self.knobs[n].render(values[n])}" for n in experiments))
return KnobValues(self.prefix, values, sources)
def serve_env(self, values: Optional[Mapping[str, Any]] = None) -> Dict[str, str]:
"""``{"<PREFIX>_<NAME>": "<value>"}`` for pinning in ``tt-model.yaml serve.env`` (defaults if no values)."""
vals = values if values is not None else self.defaults()
return {self.env_name(n): k.render(vals[n]) for n, k in self.knobs.items()}
def doc_table(self) -> str:
"""A Markdown table of the knobs (for SERVING.md / OPT_REPORT.md)."""
rows = ["| variable | default | meaning |", "|---|---|---|"]
for name, knob in self.knobs.items():
extra = f" (one of {', '.join(map(str, knob.choices))})" if knob.choices else ""
tag = " **experiment**" if knob.experiment else ""
rows.append(f"| `{self.env_name(name)}` | `{knob.render(knob.default)}` | {knob.doc}{extra}{tag} |")
return "\n".join(rows)