File size: 8,954 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 179 180 181 182 183 184 185 186 187 188 189 190 191 192 | # SPDX-License-Identifier: Apache-2.0
"""C05: explicit compute-kernel configs and per-module precision policies.
Silent defaults this module exists to avoid (TT_PLATFORM.md section 0 item 7, PLAN.md section 0.2):
- ``ttnn.matmul`` / ``ttnn.linear`` fall back to **LoFi** when a ``program_config`` or ``core_grid`` is given without
a ``compute_kernel_config``;
- ``ttnn.WormholeComputeKernelConfig()`` built without ``math_fidelity`` carries ``MathFidelity.Invalid``.
So every op that takes a compute config gets one built here, with the fidelity spelled out. The default precision
is HiFi2 + fp32 accumulation, no approximations (the reference bundles' safe start; LoFi failed their gates almost
everywhere, RP section 2.9). A :class:`PrecisionPolicy` maps module names (globs, first match wins) to a
:class:`Precision`; ``<PREFIX>_PRECISION`` overrides rules at build time (the A/B switch), e.g.
``CENTERPOINT_PRECISION="backbone.*=HiFi4+fp32;head.*=LoFi:w=bfp8"``.
Other silent defaults to override by hand (not compute configs): ``ttnn.layer_norm`` epsilon 1e-12, SDPA
``is_causal=True``, ``ttnn.embedding`` PADDED returning the cached pad row, fused HARDSWISH skipped in conv2d.
"""
from __future__ import annotations
import fnmatch
import os
from dataclasses import asdict, dataclass, replace
from typing import Any, Dict, List, Mapping, Optional, Sequence, Tuple, Union
from .tensors import dtype_name
__all__ = ["FIDELITIES", "Precision", "PRESETS", "compute_kernel_config", "PrecisionPolicy"]
FIDELITIES = ("LoFi", "HiFi2", "HiFi3", "HiFi4")
_FIDELITY_KEY = {f.lower(): f for f in FIDELITIES}
def _fidelity(name: str) -> str:
key = _FIDELITY_KEY.get(str(name).strip().lower())
if key is None:
raise ValueError(f"math fidelity {name!r}: expected one of {FIDELITIES}")
return key
def compute_kernel_config(fidelity: str = "HiFi2", *, fp32_acc: bool = True, approx: bool = False,
packer_l1_acc: bool = False, dst_full_sync: bool = False):
"""A ``ttnn.WormholeComputeKernelConfig`` (the same class as ``BlackholeComputeKernelConfig``) with every
field explicit. ``fp32_acc`` = ``fp32_dest_acc_en`` (halves DST capacity: 4 tiles in half-sync)."""
import ttnn
return ttnn.WormholeComputeKernelConfig(math_fidelity=getattr(ttnn.MathFidelity, _fidelity(fidelity)),
math_approx_mode=bool(approx), fp32_dest_acc_en=bool(fp32_acc),
packer_l1_acc=bool(packer_l1_acc), dst_full_sync_en=bool(dst_full_sync))
@dataclass(frozen=True)
class Precision:
"""Fidelity / accumulation / dtype choice of one module (an op or a group of ops)."""
fidelity: str = "HiFi2"
fp32_acc: bool = True
approx: bool = False
packer_l1_acc: bool = False
dst_full_sync: bool = False
weights: str = "bfloat16"
activations: str = "bfloat16"
def __post_init__(self) -> None:
object.__setattr__(self, "fidelity", _fidelity(self.fidelity))
object.__setattr__(self, "weights", dtype_name(self.weights))
object.__setattr__(self, "activations", dtype_name(self.activations))
@classmethod
def parse(cls, spec: Union[str, "Precision"]) -> "Precision":
"""``"HiFi4+fp32"``, ``"LoFi"``, ``"HiFi2+fp32+approx+l1acc:w=bfp8:a=bf16"`` or a preset name
(``accurate`` / ``balanced`` / ``fast``). Flags: ``fp32`` (fp32 accumulation; absent = bf16 DST),
``approx``, ``l1acc``, ``fullsync``; ``w=`` / ``a=`` set weight / activation dtypes."""
if isinstance(spec, Precision):
return spec
text = spec.strip()
if text.lower() in PRESETS:
return PRESETS[text.lower()]
head, *opts = text.split(":")
fid, *flags = [p.strip() for p in head.split("+")]
kw: Dict[str, Any] = {"fidelity": fid, "fp32_acc": False}
for flag in (f.lower() for f in flags):
if flag == "fp32":
kw["fp32_acc"] = True
elif flag == "approx":
kw["approx"] = True
elif flag == "l1acc":
kw["packer_l1_acc"] = True
elif flag == "fullsync":
kw["dst_full_sync"] = True
else:
raise ValueError(f"precision {spec!r}: unknown flag {flag!r}")
for opt in opts:
k, _, v = opt.partition("=")
k = k.strip().lower()
if k in ("w", "weights"):
kw["weights"] = v.strip()
elif k in ("a", "act", "activations"):
kw["activations"] = v.strip()
else:
raise ValueError(f"precision {spec!r}: unknown option {k!r}")
return cls(**kw)
@property
def label(self) -> str:
"""Round-trips through :meth:`parse`."""
flags = "".join(f"+{f}" for f, on in (("fp32", self.fp32_acc), ("approx", self.approx),
("l1acc", self.packer_l1_acc), ("fullsync", self.dst_full_sync)) if on)
return f"{self.fidelity}{flags}:w={self.weights}:a={self.activations}"
def with_(self, **changes: Any) -> "Precision":
return replace(self, **changes)
def compute_kernel_config(self):
return compute_kernel_config(self.fidelity, fp32_acc=self.fp32_acc, approx=self.approx,
packer_l1_acc=self.packer_l1_acc, dst_full_sync=self.dst_full_sync)
def weights_dtype(self):
from .tensors import ttnn_dtype
return ttnn_dtype(self.weights)
def activations_dtype(self):
from .tensors import ttnn_dtype
return ttnn_dtype(self.activations)
def to_dict(self) -> Dict[str, Any]:
return asdict(self)
PRESETS: Dict[str, Precision] = {
"accurate": Precision("HiFi4", fp32_acc=True), # norms, softmax logits, box regression, grid sampling
"balanced": Precision("HiFi2", fp32_acc=True), # default for big matmuls / convs
"fast": Precision("LoFi", fp32_acc=False, weights="bfloat8_b"), # only after the gates pass with it
}
RuleSpec = Union[Mapping[str, Union[str, Precision]], Sequence[Tuple[str, Union[str, Precision]]]]
class PrecisionPolicy:
"""Ordered ``pattern -> Precision`` rules (``fnmatch`` globs on dotted module names, first match wins) plus a
default. Typical use, once at model build::
POLICY = PrecisionPolicy({"backbone.*": "balanced", "head.reg*": "accurate"}, default="balanced")
policy = POLICY.with_env("CENTERPOINT") # <PREFIX>_PRECISION overrides, read once
cfg = policy.compute_kernel_config("backbone.block3.conv2")
w_dtype = policy.resolve("backbone.block3.conv2").weights_dtype()
``resolve`` records each module it answered for, so ``describe()`` shows the policy that reached the ops."""
def __init__(self, rules: Optional[RuleSpec] = None, default: Union[str, Precision] = "balanced"):
items = rules.items() if isinstance(rules, Mapping) else (rules or ())
self.rules: List[Tuple[str, Precision]] = [(str(p), Precision.parse(v)) for p, v in items]
self.default = Precision.parse(default)
self.used: Dict[str, str] = {}
self._configs: Dict[Precision, Any] = {}
def resolve(self, module: str) -> Precision:
for pattern, prec in self.rules:
if fnmatch.fnmatchcase(module, pattern):
self.used[module] = prec.label
return prec
self.used[module] = self.default.label
return self.default
def compute_kernel_config(self, module: str):
prec = self.resolve(module)
if prec not in self._configs:
self._configs[prec] = prec.compute_kernel_config()
return self._configs[prec]
def override(self, spec: str) -> "PrecisionPolicy":
"""A new policy with ``"pattern=precision;pattern=precision"`` rules placed before the existing ones
(``*=...`` effectively replaces the default for unmatched modules)."""
extra = []
for part in filter(None, (p.strip() for p in spec.split(";"))):
pattern, sep, prec = part.partition("=")
if not sep:
raise ValueError(f"precision override {part!r}: expected pattern=precision")
extra.append((pattern.strip(), Precision.parse(prec)))
return PrecisionPolicy(extra + self.rules, self.default)
def with_env(self, prefix: str, env: Optional[Mapping[str, str]] = None) -> "PrecisionPolicy":
"""Apply ``<PREFIX>_PRECISION`` if set (else return ``self``). Call once at build."""
spec = (os.environ if env is None else env).get(f"{prefix}_PRECISION", "").strip()
return self.override(spec) if spec else self
def describe(self) -> Dict[str, Any]:
return {"default": self.default.label, "rules": [[p, prec.label] for p, prec in self.rules],
"used": dict(sorted(self.used.items()))}
|