superpoint-p150 / code /models /tt /row_conv.py
changh95's picture
Optimized build (2026-10-03): trace 3.73 -> 0.50 ms
c699c4c verified
Raw History Blame Contribute Delete
14.2 kB
# SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc.
# SPDX-License-Identifier: Apache-2.0
"""ABANDONED (round 7): with SP_RC_WPROC=1 the kernel still hung chip 6 after ~700 back-to-back launches
(chipstate/chip6/FAULT, OPT_REPORT.md "Round 7"); do not run it. Kept for reference only.
EXPERIMENTAL, NOT WIRED INTO THE MODEL (round 5): one eager call per Cout=256 was bit-identical to
ttnn.conv2d on chip 10; the next test run (Cout=128 eager + 10 calls in one trace) hung the chip
(chipstate/chip10/FAULT, OPT_REPORT.md "Round 5"). Round 6 (chip 4, end-of-op 'done' barrier on semaphore 2
added): Cout 128 and 256 bit-identical to ttnn.conv2d, 1 and 2 eager calls, traces of 2 and 10 calls (3 replays)
all passed; the timing stage (~500 back-to-back launches) hung chip 4 (chipstate/chip4/FAULT). Leading
hypothesis: the two groups' weight multicasts run concurrently on NOC 1 (PROC 0, start/end swapped), so each
path from its sender in row 0 to the far corner crosses the other group's rectangle (path-reservation
deadlock, rare per launch). Next: SP_RC_WPROC=1 (sender on NOC 0, multicast starts at the sender's own
corner, the two rectangles share no router), staged on a fresh chip. Validate step by step before use.
3x3 / pad-1 conv (+ bias + ReLU) for the 60x80 layers (block 3, heads) as ONE custom generic_op
without halo / im2col / tilize, bit-identical to ttnn.conv2d's height-sharded program.
ttnn's conv for these layers is a halo op (7.7-8.9 us) + a conv on 75 of the 120 cores whose
per-core work is dominated by the im2col reader and tilize (the math is 2 M tiles x 36 K x N).
Here each image row r (80 px, padded to 96 = 3 tile rows m) is computed by two cores, one per half
of the output channels (group g): core (r, g) = logical (x = 6 g + r % 6, y = r // 6).
Input / output: the TILE [H*W, C] layout ttnn's conv uses, height-sharded [64, C] over the grid
(row-major; 75 shards hold data). The data-movement RISCs build the in0 tiles of the core directly
from the (neighbour) shards with a host-built NoC copy schedule (kernels/sp_conv/rc_dm.cpp):
T(ky, kx, m, c) row i = pixel 32 m + i + kx - 1 of image row r + ky - 1, channels 32 c .. 32 c + 31
(zero outside the image), i.e. exactly ttnn's im2col row for the output pixel, K block ky, K tiles
(kx, c) -- 12 tiles per K block for Cin = 128. Each RISC builds half of the channel tiles c.
Compute (kernels/sp_conv/rc_compute.cpp) reproduces ttnn's conv_bmm_tilize.cpp order for
height sharding with packer_l1_acc + fused bias: per K block ky (= ttnn's in0_block_w_i, 12 tiles)
the fp32 DST accumulates the 12 tile products in K order, the block result is packed to fp32
partials (ky = 0 overwrite, ky > 0 packer L1 accumulate), then partials + bias
(add_tiles_bcast_rows) with ReLU on pack -> bf16. Same values, same order => same bits.
Weights: per group the [1152, Cout/2] bf16 matrix (K order ky, kx, ci as ttnn's prepared weights)
is read from DRAM by one sender core (logical (6 g, 0)) in three K-block chunks and multicast to the
group's 6 x 10 rectangle (the two groups use disjoint rectangles); a semaphore carries the number
of chunks delivered. Receivers first signal the sender that they run (no multicast into a core
that has not started the program).
Output rows px 0..79 of tile rows m are written (1 KB = 16-row halves) to the output shards.
"""
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")
H, W = 60, 80
WP = 96 # padded row: 3 tile rows
SHARD = 64 # pixels per input / output shard (ttnn conv layout for 4800 px on the 12 x 10 grid)
NG = 2 # output-channel groups (cores per image row)
GX = 6 # grid columns per group
ZSEL = 7 # schedule source index of the local zero page
NSRC = 6 # source core slots per core (input)
NDST = 3 # destination core slots per core (output)
def _src(name):
with open(os.path.join(_KDIR, name)) as f:
return f.read()
def _unit(row: int, fc: int) -> int:
"""32-byte unit of tile row `row` (0..31), column half fc, inside a tile (face layout)."""
return ((row >> 4) * 2 + fc) * 16 + (row & 15)
def _merge(ent):
"""[(sel, su, du, L)] sorted by du -> merged runs contiguous in source and destination."""
out = []
for e in ent:
if out:
sel, su, du, L = out[-1]
if e[0] == sel and e[2] == du + L and (sel == ZSEL or e[1] == su + L) and (L + e[3] <= 32 if sel == ZSEL else True):
out[-1] = (sel, su, du, L + e[3])
continue
out.append(e)
return out
def core_of(r: int, g: int):
return GX * g + r % GX, r // GX
def input_schedule(r: int, ct: int, ch: int):
"""Copy entries of core row r per K block ky: [3][(sel, src_unit, dst_unit, n_units)] for channel
tile 0 (the kernel adds 64 units per channel tile on both sides), plus the source shard list."""
shards = []
blocks = []
blk_h = 9 * ch
for ky in range(3):
R = r + ky - 1
ent = []
for m in range(3):
for kx in range(3):
dst_tile = (m * 3 + kx) * ch # within this K block's CB_A reservation
for i in range(32):
q = 32 * m + i + kx - 1
valid = 0 <= R < H and 0 <= q < W
for fc in (0, 1):
du = dst_tile * 64 + _unit(i, fc)
if not valid:
ent.append((ZSEL, 0, du, 1))
continue
p = W * R + q
s, lr = divmod(p, SHARD)
if s not in shards:
shards.append(s)
tr, ri = divmod(lr, 32)
su = tr * ct * 64 + _unit(ri, fc)
ent.append((shards.index(s), su, du, 1))
ent.sort(key=lambda e: e[2])
blocks.append(_merge(ent))
assert len(shards) <= NSRC, shards
return blocks, shards
def output_schedule(r: int, g: int, nl: int, cto: int):
"""Output entries of core (r, g): [(sel, src_unit, dst_unit, n_units)] for local n = 0 (the kernel adds
64 units per output channel tile on both sides) + destination shard list."""
shards = []
ent = []
for m in range(3):
for h in range(2):
q0 = 32 * m + 16 * h
if q0 >= W:
continue
p = W * r + q0
s, lr = divmod(p, SHARD)
assert lr % 16 == 0
if s not in shards:
shards.append(s)
tr, ri = divmod(lr, 32)
su = m * nl * 64 + h * 32 # rows 16 h .. 16 h + 15 of out tile (m, n = 0) = units 32 h .. 32 h + 31
du = (tr * cto + g * nl) * 64 + (ri // 16) * 32
ent.append((shards.index(s), su, du, 32))
assert len(shards) <= NDST
return ent, shards
class RowConv:
"""conv3x3(pad 1) + bias (+ ReLU) on [1, 1, 4800, Cin] TILE height-sharded [64, Cin] (60 x 80)."""
def __init__(self, device, weight: torch.Tensor, bias: torch.Tensor, relu: bool = True):
self.device = device
co, ci = int(weight.shape[0]), int(weight.shape[1])
assert tuple(weight.shape[2:]) == (3, 3) and ci % 64 == 0 and co % (32 * NG) == 0
self.ci, self.co, self.relu = ci, co, relu
self.ct = ci // 32 # input channel tiles
self.ch = self.ct // 2 # channel tiles per data-movement RISC
self.nl = co // 32 // NG # output channel tiles per core
self.cto = co // 32
# K order (ky, kx, ci) as ttnn's prepared conv weights; group g = output channels g * co/2 ..
wk = weight.detach().float().permute(2, 3, 1, 0).reshape(9 * ci, co) # [ky, kx, ci] x co
wg = torch.cat([wk[:, g * co // NG:(g + 1) * co // NG] for g in range(NG)], 0) # [NG * 9ci, co/NG]
self.w = ttnn.from_torch(wg.to(torch.bfloat16), dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT, device=device,
memory_config=ttnn.DRAM_MEMORY_CONFIG)
bm = torch.zeros(32, co)
bm[0] = bias.detach().float()
self.b = ttnn.from_torch(bm.to(torch.bfloat16), dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT, device=device,
memory_config=ttnn.DRAM_MEMORY_CONFIG)
# per-core schedule pages (core order = row-major logical grid of the 120 cores)
self.cores = [ttnn.CoreCoord(x, y) for y in range(10) for x in range(12)]
pages = []
self.src_shards, self.dst_shards = {}, {}
for c in self.cores:
g, r = c.x // GX, c.y * GX + c.x % GX
blocks, ss = input_schedule(r, self.ct, self.ch)
oent, ds = output_schedule(r, g, self.nl, self.cto)
self.src_shards[(c.x, c.y)], self.dst_shards[(c.x, c.y)] = ss, ds
words = [len(b) for b in blocks] + [len(oent)]
for b in blocks + [oent]:
for sel, su, du, L in b:
assert su < 1 << 16 and du < 1 << 16 and L < 1 << 16
words += [(sel << 16) | su, (du << 16) | L]
pages.append(words)
self.page_words = (max(len(p) for p in pages) + 7) // 8 * 8
tab = np.zeros((len(pages), self.page_words), np.uint32)
for i, p in enumerate(pages):
tab[i, :len(p)] = p
self.sched = ttnn.from_torch(torch.from_numpy(tab.view(np.int32)), dtype=ttnn.uint32, layout=ttnn.ROW_MAJOR_LAYOUT,
device=device, memory_config=ttnn.DRAM_MEMORY_CONFIG)
self._dm = _src("rc_dm.cpp")
self._cc = _src("rc_compute.cpp")
self.wproc = int(os.environ.get("SP_RC_WPROC", "0"))
def supports(self, x: ttnn.Tensor) -> bool:
if not x.is_sharded() or x.layout != ttnn.TILE_LAYOUT or x.dtype != ttnn.bfloat16:
return False
mc = x.memory_config()
return (mc.memory_layout == ttnn.TensorMemoryLayout.HEIGHT_SHARDED and mc.buffer_type == ttnn.BufferType.L1
and list(mc.shard_spec.shape) == [SHARD, self.ci] and list(x.shape)[-2:] == [H * W, self.ci]
and mc.shard_spec.orientation == ttnn.ShardOrientation.ROW_MAJOR)
def __call__(self, x: ttnn.Tensor) -> ttnn.Tensor:
dev = self.device
mc = x.memory_config()
in_cores = ttnn.corerange_to_cores(mc.shard_spec.grid, row_wise=True)
omc = ttnn.MemoryConfig(ttnn.TensorMemoryLayout.HEIGHT_SHARDED, ttnn.BufferType.L1,
ttnn.ShardSpec(mc.shard_spec.grid, [SHARD, self.co], ttnn.ShardOrientation.ROW_MAJOR))
out = ttnn.allocate_tensor_on_device(ttnn.Shape([1, 1, H * W, self.co]), ttnn.bfloat16, ttnn.TILE_LAYOUT, dev, omc)
grid = ttnn.CoreRangeSet([ttnn.CoreRange(ttnn.CoreCoord(0, 0), ttnn.CoreCoord(11, 9))])
CB_A0, CB_A1, CB_W, CB_B, CB_P, CB_O, CB_T = 0, 1, 2, 3, 4, 5, 6
bf = ttnn.bfloat16
blk_h = 9 * self.ch
def cb(i, pages, page=2048, fmt=bf):
f = ttnn.CBFormatDescriptor(buffer_index=i, data_format=fmt, page_size=page)
return ttnn.CBDescriptor(total_size=page * pages, core_ranges=grid, format_descriptors=[f])
page_bytes = self.page_words * 4
cbs = [cb(CB_A0, 3 * blk_h), cb(CB_A1, 3 * blk_h), cb(CB_W, 36 * self.nl), cb(CB_B, self.nl),
cb(CB_P, 3 * self.nl, 4096, ttnn.float32), cb(CB_O, 3 * self.nl), cb(CB_T, 1, 2 * (page_bytes + 1024))]
acc = (ttnn.TensorAccessorArgs(self.w).get_compile_time_args() + ttnn.TensorAccessorArgs(self.b).get_compile_time_args()
+ ttnn.TensorAccessorArgs(self.sched).get_compile_time_args())
xa, oa, wa, ba, sa = x.buffer_address(), out.buffer_address(), self.w.buffer_address(), self.b.buffer_address(), self.sched.buffer_address()
SC = ttnn.KernelDescriptor.SourceType.SOURCE_CODE
defs = [tuple((d + "=1").split("=")[:2]) for d in os.environ.get("SP_RC_DEFS", "").split(",") if d]
noc = lambda c: dev.worker_core_from_logical_core(c)
ks = []
for proc, cfg, cba in ((0, ttnn.ReaderConfigDescriptor(), CB_A0), (1, ttnn.WriterConfigDescriptor(), CB_A1)):
rt = ttnn.RuntimeArgs()
for k, c in enumerate(self.cores):
g = c.x // GX
snd = noc(ttnn.CoreCoord(GX * g, 0))
m0, m1 = noc(ttnn.CoreCoord(GX * g, 0)), noc(ttnn.CoreCoord(GX * g + GX - 1, 9))
is_snd = int(c.x == GX * g and c.y == 0)
src = []
for s in self.src_shards[(c.x, c.y)]:
n = noc(in_cores[s])
src += [n.x, n.y]
src += [0, 0] * (NSRC - len(src) // 2)
dst = []
for s in self.dst_shards[(c.x, c.y)]:
n = noc(in_cores[s])
dst += [n.x, n.y]
dst += [0, 0] * (NDST - len(dst) // 2)
rt[c.x][c.y] = [xa, oa, wa, ba, sa, k, is_snd, snd.x, snd.y, m0.x, m0.y, m1.x, m1.y, GX * 10 - 1,
g * 36 * self.nl, g * self.nl] + src + dst
ks.append(ttnn.KernelDescriptor(
kernel_source=self._dm, source_type=SC, core_ranges=grid,
compile_time_args=[proc, cba, CB_W, CB_B, CB_O, CB_T, self.nl, self.ch, self.ct, page_bytes, blk_h, self.wproc] + acc,
runtime_args=rt, config=cfg, defines=defs))
ccfg = ttnn.ComputeConfigDescriptor(math_fidelity=ttnn.MathFidelity.HiFi2, fp32_dest_acc_en=True, dst_full_sync_en=False)
ks.append(ttnn.KernelDescriptor(kernel_source=self._cc, source_type=SC, core_ranges=grid,
compile_time_args=[CB_A0, CB_A1, CB_W, CB_B, CB_P, CB_O, self.nl, self.ch, self.ct, blk_h, int(self.relu)],
runtime_args=[], config=ccfg, defines=defs))
sems = [ttnn.SemaphoreDescriptor(id=i, core_ranges=grid, initial_value=0) for i in (0, 1, 2)]
return ttnn.generic_op([x, self.w, self.b, self.sched, out], ttnn.ProgramDescriptor(kernels=ks, semaphores=sems, cbs=cbs))