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()))}