Download code/tt_diffusion_planner/ttaw/precision.py from changh95/diffusion-planner-p150: direct link, hf CLI and curl.
- Browser
- Download file 8.95 kB
-
https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/ttaw/precision.py
- Command line
-
hf download hf://changh95/diffusion-planner-p150/code/tt_diffusion_planner/ttaw/precision.py
-
curl -L -o precision.py https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/ttaw/precision.py
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)) | |
| 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)) | |
| 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) | |
| 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()))} | |