changh95's picture
tt-model push diffusion-planner-p150 (container)
4d9b003 verified
Raw History Blame Contribute Delete
8.95 kB
# 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()))}