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