# 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`; ``_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") # _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 ``_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()))}