tt-hous's picture
Add files using upload-large-folder tool
b025706 verified
Raw History Blame Contribute Delete
53.9 kB
# SPDX-FileCopyrightText: © 2024 Tenstorrent USA, Inc.
# SPDX-License-Identifier: Apache-2.0
import math
import torch
from tqdm import tqdm
import ttnn
from models.common.lightweightmodule import LightweightModule
from models.common.rmsnorm import RMSNorm
from models.common.sampling.generator import SamplingGenerator
from models.common.sampling.tt_sampling import TOPK_MAX_WIDTH, TTSampling
from models.tt_transformers.tt.ccl import TT_CCL
from models.tt_transformers.tt.common import Mode, copy_host_to_device
from models.tt_transformers.tt.decoder import TransformerBlock
from models.tt_transformers.tt.distributed_norm import DistributedNorm
from models.tt_transformers.tt.embedding import Embedding, ScaledEmbedding
from models.tt_transformers.tt.lm_head import LMHead
from models.tt_transformers.tt.model_config import TensorGroup
from models.tt_transformers.tt.rope import HfRotarySetup, RotarySetup
def _get_trace_rope_table_len(max_seq_len, trace_prefill_seq_lens):
if not trace_prefill_seq_lens:
return max_seq_len
slice_alignment = math.lcm(*trace_prefill_seq_lens)
min_table_len = max_seq_len + max(trace_prefill_seq_lens)
return ((min_table_len + slice_alignment - 1) // slice_alignment) * slice_alignment
def _prefill_rope_setups_to_pad(rope_setup, rope_local_setup, rope_setup_class):
"""Select the RoPE setups whose shared prefill tables Transformer slices.
A caller-supplied rope_setup_class builds its prefill cosine and sine mats
per request on the host, inside its own prepare_inputs_prefill, and its
forward never calls Transformer._slice_prefill_rot_mats. Such a setup owns
no cos_matrix_prefill to pad. rope_local_setup always comes from the
built-in classes, so it always owns one.
"""
rope_setups = [] if rope_setup_class is not None else [rope_setup]
if rope_local_setup is not None:
rope_setups.append(rope_local_setup)
return rope_setups
def _pad_prefill_rope_tables(rope_setups, max_seq_len, trace_prefill_seq_lens):
table_len = _get_trace_rope_table_len(max_seq_len, trace_prefill_seq_lens)
pad_len = table_len - max_seq_len
if pad_len == 0:
return
padding = [(0, 0), (0, 0), (0, pad_len), (0, 0)]
for rope_setup in rope_setups:
rope_setup.cos_matrix_prefill = ttnn.pad(
rope_setup.cos_matrix_prefill,
padding=padding,
value=0.0,
)
rope_setup.sin_matrix_prefill = ttnn.pad(
rope_setup.sin_matrix_prefill,
padding=padding,
value=0.0,
)
class Transformer(LightweightModule):
def __init__(
self,
args,
dtype,
mesh_device,
state_dict,
weight_cache_path,
paged_attention_config=None,
use_paged_kv_cache=False,
attention_class=None,
rope_setup_class=None,
block_class=None,
lm_head_cls=None,
final_norm_builder=None,
prefetcher=None,
):
super().__init__()
self.args = args
self.vocab_size = args.vocab_size
assert self.vocab_size > 0
self.n_layers = args.n_layers
self.mesh_device = mesh_device
self.dtype = dtype
self.model_config = args.get_model_config()
self.grid_size = self.args.max_grid_size
state_dict_prefix = args.get_state_dict_prefix("", None)
self.decoders_optimizations = args.decoders_optimizations
self.prefetcher = prefetcher
self.tt_ccl = TT_CCL(self.mesh_device)
# Runtime bounds for the post-prefill tail's slice. Allocated here, before any trace exists,
# and rewritten in place per call - see process_logits_after_prefill_trace.
# These buffers belong to this model/DP lane. Calls on a lane enqueue the
# copy and slice on the same command queue, in order. Concurrent host
# calls on the same model instance are not supported.
self._tail_slice_start = ttnn.from_torch(
torch.zeros(4, dtype=torch.int32),
device=self.mesh_device,
mesh_mapper=ttnn.ReplicateTensorToMesh(self.mesh_device),
)
self._tail_slice_end = ttnn.from_torch(
torch.zeros(4, dtype=torch.int32),
device=self.mesh_device,
mesh_mapper=ttnn.ReplicateTensorToMesh(self.mesh_device),
)
embd_kwargs = {
"mesh_device": mesh_device,
"args": args,
"weight_cache_path": args.weight_cache_path(dtype),
"state_dict": state_dict,
"dtype": ttnn.bfloat16, # Row major layout requires bfloat16
}
if self.args.embed_scale is not None:
embd_cls = ScaledEmbedding
embd_kwargs["embed_scale"] = self.args.embed_scale
else:
embd_cls = Embedding
self.embd = embd_cls(**embd_kwargs)
DefaultRopeSetup = HfRotarySetup if self.args.use_hf_rope else RotarySetup
ActualRopeSetupClass = rope_setup_class if rope_setup_class is not None else DefaultRopeSetup
# NoPE global layers (EXAONE-4.x): full-attention layers apply no rotary at
# all, so the global setup's cos/sin are neutralized to the identity. Only
# the Meta-style RotarySetup implements this.
use_global_nope = getattr(args, "use_global_nope", False)
if use_global_nope and self.args.use_hf_rope:
raise NotImplementedError("use_global_nope (NoPE global layers) requires the Meta-style RotarySetup")
global_rope_kwargs = {"nope": True} if use_global_nope else {}
self.rope_setup = ActualRopeSetupClass(
device=mesh_device,
batch_size=args.max_batch_size,
head_dim=args.head_dim,
max_seq_len=args.max_seq_len,
rope_theta=args.rope_theta,
rope_scaling=args.rope_scaling,
use_qk_fused=args.use_qk_fused,
prefetcher=prefetcher,
**global_rope_kwargs,
)
if args.rope_theta_local:
self.rope_local_setup = DefaultRopeSetup(
mesh_device,
args.max_batch_size,
args.head_dim,
args.max_seq_len,
args.rope_theta_local,
# Most hybrid models (Gemma-3) use unscaled rope on sliding layers;
# EXAONE-4.x instead applies its llama3-scaled rope there and sets
# rope_scaling_local (the global layers being NoPE).
rope_scaling=getattr(args, "rope_scaling_local", None),
use_qk_fused=args.use_qk_fused,
prefetcher=None,
)
# Dynamic starts share one table across fixed-width trace buckets. The
# tail prevents out-of-range reads and the common multiple preserves
# the tensor-bound slice partition geometry for every traced length.
_pad_prefill_rope_tables(
_prefill_rope_setups_to_pad(
self.rope_setup,
getattr(self, "rope_local_setup", None),
rope_setup_class,
),
args.max_seq_len,
args.trace_prefill_supported_seq_lens,
)
self.trans_mats_dict = self.rope_setup.get_both_trans_mats()
# Device tensors used to build dynamic slice params for prefill RoPE slicing.
# Keeps chunk_start_idx-driven slicing inside the traced graph.
self._tt_seq_len_buffer = ttnn.from_torch(
torch.tensor([1, 1, self.args.max_seq_len, self.args.head_dim], dtype=torch.int32),
device=self.mesh_device,
mesh_mapper=ttnn.ReplicateTensorToMesh(self.mesh_device),
)
self._tt_slice_start_zeros_4 = ttnn.from_torch(
torch.tensor([0, 0, 0, 0], dtype=torch.int32),
device=self.mesh_device,
mesh_mapper=ttnn.ReplicateTensorToMesh(self.mesh_device),
)
# Model-family dispatch (Command-R / cohere): swap the decoder block, final
# norm and LM head. Lazy imports keep the experimental cohere module out of the default path.
ActualBlockClass = block_class
ActualLMHeadCls = lm_head_cls
final_norm_builder_resolved = final_norm_builder
if str(getattr(self.args, "model_type", None) or "").lower() == "cohere":
from models.experimental.cohere.tt.cohere_decoder import CohereDecoderLayer
from models.experimental.cohere.tt.cohere_lm_head import CohereLMHead
from models.experimental.cohere.tt.cohere_norm import build_cohere_final_norm
ActualBlockClass = ActualBlockClass or CohereDecoderLayer
ActualLMHeadCls = ActualLMHeadCls or CohereLMHead
final_norm_builder_resolved = final_norm_builder_resolved or build_cohere_final_norm
ActualBlockClass = ActualBlockClass or TransformerBlock
ActualLMHeadCls = ActualLMHeadCls or LMHead
self.layers = [
ActualBlockClass(
args=args,
mesh_device=mesh_device,
tt_ccl=self.tt_ccl,
dtype=dtype,
state_dict=state_dict,
weight_cache_path=weight_cache_path,
layer_num=i,
transformation_mats=self.trans_mats_dict,
paged_attention_config=paged_attention_config,
use_paged_kv_cache=use_paged_kv_cache,
attention_class=attention_class,
prefetcher=prefetcher,
)
for i in tqdm(range(self.n_layers))
]
self.norm = (
final_norm_builder_resolved(
args=args,
mesh_device=mesh_device,
state_dict=state_dict,
weight_cache_path=weight_cache_path,
dtype=dtype,
tt_ccl=self.tt_ccl,
)
if final_norm_builder_resolved is not None
else DistributedNorm(
RMSNorm(
device=mesh_device,
dim=args.dim,
eps=args.norm_eps,
state_dict=state_dict,
state_dict_prefix=args.get_state_dict_prefix("", None),
weight_cache_path=None if args.dummy_weights else weight_cache_path,
weight_dtype=ttnn.bfloat16,
weight_key="norm",
add_unit_offset=self.args.rms_norm_add_unit_offset,
is_distributed=self.args.is_distributed_norm,
ccl_topology=self.args.ccl_topology(),
tt_ccl=self.tt_ccl,
),
args,
tt_ccl=self.tt_ccl,
prefetcher=prefetcher,
TG=args.is_galaxy,
)
) # close the final_norm_builder_resolved conditional-expression paren
self.lm_head = ActualLMHeadCls(
args=args,
mesh_device=mesh_device,
tt_ccl=self.tt_ccl,
dtype=dtype,
state_dict=state_dict,
state_dict_prefix=state_dict_prefix,
weight_cache_path=weight_cache_path,
max_columns_per_device=self.args.max_columns_per_device_lm_head,
prefetcher=prefetcher,
)
# Initialize on-device sampling if supported
# Sampling on device is supported only if each device holds at most TOPK_MAX_WIDTH logits.
# On a single device TTSampling cuts the padded vocab into as many same-device chunks as
# needed (power-of-two, each <= TOPK_MAX_WIDTH), so any vocab it can cut tile-aligned is
# supported (#53064); anything it cannot falls back to host sampling.
padded_vocab_size = getattr(self.args, "padded_vocab_size", None) or self.args.vocab_size
if list(self.mesh_device.shape) != [1, 1]:
vocab_fits_on_device = padded_vocab_size // self.args.num_devices <= TOPK_MAX_WIDTH
else:
vocab_fits_on_device = TTSampling.num_single_device_vocab_splits(padded_vocab_size) is not None
self._supports_on_device_sampling = prefetcher is None and vocab_fits_on_device
if self._supports_on_device_sampling:
self.sampling = SamplingGenerator(
args=args,
mesh_device=mesh_device,
tt_ccl=self.tt_ccl,
)
else:
self.sampling = None
def update_weights(
self,
hf_state_dict: dict[str, ttnn.Tensor],
*,
hf_rope: bool = False,
) -> None:
"""In-place replace every weight from an HF-keyed dict of on-device 4D
ttnn tensors (replicated, DRAM-interleaved, TILE, bf16). Keys follow HF
safetensors naming; shapes are HF Linear/gamma/embedding wrapped in two
leading unit dims.
Strict by construction: every required key must be present (missing ->
``KeyError``) and every provided key consumed by exactly one leaf
``.update()`` (extras -> ``ValueError``). No "loose" mode -- silent
partial updates are an expensive class of bug.
``hf_rope=False`` (default): caller has already permuted Q/K rows into
this model's convention (right for the ttml -> TTT transfer, both store
Meta-permuted rows). ``hf_rope=True`` defers HF -> Meta permutation to
``Attention.update`` (currently raises -- kernel not wired up).
Tied embeddings: the protocol still requires both
``model.embed_tokens.weight`` and ``lm_head.weight`` (typically the same
source tensor), keeping dispatch one-to-one with device buffers.
Every existing buffer keeps its device allocation, so captured traces
and the prefetcher's recorded addresses stay valid.
"""
unconsumed = set(hf_state_dict.keys())
def consume(key: str) -> ttnn.Tensor:
if key not in hf_state_dict:
raise KeyError(f"Transformer.update_weights: missing required HF key {key!r}")
unconsumed.discard(key)
return hf_state_dict[key]
# Top-level (always required).
self.embd.update(embed_tokens=consume("model.embed_tokens.weight"))
self.norm.update(weight=consume("model.norm.weight"))
self.lm_head.update(weight=consume("lm_head.weight"))
# Per-layer: prefix-strip into a layer-local dict, dispatch.
for i, block in enumerate(self.layers):
prefix = f"model.layers.{i}."
layer_dict = {}
for key in list(hf_state_dict.keys()):
if key.startswith(prefix):
layer_dict[key[len(prefix) :]] = hf_state_dict[key]
unconsumed.discard(key)
block.update_weights(layer_dict, hf_rope=hf_rope)
if unconsumed:
sample = sorted(unconsumed)[:10]
raise ValueError(
f"Transformer.update_weights: {len(unconsumed)} HF key(s) not "
f"consumed by any leaf .update(). This usually means a typo, "
f"a stray weight, or a layer-index off-by-one. "
f"Showing up to 10: {sample}"
)
def process_logits_after_prefill_trace(self, logits, last_token_idx):
get_last_token = (last_token_idx // 32) * 32
seq_len = int(logits.shape[-2])
# Pass the offset as a runtime argument rather than a compile-time attribute. With literal
# bounds every distinct prompt offset compiles its own slice program, and since this runs
# after the prefill traces are captured - once per data-parallel group - that was the single
# largest source of buffers left live across trace replays on a DP run. Warmup cannot cover
# it either: it only ever sees bucket-length mock prompts, and real prompts are shorter.
#
# The tensor-args path needs the slice tile-aligned, which this one already is: it takes 32
# rows starting at a multiple of 32. num_devices splits the sequence into equal parts, so
# seq_len // 32 gives exactly the 32-row window, and the program then keys on the padded
# prefill bucket instead of the offset.
if seq_len % 32 == 0:
for device_tensor, values in (
(self._tail_slice_start, [0, 0, get_last_token, 0]),
(self._tail_slice_end, [1, 1, get_last_token + 32, int(logits.shape[-1])]),
):
ttnn.copy_host_to_device_tensor(
ttnn.from_torch(
torch.tensor(values, dtype=torch.int32),
mesh_mapper=ttnn.ReplicateTensorToMesh(self.mesh_device),
),
device_tensor,
)
logits = ttnn.slice(
input_tensor=logits,
starts=self._tail_slice_start,
ends=self._tail_slice_end,
slice_dim=2,
num_devices=seq_len // 32,
)
else:
logits = ttnn.slice(
logits,
(0, 0, get_last_token, 0),
(1, 1, get_last_token + 32, logits.shape[-1]),
)
logits = self._apply_norm_and_lm_head(logits)
return logits
def extract_last_tokens_batched_prefill(
self, hidden_states, last_token_idx_list, padded_batch, prefill_seq_len, target_batch=None
):
"""Extract each user's last-token hidden state from batched prefill output.
Reads hidden states to host, extracts the relevant row for each user,
and sends the combined tensor back to device with the correct column-sharded
mesh mapping (ShardTensorToMesh dim=-1) so the DistributedNorm all-gather
produces the correct full hidden dim.
Args:
hidden_states: [padded_batch, 1, prefill_seq_len, dim_per_device] on device (column-sharded, TILE_LAYOUT)
last_token_idx_list: list of length padded_batch with per-user last token positions
padded_batch: number of slots (typically 32)
prefill_seq_len: padded sequence length per user
Returns:
user_tokens: [1, 1, target_batch or padded_batch, dim_per_device] per device,
column-sharded, TILE_LAYOUT
"""
active_indices = [lt for lt in last_token_idx_list if lt > 0]
all_same = len(set(active_indices)) <= 1
if all_same and active_indices:
common_last = active_indices[0]
get_last = (common_last // 32) * 32
R = common_last % 32
block = ttnn.slice(
hidden_states,
(0, 0, get_last, 0),
(padded_batch, 1, get_last + 32, hidden_states.shape[-1]),
)
else:
block = hidden_states
R = None
host_tensors = [ttnn.to_torch(dt) for dt in ttnn.get_device_tensors(block)]
host_full = torch.cat(host_tensors, dim=-1)
if R is not None:
combined = host_full[:, :, R : R + 1, :].reshape(1, 1, padded_batch, -1).contiguous()
else:
rows = []
for slot in range(padded_batch):
lt_idx = last_token_idx_list[slot]
rows.append(host_full[slot : slot + 1, :, lt_idx : lt_idx + 1, :])
combined = torch.cat(rows, dim=0).reshape(1, 1, padded_batch, -1).contiguous()
target_batch = padded_batch if target_batch is None else target_batch
if target_batch < padded_batch:
raise ValueError(f"target_batch {target_batch} must be >= padded_batch {padded_batch}")
if target_batch > padded_batch:
padded_combined = torch.zeros(
1,
1,
target_batch,
combined.shape[-1],
dtype=combined.dtype,
)
padded_combined[:, :, :padded_batch, :] = combined
combined = padded_combined
user_tokens = ttnn.from_torch(
combined,
device=self.mesh_device,
dtype=ttnn.bfloat16,
layout=ttnn.TILE_LAYOUT,
mesh_mapper=ttnn.ShardTensorToMesh(self.mesh_device, dim=-1),
)
return user_tokens
def process_logits_after_batched_prefill(self, hidden_states, last_token_idx_list, padded_batch, prefill_seq_len):
"""Extract last tokens and run norm + lm_head once for all users."""
user_tokens = self.extract_last_tokens_batched_prefill(
hidden_states, last_token_idx_list, padded_batch, prefill_seq_len
)
return self._apply_norm_and_lm_head(user_tokens)
def _apply_final_logit_softcapping(self, logits):
"""Gemma-2 final logit soft-capping: logits -> tanh(logits / cap) * cap.
No-op unless args.final_logit_softcapping is set (only Gemma-2 sets it), so
this leaves every other model's output path unchanged.
"""
cap = self.args.final_logit_softcapping
if cap is None or cap <= 0:
return logits
logits = ttnn.multiply(logits, 1.0 / cap)
logits = ttnn.tanh(logits)
logits = ttnn.multiply(logits, cap)
return logits
def _apply_norm_and_lm_head(self, x):
"""Shared norm + lm_head for prefill logit processing. Input: [1, 1, 32, hidden_dim]."""
x = self.norm(
x, mode=Mode.PREFILL, norm_config=self.args.get_norm_config("lm_head", Mode.PREFILL, self.prefetcher)
)
lm_head_input_mem_cfg = self.args.get_lm_head_input_mem_config(Mode.PREFILL, None)
if lm_head_input_mem_cfg.is_sharded():
x = ttnn.interleaved_to_sharded(x, lm_head_input_mem_cfg)
logits = self.lm_head(x)
logits = self._apply_final_logit_softcapping(logits)
logits = ttnn.to_memory_config(logits, memory_config=ttnn.DRAM_MEMORY_CONFIG)
return logits
def process_hidden_states_after_prefill_trace(self, hidden_states, last_token_idx):
"""
Process hidden states after prefill trace, stopping before LM head.
Returns hidden states (after norm) instead of logits.
Used for embedding models that need hidden states rather than logits.
"""
get_last_token = (last_token_idx // 32) * 32
hidden_states = ttnn.slice(
hidden_states,
(0, 0, get_last_token, 0),
(1, 1, get_last_token + 32, hidden_states.shape[-1]),
)
# Apply norm (this is the final layer norm before LM head)
hidden_states = self.norm(hidden_states, mode="prefill")
# Convert to row major layout for output (but don't apply LM head)
hidden_states = ttnn.to_layout(
hidden_states, layout=ttnn.ROW_MAJOR_LAYOUT, memory_config=ttnn.DRAM_MEMORY_CONFIG
)
return hidden_states
def prepare_prefill_inputs_trace(
self,
tokens,
page_table=None,
chunk_page_table=None,
chunk_start_idx=0,
batch_size=1,
user_id=0,
**kwargs,
):
"""
Inputs are torch tensors or python types. This function returns ttnn
tensors on host.
"""
host_inputs = self.prepare_inputs_prefill(
tokens,
page_table=page_table,
chunk_page_table=chunk_page_table,
chunk_start_idx=chunk_start_idx,
trace_enabled=True,
batch_size=batch_size,
user_id=user_id,
)
return host_inputs
def transform_and_embed_prefill_inputs_device(
self,
tokens,
tt_page_table,
tt_chunk_page_table,
tt_chunk_start_idx,
):
tt_tokens = self.embd(tokens)
tt_tokens = ttnn.unsqueeze_to_4D(tt_tokens)
return tt_tokens, tt_page_table, tt_chunk_page_table, tt_chunk_start_idx
def prepare_inputs_prefill(
self,
tokens,
start_pos=0,
page_table=None,
chunk_page_table=None,
chunk_start_idx=None,
trace_enabled=False,
last_token_idx=None,
global_user_id=None,
batch_size=1,
user_id=0,
**kwargs,
):
"""
Inputs are torch tensors or python types. This function returns ttnn
tensors on device if trace is disabled or on host if trace is enabled.
TODO: Debate whether this function is responsible for padding
"""
# We set the device to None if trace is enabled so we keep the tensors on host instead of sending it to the device (None - keeps on host, device - sends to specified device)
# We will send them to device later (copy_host_to_device)
device = None if trace_enabled else self.mesh_device
assert tokens.dim() == 2, "tokens must be a 2D tensor"
# For batched prefill, tokens come in as [padded_batch, S]
# Each user's tokens are at their slot index in dimension 0
# Reshape to [1, 1, 1, padded_batch * S] for embedding
if batch_size > 1:
# Tokens are in slot-based format [padded_batch, S_per_user]
S = tokens.shape[-1] # Per-user sequence length
tokens = tokens.reshape(1, 1, 1, -1) # Flatten to [1, 1, 1, padded_batch * S]
else:
tokens = tokens.reshape(1, 1, 1, -1)
S = tokens.shape[-1]
tokens = ttnn.from_torch(
tokens,
device=device,
dtype=ttnn.uint32,
layout=ttnn.ROW_MAJOR_LAYOUT,
mesh_mapper=ttnn.ReplicateTensorToMesh(self.mesh_device),
)
# self.embd expects that tokens are on device ; if trace is enabled, the tensors will be later on device, so we will do these 2 steps when we copy the tokens to the device
if not trace_enabled:
tokens_embd = self.embd(tokens)
tokens_embd = ttnn.unsqueeze_to_4D(tokens_embd)
# Slice the rot mats to the prefill seqlen
trace_mat_len = self.rope_setup.cos_matrix_prefill.shape[2]
mat_len = self.args.max_seq_len
seq_len = last_token_idx + 1 if last_token_idx is not None else S
assert mat_len >= seq_len, f"Sequence length {seq_len} exceeds max seq len {mat_len}"
required_end = start_pos + S
pad_len = max(0, required_end - mat_len)
# We set the end_pos to max_seq_len so that we don't create a new tensor for the whole cos_matrix and sin_matrix
# In case of trace, we will use the whole matrix for all seq_lens supported by trace
prefill_start_pos = 0 if trace_enabled else start_pos
slice_end = trace_mat_len if trace_enabled else min(mat_len, required_end)
cos_slice = self.rope_setup.cos_matrix_prefill[:, :, prefill_start_pos:slice_end, :]
sin_slice = self.rope_setup.sin_matrix_prefill[:, :, prefill_start_pos:slice_end, :]
if pad_len > 0:
# Padding: [(before, after), ...] for each dim; pad at end of 3rd dim (dim=2) by pad_len
padding = [(0, 0)] * 4
padding[2] = (0, pad_len)
cos_slice = ttnn.pad(cos_slice, padding=padding, value=0.0)
sin_slice = ttnn.pad(sin_slice, padding=padding, value=0.0)
tt_rot_mats_prefill_global = [cos_slice, sin_slice]
if hasattr(self, "rope_local_setup"):
local_trace_mat_len = self.rope_local_setup.cos_matrix_prefill.shape[2]
local_mat_len = self.args.max_seq_len
local_required_end = start_pos + S
local_pad_len = max(0, local_required_end - local_mat_len)
local_slice_end = local_trace_mat_len if trace_enabled else min(local_mat_len, local_required_end)
local_cos_slice = self.rope_local_setup.cos_matrix_prefill[:, :, prefill_start_pos:local_slice_end, :]
local_sin_slice = self.rope_local_setup.sin_matrix_prefill[:, :, prefill_start_pos:local_slice_end, :]
if local_pad_len > 0:
# Pad at end of 3rd dim (dim=2) by local_pad_len
local_padding = [(0, 0)] * 4
local_padding[2] = (0, local_pad_len)
local_cos_slice = ttnn.pad(local_cos_slice, padding=local_padding, value=0.0)
local_sin_slice = ttnn.pad(local_sin_slice, padding=local_padding, value=0.0)
tt_rot_mats_prefill_local = [local_cos_slice, local_sin_slice]
else:
tt_rot_mats_prefill_local = None
if page_table is not None:
# For batched prefill, replicate page_table to all devices (same as single-user path)
# The KV cache fill will loop over users and use batch_idx=user_id for each
tt_page_table = ttnn.from_torch(
page_table,
device=device,
dtype=ttnn.int32,
layout=ttnn.ROW_MAJOR_LAYOUT,
mesh_mapper=ttnn.ReplicateTensorToMesh(self.mesh_device),
)
else:
tt_page_table = None
if chunk_page_table is not None:
tt_chunk_page_table = ttnn.from_torch(
chunk_page_table,
device=device,
dtype=ttnn.int32,
layout=ttnn.ROW_MAJOR_LAYOUT,
mesh_mapper=ttnn.ReplicateTensorToMesh(self.mesh_device),
)
else:
tt_chunk_page_table = None
if chunk_start_idx is not None and int(chunk_start_idx) > 0:
chunk_start_idx_tensor = torch.tensor([chunk_start_idx], dtype=torch.int32)
tt_chunk_start_idx = ttnn.from_torch(
chunk_start_idx_tensor,
device=device,
dtype=ttnn.int32,
mesh_mapper=ttnn.ReplicateTensorToMesh(self.mesh_device),
)
else:
tt_chunk_start_idx = None
return (
tokens if trace_enabled else tokens_embd,
tt_rot_mats_prefill_global,
tt_rot_mats_prefill_local,
tt_page_table,
tt_chunk_page_table,
tt_chunk_start_idx,
)
def prepare_inputs_decode(self, *inputs):
"""
Inputs are torch tensors or python types. This function returns ttnn
tensors on device.
Its implementation can take advantage of a few other functions which the
model must implement.
"""
host_inputs = self.prepare_decode_inputs_host(*inputs)
device_inputs = copy_host_to_device(host_inputs, mesh_device=self.mesh_device) # Helper function
return device_inputs
def prepare_decode_inputs_host(self, tokens, current_pos, page_table=None):
"""
Inputs are torch tensors or python types. Outputs are ttnn tensors on host.
NOTE: Tokens and current_pos are padded to batch
"""
B = tokens.shape[0]
assert current_pos.shape[0] == B, "Batch size mismatch"
assert (
B == self.args.max_batch_size
), f"Batch size {B} must be equal to max_batch_size {self.args.max_batch_size}"
# Necessary padding to be full tile sized when on device
tokens = torch.nn.functional.pad(tokens.view(-1), (0, 32 - len(tokens)), "constant", 0)
tokens = ttnn.from_torch(
tokens,
device=None,
dtype=ttnn.uint32,
mesh_mapper=ttnn.ReplicateTensorToMesh(self.mesh_device),
)
tokens = ttnn.unsqueeze_to_4D(tokens)
rot_current_pos = torch.maximum(
current_pos, torch.tensor(0, dtype=torch.int64)
) # Ensure position indices are non-negative
rope_idxs = self.rope_setup.get_rot_idxs(rot_current_pos, on_host=True)
current_pos_tt = ttnn.from_torch(
current_pos,
device=None,
dtype=ttnn.int32,
mesh_mapper=ttnn.ShardTensor2dMesh(
self.mesh_device,
dims=(None, 0) if (self.args.is_galaxy and B > 1) else (None, None),
mesh_shape=self.args.cluster_shape,
),
)
if page_table is not None:
page_table = ttnn.from_torch(
page_table,
device=None,
dtype=ttnn.int32,
mesh_mapper=ttnn.ShardTensor2dMesh(
self.mesh_device,
dims=(None, -2) if (self.args.is_galaxy and B > 1) else (None, None),
mesh_shape=self.args.cluster_shape,
),
)
return tokens, current_pos_tt, rope_idxs, page_table
def _transform_decode_inputs_device(
self,
tokens,
):
"""
Inputs are ttnn tensors on device. This function applies any on-device
transformations which should happen before forward decode.
For example: tilize, reshape, shard.
Return transformed device tensors
Embed tokens
"""
decode_residual_mem_cfg = self.args.get_residual_mem_config(Mode.DECODE, self.prefetcher)
tt_tokens = self.embd(
tokens,
memory_config=ttnn.DRAM_MEMORY_CONFIG if self.prefetcher is None else decode_residual_mem_cfg,
)
tt_tokens = ttnn.unsqueeze_to_4D(tt_tokens)
tt_tokens = ttnn.to_memory_config(tt_tokens, decode_residual_mem_cfg)
return tt_tokens
def concat_host_output(self, tt_out, is_log_probs=False):
"""
Concatenate the output of the devices into a single host tensor.
"""
torch_out_tensors = [ttnn.to_torch(x) for x in ttnn.get_device_tensors(tt_out)]
if self.args.is_galaxy:
row_dim, col_dim = (3, 1)
else:
row_dim, col_dim = (1, -1)
rows, cols = self.args.cluster_shape
mesh_shape = [torch_out_tensors[i : i + cols] for i in range(0, len(torch_out_tensors), cols)]
if is_log_probs:
row_concatenated = []
for row in mesh_shape:
row_reshaped = [tensor.reshape(1, 1, -1, 1) for tensor in row]
row_concatenated.append(torch.cat(row_reshaped, dim=col_dim))
else:
row_concatenated = [torch.cat(row, dim=col_dim) for row in mesh_shape]
return torch.cat(row_concatenated, dim=row_dim)
def process_output_prefill(self, tt_out, last_token_idx):
"""
Input is ttnn host tensor of logits. Output is torch logits tensor.
NOTE: In this model, prefill always uses get_last_token
"""
assert tt_out.storage_type() == ttnn.StorageType.HOST, "Expected host tensor"
return self.concat_host_output(tt_out)[0, 0, last_token_idx, : self.vocab_size]
def process_output_prefill_hidden_states(self, tt_out, last_token_idx):
"""
Input is ttnn host tensor of hidden states (after norm, before LM head).
Output is torch hidden states tensor of shape [hidden_size].
Used for embedding models.
"""
assert tt_out.storage_type() == ttnn.StorageType.HOST, "Expected host tensor"
# Extract the last token's hidden state
# Shape: [batch=1, head=1, seq, hidden_dim] -> [hidden_dim]
# For hidden states, if they're replicated across devices (not sharded),
# we should take just the first device's output to avoid incorrect concatenation.
# If sharded, concat_host_output will properly concatenate them.
concatenated = self.concat_host_output(tt_out)
# Check if concatenation resulted in oversized tensor (replicated case)
# If so, take only the first device's portion (first self.args.dim elements)
if concatenated.shape[-1] > self.args.dim:
# Hidden states are replicated, take first device's output
return concatenated[0, 0, last_token_idx, : self.args.dim]
else:
# Hidden states are sharded, concatenation is correct
return concatenated[0, 0, last_token_idx, :]
def process_output_decode(self, tt_out, B, S=1, is_tokens=False, is_log_probs=False):
"""
Input is ttnn host tensor of logits if is_tokens=False, otherwise tokens. Output is the corresponding torch tensor.
"""
if is_tokens or is_log_probs:
# Pad to 32 to match the expected batch size for decode operations (tiles are 32x32)
padded_batch_size = 32
if not is_log_probs:
tt_out = ttnn.reshape(tt_out, ttnn.Shape([1, 1, padded_batch_size, 1]))
return self.concat_host_output(tt_out, is_log_probs)[0, 0, :B, 0]
if self.args.num_devices > 1:
tt_out = ttnn.to_torch(ttnn.get_device_tensors(tt_out)[0]).float()
else:
tt_out = ttnn.to_torch(tt_out).float()
tt_out = tt_out[:, :, :B, : self.vocab_size].view(B, S, -1)
return tt_out
def ttnn_prefill_forward(
self,
x,
rot_mats_global=None,
rot_mats_local=None,
user_id=0,
page_table=None,
chunk_page_table=None,
chunk_start_idx=None,
get_last_token=-1,
kv_cache=None,
batch_size=1,
page_tables_per_layer=None,
):
"""
This method will take device tensors and any other args to run forward.
It returns ttnn device tensors.
"""
if page_tables_per_layer is None:
# vLLM hybrid bridges (HybridAttentionForCausalLM subclasses) stash
# the per-layer list on the model handle for the duration of a
# forward call rather than threading the kwarg through Generator's
# many ttnn_prefill_forward call sites. Pick it up here when set.
page_tables_per_layer = getattr(self, "_active_page_tables_per_layer", None)
page_tables_per_layer = self._page_tables_to_ttnn(page_tables_per_layer)
return self.forward(
x,
current_pos=None,
rot_mats_global=rot_mats_global,
rot_mats_local=rot_mats_local,
user_id=user_id,
mode=Mode.PREFILL,
page_table=page_table,
chunk_page_table=chunk_page_table,
chunk_start_idx=chunk_start_idx,
get_last_token=get_last_token,
kv_cache=kv_cache,
batch_size=batch_size,
page_tables_per_layer=page_tables_per_layer,
)
def _page_table_mesh_mapper(self, B):
"""Mesh mapper for per-layer page tables, matching the layout that
:meth:`prepare_decode_inputs_host` uses for the legacy single
``page_table`` kwarg: shard the batch dim across mesh axis 1 on
Galaxy when ``B>1``, replicate otherwise. The hybrid bridge
chunks the global page table per-DP before calling into a
submesh, so ``B`` here is the per-DP batch — same value the
legacy path sees on entry to ``prepare_decode_inputs_host``.
"""
return ttnn.ShardTensor2dMesh(
self.mesh_device,
dims=(None, -2) if (self.args.is_galaxy and B > 1) else (None, None),
mesh_shape=self.args.cluster_shape,
)
def _page_tables_to_ttnn(self, page_tables_per_layer):
"""Resolve a per-layer list of ``torch.Tensor`` page tables to a
list of *persistent* ttnn device tensors (allocate-only).
Tracing bakes each input tensor's device address into the captured
graph; replaying the trace reads from those exact addresses
regardless of any new ttnn objects created on the Python side.
Allocating fresh device tensors on every call would therefore
make traced inference read stale memory at the original
addresses, so we lazily allocate one persistent device tensor per
layer on first use and *only* update contents from outside the
traced ``ttnn_*_forward`` calls (writes are forbidden during trace
capture). The hybrid bridge calls
:meth:`update_persistent_per_layer_page_tables` *before* invoking
``Generator``'s decode/prefill which executes traces — that's
where content updates happen.
First call (warmup compile) populates the persistent buffers from
the input torch tensors; subsequent calls return the existing
buffers unchanged. ``None`` entries propagate; already-ttnn
entries pass through.
"""
if page_tables_per_layer is None:
return None
persistent = getattr(self, "_persistent_per_layer_page_tables", None)
n = len(page_tables_per_layer)
if persistent is None or len(persistent) != n:
persistent = []
for pt in page_tables_per_layer:
if pt is None:
persistent.append(None)
continue
if isinstance(pt, ttnn.Tensor):
persistent.append(pt)
continue
persistent.append(
ttnn.from_torch(
pt,
device=self.mesh_device,
dtype=ttnn.int32,
layout=ttnn.ROW_MAJOR_LAYOUT,
mesh_mapper=self._page_table_mesh_mapper(pt.shape[0]),
)
)
self._persistent_per_layer_page_tables = persistent
return persistent
def update_persistent_per_layer_page_tables(self, page_tables_per_layer):
"""Update content of persistent per-layer page_table device
tensors in place. Called by the hybrid bridge *before* invoking
``Generator``'s decode/prefill so traced replay observes the new
block IDs at the captured addresses. Must be called outside trace
capture (writes forbidden inside).
No-op if persistent tensors haven't been allocated yet (first
call goes through :meth:`_page_tables_to_ttnn`'s allocation).
"""
if page_tables_per_layer is None:
return
persistent = getattr(self, "_persistent_per_layer_page_tables", None)
if persistent is None or len(persistent) != len(page_tables_per_layer):
return
for i, pt in enumerate(page_tables_per_layer):
if pt is None or persistent[i] is None or isinstance(pt, ttnn.Tensor):
continue
host_pt = ttnn.from_torch(
pt,
device=None,
dtype=ttnn.int32,
layout=ttnn.ROW_MAJOR_LAYOUT,
mesh_mapper=self._page_table_mesh_mapper(pt.shape[0]),
)
ttnn.copy_host_to_device_tensor(host_pt, persistent[i])
def _increment_decode_positions_device(self, current_pos, rot_mat_idxs):
ttnn.plus_one(current_pos, skip_negative_entries=True)
ttnn.plus_one(rot_mat_idxs)
def _slice_prefill_rot_mats(self, rot_mats, chunk_start_idx, prefill_seq_len):
"""Slice full prefill RoPE mats to the traced prefill sequence length."""
if rot_mats is None or chunk_start_idx is None or not isinstance(chunk_start_idx, ttnn.Tensor):
return rot_mats
full_rot_cos, full_rot_sin = rot_mats[0], rot_mats[1]
full_seq_len = full_rot_cos.shape[2]
if prefill_seq_len <= 0:
raise ValueError(f"Prefill sequence length must be positive, got {prefill_seq_len}")
if full_rot_sin.shape[2] != full_seq_len:
raise ValueError(
f"Prefill RoPE cosine and sine sequence lengths must match, got "
f"{full_seq_len} and {full_rot_sin.shape[2]}"
)
if full_seq_len == prefill_seq_len:
return rot_mats
if full_seq_len % prefill_seq_len != 0:
raise ValueError(
f"Full RoPE sequence length {full_seq_len} must be evenly divisible by "
f"prefill sequence length {prefill_seq_len}"
)
# Tensor-bound slice fixes output geometry as input length divided by
# num_devices; this argument is a partition count, not the mesh width.
num_partitions = full_seq_len // prefill_seq_len
z = self._tt_slice_start_zeros_4
tt_slice_starts = ttnn.concat([z[0:2], chunk_start_idx, z[3:4]], dim=0)
rot_cos_slice = ttnn.slice(
input_tensor=full_rot_cos,
starts=tt_slice_starts,
ends=self._tt_seq_len_buffer,
slice_dim=2,
num_devices=num_partitions,
)
rot_sin_slice = ttnn.slice(
input_tensor=full_rot_sin,
starts=tt_slice_starts,
ends=self._tt_seq_len_buffer,
slice_dim=2,
num_devices=num_partitions,
)
return (rot_cos_slice, rot_sin_slice)
def ttnn_decode_forward(
self,
x,
current_pos,
rot_mat_idxs=None,
page_table=None,
kv_cache=None,
on_device_logits=False,
page_tables_per_layer=None,
):
"""
This method will take device tensors and any other args to run forward.
It returns ttnn device tensors.
"""
rot_mats_global = self.rope_setup.get_rot_mats(rot_mat_idxs)
rot_mats_local = self.rope_local_setup.get_rot_mats(rot_mat_idxs) if hasattr(self, "rope_local_setup") else None
x_embed = self._transform_decode_inputs_device(x)
if page_tables_per_layer is None:
# See ttnn_prefill_forward: hybrid bridges stash the per-layer list
# on the model when active, since Generator doesn't thread the kwarg.
page_tables_per_layer = getattr(self, "_active_page_tables_per_layer", None)
page_tables_per_layer = self._page_tables_to_ttnn(page_tables_per_layer)
tt_logits = self.forward(
x_embed,
current_pos,
rot_mats_global=rot_mats_global,
rot_mats_local=rot_mats_local,
mode=Mode.DECODE,
page_table=page_table,
kv_cache=kv_cache,
page_tables_per_layer=page_tables_per_layer,
)
if on_device_logits:
assert self.sampling is not None, (
"ttnn_decode_forward got on_device_logits=True but no on-device sampling "
"module exists (self.sampling is None)."
)
self._increment_decode_positions_device(current_pos, rot_mat_idxs)
return tt_logits
# Gather the output across all devices and untilize the tensor (for argmax)
if self.args.num_devices > 1:
cluster_axis = 0 if self.args.is_galaxy else None
num_links = 2 if self.args.is_galaxy else 1
tt_logits = ttnn.experimental.all_gather_async(
tt_logits,
persistent_output_buffer=None,
dim=3,
multi_device_global_semaphore=self.tt_ccl.get_and_cycle_ag_semaphore_handles(cluster_axis),
num_links=num_links,
memory_config=tt_logits.memory_config() if self.prefetcher is None else ttnn.DRAM_MEMORY_CONFIG,
cluster_axis=cluster_axis,
topology=self.args.ccl_topology(),
barrier_semaphore=self.tt_ccl.get_and_cycle_barrier_semaphore_handle(cluster_axis),
chunks_per_sync=10,
num_workers_per_link=2,
num_buffers_per_channel=2,
subdevice_id=self.prefetcher.worker_sub_device_id if self.prefetcher is not None else None,
)
tt_logits = ttnn.untilize(
tt_logits,
use_multicore=True,
memory_config=ttnn.DRAM_MEMORY_CONFIG,
sub_core_grids=self.prefetcher.all_worker_cores_range_set if self.prefetcher is not None else None,
)
return tt_logits, None
def switch_mode(self, mode: Mode):
if self.prefetcher is not None:
self.prefetcher.init(mode)
self.prefetcher.prefetch()
def forward(
self,
x: ttnn.Tensor,
current_pos,
rot_mats_global=None,
rot_mats_local=None,
user_id=0,
mode: Mode = Mode.DECODE,
page_table=None,
chunk_page_table=None,
chunk_start_idx=None,
get_last_token=-1,
kv_cache=None,
batch_size=1,
page_tables_per_layer=None,
):
if mode == Mode.DECODE:
# Run prefetcher if it is enabled
if self.prefetcher is not None:
self.prefetcher.run()
if mode == Mode.PREFILL:
# For traced prefill, keep RoPE slicing in-graph and driven by the
# on-device chunk_start_idx input. Batched prefill arrives flattened
# to [1, 1, batch_size * S_per_user, dim] and each TransformerBlock
# restores the batch dimension before attention, so the RoPE slice
# width is the per-user length, not the flattened one.
prefill_seq_len = x.shape[2] // batch_size
rot_mats_global = self._slice_prefill_rot_mats(rot_mats_global, chunk_start_idx, prefill_seq_len)
if rot_mats_local is not None:
rot_mats_local = self._slice_prefill_rot_mats(rot_mats_local, chunk_start_idx, prefill_seq_len)
if page_tables_per_layer is not None and len(page_tables_per_layer) != len(self.layers):
raise ValueError(
f"page_tables_per_layer has {len(page_tables_per_layer)} entries "
f"but model has {len(self.layers)} layers"
)
for i, layer in enumerate(self.layers):
# No-op if callers already provide the right memory config
activation_dtype = self.args.decoders_optimizations.get_tensor_dtype(
decoder_id=i, tensor=TensorGroup.ACTIVATION
)
if mode == Mode.DECODE and not self.args.is_galaxy:
x = ttnn.to_memory_config(
x,
self.args.get_residual_mem_config(mode, self.prefetcher),
activation_dtype,
)
elif activation_dtype is not None and x.dtype != activation_dtype:
x = ttnn.typecast(x, activation_dtype)
# vLLM hybrid kv-cache-groups: each attention layer gets its own
# paged pool (sliding-window vs full-attention have different
# block counts). When ``page_tables_per_layer`` is None we fall
# back to broadcasting the single ``page_table`` to every layer
# — byte-equivalent to the pre-hybrid path used by every legacy
# caller (demos, unit tests, non-hybrid vLLM bridges).
layer_page_table = page_tables_per_layer[i] if page_tables_per_layer is not None else page_table
x = layer(
x,
current_pos,
rot_mats_global=rot_mats_global,
rot_mats_local=rot_mats_local,
user_id=user_id,
mode=mode,
page_table=layer_page_table,
chunk_page_table=chunk_page_table,
chunk_start_idx=chunk_start_idx,
kv_cache=kv_cache[i] if kv_cache is not None else None,
batch_size=batch_size,
)
if mode == Mode.DECODE:
if self.prefetcher is not None:
self.prefetcher.stop()
if mode == Mode.PREFILL and get_last_token == -1:
return x
# Slicing the tensor to the nearest ceiling/floor multiples of 32 for the prefill_len, to get the last token
if get_last_token != -1:
seq_len = int(x.shape[2])
if seq_len % 32 == 0:
# Runtime bounds, as in process_logits_after_prefill_trace: with literal bounds every
# distinct prompt offset compiles its own slice program. Untraced-prefill models
# (Gemma-3) reach this slice on every real request, after warmup has recorded their
# decode traces, and warmup only ever sees bucket-length mock prompts - measured on
# Gemma-3-27B DP-4 as the last 8 buffers left live across trace replays.
for device_tensor, values in (
(self._tail_slice_start, [0, 0, get_last_token, 0]),
(self._tail_slice_end, [1, 1, get_last_token + 32, int(x.shape[-1])]),
):
ttnn.copy_host_to_device_tensor(
ttnn.from_torch(
torch.tensor(values, dtype=torch.int32),
mesh_mapper=ttnn.ReplicateTensorToMesh(self.mesh_device),
),
device_tensor,
)
x = ttnn.slice(
input_tensor=x,
starts=self._tail_slice_start,
ends=self._tail_slice_end,
slice_dim=2,
num_devices=seq_len // 32,
)
else:
x = ttnn.slice(x, (0, 0, get_last_token, 0), (1, 1, get_last_token + 32, x.shape[-1]))
# Output norm
x = self.norm(x, mode=mode, norm_config=self.args.get_norm_config("lm_head", mode, self.prefetcher))
lm_head_input_mem_cfg = self.args.get_lm_head_input_mem_config(
mode, None if mode == Mode.PREFILL else self.prefetcher
)
if mode == Mode.PREFILL and lm_head_input_mem_cfg.is_sharded():
x = ttnn.interleaved_to_sharded(x, lm_head_input_mem_cfg)
if mode == Mode.DECODE and self.prefetcher is not None:
x = ttnn.to_memory_config(x, self.args.get_lm_head_input_mem_config(mode, self.prefetcher))
x = self.lm_head(x)
x = self._apply_final_logit_softcapping(x)
if mode == Mode.PREFILL:
x = ttnn.to_memory_config(x, memory_config=ttnn.DRAM_MEMORY_CONFIG)
return x