changh95's picture
Add files using upload-large-folder tool
0190e6b verified
Raw History Blame Contribute Delete
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
# ------------------------------------------------------------------------------------------
@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}. "<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)
# ------------------------------------------------------------------------------------------
@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