superpoint-p150 / code /models /tt /conv_cell.py
changh95's picture
Optimized build (2026-10-03): trace 3.73 -> 0.50 ms
c699c4c verified
Raw History Blame Contribute Delete
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")