changh95's picture
tt-model push diffusion-planner-p150 (container)
be62f78 verified
Raw History Blame Contribute Delete
41.8 kB
# SPDX-License-Identifier: Apache-2.0
"""Device building blocks of the planner graph: weights uploaded once with the module's precision, ops that always
carry an explicit compute config (``ttaw.precision``) and LayerNorm epsilon 1e-5 (ttnn's default is 1e-12).
``Build`` resolves a module name against the precision policy (``tt/config.py`` :data:`DEFAULT_PRECISION` +
``DIFFUSION_PLANNER_PRECISION``): ``w=`` is the weight dtype, ``a=`` the module's residual-stream dtype; hidden MLP
activations are bf16 (``HIDDEN``) unless a module says otherwise.
"""
from __future__ import annotations
import fnmatch
from typing import Any, Dict, Optional, Sequence
import numpy as np
from ..reference import config as C
from ..ttaw.precision import Precision, PrecisionPolicy
from ..ttaw.tensors import to_device, ttnn_dtype
from . import config as T
from .params import Lin, Norm
__all__ = ["Build", "Linear", "SplitLinear", "make_linear", "chain2", "LayerNorm", "Const", "ATTN", "policy",
"layer_norm_fp32", "fold_batch"]
ATTN = "bfloat16" # dtype of the Q / K / V projections (SDPA takes bf16)
def policy(env: Optional[Dict[str, str]] = None, spec: Optional[str] = None) -> PrecisionPolicy:
"""The planner's precision policy: :data:`tt.config.DEFAULT_PRECISION`, then ``DIFFUSION_PLANNER_PRECISION`` (or
``spec``) rules on top."""
pol = PrecisionPolicy(dict(T.DEFAULT_PRECISION), default="HiFi4+fp32:w=bf16:a=fp32").with_env(
"DIFFUSION_PLANNER", env)
return pol.override(spec) if spec else pol
class Build:
"""Device + precision policy + graph options shared by the modules while they upload their weights.
``ln_fp32`` / ``hidden_fp32``: module globs (``tt.config.KNOBS`` ``LN_FP32`` / ``HIDDEN_FP32`` by default) whose
LayerNorms use :func:`layer_norm_fp32` / whose hidden MLP activations are fp32."""
def __init__(self, device: Any, pol: Optional[PrecisionPolicy] = None, *, ln_fp32: Optional[Sequence[str]] = None,
hidden_fp32: Optional[Sequence[str]] = None, split: Optional[Sequence[str]] = None,
attn_fp32_acc: Optional[Sequence[str]] = None, attn_matmul: Optional[Sequence[str]] = None):
self.device = device
self.policy = pol or policy()
knobs = T.KNOBS.read()
self.ln_fp32 = tuple(T.globs(knobs.LN_FP32) if ln_fp32 is None else ln_fp32)
self.hidden_fp32 = tuple(T.globs(knobs.HIDDEN_FP32) if hidden_fp32 is None else hidden_fp32)
self.split = tuple(T.globs(knobs.SPLIT_MATMUL) if split is None else split)
self.attn_fp32 = tuple(T.globs(knobs.ATTN_FP32_ACC) if attn_fp32_acc is None else attn_fp32_acc)
self.attn_mm = tuple(T.globs(knobs.ATTN_MATMUL) if attn_matmul is None else attn_matmul)
self.ch2d = bool(knobs.ENC_CH2D)
self.attn_fast = int(knobs.ATTN_FAST)
self.dec_mmcfg = bool(knobs.DEC_MMCFG) and hasattr(device, "compute_with_storage_grid_size")
on_device = hasattr(device, "compute_with_storage_grid_size")
self.ln_kernel = int(knobs.LN_KERNEL) if on_device else 0
self.ln_resid = bool(knobs.LN_RESID)
self.split_kcat = int(knobs.SPLIT_KCAT)
self.attn_smask = bool(knobs.ATTN_SMASK)
self.attn_smsm = int(knobs.ATTN_SMSM)
self.attn_fused = bool(knobs.ATTN_FUSED)
self.kcat_emit = bool(knobs.KCAT_EMIT)
self.ln_tr = bool(knobs.LN_TR)
self.enc_kcat = bool(knobs.ENC_KCAT)
self.lin_act = bool(knobs.LIN_ACT)
self.kcat_act = bool(knobs.KCAT_ACT)
self.ln_split = bool(knobs.LN_SPLIT)
self.ln_sfpu_bcast = bool(knobs.LN_SFPU_BCAST)
self.kcat_l1 = bool(knobs.KCAT_L1)
self.kcat_act_once = bool(knobs.KCAT_ACT_ONCE)
self.attn_l1 = int(knobs.ATTN_L1) if on_device else 0
self.dec_l1 = int(knobs.DEC_L1) if on_device else 0
self.enc_l1 = bool(knobs.ENC_L1) and on_device
self.fus_l1 = bool(knobs.FUS_L1) and on_device
self.compact_enc = int(knobs.COMPACT) >= 2
self.uploaded_bytes = 0
@staticmethod
def _match(module: str, patterns: Sequence[str]) -> bool:
return any(fnmatch.fnmatchcase(module, p) for p in patterns)
def ln_mode(self, module: str) -> str:
return "fp32" if self._match(module, self.ln_fp32) else "device"
def hidden(self, module: str) -> str:
"""dtype of the hidden MLP activations of ``module``: fp32 for the ``HIDDEN_FP32`` globs and for split-matmul
modules (a bf16 hidden would undo the split), else bf16."""
return "float32" if self._match(module, self.hidden_fp32 + self.split) else "bfloat16"
def split_matmul(self, module: str) -> bool:
"""``module``'s fp32 matmuls use :class:`SplitLinear` (~1e-5 instead of the TF32-like ~1e-3)."""
return self._match(module, self.split)
def attn_fp32_acc(self, module: str) -> bool:
return self._match(module, self.attn_fp32)
def attn_matmul(self, module: str) -> bool:
"""``module``'s attention runs as fp32 matmuls + softmax (C20 ``attention_matmul``)."""
return self._match(module, self.attn_mm)
def options(self) -> Dict[str, Any]:
return {"ln_fp32": list(self.ln_fp32), "hidden_fp32": list(self.hidden_fp32), "split": list(self.split),
"attn_fp32_acc": list(self.attn_fp32), "attn_matmul": list(self.attn_mm), "enc_ch2d": self.ch2d,
"attn_fast": self.attn_fast,
"dec_mmcfg": self.dec_mmcfg, "ln_kernel": self.ln_kernel, "ln_resid": self.ln_resid,
"split_kcat": self.split_kcat, "attn_smask": self.attn_smask, "attn_smsm": self.attn_smsm, "attn_fused": self.attn_fused, "kcat_emit": self.kcat_emit, "ln_tr": self.ln_tr,
"enc_kcat": self.enc_kcat, "lin_act": self.lin_act,
"kcat_act": self.kcat_act, "ln_split": self.ln_split,
"ln_sfpu_bcast": self.ln_sfpu_bcast, "kcat_l1": self.kcat_l1,
"kcat_act_once": self.kcat_act_once, "attn_l1": self.attn_l1,
"dec_l1": self.dec_l1, "enc_l1": self.enc_l1,
"fus_l1": self.fus_l1}
def prec(self, module: str) -> Precision:
return self.policy.resolve(module)
def cfg(self, module: str):
return self.policy.compute_kernel_config(module)
def stream(self, module: str) -> str:
"""The residual-stream dtype of ``module``."""
return self.prec(module).activations
def attn_mem(self):
"""``ATTN_L1``: the memory config of the fused attention's inputs (Q / K / V projections, hoisted cross K / V,
masks): L1 interleaved, else None (DRAM)."""
if not self.attn_l1:
return None
import ttnn
return ttnn.L1_MEMORY_CONFIG
def dec_mem(self):
"""``DEC_L1``: the memory config of the decoder blocks' intermediates (stream, LN operands, linear outputs,
attention outputs): L1 interleaved, else None (DRAM)."""
if not self.dec_l1:
return None
import ttnn
return ttnn.L1_MEMORY_CONFIG
def enc_mem(self):
"""``ENC_L1``: the memory config of the mixer blocks' intermediates (LN outputs and stream, token / channel
MLP outputs): L1 interleaved, else None (DRAM)."""
if not self.enc_l1:
return None
import ttnn
return ttnn.L1_MEMORY_CONFIG
def upload(self, array: np.ndarray, dtype: str, *, shape4: bool = True):
"""Host float array -> DRAM TILE device tensor (rank padded to 4 with leading 1s)."""
a = np.ascontiguousarray(np.asarray(array, np.float32))
if shape4 and a.ndim < 4:
a = a.reshape((1,) * (4 - a.ndim) + a.shape)
self.uploaded_bytes += a.size * (4 if dtype in ("float32", "fp32") else 2)
return to_device(a, self.device, dtype)
# OPT round 1 item 2a: explicit 2-D multicast program configs for the decoder's 352-row matmuls, the fastest
# bit-identical candidate per (K, N, split pass) of the device sweep (logs/diffusion-planner/opt_r1/mm_sweep.json;
# "hh" = x_hi @ w_hi, "hl" = x_hi @ w_lo, "lh" = x_lo @ w_hi + bias; a plain Linear uses "hh").
# value: (transpose_mcast, per_core_M, per_core_N, in0_block_w, out_subblock_w); every candidate was bit-identical to
# the auto config (fp32 DEST accumulation over K either way).
DEC_ROWS = 352
DEC_MM_CONFIGS = {
(324, 512, "hh"): (True, 1, 3, 11, 1), (324, 512, "hl"): (True, 2, 2, 11, 2), (324, 512, "lh"): (True, 1, 3, 11, 1),
(512, 256, "hh"): (False, 2, 1, 8, 1), (512, 256, "hl"): (False, 2, 1, 8, 1), (512, 256, "lh"): (False, 2, 1, 8, 1),
(256, 768, "hh"): (False, 2, 3, 8, 1), (256, 768, "hl"): (False, 2, 2, 8, 2), (256, 768, "lh"): (True, 1, 3, 8, 1),
(256, 256, "hh"): (False, 2, 2, 8, 2), (256, 256, "hl"): (False, 2, 2, 8, 2), (256, 256, "lh"): (False, 2, 1, 8, 1),
(256, 1024, "hh"): (False, 2, 3, 8, 1), (256, 1024, "hl"): (False, 2, 3, 8, 1),
(256, 1024, "lh"): (False, 2, 3, 8, 1),
(1024, 256, "hh"): (True, 1, 1, 8, 1), (1024, 256, "hl"): (True, 1, 1, 8, 1), (1024, 256, "lh"): (False, 2, 1, 8, 1),
(1024, 324, "hh"): (False, 2, 1, 8, 1), (1024, 324, "hl"): (False, 2, 1, 8, 1),
(1024, 324, "lh"): (False, 2, 1, 8, 1),
}
def compact_rows(rows: int) -> bool:
"""``COMPACT`` (OPT round 5): a decoder row count other than :data:`DEC_ROWS` that the derived configs serve (a
tile-aligned agent bucket below the full capacity)."""
return rows % 32 == 0 and 32 <= rows < DEC_ROWS
def fit_pcm(tr: bool, pcm: int, rows: int, grid) -> int:
"""``per_core_M`` of a 2-D multicast config for ``rows`` rows: the swept value at :data:`DEC_ROWS`, else the
smallest value whose M blocks fit the grid axis that carries them (the in0 block width, ``per_core_N`` and the
subblocks stay as swept: the K accumulation order of every output element is unchanged)."""
if rows == DEC_ROWS:
return pcm
mt = -(-rows // 32)
lim = grid.x if tr else grid.y
p = 1
while -(-mt // p) > lim:
p += 1
return p
def mcast_fits(tr: bool, pcm: int, pcn: int, rows: int, n: int, grid) -> bool:
"""True when a 2-D multicast config's output blocks fit the grid: the swept configs assume the 12x10 ETH-dispatch
grid (N / per_core_N up to 12 blocks on x); on a smaller grid (WORKER dispatch, 11x10) the callers fall back to
the auto config instead of failing at capture (never hard-code the grid: CLAUDE.md)."""
mb, nb = -(-(-(-rows // 32)) // pcm), -(-(-(-n // 32)) // pcn)
gx, gy = (mb, nb) if tr else (nb, mb)
return gx <= grid.x and gy <= grid.y
class DecConfigs:
"""``DEC_MMCFG`` configs of one decoder linear: ``get(rows)`` -> ``{"hh" | "hl" | "lh": program config}`` for
:data:`DEC_ROWS` or (``COMPACT``) a smaller agent bucket, else None."""
def __init__(self, build: "Build", k: int, n: int):
self.build, self.k, self.n, self._cache = build, k, n, {}
def get(self, rows: int) -> Optional[Dict[str, Any]]:
rows = int(rows)
if rows != DEC_ROWS and not compact_rows(rows):
return None
if rows not in self._cache:
import ttnn
grid = self.build.device.compute_with_storage_grid_size()
out = {}
for ps in ("hh", "hl", "lh"):
tr, pcm, pcn, kb, sw = DEC_MM_CONFIGS[(self.k, self.n, ps)]
if not mcast_fits(tr, fit_pcm(tr, pcm, rows, grid), pcn, rows, self.n, grid):
out = None # a grid smaller than the swept one: the auto configs
break
out[ps] = ttnn.MatmulMultiCoreReuseMultiCastProgramConfig(
compute_with_storage_grid_size=grid, in0_block_w=kb, out_subblock_h=1, out_subblock_w=sw,
per_core_M=fit_pcm(tr, pcm, rows, grid), per_core_N=pcn, transpose_mcast=tr,
fused_activation=None, fuse_batch=True)
self._cache[rows] = out
return self._cache[rows]
def dec_configs(build: "Build", module: str, shape) -> Optional[DecConfigs]:
"""The :class:`DecConfigs` of a decoder linear of weight ``shape`` (``DEC_MMCFG``), else None."""
if not (build.dec_mmcfg and module.startswith("dec.")):
return None
k, n = int(shape[0]), int(shape[1])
if (k, n, "hh") not in DEC_MM_CONFIGS:
return None
import ttnn
if not hasattr(ttnn, "MatmulMultiCoreReuseMultiCastProgramConfig"):
return None # the host fake ttnn
return DecConfigs(build, k, n)
# OPT round 2 item 2 (SPLIT_KCAT): X' row tiles 3 Kt + 1 rounded up to a multiple of kcat_pad(K) (the K block of
# the program config must divide them), and the fastest 2-D multicast config per (rows, K' tiles, N) of the device
# sweep (logs/diffusion-planner/opt_r2/kcat_sweep.json): value (transpose_mcast, per_core_M, per_core_N, in0_block_w,
# out_subblock_w); shapes not listed use the auto config.
def kcat_pad(k: int) -> int:
"""Sweep: K = 1024 (3 Kt + 1 = 97, prime) pads to 104 = 8 x 13; the other K keep 3 Kt + 1 (25 = 5 x 5, 49)."""
return 8 if k >= 1024 else 1
KCAT_CONFIGS: Dict[tuple, tuple] = {
(352, 49, 256): (False, 2, 1, 7, 1), (352, 25, 768): (False, 2, 2, 5, 2), (352, 25, 256): (False, 2, 1, 5, 1),
(352, 25, 1024): (False, 2, 3, 5, 1), (352, 104, 256): (False, 2, 1, 13, 1),
(352, 104, 324): (False, 2, 1, 13, 1), (576, 25, 256): (False, 2, 1, 5, 1),
}
def kcat_config(build: "Build", ktp: int, n: int, rows: int):
"""The program config of the K-concatenated matmul ``[rows, 32 ktp] @ [32 ktp, n]``: the swept entry, or
(``COMPACT``) the :data:`DEC_ROWS` entry with ``per_core_M`` fitted to a smaller agent bucket; None = auto."""
import ttnn
if not hasattr(ttnn, "MatmulMultiCoreReuseMultiCastProgramConfig") or not hasattr(
build.device, "compute_with_storage_grid_size"):
return None
grid = build.device.compute_with_storage_grid_size()
ent = KCAT_CONFIGS.get((rows, ktp, n))
if ent is None and compact_rows(rows):
ent = KCAT_CONFIGS.get((DEC_ROWS, ktp, n))
if ent is None:
return None
tr, pcm, pcn, kb, sw = ent
pcm = fit_pcm(tr, pcm, rows, grid) if (rows, ktp, n) not in KCAT_CONFIGS else pcm
if not mcast_fits(tr, pcm, pcn, rows, n, grid):
return None # a grid smaller than the swept one: the auto config
return ttnn.MatmulMultiCoreReuseMultiCastProgramConfig(
compute_with_storage_grid_size=grid, in0_block_w=kb, out_subblock_h=1, out_subblock_w=sw,
per_core_M=pcm, per_core_N=pcn,
transpose_mcast=tr, fused_activation=None, fuse_batch=True)
# OPT round 4 item 1 (KCAT_L1): the K = 1024 split operands X' [352, 32 Ktp] written L1 block-sharded by the operand
# build and read in place by a 2-D multicast matmul with in0 sharded (no 4.7 MB DRAM round trip). The K slices of
# X' are the matmul's N blocks, so Ktp is padded to a multiple of them (zero tiles, exact). Value per (rows, K, N):
# (transpose_mcast, per_core_M, per_core_N, pad, in0_block_w), the fastest bit-identical candidate of the device
# check (logs/diffusion-planner/opt_r4/kl1_real.log: operand + matmul 62.6 -> 41.3 us for N = 256, 62.3 -> 41.8 us
# for N = 324).
KCAT_L1_CONFIGS: Dict[tuple, tuple] = {
(352, 1024, 256): (False, 2, 1, 8, 13),
(352, 1024, 324): (False, 2, 1, 11, 9),
}
def kcat_l1_layout(build: "Build", rows: int, k: int, n: int, ktp: int):
"""``(memory_config, program_config)`` of the L1-sharded operand path for ``[rows, 32 ktp] @ [32 ktp, n]``."""
import ttnn
tr, pcm, pcn, _pad, kb = KCAT_L1_CONFIGS[(DEC_ROWS, k, n)]
grid = build.device.compute_with_storage_grid_size()
pcm = fit_pcm(tr, pcm, rows, grid) # COMPACT: a smaller agent bucket (the K blocking unchanged)
mt, nt = -(-rows // 32), -(-n // 32)
nb, mb = -(-nt // pcn), -(-mt // pcm)
assert ktp % nb == 0 and (ktp // nb) % kb == 0, (ktp, nb, kb)
gx, gy = (mb, nb) if tr else (nb, mb)
if gx > grid.x or gy > grid.y:
return None # a grid smaller than the swept one: the DRAM operand path
crs = ttnn.CoreRangeSet({ttnn.CoreRange(ttnn.CoreCoord(0, 0), ttnn.CoreCoord(gx - 1, gy - 1))})
spec = ttnn.ShardSpec(crs, [32 * pcm, 32 * (ktp // nb)],
ttnn.ShardOrientation.COL_MAJOR if tr else ttnn.ShardOrientation.ROW_MAJOR)
mem = ttnn.MemoryConfig(ttnn.TensorMemoryLayout.BLOCK_SHARDED, ttnn.BufferType.L1, spec)
pc = ttnn.MatmulMultiCoreReuseMultiCastProgramConfig(
compute_with_storage_grid_size=grid, in0_block_w=kb, out_subblock_h=1,
out_subblock_w=max(d for d in (1, 2, 4) if pcn % d == 0), per_core_M=pcm, per_core_N=pcn,
transpose_mcast=tr, fused_activation=None, fuse_batch=True)
return mem, pc
def _activate(y, act):
"""The stock activation programs of the linears (``None`` / ``"gelu"`` / ``"gelu_tanh"``)."""
import ttnn
if act == "gelu":
return ttnn.gelu(y, fast_and_approximate_mode=False)
if act == "gelu_tanh":
return ttnn.gelu(y, variant=ttnn.GeluVariant.Tanh)
return y
def _identity(t):
return t
def fold_batch(x, n_out: int, enabled: bool):
"""Run a linear on a batched activation ``[B0, B1, T, K] @ [K, N]`` as one 2-D matmul over all ``B0*B1*T`` rows
(``ENC_CH2D``, OPT round 1 item 1). The stock auto-config tiles M per batch element (``fuse_batch=False``), which
puts the encoder's ``[1, E, T, C]`` channel MLPs on 4-8 cores (970 us at E = 320). Returns ``(x2, back, pc)``:
- ``T % 32 == 0`` (the MixerBlock channel MLPs, T = 64): the free TILE view ``[1, 1, B*T, K]`` (auto config,
2-D multicast over the grid), ``back`` reshapes the output to ``[B0, B1, T, N]``;
- otherwise (the pre-projections, T = 6 / 20 / 40 padded per element): an explicit 1-D in1-multicast program config
with ``fuse_batch=True`` (rows of tiles split over the grid, the full N per core, ``in0_block_w = Kt``);
``back`` is the identity.
The products and their fp32 DEST accumulation over K are the same; only the core assignment changes."""
import ttnn
shape = tuple(x.shape)
if not enabled or len(shape) != 4 or shape[0] * shape[1] == 1:
return x, _identity, None
b0, b1, t, k = shape
if t % 32 == 0:
x2 = ttnn.reshape(x, (1, 1, b0 * b1 * t, k))
return x2, (lambda y: ttnn.reshape(y, (b0, b1, t, y.shape[-1]))), None
cfg_cls = getattr(ttnn, "MatmulMultiCoreReuseMultiCast1DProgramConfig", None)
dev = x.device() if callable(getattr(x, "device", None)) else None
if cfg_cls is None or dev is None or not hasattr(dev, "compute_with_storage_grid_size"):
return x, _identity, None # the host fake: numerics are the same
grid = dev.compute_with_storage_grid_size()
cores = grid.x * grid.y
m_tiles = b0 * b1 * (-(-t // 32))
kt, nt = -(-k // 32), -(-n_out // 32)
per_core_m = -(-m_tiles // cores)
sub_w = max(d for d in (1, 2, 4) if nt % d == 0) # fp32 DEST: subblock h * w <= 4
pc = cfg_cls(compute_with_storage_grid_size=grid, in0_block_w=kt, out_subblock_h=1, out_subblock_w=sub_w,
per_core_M=per_core_m, per_core_N=nt, fuse_batch=True, fused_activation=None, mcast_in0=False)
return x, _identity, pc
def tall_config(x, wshape, fused_activation=None):
"""The stock auto config of a tall 2-D matmul ``[1, 1, M, K] @ [K, N]`` (M / 32 >= the grid, K, N <= 128): the
1-D in1-multicast config the profile shows for the mixer linears (``MatmulMultiCoreReuseMultiCast1DProgramConfig``,
``per_core_M = ceil(Mt / cores)``, ``in0_block_w = min(Kt, 2)``, the full N per core, ``mcast_in0 = False``), so a
fused activation can be added without changing the matmul (``LIN_ACT``). None for any other shape."""
import ttnn
shp = tuple(x.shape)
dev = x.device() if callable(getattr(x, "device", None)) else None
cls = getattr(ttnn, "MatmulMultiCoreReuseMultiCast1DProgramConfig", None)
if cls is None or dev is None or not hasattr(dev, "compute_with_storage_grid_size") or len(shp) != 4:
return None
if shp[0] * shp[1] != 1 or shp[-2] % 32 or wshape[0] % 32 or wshape[1] % 32:
return None
grid = dev.compute_with_storage_grid_size()
cores = grid.x * grid.y
mt, kt, nt = shp[-2] // 32, wshape[0] // 32, wshape[1] // 32
if mt < cores or nt > 4 or kt > 4:
return None
pcm = -(-mt // cores)
sw = max(d for d in (1, 2, 4) if nt % d == 0 and d <= 4)
sh = max(d for d in (1, 2, 4) if pcm % d == 0 and d * sw <= 4)
return cls(compute_with_storage_grid_size=grid, in0_block_w=min(kt, 2), out_subblock_h=sh, out_subblock_w=sw,
out_block_h=pcm, out_block_w=nt, per_core_M=pcm, per_core_N=nt, fuse_batch=False,
fused_activation=fused_activation, mcast_in0=False)
class Linear:
"""``y = x @ w (+ b)`` with fused ``activation`` ("gelu" = exact erf, "gelu_tanh"); output dtype ``out``."""
def __init__(self, build: Build, lin: Lin, module: str, *, out: str, activation: Optional[str] = None,
w_dtype: Optional[str] = None):
p = build.prec(module)
wd = w_dtype or p.weights
self.module, self.activation, self.out = module, activation, ttnn_dtype(out)
self.w = build.upload(lin.w, wd)
self.b = None if lin.b is None else build.upload(np.asarray(lin.b).reshape(1, -1), wd)
self.cfg = build.cfg(module)
self.shape = tuple(lin.w.shape)
self.fold = build.ch2d and module.startswith("enc.")
self.dec_pc = dec_configs(build, module, self.shape)
self.act_fuse = build.lin_act and module.startswith("enc.mixer.")
self.out_mem = None # ATTN_L1: output memory config (None = DRAM)
def __call__(self, x):
import ttnn
x, back, pc = fold_batch(x, self.shape[1], self.fold)
dpc = None if self.dec_pc is None else self.dec_pc.get(tuple(x.shape)[-2])
if dpc is not None and x.dtype == ttnn.bfloat16:
pc = dpc["hh"]
act = self.activation
if self.act_fuse and act == "gelu" and pc is None:
pc = tall_config(x, self.shape, ttnn.UnaryWithParam(ttnn.UnaryOpType.GELU, 0.0))
if pc is not None: # LIN_ACT: GELU in the matmul epilogue (fp32 DEST, before rounding)
act = None
kw = {} if pc is None else {"program_config": pc}
if self.out_mem is not None:
kw["memory_config"] = self.out_mem
return back(ttnn.linear(x, self.w, bias=self.b, activation=act, dtype=self.out,
compute_kernel_config=self.cfg, **kw))
def layer_norm_fp32(x, gamma=None, beta=None, *, eps: float = C.LN_EPS):
"""LayerNorm over the last dim as fp32 element-wise / reduction ops: ``(x - mean) * rsqrt(var + eps) * gamma +
beta`` (7-8 programs). The fused ``ttnn.layer_norm`` is ~2.5e-3 relative even for fp32 input (probe P10), which
swamps the small per-entity deviations of the offset-dominated mixer rows."""
import ttnn
if x.dtype != ttnn.float32:
x = ttnn.typecast(x, ttnn.float32)
xc = ttnn.subtract(x, ttnn.mean(x, dim=-1, keepdim=True))
var = ttnn.mean(ttnn.multiply(xc, xc), dim=-1, keepdim=True)
y = ttnn.multiply(xc, ttnn.rsqrt(ttnn.add(var, float(eps)), fast_and_approximate_mode=False))
if gamma is not None:
y = ttnn.multiply(y, gamma)
if beta is not None:
y = ttnn.add(y, beta)
return y
class KcatOperand:
"""A split operand ``[x_hi | x_hi | x_lo | 1 | 0..]`` written by its producer (``KCAT_EMIT``: the split
LayerNorm, the fused attention) for a K-concatenated :class:`SplitLinear` (instead of x + a ``kcat_operand``
program). ``t``: the device tensor ``[..., M, 32 * ktp]``."""
def __init__(self, t, ktp: int):
self.t, self.ktp = t, int(ktp)
@property
def shape(self):
return self.t.shape
def operand_ktp(build: "Build", consumer) -> int:
"""Row tiles of the split operand ``consumer`` takes when its producer may emit it (``KCAT_EMIT``), else 0."""
if not getattr(build, "kcat_emit", False) or not isinstance(consumer, SplitLinear):
return 0
return consumer.operand_ktp()
class SplitLinear:
"""``y = x @ w (+ b)`` to ~1e-5 relative from three device matmuls on split operands (the "bf16x3" scheme):
``x_hi @ w_hi + x_hi @ w_lo + x_lo @ w_hi`` with ``*_hi`` the bf16 roundings (bf16 x bf16 products are exact
under fp32 accumulation) and ``*_lo = * - *_hi`` the fp32 remainders (exact by Sterbenz), whose TF32-like operand
truncation (probe P12) then costs ~2^-18 of the full value; ``x_lo @ w_lo`` (~2^-16) is dropped. Output fp32,
activation applied to the sum. 8-9 programs instead of 1: for the small pre-projection island only."""
def __init__(self, build: Build, lin: Lin, module: str, *, activation: Optional[str] = None,
out: str = "float32"):
from ..ttaw.tensors import round_to_bf16
if activation not in (None, "gelu", "gelu_tanh"):
raise ValueError(f"SplitLinear: unsupported activation {activation!r}")
w = np.asarray(lin.w, np.float32)
w_hi = round_to_bf16(w)
self.w_hi = build.upload(w_hi, "bfloat16")
self.w_lo = build.upload((w - w_hi).astype(np.float32), "float32")
self.b = None if lin.b is None else build.upload(np.asarray(lin.b, np.float32).reshape(1, -1), "float32")
self.cfg = build.cfg(module)
self.activation = activation
self.out = ttnn_dtype(out)
self.shape = tuple(w.shape)
self.fold = build.ch2d and module.startswith("enc.")
self.dec_pc = dec_configs(build, module, self.shape)
# OPT round 2 item 2 (SPLIT_KCAT): the three passes as ONE fp32 matmul over the concatenated K axis,
# [x_hi | x_hi | x_lo | ones] @ [w_hi; w_lo; w_hi; (b_hi, b_lo, 0...)]: the same exact / TF32-truncated
# products, accumulated in one fp32 DEST sum (the bias exact as two K rows instead of the packer epilogue).
# Decoder modules with a tile-aligned K only (the 352-row linears); others keep the 3-pass form.
# OPT round 3 item 2 (ENC_KCAT): the encoder's split linears (pre-projections, the ego / neighbour island)
# in the same form, any K: each weight block padded to whole tiles with zero rows (the operand keeps the
# input's tile padding, which meets those zero rows as in the 3-pass form).
self.kcat = None
self.out_mem = None # ATTN_L1: output memory config of the K-concatenated matmul
k = self.shape[0]
enc_kcat = build.enc_kcat and (module.startswith("enc.pre.") or module.startswith("enc.island."))
if build.split_kcat and ((module.startswith("dec.") and k % 32 == 0) or enc_kcat):
rows = np.zeros((32, self.shape[1]), np.float32)
if lin.b is not None:
b = np.asarray(lin.b, np.float32).reshape(-1)
b_hi = round_to_bf16(b)
rows[0], rows[1] = b_hi, b - b_hi
from .kcat_kernel import kcat_tiles
kp = -(-k // 32) * 32
self.kcat_kp = kp
def padk(a):
return np.concatenate([a, np.zeros((kp - k, self.shape[1]), np.float32)], axis=0)
l1_key = (DEC_ROWS, kp, self.shape[1])
use_l1 = (build.kcat_l1 and build.split_kcat == 2 and module.startswith("dec.")
and l1_key in KCAT_L1_CONFIGS and hasattr(build.device, "compute_with_storage_grid_size"))
self.kcat_padv = KCAT_L1_CONFIGS[l1_key][3] if use_l1 else kcat_pad(kp)
self.kcat_ktp = kcat_tiles(kp, self.kcat_padv)
# KCAT_L1: (memory config, program config) of the L1-sharded operand path for DEC_ROWS rows
self.kcat_l1 = {} if use_l1 else None # {rows: layout}, filled per row count on first use
pad_rows = np.zeros((32 * self.kcat_ktp - 3 * kp - 32, self.shape[1]), np.float32)
wk = np.concatenate([padk(w_hi), padk((w - w_hi).astype(np.float32)), padk(w_hi), rows, pad_rows],
axis=0)
self.kcat = build.upload(wk, "float32")
self.kcat_pc = {} # {rows: program config or None}
self.kcat_mode = build.split_kcat
self._act_once = build.kcat_act_once
self._ones = {}
self._build = build
def _ones_tile(self, x):
"""``[.., M, 32 (Kt' - 3 Kt)]`` fp32 with columns 0 and 1 = 1 (the bias rows of the concatenated weight), the
rest 0 (the K padding)."""
key = tuple(tuple(x.shape)[:-1])
if key not in self._ones:
o = np.zeros(key + (32 * self.kcat_ktp - 3 * self.kcat_kp,), np.float32)
o[..., :2] = 1.0
self._ones[key] = self._build.upload(o, "float32")
return self._ones[key]
def _call_kcat(self, x, pre_act=None):
import ttnn
f32 = ttnn.float32
if isinstance(x, KcatOperand): # KCAT_EMIT: the producer wrote the operand
assert pre_act is None and x.ktp == self.kcat_ktp
return self._kcat_mm(x.t, x.t)
if x.dtype != f32:
x = ttnn.typecast(x, f32)
from .kcat_kernel import kcat_operand, supported
if self.kcat_mode == 2 and supported(x): # else (the host fake ttnn) the stock build: the same values
rows = tuple(x.shape)[-2]
if self.kcat_l1 is not None and (rows == DEC_ROWS or compact_rows(rows)) and not self.fold:
if rows not in self.kcat_l1:
self.kcat_l1[rows] = kcat_l1_layout(self._build, rows, self.kcat_kp, self.shape[1],
self.kcat_ktp)
if self.kcat_l1 is not None and self.kcat_l1.get(rows) is not None and not self.fold:
mem, pc = self.kcat_l1[rows] # KCAT_L1: X' in L1, read in place by the sharded-in0 matmul
xk = kcat_operand(x, self.kcat_padv, pre_act, memory_config=mem, act_once=self._act_once)
y = ttnn.matmul(xk, self.kcat, dtype=f32, compute_kernel_config=self.cfg, program_config=pc,
memory_config=self.out_mem or ttnn.DRAM_MEMORY_CONFIG)
xk.deallocate()
return y
return self._kcat_mm(kcat_operand(x, self.kcat_padv, pre_act, act_once=self._act_once), x)
x = _activate(x, pre_act)
x_hi = ttnn.typecast(ttnn.typecast(x, ttnn.bfloat16), f32)
x_lo = ttnn.subtract(x, x_hi)
xk = ttnn.concat([x_hi, x_hi, x_lo, self._ones_tile(x)], dim=-1)
return self._kcat_mm(xk, x)
def _kcat_mm(self, xk, x):
import ttnn
if self.fold: # encoder [1, E, T, K'] operands: one 2-D matmul (ENC_CH2D)
xk, back, pc = fold_batch(xk, self.shape[1], True)
kw = {} if pc is None else {"program_config": pc}
if self.out_mem is not None:
kw["memory_config"] = self.out_mem
return back(ttnn.matmul(xk, self.kcat, dtype=ttnn.float32, compute_kernel_config=self.cfg, **kw))
kw = self._kcat_kw(x)
if self.out_mem is not None:
kw["memory_config"] = self.out_mem
return ttnn.matmul(xk, self.kcat, dtype=ttnn.float32, compute_kernel_config=self.cfg, **kw)
def _kcat_ok(self, x) -> bool:
"""The K-concatenated form applies: tile-aligned K (any operand build), or the generic_op operand build
(which keeps the input's tile padding for an unaligned K)."""
if self.kcat is None:
return False
if isinstance(x, KcatOperand) or self.kcat_kp == self.shape[0]:
return True
from .kcat_kernel import supported
return self.kcat_mode == 2 and supported(x)
def _kcat_kw(self, x):
rows = int(tuple(x.shape)[-2])
if rows not in self.kcat_pc:
self.kcat_pc[rows] = kcat_config(self._build, self.kcat_ktp, self.shape[1], rows)
pc = self.kcat_pc[rows]
return {} if pc is None else {"program_config": pc}
def operand_ktp(self) -> int:
"""Row tiles of the operand a producer may write for this linear (``KCAT_EMIT``): the generic_op form with a
tile-aligned K and no K-block padding of a different layout; else 0."""
if self.kcat is None or self.kcat_mode != 2 or self.kcat_kp != self.shape[0]:
return 0
return self.kcat_ktp
def fuses_input_act(self) -> bool:
"""True when this linear can take its input before the previous linear's activation (``KCAT_ACT``)."""
return self.kcat is not None and self.kcat_mode == 2 and self.kcat_kp == self.shape[0]
def __call__(self, x, *, pre_act=None, defer_act: bool = False):
"""``pre_act``: apply ``"gelu"`` / ``"gelu_tanh"`` to ``x`` first (fused into the operand build when
:meth:`fuses_input_act`); ``defer_act``: return the pre-activation fp32 output (the next linear applies
:attr:`activation` through ``pre_act``)."""
import ttnn
f32 = ttnn.float32
if self._kcat_ok(x):
y = self._call_kcat(x, pre_act)
if defer_act:
return y
y = _activate(y, self.activation)
if self.out != f32:
y = ttnn.typecast(y, self.out)
return y
x = _activate(x, pre_act)
if defer_act:
assert self.out == f32
act, self.activation = self.activation, None
try:
return self(x)
finally:
self.activation = act
x, back, pc = fold_batch(x, self.shape[1], self.fold)
pcs = {"hh": pc, "hl": pc, "lh": pc}
dpc = None if self.dec_pc is None else self.dec_pc.get(tuple(x.shape)[-2])
if dpc is not None:
pcs = dpc
def kw(ps):
p = pcs[ps]
return {"compute_kernel_config": self.cfg, **({} if p is None else {"program_config": p})}
if x.dtype != f32: # a bf16 input is its own hi part: two matmuls
y = ttnn.matmul(x, self.w_hi, dtype=f32, **kw("hh"))
y = ttnn.add(y, ttnn.linear(x, self.w_lo, bias=self.b, dtype=f32, **kw("hl")))
else:
x_hi = ttnn.typecast(x, ttnn.bfloat16)
x_lo = ttnn.subtract(x, ttnn.typecast(x_hi, f32))
y = ttnn.matmul(x_hi, self.w_hi, dtype=f32, **kw("hh"))
y = ttnn.add(y, ttnn.matmul(x_hi, self.w_lo, dtype=f32, **kw("hl")))
y = ttnn.add(y, ttnn.linear(x_lo, self.w_hi, bias=self.b, dtype=f32, **kw("lh")))
if self.activation == "gelu":
y = ttnn.gelu(y, fast_and_approximate_mode=False)
elif self.activation == "gelu_tanh":
y = ttnn.gelu(y, variant=ttnn.GeluVariant.Tanh)
if self.out != f32:
y = ttnn.typecast(y, self.out)
return back(y)
def chain2(first, second, x, enabled: bool):
"""``second(first(x))``; with ``enabled`` (``KCAT_ACT``) and a K-concatenated ``second``, ``first``'s activation
runs inside ``second``'s operand build instead of as its own program (same LLK, same values)."""
import ttnn
if (enabled and isinstance(first, SplitLinear) and isinstance(second, SplitLinear) and first.activation
and first.out == ttnn.float32 and second.fuses_input_act()):
return second(first(x, defer_act=True), pre_act=first.activation)
return second(first(x))
def make_linear(build: Build, lin: Lin, module: str, *, out: str, activation: Optional[str] = None):
""":class:`SplitLinear` when ``module`` is in the build's ``SPLIT_MATMUL`` globs, else :class:`Linear`."""
if build.split_matmul(module):
return SplitLinear(build, lin, module, activation=activation, out=out)
return Linear(build, lin, module, out=out, activation=activation)
class LayerNorm:
"""LayerNorm with ``epsilon=1e-5`` and fp32 affine rows (``[1, 1, 1, W]`` TILE): ``ttnn.layer_norm``, or
:func:`layer_norm_fp32` when the module is in the build's ``ln_fp32`` globs. ``gamma`` / ``beta`` may be given
as device tensors (the per-step folded adaLN rows of the decoder)."""
def __init__(self, build: Build, nrm: Optional[Norm], module: str):
self.cfg = build.cfg(module)
self.mode = build.ln_mode(module)
self.fused = build.ln_kernel
self.resid = build.ln_resid
self.split = build.ln_split
self.sbc = build.ln_sfpu_bcast
self.tr = build.ln_tr
self.mem = None # DEC_L1: memory config of the split-row LN's outputs
self.g = self.b = None
if nrm is not None:
self.g = build.upload(np.asarray(nrm.gamma).reshape(1, -1), "float32")
self.b = build.upload(np.asarray(nrm.beta).reshape(1, -1), "float32")
def __call__(self, x, gamma=None, beta=None, *, kcat: int = 0):
"""``kcat`` (``KCAT_EMIT``, the consumer's :func:`operand_ktp`): return the consumer's split operand as a
:class:`KcatOperand` when the split-row kernel runs (else y as usual)."""
import ttnn
g = self.g if gamma is None else gamma
b = self.b if beta is None else beta
if self.mode == "fp32":
if self.fused:
from .ln_kernel import layer_norm_fp32_fused, supported
if self.split:
from .ln_kernel import layer_norm_fp32_split, split_supported
if split_supported(x):
y = layer_norm_fp32_split(x, g, b, eps=C.LN_EPS, sfpu_bcast=self.sbc, kcat_ktp=kcat,
memory_config=self.mem)
return KcatOperand(y, kcat) if kcat else y
if supported(x):
return layer_norm_fp32_fused(x, g, b, eps=C.LN_EPS, lean=self.fused == 2, sfpu_bcast=self.sbc,
memory_config=self.mem)
return layer_norm_fp32(x, g, b)
kw = {} if self.mem is None else {"memory_config": self.mem}
return ttnn.layer_norm(x, epsilon=C.LN_EPS, weight=g, bias=b, compute_kernel_config=self.cfg, **kw)
def can_transpose(self, x) -> bool:
"""``LN_TR``: the fused one-core-per-row kernel runs for ``x`` (not the split-row form), so it can read the
residual / write the output transposed per entity."""
if not (self.tr and self.mode == "fp32" and self.fused and self.resid):
return False
from .ln_kernel import split_supported, supported
return supported(x) and not (self.split and split_supported(x))
def transposed(self, x, *, res=None):
"""``LN_TR``: ``LN(x)`` (or, with ``res``, ``h = x + res`` and ``LN(h)``) with the output written as
``[.., E, W, T]`` (the per-entity transpose); returns ``y^T`` or ``(h, y^T)``."""
from .ln_kernel import layer_norm_fp32_fused
return layer_norm_fp32_fused(x, self.g, self.b, eps=C.LN_EPS, residual=res, lean=self.fused == 2,
sfpu_bcast=self.sbc, out_t=True, memory_config=self.mem)
def residual_t(self, x, res_t_tensor):
"""``LN_TR``: ``h = x + res^T`` (``res`` ``[.., E, W, T]``) and ``LN(h)`` -> ``(h, y)``."""
from .ln_kernel import layer_norm_fp32_fused
return layer_norm_fp32_fused(x, self.g, self.b, eps=C.LN_EPS, residual=res_t_tensor, lean=self.fused == 2,
sfpu_bcast=self.sbc, res_t=True, memory_config=self.mem)
def residual(self, x, res, rgate=None, gamma=None, beta=None, *, write_h: bool = True, kcat: int = 0):
"""``h = x + res (* rgate)`` then ``(h, LN(h))``: one fused program (``LN_RESID``, ``tt/ln_kernel.py``)
when this LN runs the fused fp32 kernel, else the stock ``add`` (+ ``multiply``) and :meth:`__call__`.
``h`` is None when ``write_h`` is False and the fused program runs."""
import ttnn
g = self.g if gamma is None else gamma
b = self.b if beta is None else beta
if self.mode == "fp32" and self.fused and self.resid:
from .ln_kernel import layer_norm_fp32_fused, supported
if supported(x) and supported(res) and tuple(x.shape) == tuple(res.shape):
if self.split:
from .ln_kernel import layer_norm_fp32_split, split_supported
if split_supported(x):
h, y = layer_norm_fp32_split(x, g, b, eps=C.LN_EPS, residual=res, rgate=rgate,
write_h=write_h, sfpu_bcast=self.sbc, kcat_ktp=kcat,
memory_config=self.mem)
return h, (KcatOperand(y, kcat) if kcat else y)
return layer_norm_fp32_fused(x, g, b, eps=C.LN_EPS, residual=res, rgate=rgate, write_h=write_h,
lean=self.fused == 2, sfpu_bcast=self.sbc, memory_config=self.mem)
h = ttnn.add(x, res if rgate is None else ttnn.multiply(res, rgate))
return h, self(h, gamma, beta)
class Const:
"""A constant device tensor (upload once)."""
def __init__(self, build: Build, array: np.ndarray, dtype: str, memory_config=None):
self.t = build.upload(array, dtype)
if memory_config is not None:
import ttnn
self.t = ttnn.to_memory_config(self.t, memory_config)
def __call__(self):
return self.t