# SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc. # SPDX-License-Identifier: Apache-2.0 """Block-0 conv_b (3x3, 64 -> 64, ReLU) as ONE custom generic_op on the "cell tile" layout. ttnn.conv2d on the 64-channel 480x640 activation is reader/tilize bound (halo 24 us + Move 6 us + conv 148 us). The cell-conv conv_a already produces, per 8-pixel cell, the 8 x 64 channels of its pixels; written as TILE ([N/8 cells, 512]: tile (tr, q) = cells 32tr..32tr+31, pixel p = q // 2, channel half h = q % 2) the tap of pixel p + kx of the same cell is just another tile column, so conv_b becomes a sum of tile matmuls with no im2col and no tilize: out(TR, p) = relu( sum_{ky, kx, h} X(ky, kx, h) @ W[ky, kx, h] + bias ) Taps that leave the tile grid are rebuilt per output tile row by the two data-movement RISCs from a host-built copy schedule (kernels/sp_conv/cb0_reader.cpp): half-tile row shifts for ky = +-1 (an image row is 80 cells = 2.5 tile rows), one-row shifts for the pixel taps that cross a cell, zero rows at the image borders, neighbour-core cells over the NoC. The compute kernel (cb0_compute.cpp) runs one output pixel x 2 output tiles per matmul_block call (rt = 1, ct = 2; measured faster than rt = 2 x ct = 2 and rt = 8 x ct = 1 blocks). Group layout (GROUP = 44 tiles per output tile row TR, row i = output cell c = 32 TR + i): [ 0..15] H(-80) q0..q15 [16..31] H(+80) q0..q15 [32..43] (s, p') = (-1, 7), (+1, 0), (-81, 7), (-79, 0), (+79, 7), (+81, 0), each h0 h1 where S(s, p') / H(s) row i = source cell c + s (pixel p', or all pixels), zero when that cell is outside the image or (for the +-1 pixel shifts) in another image row. """ from __future__ import annotations import os import numpy as np import torch import ttnn _KDIR = os.path.join(os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))), "kernels", "sp_conv") W_TILES = 36 # weight tiles in CB_W (tiles 36, 37 of the weight tensor = bias -> CB_B) CPR = 80 # cells per image row (block 0: 640 px / 8, block 1: 320 px / 4) UNIT = 16 # elements per 32-byte unit (one face row) MAXE = 256 # schedule entries per (core class, TR) class CellGeom: """P pixels per cell, `cells` cells per core (a multiple of 80 and 32), `rows` image rows in total. QT = 2 P tiles per tile row (pixel p, channel half h at column 2 p + h), TRS tile rows per core, GROUP = 2 QT + 12 shifted-tap tiles per output tile row.""" def __init__(self, P: int, cells: int, rows: int): assert cells % CPR == 0 and cells % 32 == 0 and P % 2 == 0 self.P, self.cells, self.rows = P, cells, rows self.QT = 2 * P self.TRS = cells // 32 self.GROUP = 2 * self.QT + 12 self.rows_core = cells // CPR G0 = CellGeom(8, 320, 480) # block 0: 480 x 640, 4 image rows per core G1 = CellGeom(4, 160, 240) # block 1: 240 x 320, 2 image rows per core # block-0 names (tests / scripts) CELLS, TRS, QT, GROUP = G0.cells, G0.TRS, G0.QT, G0.GROUP def _src(name): with open(os.path.join(_KDIR, name)) as f: return f.read() def weight_matrix(w: torch.Tensor, b: torch.Tensor, order: str = "kxrev") -> torch.Tensor: """[704, 64] bf16 (22 x 2 tiles). order "kxrev" (cb0_compute.cpp): row tile R = 6 ky + 3 h + (2 - kx) (ky, kx = 0..2) holds w[:, 32 h + ci, ky, kx] in row ci, i.e. in1 tile 12 ky + 6 h + 2 (2 - kx) + n. order "tap" (af7b261 kernel): K rows t*64 + ci (t = 3*ky + kx) = w[:, ci, ky, kx]; tiles 36, 37 = bias row (row 0); 38, 39 = 0; tile 40 = ones in column 0 (in0 of the bias product); 41..43 = 0.""" assert tuple(w.shape) == (64, 64, 3, 3) m = torch.zeros(704, 64, dtype=torch.float32) for ky in range(3): for kx in range(3): if order == "tap": t = 3 * ky + kx m[t * 64:(t + 1) * 64, :] = w[:, :, ky, kx].t() else: for h in range(2): r = 6 * ky + 3 * h + (2 - kx) m[r * 32:(r + 1) * 32, :] = w[:, 32 * h:32 * (h + 1), ky, kx].t() m[576, :] = b m[640:672, 0] = 1.0 return m.to(torch.bfloat16) # ---- copy schedule --------------------------------------------------------------------------- SEL_LOCAL, SEL_PREV, SEL_NEXT, SEL_ZERO = 0, 1, 2, 3 def _row_units(row: int): """(left face unit, right face unit) offsets of tile row `row` inside a tile (32-B units).""" f = (row >> 4) * 2 return f * 16 + (row & 15), (f + 1) * 16 + (row & 15) def _builds(g: CellGeom = G0): """[(slot0, s, q0, nq)]: s = source cell shift, q0..q0+nq-1 source tile columns.""" b = [(0, -80, 0, g.QT), (g.QT, 80, 0, g.QT)] for e, (s, pp) in enumerate(((-1, g.P - 1), (1, 0), (-81, g.P - 1), (-79, 0), (79, g.P - 1), (81, 0))): b.append((2 * g.QT + 2 * e, s, 2 * pp, 2)) return b def schedule(core_class: str, g: CellGeom = G0): """Copy list per TR for a core of class 'first' / 'mid' / 'last' (4 image rows per core): [TRS][entries (sel, src_unit, dst_unit, n_units)], 32-byte units; zero runs read the RISC's local zero page (sel 3, src 0).""" first, last = core_class == "first", core_class == "last" out = [] for tr in range(g.TRS): ent = [] for slot0, s, q0, nq in _builds(g): ky = (s + 81) // 80 - 1 if s != 0 else 0 delta = s - 80 * ky # per target row: (sel, src tile row index within source shard, src row) or zero rows = [] for i in range(32): c = 32 * tr + i cs = c + s cm = c % CPR if (delta == -1 and cm == 0) or (delta == 1 and cm == CPR - 1): rows.append(None) continue if cs < 0: if first: rows.append(None) continue sel, lc = SEL_PREV, cs + g.cells elif cs >= g.cells: if last: rows.append(None) continue sel, lc = SEL_NEXT, cs - g.cells else: sel, lc = SEL_LOCAL, cs rows.append((sel, lc >> 5, lc & 31)) for q in range(nq): dst_tile = slot0 + q i = 0 while i < 32: r = rows[i] # maximal run inside one target 16-row half j = i + 1 while j < 32 and (j & 15) and _cont(rows[i], rows[j], j - i): j += 1 L = j - i for face in (0, 1): du = dst_tile * 64 + _row_units(i)[face] if r is None: ent.append((SEL_ZERO, 0, du, L)) else: sel, ts, rs = r su = (ts * g.QT + q0 + q) * 64 + _row_units(rs)[face] ent.append((sel, su, du, L)) i = j out.append(_merge(ent)) return out def _cont(a, b, k): """row b continues run a (k rows later) within one source 16-row half / one zero run.""" if a is None or b is None: return a is None and b is None return a[0] == b[0] and a[1] == b[1] and b[2] == a[2] + k and (a[2] >> 4) == (b[2] >> 4) def _merge(ent): """Merge consecutive entries that are contiguous in both source and destination (e.g. the two face halves of an aligned half-tile copy, or whole tiles).""" out = [] for e in ent: if out: sel, su, du, L = out[-1] if e[0] == sel and e[2] == du + L and (sel == SEL_ZERO or e[1] == su + L) and L + e[3] <= (32 if sel == SEL_ZERO else 1 << 15): out[-1] = (sel, su if sel != SEL_ZERO else 0, du, L + e[3]) continue out.append(e) return out def schedule_table(g: CellGeom = G0): """uint32 [3 classes * TRS, MAXE, 2]: entry 0 word 0 = count; entries 1.. = (sel << 16 | src, dst << 16 | len) in 32-B units. Entries are sorted by (source, length) so that the reader re-programs its NoC read state (source core, size) only when that pair changes (the copies are independent).""" tab = np.zeros((3, g.TRS, MAXE, 2), np.uint32) for ci, cls in enumerate(("first", "mid", "last")): for tr, ent in enumerate(schedule(cls, g)): ent = sorted(ent, key=lambda e: (e[0], e[3], e[2])) assert len(ent) < MAXE, len(ent) tab[ci, tr, 0, 0] = len(ent) for k, (sel, su, du, L) in enumerate(ent): assert su < 1 << 16 and du < 1 << 16 and L < 1 << 16 tab[ci, tr, 1 + k, 0] = (sel << 16) | su tab[ci, tr, 1 + k, 1] = (du << 16) | L return tab def simulate(x_tiles: np.ndarray, core: int, ncores: int, g: CellGeom = G0) -> np.ndarray: """Host model of the reader for one core: x_tiles [ncores, TRS*QT, 1024] (tile memory order) -> groups [TRS, GROUP, 1024].""" cls = "first" if core == 0 else ("last" if core == ncores - 1 else "mid") flat = x_tiles.reshape(ncores, -1, UNIT) out = np.zeros((g.TRS, g.GROUP * 64, UNIT), x_tiles.dtype) for tr, ent in enumerate(schedule(cls, g)): for sel, su, du, L in ent: if sel == SEL_ZERO: out[tr, du:du + L] = 0 else: src = flat[core + {SEL_LOCAL: 0, SEL_PREV: -1, SEL_NEXT: 1}[sel]] out[tr, du:du + L] = src[su:su + L] return out.reshape(g.TRS, g.GROUP, 1024) def in0_index_pp(tr: int, pp: int, ky: int, h: int, g: CellGeom = G0): """(source, tile index) of in0 for source pixel pp (-1..P of this cell), tap row ky, channel half h, exactly as cb0_compute.cpp computes it. source 'x' = local shard, 'g' = this TR's group.""" if 0 <= pp <= g.P - 1: if ky == 0: return "x", tr * g.QT + 2 * pp + h return "g", (0 if ky < 0 else g.QT) + 2 * pp + h e = {(0, -1): 0, (0, 1): 1, (-1, -1): 2, (-1, 1): 3, (1, -1): 4, (1, 1): 5}[(ky, -1 if pp < 0 else 1)] return "g", 2 * g.QT + 2 * e + h def in0_index(tr: int, p: int, ky: int, kx: int, h: int, g: CellGeom = G0): """(source, tile index) of in0 for output pixel p, tap (ky, kx), channel half h.""" return in0_index_pp(tr, p + kx, ky, h, g) def compute_calls(pix: int = 2, P: int = 8): """The matmul_block calls of cb0_compute.cpp for one output tile row, in issue order: [(p0, ky, pp, h, in1 tile, dst tile, ct)] (dst tile d + c = pixel p0 + (d + c) // 2, half (d + c) % 2, in1 tile w + c).""" calls = [] for p0 in range(0, P, pix): for ky in (-1, 0, 1): for pp in range(p0 - 1, p0 + pix + 1): lo, hi = max(pp - 1, p0), min(pp + 1, p0 + pix - 1) ct = 2 * (hi - lo + 1) w0 = 12 * (ky + 1) + 2 * (1 - pp + lo) d = 2 * (lo - p0) for h in (0, 1): calls.append((p0, ky, pp, h, w0 + 6 * h, d, ct)) return calls def cell_memory_config(grid, g: CellGeom, px_per_cell: int): return ttnn.MemoryConfig(ttnn.TensorMemoryLayout.HEIGHT_SHARDED, ttnn.BufferType.L1, ttnn.ShardSpec(grid, [g.cells, px_per_cell * 64], ttnn.ShardOrientation.ROW_MAJOR)) class CellConv: """3x3 / pad-1 conv 64 -> 64 + bias + ReLU on the cell tile layout of geometry g, ONE generic_op (kernels/sp_conv/cb0_reader.cpp + cb0_compute.cpp). Input: TILE [N/P cells, P * 64] height-sharded [g.cells, P * 64]. Output: the same layout, or with hmax=True the horizontally half-pooled map [cells, P/2 * 64] (pixel pairs max-reduced: the first half of the following 2x2 pool).""" def __init__(self, device, weight: torch.Tensor, bias: torch.Tensor, g: CellGeom = G0, hmax: bool = True, trb: int | None = None, compute_src: str | None = None, pix: int | None = None, worder: str = "kxrev"): self.device, self.g, self.hmax = device, g, hmax # output pixels per DST block: 2 (half-sync DST, pack overlaps math) or 4 (full-sync DST) self.pix = int(pix if pix is not None else os.environ.get("SP_CB0_PIX", "2")) assert self.pix in (2, 4) self.mcast = os.environ.get("SP_CB0_MCAST", "1") == "1" # weights multicast (see cb0_reader.cpp) # output tile rows per bias phase (divides TRS); SP_CB0_TRB overrides for measurements self.trb = int(trb if trb is not None else os.environ.get("SP_CB0_TRB", "1")) assert g.TRS % self.trb == 0, self.trb wm = weight_matrix(weight.detach().float(), bias.detach().float(), worder) self.w = ttnn.from_torch(wm, dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT, device=device, memory_config=ttnn.DRAM_MEMORY_CONFIG) tab = schedule_table(g) self._tab = tab self._sched_rep = {} # n cores -> per-core replicated schedule tensor (SP_SCHED_REP) self.nent = int(tab[:, :, 0, 0].max()) self.sched = ttnn.from_torch(torch.from_numpy(tab.reshape(3 * g.TRS, MAXE * 2).view(np.int32)), dtype=ttnn.uint32, layout=ttnn.ROW_MAJOR_LAYOUT, device=device, memory_config=ttnn.DRAM_MEMORY_CONFIG) self._r = _src("cb0_reader.cpp") self._c = _src("cb0_compute.cpp") if compute_src is None else open(compute_src).read() self._c_ref = compute_src is not None # af7b261 compute kernel (measurements): 16-tile CB_P # tile rows whose group the reader's PROC 1 (NOC 0) builds (bit tr): rows 3, 6, 7, 8, 9. Rows 7-9 # read the next core's cells, rows 0-2 the previous core's: those reads are cheap on NOC 0 / NOC 1 # respectively and congest the other NoC (measured 117 -> 103 us vs the even/odd split; the # wrong way round 150+ us) # block 1 (TRS 5): rows 0-2 read the previous core, rows 2-4 the next one -> PROC 1 rows 3, 4 self.bmask = int(os.environ.get("SP_CB0_BMASK" if g is G0 else "SP_CB1_BMASK", "0x3C8" if g is G0 else "0x18"), 0) if self._c_ref: self.bmask = 0x2AA # the af7b261 compute kernel takes odd rows from PROC 1 (tr & 1) def supports(self, x: ttnn.Tensor) -> bool: g = self.g if not x.is_sharded() or x.layout != ttnn.TILE_LAYOUT or x.dtype != ttnn.bfloat16: return False mc = x.memory_config() sh = mc.shard_spec.shape return (mc.memory_layout == ttnn.TensorMemoryLayout.HEIGHT_SHARDED and mc.buffer_type == ttnn.BufferType.L1 and list(sh) == [g.cells, g.P * 64] and x.shape[-1] == g.P * 64 and x.shape[-2] == g.cells * mc.shard_spec.grid.num_cores() and g.rows == g.rows_core * mc.shard_spec.grid.num_cores()) def release(self): ttnn.deallocate(self.w) ttnn.deallocate(self.sched) for t in self._sched_rep.values(): ttnn.deallocate(t) self._sched_rep = {} def alloc_weight_buffers(self, grid): """(W, bias) in-graph L1 tensors on ``grid`` with this op's weight tiles (36) and bias tiles (2) per core: the target of the previous op's weight prefetch (SP_CW_PF) and the backing of this op's CB_W / CB_B (``pre``).""" n = grid.num_cores() out = [] for nt in (W_TILES, 2): mc = ttnn.MemoryConfig(ttnn.TensorMemoryLayout.HEIGHT_SHARDED, ttnn.BufferType.L1, ttnn.ShardSpec(grid, [nt * 32, 32], ttnn.ShardOrientation.ROW_MAJOR)) out.append(ttnn.allocate_tensor_on_device(ttnn.Shape([1, 1, n * nt * 32, 32]), ttnn.bfloat16, ttnn.TILE_LAYOUT, self.device, mc)) return tuple(out) def __call__(self, x: ttnn.Tensor, pre=None, pf=None) -> ttnn.Tensor: """-> TILE [N/P cells, opx * 64] height-sharded [g.cells, opx * 64] on the same cores; opx = P / 2 with hmax (pixel pairs (2j, 2j + 1) max-reduced), else P. SP_CW_PF: ``pre`` = (W, bias) L1 tensors (``alloc_weight_buffers``) already holding this op's weights (CB_W / CB_B are backed by them; no DRAM read, no multicast); ``pf`` = [(later CellConv, W, bias), ...]: after its own work the sender core reads those ops' weights from DRAM and multicasts them into those tensors.""" g = self.g mc = x.memory_config() grid = mc.shard_spec.grid cores = ttnn.corerange_to_cores(grid, row_wise=True) n = len(cores) opx = g.P // 2 if self.hmax else g.P omc = ttnn.MemoryConfig(ttnn.TensorMemoryLayout.HEIGHT_SHARDED, ttnn.BufferType.L1, ttnn.ShardSpec(grid, [g.cells, opx * 64], mc.shard_spec.orientation)) out = ttnn.allocate_tensor_on_device(ttnn.Shape([1, 1, x.shape[-2], opx * 64]), ttnn.bfloat16, ttnn.TILE_LAYOUT, self.device, omc) CB_X, CB_S0, CB_S1, CB_W, CB_OUT, CB_T, CB_P, CB_B = 0, 1, 2, 3, 4, 5, 6, 7 bf = ttnn.bfloat16 cb_x = ttnn.cb_descriptor_from_sharded_tensor(CB_X, x) cb_x.format_descriptors = [ttnn.CBFormatDescriptor(buffer_index=CB_X, data_format=bf, page_size=2048)] cb_o = ttnn.cb_descriptor_from_sharded_tensor(CB_OUT, out) cb_o.format_descriptors = [ttnn.CBFormatDescriptor(buffer_index=CB_OUT, data_format=bf, page_size=2048)] def cb(i, pages, page=2048): f = ttnn.CBFormatDescriptor(buffer_index=i, data_format=bf, page_size=page) return ttnn.CBDescriptor(total_size=page * pages, core_ranges=grid, format_descriptors=[f]) # per RISC scratch: schedule of its TRs + 1 KB zero page n1 = bin(self.bmask).count("1") sched_bytes = max(n1, g.TRS - n1) * MAXE * 8 scratch = 2 * (sched_bytes + 1024) fp = ttnn.CBFormatDescriptor(buffer_index=CB_P, data_format=ttnn.float32, page_size=4096) cb_p = ttnn.CBDescriptor(total_size=4096 * (16 if self._c_ref else 2 * opx * self.trb), core_ranges=grid, format_descriptors=[fp]) if pre is not None: cb_w = ttnn.cb_descriptor_from_sharded_tensor(CB_W, pre[0]) cb_w.format_descriptors = [ttnn.CBFormatDescriptor(buffer_index=CB_W, data_format=bf, page_size=2048)] cb_b = ttnn.cb_descriptor_from_sharded_tensor(CB_B, pre[1]) cb_b.format_descriptors = [ttnn.CBFormatDescriptor(buffer_index=CB_B, data_format=bf, page_size=2048)] else: cb_w, cb_b = cb(CB_W, W_TILES), cb(CB_B, 2) cbs = [cb_x, cb(CB_S0, 2 * g.GROUP), cb(CB_S1, 2 * g.GROUP), cb_w, cb_o, cb(CB_T, 1, scratch), cb_p, cb_b] sched = self.sched rep = os.environ.get("SP_SCHED_REP", "1") == "1" if rep: # one copy of the class schedule per core (pages spread over the DRAM banks instead of 120 cores reading # the same 3 x TRS pages) and only the used entries read (SCHED_RD) if n not in self._sched_rep: cls_of = [0 if i == 0 else (2 if i == n - 1 else 1) for i in range(n)] t = self._tab.reshape(3, g.TRS, MAXE * 2)[cls_of].reshape(n * g.TRS, MAXE * 2) self._sched_rep[n] = ttnn.from_torch(torch.from_numpy(np.ascontiguousarray(t).view(np.int32)), dtype=ttnn.uint32, layout=ttnn.ROW_MAJOR_LAYOUT, device=self.device, memory_config=ttnn.DRAM_MEMORY_CONFIG) sched = self._sched_rep[n] addr, waddr, saddr = x.buffer_address(), self.w.buffer_address(), sched.buffer_address() SC = ttnn.KernelDescriptor.SourceType.SOURCE_CODE defs = [tuple((d + "=1").split("=")[:2]) for d in os.environ.get("SP_CB0_DEFS", "").split(",") if d] # weights: DRAM -> sender core (logical core 0) -> multicast to the other cores (one rectangle) lx = [c.x for c in cores] ly = [c.y for c in cores] rect = (max(lx) - min(lx) + 1) * (max(ly) - min(ly) + 1) == n mcast = self.mcast and rect and n > 1 and pre is None if mcast: defs.append(("MCAST", "1")) if pre is not None: defs.append(("W_PRE", "1")) if pf is not None: if not rect: raise RuntimeError("SP_CW_PF needs a rectangular core grid (one multicast rectangle)") defs.append(("PF_NEXT", "1")) pf_args = [len(pf)] + [a for (op, tw, tb) in pf for a in (op.w.buffer_address(), tw.buffer_address(), tb.buffer_address())] else: pf_args = [] defs += [("BMASK_DEF", str(self.bmask)), ("GROUP_DEF", str(g.GROUP)), ("QT_DEF", str(g.QT))] if rep: defs.append(("SCHED_RD", str((self.nent + 1) * 8 + 63 & ~63))) snd = self.device.worker_core_from_logical_core(cores[0]) m0 = self.device.worker_core_from_logical_core(ttnn.CoreCoord(min(lx), min(ly))) m1 = self.device.worker_core_from_logical_core(ttnn.CoreCoord(max(lx), max(ly))) acc = ttnn.TensorAccessorArgs(self.w).get_compile_time_args() + ttnn.TensorAccessorArgs(sched).get_compile_time_args() ks = [] for proc, cfg, cbs_ in ((0, ttnn.ReaderConfigDescriptor(), CB_S0), (1, ttnn.WriterConfigDescriptor(), CB_S1)): rt = ttnn.RuntimeArgs() for i, core in enumerate(cores): pc = self.device.worker_core_from_logical_core(cores[max(i - 1, 0)]) nc = self.device.worker_core_from_logical_core(cores[min(i + 1, n - 1)]) cls = (0 if i == 0 else (2 if i == n - 1 else 1)) if not rep else i rt[core.x][core.y] = [addr, cls, pc.x, pc.y, nc.x, nc.y, waddr, saddr, int(i == 0), snd.x, snd.y, m0.x, m0.y, m1.x, m1.y, n - 1] + pf_args ks.append(ttnn.KernelDescriptor(kernel_source=self._r, source_type=SC, core_ranges=grid, compile_time_args=[proc, cbs_, CB_X, CB_W, CB_T, W_TILES, MAXE, g.TRS, sched_bytes, CB_B] + acc, runtime_args=rt, config=cfg, defines=defs)) ccfg = ttnn.ComputeConfigDescriptor(math_fidelity=getattr(ttnn.MathFidelity, os.environ.get("SP_CB0_FID", "HiFi2")), fp32_dest_acc_en=os.environ.get("SP_CB0_FP32", "1") == "1", dst_full_sync_en=self.pix == 4) ks.append(ttnn.KernelDescriptor(kernel_source=self._c, source_type=SC, core_ranges=grid, compile_time_args=[CB_X, CB_S0, CB_S1, CB_W, CB_OUT, W_TILES, CB_P, CB_B, self.trb, self.pix, g.P, g.TRS, g.GROUP, int(self.hmax)], runtime_args=[], config=ccfg, defines=defs)) sems = [ttnn.SemaphoreDescriptor(id=i, core_ranges=grid, initial_value=0) for i in (0, 1, 2)] if mcast else [] extra = (list(pre) if pre is not None else []) + ([t for e in pf for t in (e[0].w, e[1], e[2])] if pf is not None else []) return ttnn.generic_op([x, self.w, sched] + extra + [out], ttnn.ProgramDescriptor(kernels=ks, semaphores=sems, cbs=cbs)) def ConvCell0(device, weight, bias, **kw): """Block-0 conv_b: P = 8, fused horizontal pool max.""" return CellConv(device, weight, bias, G0, hmax=True, **kw) class PoolCell: """Vertical half of a 2x2 max pool on a half-pooled cell map (CellConv(hmax=True) output: TILE [cells, P/2 px * 64], QTI = P tiles per tile row, g.rows_core image rows per core), ONE generic_op (kernels/sp_conv/pc0_{reader,compute}.cpp). Pooled row Y = max(cells 160Y + i, cells 160Y + 80 + i), i < 80: the second operand is half a tile row away, rebuilt as tiles by the data-movement RISCs (2 x 1 KB face copies per tile); SFPU max per tile pair. Output: out="rm": pack-untilized ([32 cells, P/2 px * 64] row-major == 32 * P/2 pooled pixels x 64 ch), valid rows copied into ROW_MAJOR [N_pooled_px, 64] height-sharded (ttnn conv input); out="cell": the pooled cell tiles reassembled (1 KB half-tile copies) into TILE [pooled cells, P/2 px * 64] height-sharded [g.cells / 2, P/2 * 64] (next CellConv).""" def __init__(self, device, g: CellGeom, out: str): assert out in ("rm", "cell") self.device, self.g, self.out = device, g, out self._r = _src("pc0_reader.cpp") self._c = _src("pc0_compute.cpp") def __call__(self, y: ttnn.Tensor) -> ttnn.Tensor: g = self.g qti = g.P # tiles per tile row of the half-pooled input (P/2 px x 2 halves) mc = y.memory_config() grid = mc.shard_spec.grid prow = g.rows_core // 2 # pooled rows per core if self.out == "rm": px_core = prow * CPR * (g.P // 2) omc = ttnn.MemoryConfig(ttnn.TensorMemoryLayout.HEIGHT_SHARDED, ttnn.BufferType.L1, ttnn.ShardSpec(grid, [px_core, 64], mc.shard_spec.orientation)) out = ttnn.allocate_tensor_on_device(ttnn.Shape([1, 1, grid.num_cores() * px_core, 64]), ttnn.bfloat16, ttnn.ROW_MAJOR_LAYOUT, self.device, omc) opage = qti * 64 # bytes per cell row of the untilized block else: ocells = prow * CPR assert ocells % 32 == 0 omc = ttnn.MemoryConfig(ttnn.TensorMemoryLayout.HEIGHT_SHARDED, ttnn.BufferType.L1, ttnn.ShardSpec(grid, [ocells, qti * 32], mc.shard_spec.orientation)) out = ttnn.allocate_tensor_on_device(ttnn.Shape([1, 1, grid.num_cores() * ocells, qti * 32]), ttnn.bfloat16, ttnn.TILE_LAYOUT, self.device, omc) opage = 2048 CB_A, CB_B0, CB_B1, CB_U, CB_O = 0, 1, 2, 3, 4 bf = ttnn.bfloat16 cb_a = ttnn.cb_descriptor_from_sharded_tensor(CB_A, y) cb_a.format_descriptors = [ttnn.CBFormatDescriptor(buffer_index=CB_A, data_format=bf, page_size=2048)] cb_o = ttnn.cb_descriptor_from_sharded_tensor(CB_O, out) cb_o.format_descriptors = [ttnn.CBFormatDescriptor(buffer_index=CB_O, data_format=bf, page_size=opage)] def cb(i, pages, page=2048): f = ttnn.CBFormatDescriptor(buffer_index=i, data_format=bf, page_size=page) return ttnn.CBDescriptor(total_size=page * pages, core_ranges=grid, format_descriptors=[f]) nb = (prow + 1) // 2 # pooled rows per RISC (Y = 2m + PROC) cbs = [cb_a, cb(CB_B0, 3 * qti * max(nb, 1)), cb(CB_B1, 3 * qti * max(nb, 1)), cb(CB_U, 2 * qti), cb_o] SC = ttnn.KernelDescriptor.SourceType.SOURCE_CODE mode = 1 if self.out == "cell" else 0 ks = [] for proc, cfg in ((0, ttnn.ReaderConfigDescriptor()), (1, ttnn.WriterConfigDescriptor())): ks.append(ttnn.KernelDescriptor(kernel_source=self._r, source_type=SC, core_ranges=grid, compile_time_args=[proc, CB_A, CB_B0 if proc == 0 else CB_B1, CB_U, CB_O, qti, g.TRS, prow, mode], runtime_args=ttnn.RuntimeArgs(), config=cfg)) half = mode == 1 and os.environ.get("SP_PC_HALF", "1") == "1" ccfg = ttnn.ComputeConfigDescriptor(math_fidelity=ttnn.MathFidelity.HiFi4, fp32_dest_acc_en=False, dst_full_sync_en=not half) ks.append(ttnn.KernelDescriptor(kernel_source=self._c, source_type=SC, core_ranges=grid, compile_time_args=[CB_A, CB_B0, CB_B1, CB_U, qti, g.TRS, prow, mode], runtime_args=[], config=ccfg, defines=[("PC_HALF", "1")] if half else [])) return ttnn.generic_op([y, out], ttnn.ProgramDescriptor(kernels=ks, semaphores=[], cbs=cbs)) def PoolCell0(device): """Block-0 pool -> ROW_MAJOR block-1 input (ttnn conv).""" return PoolCell(device, G0, "rm")