File size: 7,474 Bytes
4d9b003 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 | # 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)
|