"""Qwen3-VL vision tower (Alpamayo2-Super `vlm.model.visual`) on one Tenstorrent chip with ttnn. Layout decisions (see VISION_NOTES.md): * All images are processed as ONE flat token sequence [1, 1, N, 1152] (N = sum of patches, padded to a multiple of 32 only if needed). Per-image attention is enforced with the windowed SDPA (`cu_window_seqlens` = per-image boundaries, mask generated on device), so no cross-image leakage and no per-image padding. * head_dim 72 is zero-padded to 96 inside the fused QKV weight (SDPA needs a tile-aligned head dim); the softmax scale stays 72**-0.5 and the padded lanes of V are zero, so the result is exact. * 2-D rotary uses `ttnn.experimental.rotary_embedding_llama` (adjacent-pair rotation with a 32x32 transformation tile). To reproduce HF's `rotate_half` pairing (i, i+36) the Q and K output columns of the QKV weight are permuted into interleaved order (i, i+36, ...) once at load time; attention scores are invariant to that permutation of Q and K. * MLP intermediate 4304 is zero-padded to 4320 (tile multiple); fc1 has the GELU(tanh) fused. * Patch embedding is a linear (Conv3d with stride == kernel) run on device; the bilinear position embedding and rotary cos/sin tables are built on host (torch, cached per grid). """ from __future__ import annotations import math import time import torch import ttnn from .config import ModelConfig, VisionConfig from .weights import VLM_VIS, TTWeights, linear_weight from .vision_ref import cu_seqlens, patch_embed_weight, pos_embed_interpolate, rot_pos_freqs TILE = 32 def _ceil_to(x: int, m: int) -> int: return (x + m - 1) // m * m def _interleave_perm(head_dim: int) -> torch.Tensor: """HF rotate_half order [x0..x35 | x36..x71] -> interleaved [x0, x36, x1, x37, ...].""" half = head_dim // 2 perm = torch.empty(head_dim, dtype=torch.long) perm[0::2] = torch.arange(half) perm[1::2] = torch.arange(half) + half return perm def rope_transformation_mat() -> torch.Tensor: """32x32 tile used by rotary_embedding_llama: rotated = x @ T gives (-x_odd, x_even) per pair.""" t = torch.zeros(1, 1, TILE, TILE) t[..., torch.arange(0, TILE, 2), torch.arange(1, TILE, 2)] = 1.0 t[..., torch.arange(1, TILE, 2), torch.arange(0, TILE, 2)] = -1.0 return t def build_qkv_weight(w: torch.Tensor, b: torch.Tensor, vc: VisionConfig, head_dim_p: int ) -> tuple[torch.Tensor, torch.Tensor]: """HF qkv [3*hidden, hidden], bias [3*hidden] -> W' [hidden, 3*heads*head_dim_p], b' [1, ...]. Q and K head columns permuted to interleaved rotate_half order and zero-padded to head_dim_p; V head columns kept in order and zero-padded. """ H, nh, hd = vc.hidden, vc.n_heads, vc.head_dim wt = linear_weight(w) # [hidden, 3*hidden] out_w = torch.zeros(H, 3 * nh * head_dim_p, dtype=w.dtype) out_b = torch.zeros(3 * nh * head_dim_p, dtype=b.dtype) perm = _interleave_perm(hd) ident = torch.arange(hd) for s in range(3): idx = perm if s < 2 else ident for h in range(nh): src = s * H + h * hd + idx dst = slice((s * nh + h) * head_dim_p, (s * nh + h) * head_dim_p + hd) out_w[:, dst] = wt[:, src] out_b[dst] = b[src] return out_w, out_b.view(1, -1) def build_proj_weight(w: torch.Tensor, vc: VisionConfig, head_dim_p: int) -> torch.Tensor: """HF proj [hidden, hidden] -> [heads*head_dim_p, hidden] with zero rows at padded lanes.""" nh, hd = vc.n_heads, vc.head_dim wt = linear_weight(w) # [in=hidden, out] out = torch.zeros(nh * head_dim_p, wt.shape[1], dtype=w.dtype) for h in range(nh): out[h * head_dim_p:h * head_dim_p + hd] = wt[h * hd:(h + 1) * hd] return out def build_rope_tables(freqs: torch.Tensor, head_dim: int, head_dim_p: int) -> tuple[torch.Tensor, torch.Tensor]: """[N, head_dim//2] angles -> interleaved cos/sin [N, head_dim_p] (pad lanes cos=1, sin=0).""" n = freqs.shape[0] cos = torch.ones(n, head_dim_p) sin = torch.zeros(n, head_dim_p) cos[:, :head_dim] = torch.stack([freqs.cos(), freqs.cos()], -1).flatten(-2) sin[:, :head_dim] = torch.stack([freqs.sin(), freqs.sin()], -1).flatten(-2) return cos, sin class TTVision: """Vision tower: `forward(pixel_values, grid_thw) -> (image_embeds, [deepstack x3])` on device.""" def __init__(self, device, cfg: ModelConfig, ckpt, tw: TTWeights, dtype: str = "bf16", matmul_fidelity: str = "hifi2", sdpa_fidelity: str = "hifi2", sdpa_chunk: int = 128): # sdpa_chunk 128 divides 17280 (no chunk padding -> bit-exact run to run) and is as fast as 256 self.device = device self.cfg = cfg self.vc: VisionConfig = cfg.vision self.ckpt = ckpt self.tw = tw self.dtype = dtype self.head_dim_p = _ceil_to(self.vc.head_dim, TILE) # 96 self.inter_p = _ceil_to(self.vc.inter, TILE) # 4320 self.scale = self.vc.head_dim ** -0.5 self.sdpa_chunk = sdpa_chunk self.profile = False self.prof: dict[str, float] = {} self.taps: dict[str, torch.Tensor] = {} self._prep_cache: dict[tuple, dict] = {} fid = {"lofi": ttnn.MathFidelity.LoFi, "hifi2": ttnn.MathFidelity.HiFi2, "hifi4": ttnn.MathFidelity.HiFi4} self.mm_kernel = ttnn.WormholeComputeKernelConfig( math_fidelity=fid[matmul_fidelity], math_approx_mode=False, fp32_dest_acc_en=(matmul_fidelity == "hifi4"), packer_l1_acc=True) self.sdpa_kernel = ttnn.WormholeComputeKernelConfig( math_fidelity=fid[sdpa_fidelity], math_approx_mode=False, fp32_dest_acc_en=(sdpa_fidelity == "hifi4"), packer_l1_acc=True) self.mm_kernels: dict[str, object] = {} # per-linear compute kernel overrides (name -> config) self.ln_kernel = ttnn.WormholeComputeKernelConfig( math_fidelity=ttnn.MathFidelity.HiFi4, math_approx_mode=False, fp32_dest_acc_en=True, packer_l1_acc=False) self.sdpa_pc = ttnn.SDPAProgramConfig( compute_with_storage_grid_size=device.compute_with_storage_grid_size(), exp_approx_mode=False, q_chunk_size=sdpa_chunk, k_chunk_size=sdpa_chunk) # per-linear ttnn.linear kwargs (program_config / core_grid / batch split), see _default_linear_kwargs self.linear_kwargs: dict[str, dict] = self._default_linear_kwargs(device.compute_with_storage_grid_size()) self._trace = None # set by capture_trace() t0 = time.time() self._load_weights() self.load_seconds = time.time() - t0 # ------------------------------------------------------------------------------------------ # matmul program configs # ------------------------------------------------------------------------------------------ @staticmethod def _mc2d(cx, cy, per_core_m, per_core_n, sub_h, sub_w, block_w, activation=None, fuse_batch=False): fa = None if activation == "gelu_tanh": fa = ttnn.UnaryWithParam(ttnn.UnaryOpType.GELU_TANH) elif activation == "gelu": fa = ttnn.UnaryWithParam(ttnn.UnaryOpType.GELU, 0.0) return ttnn.MatmulMultiCoreReuseMultiCastProgramConfig( compute_with_storage_grid_size=(cx, cy), in0_block_w=block_w, out_subblock_h=sub_h, out_subblock_w=sub_w, per_core_M=per_core_m, per_core_N=per_core_n, transpose_mcast=False, fused_activation=fa, fuse_batch=fuse_batch) def _default_linear_kwargs(self, grid) -> dict[str, dict]: """Explicit configs for the 17280-row linears (24 images of 20x36). They bound the per-core L1 circular-buffer footprint to <= ~1 MB (the fused decode layers of the text model keep ~100-150 KB of L1 buffers resident at the top of L1, and the auto config for M = 540 tiles may use 1.4 MB+), by viewing the activations as [1, B, 17280/B, K] so per-core M is 540/(B*10) tiles. Measured on Blackhole (11x10 grid), bf16 weights, HiFi2: qkv 1.46 ms, fc1 2.14 ms, fc2 1.03 ms, proj 0.50 ms (auto configs: 1.56 / 2.49 / 1.07 / 0.61). Other M (other image counts / grids) fall back to the auto config, which fit with 96 KB reserved.""" cfg: dict[str, dict] = {} if grid.x < 9 or grid.y < 10: return cfg cfg["qkv@17280"] = {"batch": 6, "program_config": self._mc2d(9, 10, 9, 16, 1, 8, 6)} # ~950 KB cfg["fc1@17280"] = {"batch": 6, "program_config": self._mc2d(9, 10, 9, 15, 1, 5, 4, "gelu_tanh")} # ~715 KB cfg["fc2@17280"] = {"batch": 3, "program_config": self._mc2d(9, 10, 18, 4, 2, 4, 5)} # ~600 KB proj = {"batch": 3, "program_config": self._mc2d(9, 10, 18, 4, 2, 4, 8)} # ~865 KB cfg["proj@17280"] = proj cfg["patch_embed@17280"] = dict(proj) # same K=1536, N=1152 shape return cfg # ------------------------------------------------------------------------------------------ # weights # ------------------------------------------------------------------------------------------ def _get(self, name: str) -> torch.Tensor: return self.ckpt.get(VLM_VIS + name) def _dev(self, name: str, t: torch.Tensor, dtype: str | None = None): return self.tw.to_device("visual." + name, t.contiguous(), dtype or self.dtype) def _ln_params(self, name: str, dim: int): w = self._get(name + ".weight").view(1, 1, 1, dim).expand(1, 1, TILE, dim) b = self._get(name + ".bias").view(1, 1, 1, dim).expand(1, 1, TILE, dim) return self._dev(name + ".w32", w, "bf16"), self._dev(name + ".b32", b, "bf16") def _load_weights(self): vc = self.vc hp = self.head_dim_p # patch embed (Conv3d == linear over 1536) self.w_pe = self._dev("patch_embed.w", linear_weight(patch_embed_weight(self._get("patch_embed.proj.weight")))) self.b_pe = self._dev("patch_embed.b", self._get("patch_embed.proj.bias").view(1, -1), "bf16") self.pos_table = self._get("pos_embed.weight") # host, bf16 [2304, 1152] self.trans_mat = ttnn.from_torch(rope_transformation_mat().to(torch.bfloat16), device=self.device, layout=ttnn.TILE_LAYOUT, dtype=ttnn.bfloat16, memory_config=ttnn.DRAM_MEMORY_CONFIG) self.blocks = [] for i in range(vc.depth): p = f"blocks.{i}." wqkv, bqkv = build_qkv_weight(self._get(p + "attn.qkv.weight"), self._get(p + "attn.qkv.bias"), vc, hp) fc1_w = torch.zeros(vc.hidden, self.inter_p, dtype=torch.bfloat16) fc1_w[:, :vc.inter] = linear_weight(self._get(p + "mlp.linear_fc1.weight")) fc1_b = torch.zeros(1, self.inter_p, dtype=torch.bfloat16) fc1_b[0, :vc.inter] = self._get(p + "mlp.linear_fc1.bias") fc2_w = torch.zeros(self.inter_p, vc.hidden, dtype=torch.bfloat16) fc2_w[:vc.inter] = linear_weight(self._get(p + "mlp.linear_fc2.weight")) blk = { "ln1": self._ln_params(p + "norm1", vc.hidden), "ln2": self._ln_params(p + "norm2", vc.hidden), "wqkv": self._dev(p + f"attn.qkv.w_il{hp}", wqkv), "bqkv": self._dev(p + f"attn.qkv.b_il{hp}", bqkv, "bf16"), "wo": self._dev(p + f"attn.proj.w_p{hp}", build_proj_weight(self._get(p + "attn.proj.weight"), vc, hp)), "bo": self._dev(p + "attn.proj.b", self._get(p + "attn.proj.bias").view(1, -1), "bf16"), "w1": self._dev(p + f"mlp.fc1.w_p{self.inter_p}", fc1_w), "b1": self._dev(p + f"mlp.fc1.b_p{self.inter_p}", fc1_b, "bf16"), "w2": self._dev(p + f"mlp.fc2.w_p{self.inter_p}", fc2_w), "b2": self._dev(p + "mlp.fc2.b", self._get(p + "mlp.linear_fc2.bias").view(1, -1), "bf16"), } self.blocks.append(blk) self.merger = self._load_merger("merger", postshuffle=False) self.deep_mergers = [self._load_merger(f"deepstack_merger_list.{j}", postshuffle=True) for j in range(len(vc.deepstack_indexes))] def _load_merger(self, name: str, postshuffle: bool) -> dict: vc = self.vc return { "postshuffle": postshuffle, "ln": self._ln_params(name + ".norm", vc.merged_dim if postshuffle else vc.hidden), "w1": self._dev(name + ".fc1.w", linear_weight(self._get(name + ".linear_fc1.weight"))), "b1": self._dev(name + ".fc1.b", self._get(name + ".linear_fc1.bias").view(1, -1), "bf16"), "w2": self._dev(name + ".fc2.w", linear_weight(self._get(name + ".linear_fc2.weight"))), "b2": self._dev(name + ".fc2.b", self._get(name + ".linear_fc2.bias").view(1, -1), "bf16"), } # ------------------------------------------------------------------------------------------ # host-side per-grid preparation (position embeddings, rope tables, window boundaries) # ------------------------------------------------------------------------------------------ def _prepare(self, grid_thw: torch.Tensor) -> dict: key = tuple(grid_thw.flatten().tolist()) hit = self._prep_cache.get(key) if hit is not None: return hit vc = self.vc cu = cu_seqlens(grid_thw) n = cu[-1] n_pad = _ceil_to(n, TILE) if n_pad > n: cu = cu + [n_pad] # padding rows attend only to themselves pos = pos_embed_interpolate(self.pos_table, grid_thw, vc).to(torch.bfloat16) # [n, hidden] pos_p = torch.zeros(1, 1, n_pad, vc.hidden, dtype=torch.bfloat16) pos_p[0, 0, :n] = pos cos, sin = build_rope_tables(rot_pos_freqs(grid_thw, vc), vc.head_dim, self.head_dim_p) cos_p = torch.ones(1, 1, n_pad, self.head_dim_p) sin_p = torch.zeros(1, 1, n_pad, self.head_dim_p) cos_p[0, 0, :n] = cos sin_p[0, 0, :n] = sin dev = lambda t: ttnn.from_torch(t.to(torch.bfloat16), device=self.device, layout=ttnn.TILE_LAYOUT, dtype=ttnn.bfloat16, memory_config=ttnn.DRAM_MEMORY_CONFIG) prep = { "n": n, "n_pad": n_pad, "n_merged": n // (vc.merge ** 2), "cu": cu, "pos": dev(pos_p), "cos": dev(cos_p), "sin": dev(sin_p), "cu_tt": ttnn.from_torch(torch.tensor(cu, dtype=torch.int32), device=self.device, layout=ttnn.ROW_MAJOR_LAYOUT, dtype=ttnn.uint32), } self._prep_cache[key] = prep return prep # ------------------------------------------------------------------------------------------ # device graph # ------------------------------------------------------------------------------------------ def _t(self, name: str, fn): """Run fn(); when profiling, synchronize and accumulate time under `name`.""" if not self.profile: return fn() ttnn.synchronize_device(self.device) t0 = time.time() out = fn() ttnn.synchronize_device(self.device) self.prof[name] = self.prof.get(name, 0.0) + time.time() - t0 return out def _linear(self, name, x, w, b, activation=None): """ttnn.linear with an optional per-linear config from self.linear_kwargs: {"program_config": ..., "core_grid": ..., "batch": B}. "@" keys apply only to that M. "batch": B runs the matmul on x viewed as [1, B, M/B, K] (fuse_batch=False program configs), which divides the per-core output block by B and bounds the L1 circular-buffer footprint. A program_config carries its own fused activation, so `activation` is dropped in that case.""" m, k = x.shape[-2], x.shape[-1] cfg = self.linear_kwargs.get(f"{name}@{m}", self.linear_kwargs.get(name, {})) kw = {kk: v for kk, v in cfg.items() if kk != "batch"} nb = cfg.get("batch", 1) act = None if "program_config" in kw else activation kernel = self.mm_kernels.get(name, self.mm_kernel) # per-linear fidelity override def run(): xin = x if nb == 1 else ttnn.reshape(x, (1, nb, m // nb, k)) out = ttnn.linear(xin, w, bias=b, activation=act, dtype=ttnn.bfloat16, compute_kernel_config=kernel, memory_config=ttnn.DRAM_MEMORY_CONFIG, **kw) if nb != 1: out = ttnn.reshape(out, (1, 1, m, out.shape[-1])) return out return self._t(name, run) def _layer_norm(self, name, x, params): w, b = params return self._t(name, lambda: ttnn.layer_norm(x, epsilon=self.vc.ln_eps, weight=w, bias=b, compute_kernel_config=self.ln_kernel, memory_config=ttnn.DRAM_MEMORY_CONFIG)) def _attention(self, blk, h, prep): qkv = self._linear("qkv", h, blk["wqkv"], blk["bqkv"]) q, k, v = self._t("create_heads", lambda: ttnn.experimental.nlp_create_qkv_heads( qkv, num_heads=self.vc.n_heads, num_kv_heads=self.vc.n_heads, transpose_k_heads=False, memory_config=ttnn.DRAM_MEMORY_CONFIG)) ttnn.deallocate(qkv) q_r = self._t("rope", lambda: ttnn.experimental.rotary_embedding_llama( q, prep["cos"], prep["sin"], self.trans_mat, is_decode_mode=False)) ttnn.deallocate(q) k_r = self._t("rope", lambda: ttnn.experimental.rotary_embedding_llama( k, prep["cos"], prep["sin"], self.trans_mat, is_decode_mode=False)) ttnn.deallocate(k) attn = self._t("sdpa", lambda: ttnn.transformer.scaled_dot_product_attention( q_r, k_r, v, is_causal=False, scale=self.scale, program_config=self.sdpa_pc, compute_kernel_config=self.sdpa_kernel, cu_window_seqlens=prep["cu_tt"])) ttnn.deallocate(q_r) ttnn.deallocate(k_r) ttnn.deallocate(v) cat = self._t("concat_heads", lambda: ttnn.experimental.nlp_concat_heads(attn, memory_config=ttnn.DRAM_MEMORY_CONFIG)) ttnn.deallocate(attn) out = self._linear("proj", cat, blk["wo"], blk["bo"]) ttnn.deallocate(cat) return out def _block(self, blk, x, prep): h = self._layer_norm("ln", x, blk["ln1"]) a = self._attention(blk, h, prep) ttnn.deallocate(h) x2 = self._t("add", lambda: ttnn.add(x, a, memory_config=ttnn.DRAM_MEMORY_CONFIG)) ttnn.deallocate(x) ttnn.deallocate(a) h = self._layer_norm("ln", x2, blk["ln2"]) m = self._linear("fc1", h, blk["w1"], blk["b1"], activation="gelu_tanh") ttnn.deallocate(h) m2 = self._linear("fc2", m, blk["w2"], blk["b2"]) ttnn.deallocate(m) x3 = self._t("add", lambda: ttnn.add(x2, m2, memory_config=ttnn.DRAM_MEMORY_CONFIG)) ttnn.deallocate(x2) ttnn.deallocate(m2) return x3 def _merge_rows(self, x, prep): """[1,1,n_pad,hidden] -> [1,1,n/4, 4*hidden] (drop padding rows first if any).""" n, n_pad = prep["n"], prep["n_pad"] d = self.vc.merged_dim if n_pad == n and (n // (self.vc.merge ** 2)) % TILE == 0: return self._t("merge_reshape", lambda: ttnn.reshape(x, (1, 1, n // (self.vc.merge ** 2), d))) # slow generic path (row-major round trip); only hit when N is not a multiple of 128 def slow(): rm = ttnn.to_layout(x, ttnn.ROW_MAJOR_LAYOUT) if n_pad != n: rm = ttnn.slice(rm, [0, 0, 0, 0], [1, 1, n, self.vc.hidden]) rm = ttnn.reshape(rm, (1, 1, n // (self.vc.merge ** 2), d)) return ttnn.to_layout(rm, ttnn.TILE_LAYOUT) return self._t("merge_reshape_slow", slow) def _run_merger(self, mg, x, prep): if mg["postshuffle"]: xm = self._merge_rows(x, prep) h = self._layer_norm("merger_ln", xm, mg["ln"]) ttnn.deallocate(xm) else: hn = self._layer_norm("merger_ln", x, mg["ln"]) h = self._merge_rows(hn, prep) if h is not hn: ttnn.deallocate(hn) m = self._linear("merger_fc1", h, mg["w1"], mg["b1"], activation="gelu") ttnn.deallocate(h) out = self._linear("merger_fc2", m, mg["w2"], mg["b2"]) ttnn.deallocate(m) return out def _tap(self, taps, name, x, prep): if taps is not None and name in taps: self.taps[name] = ttnn.to_torch(x)[0, 0, :prep["n"]].float() def _px_host(self, pixel_values: torch.Tensor, prep: dict) -> ttnn.Tensor: """pixel_values [N, 1536] -> host bf16 ROW_MAJOR tensor [1, 1, n_pad, 1536] (zero rows for padding). Row-major so the host does no tilize work (53 MB for 24 images); the device tilizes in ~1 ms.""" n, n_pad = prep["n"], prep["n_pad"] assert pixel_values.shape[0] == n, (pixel_values.shape, n) px = pixel_values.to(torch.bfloat16) # HF casts pixels to the conv weight dtype (bf16) if n_pad != n: px = torch.cat([px, torch.zeros(n_pad - n, px.shape[1], dtype=px.dtype)], 0) return ttnn.from_torch(px.view(1, 1, n_pad, -1).contiguous(), layout=ttnn.ROW_MAJOR_LAYOUT, dtype=ttnn.bfloat16) def _embed_device(self, px_tt: ttnn.Tensor, prep: dict) -> ttnn.Tensor: """Patch embedding (linear) + position embedding, on device -> [1,1,n_pad,hidden].""" px_tile = self._t("tilize", lambda: ttnn.tilize(px_tt, memory_config=ttnn.DRAM_MEMORY_CONFIG)) pe = self._linear("patch_embed", px_tile, self.w_pe, self.b_pe) ttnn.deallocate(px_tile) x = self._t("add", lambda: ttnn.add(pe, prep["pos"], memory_config=ttnn.DRAM_MEMORY_CONFIG)) ttnn.deallocate(pe) return x def _forward_device(self, px_tt: ttnn.Tensor, prep: dict, taps: set | None = None ) -> tuple[ttnn.Tensor, list[ttnn.Tensor]]: """Whole tower on device from a resident pixel tensor. Used directly and inside trace capture.""" x = self._embed_device(px_tt, prep) self._tap(taps, "embed", x, prep) deep = [] for i, blk in enumerate(self.blocks): x = self._block(blk, x, prep) self._tap(taps, f"block{i}.out", x, prep) if i in self.vc.deepstack_indexes: j = self.vc.deepstack_indexes.index(i) deep.append(self._run_merger(self.deep_mergers[j], x, prep)) out = self._run_merger(self.merger, x, prep) ttnn.deallocate(x) return out, deep def embed(self, pixel_values: torch.Tensor, grid_thw: torch.Tensor): """Patch embedding + position embedding on device -> ([1,1,n_pad,hidden], prep).""" prep = self._prepare(grid_thw) px_tt = self._t("h2d_pixels", lambda: ttnn.to_device(self._px_host(pixel_values, prep), self.device, memory_config=ttnn.DRAM_MEMORY_CONFIG)) x = self._embed_device(px_tt, prep) ttnn.deallocate(px_tt) return x, prep def forward(self, pixel_values: torch.Tensor, grid_thw: torch.Tensor, taps: set | None = None ) -> tuple[ttnn.Tensor, list[ttnn.Tensor]]: """pixel_values [N, 1536] (any float dtype), grid_thw [n_images, 3] -> image_embeds [1,1,Nm,5120] bf16 TILE DRAM and 3 deepstack tensors of the same shape. Uses the captured trace when one exists for this grid (see capture_trace); the returned tensors are then owned by the trace and must not be deallocated by the caller.""" self.taps = {} if taps is None and self._trace is not None and self._trace["key"] == self._grid_key(grid_thw): return self.forward_traced(pixel_values, grid_thw) prep = self._prepare(grid_thw) px_tt = self._t("h2d_pixels", lambda: ttnn.to_device(self._px_host(pixel_values, prep), self.device, memory_config=ttnn.DRAM_MEMORY_CONFIG)) out, deep = self._forward_device(px_tt, prep, taps) ttnn.deallocate(px_tt) return out, deep def forward_torch(self, pixel_values: torch.Tensor, grid_thw: torch.Tensor, taps: set | None = None ) -> tuple[torch.Tensor, list[torch.Tensor]]: """Same as forward but returns float32 torch tensors [Nm, 5120] (device tensors are freed).""" out, deep = self.forward(pixel_values, grid_thw, taps) nm = self._prepare(grid_thw)["n_merged"] to = lambda t: ttnn.to_torch(t)[0, 0, :nm].float() out_t = to(out) deep_t = [to(d) for d in deep] if not self._owned_by_trace(out): ttnn.deallocate(out) for d in deep: ttnn.deallocate(d) return out_t, deep_t # ------------------------------------------------------------------------------------------ # trace (replays the whole op graph without host dispatch gaps) # ------------------------------------------------------------------------------------------ @staticmethod def _grid_key(grid_thw: torch.Tensor) -> tuple: return tuple(grid_thw.flatten().tolist()) def _owned_by_trace(self, t) -> bool: return self._trace is not None and (t is self._trace["out"] or any(t is d for d in self._trace["deep"])) def capture_trace(self, pixel_values: torch.Tensor, grid_thw: torch.Tensor) -> None: """Compile + capture the forward for this grid. Needs the device opened with a trace_region_size (>= ~32 MB). The pixel tensor is persistent on device; forward_traced() refills it and replays.""" self.release_trace() prep = self._prepare(grid_thw) host = self._px_host(pixel_values, prep) px_dev = ttnn.allocate_tensor_on_device(host.shape, ttnn.bfloat16, ttnn.ROW_MAJOR_LAYOUT, self.device, ttnn.DRAM_MEMORY_CONFIG) ttnn.copy_host_to_device_tensor(host, px_dev) out, deep = self._forward_device(px_dev, prep) # compile pass ttnn.synchronize_device(self.device) ttnn.deallocate(out) for d in deep: ttnn.deallocate(d) tid = ttnn.begin_trace_capture(self.device, cq_id=0) out, deep = self._forward_device(px_dev, prep) ttnn.end_trace_capture(self.device, tid, cq_id=0) ttnn.synchronize_device(self.device) self._trace = {"key": self._grid_key(grid_thw), "tid": tid, "px": px_dev, "out": out, "deep": deep, "prep": prep} def forward_traced(self, pixel_values: torch.Tensor, grid_thw: torch.Tensor ) -> tuple[ttnn.Tensor, list[ttnn.Tensor]]: tr = self._trace assert tr is not None and tr["key"] == self._grid_key(grid_thw), "no trace captured for this grid" host = self._t("h2d_pixels", lambda: self._px_host(pixel_values, tr["prep"])) self._t("h2d_copy", lambda: ttnn.copy_host_to_device_tensor(host, tr["px"])) self._t("execute_trace", lambda: ttnn.execute_trace(self.device, tr["tid"], cq_id=0, blocking=False)) return tr["out"], tr["deep"] def release_trace(self) -> None: if self._trace is not None: ttnn.release_trace(self.device, self._trace["tid"]) ttnn.deallocate(self._trace["px"]) self._trace = None # ------------------------------------------------------------------------------------------ # timing # ------------------------------------------------------------------------------------------ def timed_forward(self, pixel_values: torch.Tensor, grid_thw: torch.Tensor, iters: int = 1, traced: bool = False) -> tuple[float, float]: """Returns (wall seconds incl. host tilize + host->device pixel copy, device seconds after the inputs are resident) per iteration, measured with synchronize_device around the forward. Call once beforehand to warm program caches (or capture_trace() when traced=True).""" prep = self._prepare(grid_thw) ttnn.synchronize_device(self.device) wall = dev_only = 0.0 for _ in range(iters): t0 = time.time() if traced: tr = self._trace ttnn.copy_host_to_device_tensor(self._px_host(pixel_values, prep), tr["px"]) ttnn.synchronize_device(self.device) t1 = time.time() ttnn.execute_trace(self.device, tr["tid"], cq_id=0, blocking=False) ttnn.synchronize_device(self.device) t2 = time.time() else: px_tt = ttnn.to_device(self._px_host(pixel_values, prep), self.device, memory_config=ttnn.DRAM_MEMORY_CONFIG) ttnn.synchronize_device(self.device) t1 = time.time() out, deep = self._forward_device(px_tt, prep) ttnn.synchronize_device(self.device) t2 = time.time() ttnn.deallocate(px_tt) ttnn.deallocate(out) for d in deep: ttnn.deallocate(d) wall += t2 - t0 dev_only += t2 - t1 return wall / iters, dev_only / iters