Download code/alpamayo_tt/vision.py from changh95/Alpamayo2-Super-p300x2: direct link, hf CLI and curl.
- Browser
- Download file 29.4 kB
-
https://huggingface.co/changh95/Alpamayo2-Super-p300x2/resolve/main/code/alpamayo_tt/vision.py
- Command line
-
hf download hf://changh95/Alpamayo2-Super-p300x2/code/alpamayo_tt/vision.py
-
curl -L -o vision.py https://huggingface.co/changh95/Alpamayo2-Super-p300x2/resolve/main/code/alpamayo_tt/vision.py
29.4 kB
| """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 | |
| # ------------------------------------------------------------------------------------------ | |
| 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}. "<name>@<rows>" 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) | |
| # ------------------------------------------------------------------------------------------ | |
| 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 | |