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)