Download code/models/tt/conv_cell.py from changh95/superpoint-p150: direct link, hf CLI and curl.
- Browser
- Download file 28.1 kB
-
https://huggingface.co/changh95/superpoint-p150/resolve/main/code/models/tt/conv_cell.py
- Command line
-
hf download hf://changh95/superpoint-p150/code/models/tt/conv_cell.py
-
curl -L -o conv_cell.py https://huggingface.co/changh95/superpoint-p150/resolve/main/code/models/tt/conv_cell.py
28.1 kB
| # 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") | |