# SPDX-FileCopyrightText: © 2024 Tenstorrent USA, Inc. # SPDX-License-Identifier: Apache-2.0 import math import os from collections import defaultdict import torch from loguru import logger from ttnn.tools import trace_allocation_tracker import ttnn from models.common.llama_models import ( CompletionMessage, StopReason, TokenResult, create_vision_mask, encode_content, extract_images_from_messages, sample_top_p, ) from models.common.model_capabilities import ModelCapabilitiesMixin from models.common.sampling import ( SamplingParams, broadcast_sampling_params, chunk_sampling_params, format_sampling_params, scatter_sampling_params_to_slots, ) from models.common.sampling.tt_log_probs import LogProbsResult, reformat_logprobs from models.common.warmup import WarmupForwardMixin from models.tt_transformers.tt.common import ( Mode, copy_host_to_device, get_all_padded_prefill_lengths, get_block_size, get_max_prefill_chunk_size, get_padded_prefill_len, num_blocks_in_seq, ) # Maximum total tokens (batch_size * seq_len) allowed for a batched prefill pass. # Exceeding this triggers a fallback to sequential per-user prefill. MAX_BATCHED_PREFILL_SEQ_LEN = 128 * 1024 # Power-of-2 batch sizes supported by trace caching for batched prefill. SUPPORTED_PREFILL_BATCH_SIZES = (1, 2, 4, 8, 16, 32) def batched_prefill_fits_token_budget(padded_batch, seq_len, max_prefill_chunk_size): """Bound the combined activation footprint by the model/device's prefill budget. The MLP flattens the batch and sequence dimensions. Checking each sequence alone admits e.g. 32 x 1024 on Llama-8B N150, whose single-pass budget is 4096 tokens, and runs out of DRAM even before capturing a trace. """ total_tokens = padded_batch * seq_len return total_tokens <= max_prefill_chunk_size and total_tokens < MAX_BATCHED_PREFILL_SEQ_LEN def batched_prefill_padded_batch(batch_size, empty_slots, max_batch_size): """Rows the batched-prefill device batch needs for ``empty_slots``. A batched prefill places request ``i`` at device row ``empty_slots[i]``, its physical slot, and every slot-indexed buffer (``prefill_ids``, ``padded_last_token_idx``, the padded page table) is bounded by the returned value. The batch must therefore span the highest slot in use, not just the request count: vLLM hands out the slot that already owns a request's per-slot state, so a batch of N requests can legitimately land on slots above N. A span no bucket covers returns at least the span itself, so the caller's ``padded_batch > max_batch_size`` guard fires and sends the batch down the sequential path. Reporting ``max_batch_size`` there would re-enable the very out-of-bounds slot write this bound exists to prevent. """ span = batch_size if empty_slots is not None and len(empty_slots) > 0: span = max(span, max(int(s) for s in empty_slots) + 1) return next((b for b in SUPPORTED_PREFILL_BATCH_SIZES if b >= span), max(span, max_batch_size)) def gather_batched_prefill_samples( empty_slots, tokens_host, tt_log_probs, plain_log_probs_host, output_tokens, output_log_probs, ): """Move a batched prefill's sampled rows back into the caller's prefill order. The device sampled row ``empty_slots[i]`` for request ``i`` (batched prefill is laid out by physical slot), while ``output_tokens``/``output_log_probs`` are sized by the request count and returned in prefill order. Read by slot, write by position: indexing the outputs by slot overflows them once a slot reaches the request count, and silently hands one request another's token before that. """ for local_idx, slot in enumerate(empty_slots): slot = int(slot) output_tokens[local_idx] = tokens_host[slot] if isinstance(tt_log_probs, LogProbsResult): output_log_probs[local_idx] = tt_log_probs.extract_user(slot) elif plain_log_probs_host is not None: output_log_probs[local_idx] = plain_log_probs_host[slot] # Position of the page table within the decode input tuple produced by # Transformer.prepare_decode_inputs_host: (tokens, current_pos, rope_idxs, page_table). # Used to refresh only the page-table trace input when KV blocks are reallocated. DECODE_PAGE_TABLE_INPUT_IDX = 3 def _maybe_acknowledge_trace_buffers_corruptible(owner, value): """Acknowledge opt-in trace I/O that another live trace may overwrite.""" if not getattr(owner, "_tt_allow_decode_trace_buffer_reuse", False) or value is None: return if isinstance(value, (list, tuple)): for item in value: _maybe_acknowledge_trace_buffers_corruptible(owner, item) return trace_allocation_tracker.acknowledge_corruptible(value) def max_prefill_chunk_size_cutoff(sequence_length, max_prefill_chunk_size): return sequence_length > max_prefill_chunk_size def _deepseek_kvdbg_enabled() -> bool: return os.getenv("DEEPSEEK_KVDBG", "").lower() in ("1", "true", "yes", "y") def _get_max_blocks_prefill(kv_cache): first_cache_tensor = kv_cache[0][0] return int(first_cache_tensor.shape[0]) def _pad_or_create_page_table(table, target_blocks): aligned_blocks = ((target_blocks + 7) // 8) * 8 if table is not None: num_pad = aligned_blocks - table.shape[1] if num_pad > 0: padding = torch.ones(table.shape[0], num_pad, dtype=torch.int32) * -1 return torch.cat([table, padding], dim=-1) return table return torch.ones(1, aligned_blocks, dtype=torch.int32) * -1 class Generator(ModelCapabilitiesMixin, WarmupForwardMixin): def __init__(self, model, model_args, mesh_device, processor=None, tokenizer=None): """ Creating a LlamaVision wrapper requires only a mesh_device and model_args. With model_args you have the checkpoint location, can specify max batch size and max seqlen, and other model specific parameters. LlamaVision is general to text and chat. For bringup, make this class general to any backend implementation, as long as it takes torch tensors and returns torch tensors. """ self.model = model self.model_args = model_args self.mesh_device = mesh_device self.processor = processor self.tokenizer = tokenizer self.data_parallel = len(self.model) self.trace_id_prefill = defaultdict(lambda: None) self.trace_inputs_prefill = defaultdict(lambda: None) self.trace_output_prefill = defaultdict(lambda: None) self.trace_id_prefill_sampling = defaultdict(lambda: None) self.trace_input_prefill_sampling = defaultdict(lambda: None) self.trace_output_prefill_sampling = defaultdict(lambda: None) self.trace_ids_decode = defaultdict(lambda: None) # {device_sampling_bool: {device_id: trace_id}} self.trace_inputs_decode = defaultdict(lambda: None) self.trace_output_decode = defaultdict(lambda: None) self.prefill_traces_warmup = False self.already_warmed_up_prefill = False self.mode = None # Set for the duration of the first traced prefill call: its decode trace is prepared before that # call's prefill captures anything, and recorded once the prefill is done. That call's # prefill traces are also deferred until its output processing has compiled. self._defer_trace_recording = False # Prefill-side deferral, kept separate from the decode latch above: warmup arms this across # its whole sweep so no bucket compiles behind an earlier bucket's trace. self._defer_prefill_recording = False self._pending_prefill_traces = {} self._prepared_prefill_traces = {} self._prepared_prefill_sampling_traces = {} self._pending_decode_trace = None # The eager warmup phase stages decode I/O as well as programs. Keep it # until the capture phase, which may follow prefill trace recording. self._prepared_decode_traces = {} # Class-level capabilities (VLLM specific, to be overridden by subclasses). # A subclass dict replaces this one rather than merging into it, so a default # of True would be claimed by every subclass that declares nothing at all. model_capabilities = { "supports_prefix_caching": False, } def _any_trace_captured(self): """True once any trace has been captured, i.e. once allocations are no longer unconditionally safe. Used to decide whether a prefill call may still arm capture deferral: the point of deferring is to keep every compile pass and every persistent allocation ahead of the first capture, which is only achievable while nothing has been captured yet. """ return any( any(getattr(self, name, {}).values()) for name in ("trace_id_prefill", "trace_id_prefill_sampling", "trace_ids_decode") ) or any( slot["id"] is not None for model in self.model for slot in getattr(getattr(model, "sampling", None), "_trace_states", {}).values() ) def _get_sampling_contract(self, model_id: int): sampling_module = getattr(self.model[model_id], "sampling", None) sampling_dp = getattr(self.model[model_id], "sampling_dp", 1) group_batch = sampling_module.tt_sampling.max_batch_size if sampling_module is not None else None total_sampling_batch = group_batch * sampling_dp if group_batch is not None else None return sampling_module, sampling_dp, group_batch, total_sampling_batch def _mock_tokens(self, batch_size, seq_len, kv_cache, model_id): ret = dict() ret["tokens"] = torch.zeros(batch_size, seq_len, dtype=torch.long) ret["prompt_lens"] = torch.tensor([seq_len] * batch_size, dtype=torch.long) ret["empty_slots"] = list(range(batch_size)) page_table_warmup = None # second check is some tests set the kv_cache to [None] instead of None if kv_cache is not None and kv_cache[model_id] is not None: block_size = get_block_size(kv_cache[model_id]) num_blocks = num_blocks_in_seq(seq_len, block_size) page_table_warmup = torch.zeros(batch_size, num_blocks, dtype=torch.int32) ret["page_table"] = page_table_warmup return ret def warmup_model_prefill(self, kv_cache, enable_trace, can_sample_on_device, greedy_only: bool = False): self.warmup_vision_encoder() if self.already_warmed_up_prefill: return sequence_lengths_to_warmup = self.model_args[0].get_warmup_prefill_supported_seq_lens() warmup_batch_sizes = (1,) _, sampling_dp, _, _ = self._get_sampling_contract(0) if ( self.data_parallel == 1 and not getattr(self.model_args[0], "disable_batched_prefill", False) and (not can_sample_on_device or sampling_dp == 1) and not self._overrides_prefill_capture() and not self._uses_prefetcher() ): warmup_batch_sizes = tuple( batch for batch in SUPPORTED_PREFILL_BATCH_SIZES if batch <= self.model_args[0].max_batch_size ) skip_sequence_lengths = False # Sweep all sampling parameters for prefill warmup just once since it is sequence length agnostic sampling_parameters_sweeped = False if enable_trace: logger.info("Warming traced prefill batch sizes {}", warmup_batch_sizes) # Compile every bucket before recording any of them; see _easy_trace_prefill. self._defer_prefill_recording = enable_trace and not self._overrides_prefill_capture() self.already_warmed_up_prefill = True try: self._warmup_prefill_sweep( kv_cache=kv_cache, enable_trace=enable_trace, can_sample_on_device=can_sample_on_device, greedy_only=greedy_only, sequence_lengths_to_warmup=sequence_lengths_to_warmup, warmup_batch_sizes=warmup_batch_sizes, skip_sequence_lengths=skip_sequence_lengths, sampling_parameters_sweeped=sampling_parameters_sweeped, ) # Resumed buckets must also compile before any sp0 or sp1 trace is live. resumes_prefill = self.model_capabilities.get( "supports_prefix_caching", False ) or self.model_capabilities.get("supports_chunked_prefill", False) if enable_trace and resumes_prefill: self._warmup_prefill_resumed_sweep(kv_cache=kv_cache) self._defer_prefill_recording = False self._record_pending_prefill_traces() except BaseException: self.already_warmed_up_prefill = False self._prepared_prefill_traces.clear() self._prepared_prefill_sampling_traces.clear() raise finally: self._defer_prefill_recording = False self._pending_prefill_traces.clear() def _warmup_prefill_resumed_sweep(self, kv_cache): """Capture the resumed-prefill ("sp1") traces, batch 1, one per traced length. The sweep above never passes ``start_pos``, so it only ever captures the ``sp0`` half of the prefill trace key. A resumed prefill -- prefix caching, or a prompt split across engine steps -- takes the ``sp1`` half, which would otherwise be captured lazily on whichever request happens to resume first, allocating trace region nobody sized for. Reserving them here makes the requirement a function of configuration instead of traffic. Mirrors the phase-2 warmup in ``models/demos/llama3_70b_galaxy/tt/generator.py``. """ if kv_cache is None or kv_cache[0] is None: # Resumed prefill needs a page table, so there is nothing to capture. return block_size = self._paged_prefill_block_size(kv_cache[0]) for model_id in range(self.data_parallel): model_args = self.model_args[model_id] for prefill_seq_len in model_args.trace_prefill_supported_seq_lens: # A nonzero probe: ``can_enable_trace`` reads this argument only to # test it against zero, and a model that refuses a resumed prefill # refuses it for every offset. The declaration alone is not enough, # because the model args can reject what the capability allows. if not model_args.can_enable_trace(prefill_seq_len, 1): continue # The offset a real request will be floored to for this length, so # the captured program config is the one those replays need. num_cached = self._resume_offset_alignment(prefill_seq_len, block_size, model_id) # Only the suffix after the offset is padded into the bucket, so # the prompt has to clear the offset by a full bucket to land on # this key. ``capped_warmup_seq_len`` is the ceiling the rest of # warmup uses: past it the call would be split into chunks and # capture something else. total_seq_len = self._resumed_warmup_prompt_len( prefill_seq_len, num_cached, model_args.capped_warmup_seq_len ) suffix = total_seq_len - num_cached if suffix <= 0 or get_padded_prefill_len(suffix) != prefill_seq_len: logger.warning( f"Skipping resumed prefill warmup for sequence length {prefill_seq_len}: " f"offset {num_cached} leaves no suffix padding back to it within " f"{model_args.capped_warmup_seq_len} tokens. Its trace is captured on first use." ) continue num_blocks = num_blocks_in_seq(total_seq_len, block_size) logger.info( f"Warming up resumed prefill for sequence length: {prefill_seq_len}, " f"num_cached: {num_cached}" ) self.prefill_forward_text( tokens=torch.zeros(1, total_seq_len, dtype=torch.long), prompt_lens=torch.tensor([total_seq_len], dtype=torch.long), empty_slots=[0], page_table=torch.zeros(1, num_blocks, dtype=torch.int32), start_pos=[num_cached], kv_cache=kv_cache, enable_trace=True, model_id_warmup=model_id, sampling_params=None, ) def finalize_deferred_traces(self): """Record the prefill and decode traces that this call deferred. The first traced prefill call prepares decode (compile pass, persistent inputs, sampling pre-compile) before its own prefill captures anything, then records it here once that prefill is done. Prefill recording also waits for the call's output processing to compile. Recording binds only to the already prepared inputs in the pending trace stores. """ if not self._defer_trace_recording: return try: # Prefill first: its traces were prepared during this call, and the tail that runs # between preparation and here has now compiled. self._defer_prefill_recording = False self._record_pending_prefill_traces() # Deliberately no prepare fallback here: the decode trace was prepared up front by # _prefill_forward_text_impl, before this call's prefill filled the KV cache. Preparing at # this point instead would run the decode compile pass -- a real decode step at position 0 # with mock inputs -- over the prefilled cache, overwriting real K/V with mock values. If # nothing was prepared (_prepare_decode_trace_for_warmup declined and returned None), # recording is a no-op and the decode trace is set up lazily on the first decode step, as # on main. self._record_pending_traces() finally: # Cleared even if recording raised, and the stash is dropped with it: leaving the latch armed # would make every later prefill defer a capture that nothing flushes. self._defer_trace_recording = False self._pending_decode_trace = None self._defer_prefill_recording = False self._pending_prefill_traces.clear() def _warmup_prefill_sweep( self, kv_cache, enable_trace, can_sample_on_device, greedy_only, sequence_lengths_to_warmup, warmup_batch_sizes, skip_sequence_lengths, sampling_parameters_sweeped, ): # Once per data-parallel group, not once overall. The sweep is sequence-length agnostic, # but every group is a separate device with its own program cache, so sweeping only on the # first one left groups 1..N-1 compiling the sampling path on their first real request - # after warmup had recorded its traces. Measured on DP-32: 31 stranded argmax programs. swept_sampling_model_ids = set(range(self.data_parallel)) if sampling_parameters_sweeped else set() for model_id in range(self.data_parallel): for supported_length in sequence_lengths_to_warmup: if model_id != 0 and ( supported_length not in self.model_args[0].trace_prefill_supported_seq_lens or not enable_trace ): continue # Use the same combined-token budget as runtime routing. Warming # larger batches would OOM on shapes runtime handles sequentially. for batch_size in warmup_batch_sizes: if batch_size > 1 and not batched_prefill_fits_token_budget( batch_size, supported_length, self.model_args[model_id].max_prefill_chunk_size ): logger.info( f"Skipping batched prefill warmup for batch_size={batch_size}, " f"seq_len={supported_length}: exceeds model/device token budget" ) continue warmup_args = self._mock_tokens(batch_size, supported_length, kv_cache, model_id) # chunked prefill not supported without paged attention if warmup_args["page_table"] is None and max_prefill_chunk_size_cutoff( supported_length, self.model_args[0].max_prefill_chunk_size ): logger.warning( f"Skipping warmup for sequence lengths after: {supported_length} because they are greater than the max prefill chunk size and paged attention is disabled" ) skip_sequence_lengths = True break if model_id not in swept_sampling_model_ids: sampling_params = self._create_sampling_params( can_sample_on_device=can_sample_on_device, batch_size=batch_size, greedy_only=greedy_only, ) else: # Not [None]: that path skips the on-device-sampling tail, so its # programs (last-token slice, and the untilize after it - both keyed on # the prefill bucket) never compile here and land on the first real # request instead, behind the traces warmup is about to record. sampling_params = self._create_sampling_params( can_sample_on_device=can_sample_on_device, batch_size=batch_size, greedy_only=True, ) for param in sampling_params: logger.info( f"Warming up prefill for sequence length: {supported_length} for batch size: {batch_size} with sampling params: {param}" ) self.prefill_forward_text( **warmup_args, kv_cache=kv_cache, enable_trace=enable_trace, model_id_warmup=model_id, sampling_params=param, ) swept_sampling_model_ids.add(model_id) if skip_sequence_lengths: break # Vision compile for multimodal models if getattr(self.model_args[0], "is_multimodal", False): vision_chunk_size = getattr(self.model_args[0], "vision_chunk_size", 896) vision_channels = getattr(self.model_args[0], "vision_in_channels", 3) model_id = 0 # Create synthetic image for vision warmup # pixel_values is a list (one per user), each element is (num_images, C, H, W) warmup_pixel_values = [torch.zeros((1, vision_channels, vision_chunk_size, vision_chunk_size))] # Minimal text tokens for vision warmup pass, prefill expects non-empty tokens batch_size = 1 # VLMs support only batch=1 for now prefill_forward_args = self._mock_tokens(batch_size, 128, kv_cache, model_id) logger.info(f"Warming up vision encoder with image size {vision_chunk_size}x{vision_chunk_size}") self.prefill_forward_text( **prefill_forward_args, kv_cache=kv_cache, enable_trace=False, # Vision encoder warmup doesn't support trace model_id_warmup=model_id, sampling_params=None, pixel_values=warmup_pixel_values, image_sizes=[(vision_chunk_size, vision_chunk_size)], ) logger.info("Vision encoder warmup completed") def _prepare_decode_trace_for_warmup(self, kv_cache, page_table, on_device_sampling): """Full decode trace preparation (compile pass + trace inputs), run before any trace is captured. The model has to be in decode mode for the compile pass to pick the right configs. """ # Gate lives here rather than only in _prepare_decode_trace_once because this function has a # second caller that bypasses that wrapper. Returning None leaves _pending_decode_trace unset, # which is exactly the "nothing hoisted" state the original path expects. if self._uses_prefetcher(): return None if page_table is None: # Nothing here pins down the decode batch size, and guessing wrong would leave the trace inputs # the wrong shape for the first decode step. Fall back to setting it up lazily, as before. logger.info("No page table available at warmup; decode trace will be set up on first decode") return None batch_size = page_table.shape[0] # Values do not matter: the first real decode explicitly requests a # full input reload. Only the shapes have to match what # decode_forward will supply. tokens = torch.chunk(torch.zeros(batch_size, 1, dtype=torch.int64), self.data_parallel, 0) current_pos = torch.chunk(torch.zeros(batch_size, dtype=torch.int64), self.data_parallel, 0) chunked_page_table = torch.chunk(page_table, self.data_parallel, 0) previous_mode = self.mode self.mode = Mode.DECODE for i in range(len(self.model)): self.model[i].switch_mode(Mode.DECODE) try: return self._prepare_decode_trace_text( tokens, current_pos, page_table=chunked_page_table, kv_cache=kv_cache, on_device_sampling=on_device_sampling, ) finally: # Restore only the generator's own mode. No switch_mode(Mode.PREFILL): that drives # Prefetcher.init(), whose Mode.PREFILL branch has never run on main (main only ever calls # switch_mode(Mode.DECODE)) and is not functional -- it reads an unassigned all_core_range_set, # and supplying one trips "Statically allocated circular buffers ... clash with L1 buffers on # core range". Leaving it in DECODE matches main. self.mode = previous_mode def _will_row_shard_prefill(self, tokens, sampling_params): """Whether this prefill will be dispatched to the model's row-sharded batched path. Single source of truth for that decision -- both the dispatch itself and the choice of where to hoist decode trace preparation depend on it, and they must not drift apart. Only used when device sampling is active (sampling_params is not None) and the prompt uses the harmony chat template (first token is <|start|>=200006). Host sampling needs the single-user prefill path that returns full logits per user. """ is_harmony = tokens.shape[1] > 0 and int(tokens[0, 0]) == 200006 return bool( getattr(self.model[0], "users_row_sharded", False) and tokens.shape[0] > 1 and sampling_params is not None and is_harmony ) def _overrides_prefill_capture(self): """Whether a subclass replaces _capture_trace_prefill (see the warmup-less deferral gate).""" return type(self)._capture_trace_prefill is not Generator._capture_trace_prefill def _uses_prefetcher(self): """Whether any model instance drives the DRAM prefetcher. The prefetcher owns sub-device managers that differ per mode, so hoisting decode setup into the prefill phase cannot work for it: switching to DECODE and back reaches the prefetcher's Mode.PREFILL branch, which has never run on main and is not functional, while staying in DECODE leaves prefill running against decode sub-devices ("Kernel group cores do not match sub device cores", program.cpp:2205). The prefetcher is Blackhole-only and GPT-OSS does not use it, so keeping these models on the original non-hoisted path costs this change nothing it targets. """ uses = any(getattr(m, "prefetcher", None) is not None for m in self.model) if uses and not getattr(self, "_logged_prefetcher_hoist_skip", False): self._logged_prefetcher_hoist_skip = True logger.info( "DRAM prefetcher detected: keeping the original (non-hoisted) decode trace setup. " "This model does not get the hoisted trace-allocation path, because the prefetcher's " "per-mode sub-device managers are incompatible with preparing decode setup during the " "prefill phase (see _uses_prefetcher). Trace I/O buffers for this model may still be " "allocated while a trace is live." ) return uses def _prepare_decode_trace_once(self, kv_cache, page_table, on_device_sampling): """Prepare the decode trace unless it is already prepared. Safe to call from either hoist point.""" if self._uses_prefetcher(): return if self._pending_decode_trace is None: self._pending_decode_trace = self._prepare_decode_trace_for_warmup( kv_cache=kv_cache, page_table=page_table, on_device_sampling=on_device_sampling, ) def _post_prefill_tail(self, model_id, logits, sampling_enabled): """The part of the post-prefill body that runs on already-computed logits.""" if sampling_enabled: return self.model[model_id].sampling.sample(logits, enable_trace=False) return ttnn.untilize(logits, use_multicore=True) def _append_prefill_result(self, prefill_results, idx, model_id, last_token_idx, tail_out, sampling_enabled): """Queue a per-prompt prefill result for readback, in whichever shape the tail produced.""" if sampling_enabled: tt_tokens, tt_log_probs = tail_out queued = [ tt_tokens.cpu(blocking=False), tt_log_probs.cpu(blocking=False) if tt_log_probs is not None else None, ] else: queued = tail_out.cpu(blocking=False) prefill_results.append( { "idx": idx, "model_id": model_id, "last_token_idx": last_token_idx, "logits": queued, "sampling": sampling_enabled, } ) def _record_pending_traces(self): """Capture the decode trace prepared while recording was deferred. Its compile pass, persistent inputs and sampling pre-compile all ran in :meth:`_prepare_decode_trace_text` before this call's prefill captured anything, so this binds only to buffers that already existed and allocates nothing itself. """ if self._pending_decode_trace is not None: prepared, self._pending_decode_trace = self._pending_decode_trace, None on_device_sampling = prepared["on_device_sampling"] previous_mode = self.mode self.mode = Mode.DECODE for i in range(len(self.model)): self.model[i].switch_mode(Mode.DECODE) try: trace_ids, tt_out_trace, *device_inputs = self._record_decode_trace_text(prepared) finally: # See _prepare_decode_trace_for_warmup: no switch_mode(Mode.PREFILL) here either. self.mode = previous_mode self.trace_ids_decode[on_device_sampling] = trace_ids self.trace_inputs_decode[on_device_sampling] = device_inputs self.trace_output_decode[on_device_sampling] = tt_out_trace # A warmup-less call can hoist its own preparation after an eager # warmup staged this key. The hoisted trace owns the active inputs. getattr(self, "_prepared_decode_traces", {}).pop(on_device_sampling, None) def _prefill_trace_forward(self, prepared, device_inputs): """Run the prefill body for a prepared trace variant. Shared verbatim by the compile pass and the capture pass so the two can never drift. """ model_id = prepared["model_id"] transformed_inputs = self.model[model_id].transform_and_embed_prefill_inputs_device(*device_inputs) return self.model[model_id].ttnn_prefill_forward( x=transformed_inputs[0], rot_mats_global=prepared["rot_mats_global"], rot_mats_local=prepared["rot_mats_local"], page_table=transformed_inputs[1], chunk_page_table=transformed_inputs[2], chunk_start_idx=transformed_inputs[3], kv_cache=prepared["kv_cache"], **prepared["forward_kwargs"], ) def _prepare_trace_prefill( self, prefill_ids, page_table=None, chunk_page_table=None, kv_cache=None, model_id=-1, global_user_id=None, batch_size=1, user_id=0, start_pos=0, ): """Phase 1 of prefill trace setup: run the compile pass and allocate the persistent trace inputs. Both of those allocate device memory, and neither may happen while another trace is live: once a trace is captured its recorded buffer addresses are invisible to the allocator, so a buffer handed out afterwards can land inside a live trace's scratch and be clobbered on replay. Capture is split out into :meth:`_record_trace_prefill` so every trace variant can be prepared first, and only then captured back to back. """ prefill_kwargs = { "page_table": page_table, "chunk_page_table": chunk_page_table, "chunk_start_idx": start_pos, "user_id": user_id, } # Batched prefill threads the batch through the model; single-user prefill does not. forward_kwargs = {"batch_size": batch_size, "user_id": user_id} if batch_size > 1 else {} if batch_size > 1: prefill_kwargs["batch_size"] = batch_size if global_user_id is not None: prefill_kwargs["global_user_id"] = global_user_id host_inputs = self.model[model_id].prepare_prefill_inputs_trace(prefill_ids, **prefill_kwargs) prepared = { "model_id": model_id, "kv_cache": kv_cache, "forward_kwargs": forward_kwargs, # These matrices will actually be pointing to the whole cos_matrix and sin_matrix that was allocated on device in the RotarySetup class "rot_mats_global": host_inputs[1], "rot_mats_local": host_inputs[2], } host_inputs = (host_inputs[0], host_inputs[3], host_inputs[4], host_inputs[5]) mesh_device = self.model_args[model_id].mesh_device # Persistent trace inputs: these outlive the capture and are refreshed in place on every replay. # Allocated before the compile pass so that the compile pass can run over them directly. prepared["device_inputs"] = copy_host_to_device(host_inputs, mesh_device=mesh_device) # Compile run, over the *persistent* inputs -- the exact buffers the capture will bind to, so every # program it needs is cached against the specs it will actually see. Compiling over a separate # transient copy would leave any program whose cache key differs between the two uncached, and it # would then load inside the capture window, which the runtime rejects outright: "Cannot load new # binaries during trace capture" (mesh_workload.cpp). Gemma-4's LM head hit this. # The output is a real prefill result for these ids, so deferred warmup callers can consume it # instead of replaying a trace that has not been captured yet. prepared["compile_output"] = self._prefill_trace_forward(prepared, prepared["device_inputs"]) ttnn.synchronize_device(mesh_device) logger.info("Done Compiling Model") return prepared def _record_pending_prefill_traces(self): """Record single-user traces; retain prepared batched inputs for capture on first use.""" if not self._pending_prefill_traces: return logger.info(f"Recording {len(self._pending_prefill_traces)} deferred prefill trace(s)") for trace_key, prepared in self._pending_prefill_traces.items(): if prepared["forward_kwargs"].get("batch_size", 1) > 1: # Compilation and persistent allocation have already finished. # Capture this bucket only when requested so unused batch sizes # do not consume the model's reserved trace region. self._prepared_prefill_traces[trace_key] = prepared continue trace_id, tt_out_trace, *device_inputs = self._record_trace_prefill(prepared) self.trace_id_prefill[trace_key] = trace_id self.trace_inputs_prefill[trace_key] = device_inputs self.trace_output_prefill[trace_key] = tt_out_trace self._pending_prefill_traces = {} def _record_trace_prefill(self, prepared): """Phase 2 of prefill trace setup: capture the trace. Allocation-free outside the capture window by construction -- everything it binds to was allocated by :meth:`_prepare_trace_prefill` before any trace existed. """ mesh_device = self.model_args[prepared["model_id"]].mesh_device device_inputs = prepared["device_inputs"] # Release our handle on the compile-pass output before capturing, matching the pre-split behaviour. # A deferred warmup caller may still be holding it; that is its own reference to keep or drop. prepared.pop("compile_output", None) # Everything allocated between begin/end_trace_capture belongs to the trace being recorded - # the decoder's residual add and friends - and must stay allocated for replay. Recording N # traces means capture N runs while 1..N-1 are live, which reordering cannot avoid, so scope # the window instead. Acknowledgement, not elimination: it tells the checker the program is # prepared for these, it does not stop a replay writing them. Matches llama3_70b_galaxy. # No-op unless TT_METAL_TRACE_ALLOC_TRACKING=1. with trace_allocation_tracker.corruptible_allocation_scope(mesh_device): trace_id = ttnn.begin_trace_capture(mesh_device, cq_id=0) tt_out_trace = self._prefill_trace_forward(prepared, device_inputs) ttnn.end_trace_capture(mesh_device, trace_id, cq_id=0) ttnn.synchronize_device(mesh_device) logger.info("Done Capturing Prefill Trace") return trace_id, tt_out_trace, *device_inputs def _capture_trace_prefill(self, *args, **kwargs): """Prepare and immediately capture a prefill trace. Only safe when no other trace is live on the device. :meth:`warmup_model_prefill` drives the two-phase form instead so that every variant is prepared before the first capture; this single-shot form remains for callers (e.g. vLLM) that capture exactly one trace. """ return self._record_trace_prefill(self._prepare_trace_prefill(*args, **kwargs)) def _prepare_trace_prefill_sampling(self, model_id, sampling_batch): """Compile batched prefill post-processing (norm + lm_head + sampling) and allocate its trace input. Input buffer: [1, 1, sampling_batch, full_dim] host → column-sharded to [1, 1, sampling_batch, dim_per_device]. Output: (tt_tokens, tt_log_probs) from sampling. """ mesh_device = self.model_args[model_id].mesh_device full_dim = self.model_args[model_id].dim dummy_input = ttnn.from_torch( torch.zeros(1, 1, sampling_batch, full_dim, dtype=torch.bfloat16), device=mesh_device, dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT, mesh_mapper=ttnn.ShardTensorToMesh(mesh_device, dim=-1), ) logits = self.model[model_id]._apply_norm_and_lm_head(dummy_input) # count_tokens=False: this pass samples off dummy zeros. Counting it and then zeroing the # counters would also discard the caller's real output history, which sample() may have built up # before this variant was first prepared. self.model[model_id].sampling.precompile(logits, all_configs=not self._any_trace_captured()) ttnn.synchronize_device(mesh_device) logger.info("Done compiling prefill sampling") trace_input = ttnn.from_torch( torch.zeros(1, 1, sampling_batch, full_dim, dtype=torch.bfloat16), device=mesh_device, dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT, mesh_mapper=ttnn.ShardTensorToMesh(mesh_device, dim=-1), ) return {"model_id": model_id, "input": trace_input} def _record_trace_prefill_sampling(self, prepared): """Capture the batched prefill sampling trace prepared by :meth:`_prepare_trace_prefill_sampling`.""" model_id = prepared["model_id"] mesh_device = self.model_args[model_id].mesh_device trace_input = prepared["input"] # count_tokens stays on (the default) inside the capture window: capture only records the commands, # so the update runs on replays, over real sampled tokens -- exactly like SamplingGenerator's own # capture_trace. Only the eager compile pass in _prepare_trace_prefill_sampling passes # count_tokens=False, because it actually executes, over dummy logits. # As with the prefill trace, these outputs are scratch shared with # other captured graphs and are consumed immediately after replay. with trace_allocation_tracker.corruptible_allocation_scope(mesh_device): trace_id = ttnn.begin_trace_capture(mesh_device, cq_id=0) logits = self.model[model_id]._apply_norm_and_lm_head(trace_input) tt_tokens, tt_log_probs = self.model[model_id].sampling.sample(logits, enable_trace=False) ttnn.end_trace_capture(mesh_device, trace_id, cq_id=0) ttnn.synchronize_device(mesh_device) logger.info("Done capturing prefill sampling trace") return trace_id, (tt_tokens, tt_log_probs), trace_input def _capture_trace_prefill_sampling(self, model_id, sampling_batch): """Prepare and immediately capture the batched prefill sampling trace.""" return self._record_trace_prefill_sampling(self._prepare_trace_prefill_sampling(model_id, sampling_batch)) def _row_sharded_batched_prefill( self, tokens, page_table, kv_cache, prompt_lens, prefill_seq_lens, enable_trace=True, sampling_params=None, empty_slots=None, ): """Dispatch to model's row-sharded batched prefill. ``empty_slots`` is forwarded so the model can reorder users to match their decode row mapping (tenstorrent/tt-metal#44746). """ assert ( self.data_parallel == 1 ), "Row-sharded batched prefill requires data_parallel=1 (model handles DP internally)" return self.model[0].row_sharded_batched_prefill( tokens, page_table, kv_cache[0], prompt_lens, prefill_seq_lens, enable_trace=enable_trace, sampling_params=sampling_params, model_args=self.model_args[0], trace_cache={ "ids": self.trace_id_prefill, "inputs": self.trace_inputs_prefill, "outputs": self.trace_output_prefill, }, empty_slots=empty_slots, ) def _easy_trace_prefill( self, prefill_ids, page_table=None, full_page_table=None, user_id=0, last_token_idx=None, kv_cache=None, model_id=-1, prefill_seq_len=None, batch_size=1, num_cached_tokens=0, **kwargs, ): global_user_id = kwargs.get("global_user_id", None) use_start_pos = "sp1" if num_cached_tokens > 0 else "sp0" trace_key = f"{prefill_seq_len}_{model_id}_{batch_size}_{use_start_pos}" use_prefix_caching = num_cached_tokens > 0 chunk_start_idx = num_cached_tokens block_size = get_block_size(kv_cache) if page_table is not None and batch_size == 1: page_table = page_table[user_id : user_id + 1, :] if full_page_table is not None and batch_size == 1: full_page_table = full_page_table[user_id : user_id + 1, :] chunk_page_table = None max_blocks_prefill = _get_max_blocks_prefill(kv_cache) # Preserve full per-user page IDs for traced APC slicing. source_page_table = full_page_table if full_page_table is not None else page_table if source_page_table is None: raise ValueError("Traced prefill requires a page_table") page_table = _pad_or_create_page_table(source_page_table, max_blocks_prefill) if batch_size == 1: if use_prefix_caching: chunk_start_block = num_cached_tokens // block_size chunk_end_block = num_blocks_in_seq(num_cached_tokens + prefill_seq_len, block_size) chunk_page_table = source_page_table[:, chunk_start_block:chunk_end_block] chunk_blocks = num_blocks_in_seq(prefill_seq_len, block_size) chunk_page_table = _pad_or_create_page_table(chunk_page_table, chunk_blocks) if self.trace_id_prefill[trace_key] is None: if self._defer_prefill_recording: # Compile and stage only. Recording here would put this bucket's trace on device # before the remaining warmup buckets have compiled, so every one of those compiles - # and the trace inputs they stage - would land behind it. warmup_model_prefill # records the whole set afterwards. if trace_key not in self._pending_prefill_traces: self._pending_prefill_traces[trace_key] = self._prepare_trace_prefill( prefill_ids, page_table=page_table, chunk_page_table=chunk_page_table, kv_cache=kv_cache, model_id=model_id, global_user_id=global_user_id, batch_size=batch_size, user_id=user_id, start_pos=chunk_start_idx, ) else: # A pending bucket has no trace to replay yet. Refresh its inputs and # execute for this user too; the previous compile output belongs to # the previous request and did not fill this user's KV rows. prepared = self._pending_prefill_traces[trace_key] prefill_kwargs = dict( page_table=page_table, chunk_page_table=chunk_page_table, chunk_start_idx=chunk_start_idx, user_id=user_id, ) if batch_size > 1: prefill_kwargs["batch_size"] = batch_size prepared["forward_kwargs"] = {"batch_size": batch_size, "user_id": user_id} if global_user_id is not None: prefill_kwargs["global_user_id"] = global_user_id host_inputs = self.model[model_id].prepare_prefill_inputs_trace(prefill_ids, **prefill_kwargs) prepared["rot_mats_global"], prepared["rot_mats_local"] = host_inputs[1:3] prepared["device_inputs"] = copy_host_to_device( (host_inputs[0], host_inputs[3], host_inputs[4], host_inputs[5]), device_tensors=prepared["device_inputs"], mesh_device=self.model_args[model_id].mesh_device, ) prepared["compile_output"] = self._prefill_trace_forward(prepared, prepared["device_inputs"]) # The compile pass produced a real prefill result for these ids, so the caller's # output processing runs (and compiles) exactly as it would have. # Only that caller needs the result and keeping it in the pending # store would retain every bucket's activations until capture. return self._pending_prefill_traces[trace_key].pop("compile_output") if trace_key in self._prepared_prefill_traces: trace_id, tt_out_trace, *device_inputs = self._record_trace_prefill( self._prepared_prefill_traces.pop(trace_key) ) else: trace_id, tt_out_trace, *device_inputs = self._capture_trace_prefill( prefill_ids, page_table=page_table, chunk_page_table=chunk_page_table, kv_cache=kv_cache, model_id=model_id, global_user_id=global_user_id, batch_size=batch_size, user_id=user_id, start_pos=chunk_start_idx, ) self.trace_id_prefill[trace_key] = trace_id self.trace_inputs_prefill[trace_key] = device_inputs self.trace_output_prefill[trace_key] = tt_out_trace tt_out_trace = self._prefill_forward_trace( self.trace_id_prefill[trace_key], self.trace_inputs_prefill[trace_key], self.trace_output_prefill[trace_key], prefill_ids, page_table=page_table, chunk_page_table=chunk_page_table, model_id=model_id, global_user_id=global_user_id, batch_size=batch_size, user_id=user_id, start_pos=chunk_start_idx, ) return tt_out_trace def _prefill_forward_trace( self, trace_id, device_inputs, tt_out_trace, prefill_ids, user_id=0, page_table=None, chunk_page_table=None, model_id=-1, global_user_id=None, batch_size=1, start_pos=0, ): # Use actual batch_size since tokens are now in batch dimension prefill_kwargs = { "page_table": page_table, "chunk_page_table": chunk_page_table, "chunk_start_idx": start_pos, "batch_size": batch_size, "user_id": user_id, } if global_user_id is not None: prefill_kwargs["global_user_id"] = global_user_id host_inputs = self.model[model_id].prepare_prefill_inputs_trace(prefill_ids, **prefill_kwargs) host_inputs = (host_inputs[0], host_inputs[3], host_inputs[4], host_inputs[5]) device_inputs = copy_host_to_device( host_inputs, device_tensors=device_inputs, mesh_device=self.model_args[model_id].mesh_device ) ttnn.execute_trace(self.model_args[model_id].mesh_device, trace_id, cq_id=0, blocking=False) return tt_out_trace def release_request(self, slot: int) -> None: """Release finished-request seed state through the serving lifecycle hook. 'slot' is the request's current state slot, after any decode remap. KV pages and traces remain reusable; surviving requests retain their RNG streams. Hosts must notify completion before admitting replacements. """ per_model_slots = self.model_args[0].max_batch_size if not 0 <= slot < per_model_slots * self.data_parallel: raise ValueError(f"Request slot {slot} is outside the configured batch") model_id, local_slot = divmod(slot, per_model_slots) sampling = getattr(self.model[model_id], "sampling", None) if sampling is not None: sampling.seed_manager.release_slot(local_slot) getattr(self, "_slots_prefilled_since_decode", set()).discard(slot) # Note: This function is called by vLLM def prefill_forward_text(self, *args, **kwargs): """Flush deferred trace captures once the prefill that armed them ends successfully. The first traced prefill call defers only its own decode-trace recording: it prepares decode (compile pass, persistent inputs, sampling pre-compile) before its prefill captures anything, then records both prefill and decode here once output processing is done. Only the call that armed the deferral flushes it -- nested calls must not. """ already_pending = self._defer_trace_recording try: result = self._prefill_forward_text_impl(*args, **kwargs) except BaseException: # A failed prefill must not capture traces while unwinding: the capture would bind whatever # state the failure left behind, and an error raised during recording would mask the original # exception. Drop the deferred state instead; the decode trace is then set up lazily on the # first decode step, as on main. if not already_pending: self._defer_trace_recording = False self._pending_decode_trace = None self._defer_prefill_recording = False self._pending_prefill_traces.clear() if not self._any_trace_captured(): self._prepared_prefill_traces.clear() self._prepared_prefill_sampling_traces.clear() raise if not already_pending and self._defer_trace_recording: self.finalize_deferred_traces() return result def _prefill_forward_text_impl( self, tokens: torch.Tensor, # All tokens, including the cached ones page_table=None, kv_cache=None, prompt_lens=None, # Full prompt lengths, including the cached ones empty_slots=None, enable_trace=True, model_id_warmup=None, sampling_params: SamplingParams | None = None, start_pos: list[int] = None, # Cached prefixes lengths return_hidden_states=False, warmup_prefill=True, **kwargs, ): self.mode = Mode.PREFILL if page_table is not None: assert isinstance(page_table, torch.Tensor), "page_table mush be torch.Tensor" else: # Only paged attention is supported for prefill enable_trace = False on_device_sampling_requested = sampling_params is not None # we need this here because of tt-metal tests on_device_sampling_enabled = ( getattr(self.model[0], "_supports_on_device_sampling", False) and getattr(self.model[0], "sampling", None) is not None ) if warmup_prefill: # The prefill sweep records its traces before returning. Stage decode # programs and persistent inputs first, otherwise the first decode # allocates behind those traces and a later request cannot replay them. prepare_decode = ( enable_trace and not self.already_warmed_up_prefill and not self._defer_trace_recording and not self._defer_prefill_recording and not self._any_trace_captured() and not self._will_row_shard_prefill(tokens, sampling_params) and not self._overrides_prefill_capture() and not self._uses_prefetcher() ) if prepare_decode: self._prepare_decode_trace_once( kv_cache=kv_cache, page_table=page_table, on_device_sampling=on_device_sampling_requested and on_device_sampling_enabled, ) self.warmup_model_prefill( kv_cache=kv_cache, enable_trace=enable_trace, can_sample_on_device=on_device_sampling_enabled, ) if prepare_decode: # Finish deferred capture as part of warmup so real requests # can refresh and replay the prepared decode trace directly. self._record_pending_traces() elif ( enable_trace and not self._defer_trace_recording and not self._defer_prefill_recording and not self._any_trace_captured() # Excluded: the row-sharded batched path deadlocks either way round -- once a decode program is # compiled its MoE dispatch can no longer be trace-captured, and compiling decode before any # prefill deadlocks in the MoE combine. It keeps main's late decode compile. and not self._will_row_shard_prefill(tokens, sampling_params) # Excluded: models that customise the capture step. Deferred recording prepares through the # base helper and records later, which skips whatever a _capture_trace_prefill override does # between its compile pass and its capture. and not self._overrides_prefill_capture() ): # Models that skip prefill warmup (GPT-OSS sets warmup_prefill=False everywhere) would otherwise # capture their prefill trace part way through this call and then compile the whole decode graph # -- plus its long-lived trace inputs and its sampling state -- with that trace already live. # Defer this call's decode capture so it only *prepares* here; the wrapper flushes once the # prefill is done. First call only: after that the ordering is fixed and re-arming would defer a # capture later calls expect to exist. self._defer_trace_recording = True # Also hold this call's prefill trace back. Without a warmup sweep this call both # compiles and records, so _post_prefill_tail - which runs after the prefill and # compiles the tail's bucket-keyed programs - would otherwise do so behind the trace # just recorded. Deferring lets the tail compile first; the wrapper records both once # the prefill is done. self._defer_prefill_recording = True # Prepare decode before this call's prefill, not at the flush: the compile pass is a real decode # step at position 0 with mock inputs, so it writes mock K/V that the following prefill then # overwrites. At the flush it would instead corrupt the prefilled cache. The position-0 write # relies on this being the first traced prefill call (nothing captured yet), so no earlier # request can have left a cached prefix in these page-table rows -- a cached prefix # (num_cached > 0) would not be rewritten by the prefill and would stay corrupted. self._prepare_decode_trace_once( kv_cache=kv_cache, page_table=page_table, on_device_sampling=on_device_sampling_enabled, ) batch_size, batch_seq_len = tokens.shape max_batch_size_per_model = self.model_args[0].max_batch_size # Output shape depends on whether we're returning logits or hidden states if return_hidden_states: # For hidden states, output shape is [batch_size, hidden_size] # Note: dim is the hidden dimension size hidden_size = self.model_args[0].dim output_tensor = torch.zeros(batch_size, hidden_size) else: # Each model expected to run the same model, safe to use 1st vocab size output_tensor = torch.zeros(batch_size, 1, self.model_args[0].vocab_size) output_tokens = torch.zeros(batch_size, 1, dtype=torch.int64) output_log_probs = [None] * batch_size sampling_executed = False prompt_lens = prompt_lens if prompt_lens is not None else torch.tensor([batch_seq_len] * batch_size) if empty_slots is None: empty_slots = list(range(batch_size)) # For row-sharded users, use max_local_batch_size (users per row) for group_user_id local_batch_size = getattr(self.model_args[0], "max_local_batch_size", max_batch_size_per_model) if not isinstance(prompt_lens, list): prompt_lens = prompt_lens.tolist() # Pad by uncached suffix length: only (seq_len - num_cached) tokens reach the kernel. # int() normalizes numpy.int64 from vLLM callers (bit_length() requires a Python int). num_cached_per_user = [int(n) for n in start_pos] if start_pos is not None else [0] * len(prompt_lens) assert len(num_cached_per_user) == len( prompt_lens ), f"start_pos length {len(num_cached_per_user)} != prompt_lens length {len(prompt_lens)}" if start_pos is not None: num_cached_per_user = self._align_resume_offsets(num_cached_per_user, prompt_lens, kv_cache) # The per-user loop below re-reads the offset from ``start_pos``, so the # aligned list has to replace it: otherwise the padded length computed # here would not match the token slice the kernel receives. start_pos = num_cached_per_user for i, (seq_len, num_cached) in enumerate(zip(prompt_lens, num_cached_per_user)): assert 0 <= num_cached < seq_len, f"user {i}: num_cached={num_cached} must be < seq_len={seq_len}" prefill_seq_lens = [ get_padded_prefill_len(seq_len - num_cached) for seq_len, num_cached in zip(prompt_lens, num_cached_per_user) ] # Row-sharded batched prefill: process 1 user per row per iteration. # See _will_row_shard_prefill for when this path is taken. if self._will_row_shard_prefill(tokens, sampling_params): return self._row_sharded_batched_prefill( tokens, page_table, kv_cache, prompt_lens, prefill_seq_lens=prefill_seq_lens, enable_trace=enable_trace, sampling_params=sampling_params, empty_slots=empty_slots, ) # Batched prefill: all prompts share the same padded length so they can # be processed in a single forward pass. padded_batch is rounded up to # the nearest SUPPORTED_PREFILL_BATCH_SIZES entry (not max_batch_size) # to keep all_gather buffers within DRAM limits. use_batched_prefill = ( batch_size > 1 and len(set(prefill_seq_lens)) == 1 and self.data_parallel == 1 and not getattr(self.model_args[0], "disable_batched_prefill", False) and all( n == 0 for n in num_cached_per_user ) # batched path feeds full tokens; incompatible with cached prefixes ) # Batched prefill passes a per-user `last_token_idx` *list* (and a list # `user_id`) into prefill_forward_single_user_text. That function's # chunked-prefill branch only supports a single sequence with a scalar # last_token_idx -- it compares/slices it arithmetically # (`last_token_idx < seq_len`, `// chunk_size`, ...), pins user_id=0 and # slices the page table to one row. A batch whose padded length exceeds # max_prefill_chunk_size therefore reaches the chunked path with a list # and dies on the first assert with # `TypeError: '<' not supported between instances of 'list' and 'int'`. # Such prompts already require multi-pass chunked prefill (so batching # buys no single-pass win and would re-introduce the very DRAM pressure # chunking exists to relieve); keep them on the sequential per-user path # that chunks each prompt correctly. See tenstorrent/tt-metal#45234. if use_batched_prefill and any(s > self.model_args[0].max_prefill_chunk_size for s in prefill_seq_lens): logger.info( f"Batched prefill disabled: padded prefill len {prefill_seq_lens[0]} exceeds " f"max_prefill_chunk_size {self.model_args[0].max_prefill_chunk_size}; chunked " f"prefill requires the sequential prefill path (#45234)" ) use_batched_prefill = False if use_batched_prefill and on_device_sampling_requested: sampling_module, sampling_dp, _, _ = self._get_sampling_contract(0) if sampling_module is not None and sampling_dp > 1: # NOTE: Batched prefill disabled: on-device sampling # must fall back to sequential prefill until a row-sharded # batched-prefill sampling contract is implemented. use_batched_prefill = False if use_batched_prefill: padded_batch = batched_prefill_padded_batch(batch_size, empty_slots, self.model_args[0].max_batch_size) if padded_batch > self.model_args[0].max_batch_size: logger.info( f"Batched prefill disabled: padded_batch {padded_batch} exceeds " f"max_batch_size {self.model_args[0].max_batch_size}" ) use_batched_prefill = False elif not batched_prefill_fits_token_budget( padded_batch, prefill_seq_lens[0], self.model_args[0].max_prefill_chunk_size ): logger.info( f"Batched prefill disabled: {padded_batch} x {prefill_seq_lens[0]} = " f"{padded_batch * prefill_seq_lens[0]} tokens exceeds model/device token budget " f"{self.model_args[0].max_prefill_chunk_size} or reaches kernel limit {MAX_BATCHED_PREFILL_SEQ_LEN}" ) use_batched_prefill = False if not use_batched_prefill: padded_batch = self.model_args[0].max_batch_size all_users = [0] if use_batched_prefill else empty_slots sampling_params_per_out: list[SamplingParams | None] = [None] * len(empty_slots) prompt_tokens_per_out: list[torch.Tensor | None] = [None] * len(empty_slots) prefill_results: list[dict] = [] for idx, user_id in enumerate(all_users): model_id = user_id // max_batch_size_per_model if model_id_warmup is None else model_id_warmup group_user_id = user_id % local_batch_size if page_table is None else 0 if use_batched_prefill: batch_user_ids = empty_slots last_token_idx = [(seq_len - 1) for seq_len in prompt_lens] prefill_seq_len = prefill_seq_lens[0] seq_len = prompt_lens else: batch_user_ids = None seq_len = int(prompt_lens[idx]) num_cached_tokens = int(start_pos[idx]) if start_pos is not None else 0 last_token_idx = seq_len - 1 prefill_seq_len = prefill_seq_lens[idx] logger.info(f"Prefilling User {user_id + 1} up to {seq_len} tokens") local_kwargs = kwargs.copy() # Avoid modifying original kwargs if getattr(self.model[model_id], "users_row_sharded", False): local_kwargs["global_user_id"] = batch_user_ids if use_batched_prefill else user_id sampling_enabled = ( on_device_sampling_requested and getattr(self.model[model_id], "_supports_on_device_sampling", False) and getattr(self.model[model_id], "sampling", None) is not None ) if use_batched_prefill: # Galaxy 70B approach: slot-based placement with shape [padded_batch, prefill_seq_len] # Each request is placed at its corresponding slot index prefill_ids = torch.zeros(padded_batch, prefill_seq_len, dtype=torch.long, device=tokens.device) padded_last_token_idx = [0] * padded_batch # dummy idx for padded slots for local_idx, slot in enumerate(empty_slots): seq_len_local = int(seq_len[local_idx]) padded_tokens = torch.cat( [ tokens[local_idx : local_idx + 1, :seq_len_local], torch.zeros(1, prefill_seq_len - seq_len_local, dtype=torch.long, device=tokens.device), ], dim=-1, ) prefill_ids[slot : slot + 1] = padded_tokens padded_last_token_idx[slot] = last_token_idx[local_idx] last_token_idx = padded_last_token_idx else: num_cached_tokens = int(start_pos[idx]) if start_pos is not None else 0 prefill_ids = torch.cat( [ tokens[idx : idx + 1, num_cached_tokens:seq_len], torch.zeros(1, prefill_seq_len - (seq_len - num_cached_tokens)).long(), ], dim=-1, ) enable_trace_current_prompt = enable_trace and self.model_args[model_id].can_enable_trace( prefill_seq_len, num_cached_tokens if not use_batched_prefill else 0 ) logger.info( f"Prefill seq len: {prefill_seq_len}, max_prefill_chunk_size: {self.model_args[0].max_prefill_chunk_size}, trace: {enable_trace_current_prompt}" ) if page_table is not None: # For batched prefill: pass full page_table (function handles slot placement) # For non-batched prefill: pass sliced page_table for current user (like original code) page_table_for_user = page_table if use_batched_prefill else page_table[idx : idx + 1] # A resumed chunk needs a page table spanning the cached tokens # too: ``prefill_forward_single_user_text`` slices # ``chunk_page_table`` at an absolute block offset, so a # chunk-width table leaves that slice inside its own zero pad and # ``paged_fill_cache`` writes the chunk into physical block 0. # ``seq_len`` is the cumulative prompt length; ``prefill_seq_len`` # is only this chunk's padded width. page_table_use_full_len = bool( not use_batched_prefill and not enable_trace_current_prompt and num_cached_tokens ) page_table_user = self._get_prefill_user_page_table( page_table_for_user, kv_cache[model_id], seq_len, trace_enabled=enable_trace_current_prompt, prefill_seq_len=prefill_seq_len, use_batched_prefill=use_batched_prefill, user_id=batch_user_ids if use_batched_prefill else user_id, padded_batch_size=padded_batch if use_batched_prefill else None, use_full_prompt_len=page_table_use_full_len, ) full_page_table_user = None if enable_trace_current_prompt and not use_batched_prefill: # Keep the full per-user mapping for traced APC page slicing. full_page_table_user = self._get_prefill_user_page_table( page_table_for_user, kv_cache[model_id], seq_len, trace_enabled=False, prefill_seq_len=prefill_seq_len, use_batched_prefill=False, user_id=user_id, padded_batch_size=None, use_full_prompt_len=True, ) else: page_table_user = None full_page_table_user = None if page_table_user is not None and _deepseek_kvdbg_enabled(): sample = [] if page_table_user.numel(): flat = page_table_user.reshape(-1) sample = flat[: min(16, flat.numel())].tolist() logger.debug( "KVDBG deepseek prefill user global={} local={} seq_len={} cached={} page_table_shape={} sample={}", user_id, group_user_id, seq_len, num_cached_tokens, list(page_table_user.shape), sample, ) model_kv_cache = kv_cache[model_id] if kv_cache is not None else None # Check if 'pixel_values' exists and index it safely if local_kwargs.get("pixel_values", None) is not None: local_kwargs["pixel_values"] = local_kwargs["pixel_values"][idx] if "image_grid_thw" in local_kwargs: local_kwargs["image_grid_thw"] = local_kwargs["image_grid_thw"][idx] if "image_sizes" in local_kwargs and local_kwargs["image_sizes"] is not None: local_kwargs["image_sizes"] = local_kwargs["image_sizes"][idx] if sampling_enabled and not use_batched_prefill: sampling_executed = True sampling_dp = getattr(self.model[model_id], "sampling_dp", 1) total_batch = self.model[model_id].sampling.tt_sampling.max_batch_size * sampling_dp per_request_params = format_sampling_params( broadcast_sampling_params(sampling_params, idx, slot_len=total_batch), total_batch ) assert per_request_params is not None, "Sampling was executed but missing per-request sampling params" # empty_slots uses max_batch_size_per_model (not total_batch) because # the seed manager operates on per-row slots (0..31). When sampling_dp > 1 # the params are already broadcast across all rows by broadcast_sampling_params. self.model[model_id].sampling.apply_prefill_state( sampling_params=per_request_params, prompt_tokens=prefill_ids[:, :seq_len].repeat(total_batch, 1), empty_slots=[user_id % max_batch_size_per_model], ) if enable_trace_current_prompt: logits = self._easy_trace_prefill( prefill_ids, page_table=page_table_user, full_page_table=full_page_table_user, user_id=batch_user_ids if use_batched_prefill else group_user_id, last_token_idx=last_token_idx, kv_cache=model_kv_cache, model_id=model_id, prefill_seq_len=prefill_seq_len, batch_size=padded_batch if use_batched_prefill else 1, num_cached_tokens=0 if use_batched_prefill else num_cached_tokens, **local_kwargs, ) else: logits = self.prefill_forward_single_user_text( prefill_ids, page_table=page_table_user, user_id=batch_user_ids if use_batched_prefill else group_user_id, last_token_idx=last_token_idx, kv_cache=model_kv_cache, model_id=model_id, num_cached_tokens=0 if use_batched_prefill else num_cached_tokens, batch_size=padded_batch if use_batched_prefill else 1, **local_kwargs, ) if use_batched_prefill: hidden_dim = logits.shape[-1] logits = ttnn.reshape(logits, [padded_batch, 1, prefill_seq_len, hidden_dim]) if sampling_enabled: sampling_executed = True sampling_module, sampling_dp, sampling_batch, _ = self._get_sampling_contract(model_id) assert sampling_module is not None assert sampling_batch is not None max_prompt_len = max(int(prompt_lens[i]) for i in range(len(empty_slots))) combined_prompt_tokens = torch.zeros(sampling_batch, max_prompt_len, dtype=torch.long) for local_idx, slot in enumerate(empty_slots): plen = int(prompt_lens[local_idx]) combined_prompt_tokens[slot, :plen] = prefill_ids[slot, :plen] # ``combined_prompt_tokens`` above and the extracted hidden states # are both laid out by slot, so the params have to be as well. combined_params = scatter_sampling_params_to_slots( format_sampling_params(sampling_params, sampling_batch), empty_slots, sampling_batch, ) sampling_module.apply_prefill_state( sampling_params=combined_params, prompt_tokens=combined_prompt_tokens, empty_slots=empty_slots, replicate_seeds=False, ) user_hidden = self.model[model_id].extract_last_tokens_batched_prefill( logits, last_token_idx, padded_batch, prefill_seq_len, target_batch=sampling_batch, ) sampling_input_key = f"sampling_{prefill_seq_len}_{model_id}_{sampling_batch}_{sampling_dp}" sampling_trace_key = ( f"{sampling_input_key}_{sampling_module._penalties_active}_" f"{sampling_module.tt_sampling.log_probs_calculator.enable_log_probs}_" f"{sampling_module.tt_sampling.force_argmax_sampling}" ) if enable_trace_current_prompt and self._defer_prefill_recording: if sampling_input_key not in self._prepared_prefill_sampling_traces: self._prepared_prefill_sampling_traces[ sampling_input_key ] = self._prepare_trace_prefill_sampling(model_id, sampling_batch) # Consume this warmup's actual hidden states. Capture # waits until every batched extraction/sampling program # has compiled and all persistent inputs exist. batched_logits = self.model[model_id]._apply_norm_and_lm_head(user_hidden) tt_tokens, tt_log_probs = self.model[model_id].sampling.sample( batched_logits, enable_trace=False ) elif enable_trace_current_prompt: if self.trace_id_prefill_sampling[sampling_trace_key] is None: ( s_trace_id, s_trace_output, s_trace_input, ) = ( self._record_trace_prefill_sampling( self._prepared_prefill_sampling_traces[sampling_input_key] ) if sampling_input_key in self._prepared_prefill_sampling_traces else self._capture_trace_prefill_sampling(model_id, sampling_batch) ) self.trace_id_prefill_sampling[sampling_trace_key] = s_trace_id self.trace_output_prefill_sampling[sampling_trace_key] = s_trace_output self.trace_input_prefill_sampling[sampling_trace_key] = s_trace_input s_trace_input = self.trace_input_prefill_sampling[sampling_trace_key] user_hidden_host = user_hidden.cpu() # Readback is blocking; the sampling trace only needs # the host copy and its persistent input from here. del user_hidden ttnn.copy_host_to_device_tensor(user_hidden_host, s_trace_input) ttnn.execute_trace( self.model_args[model_id].mesh_device, self.trace_id_prefill_sampling[sampling_trace_key], cq_id=0, blocking=False, ) tt_tokens, tt_log_probs = self.trace_output_prefill_sampling[sampling_trace_key] else: batched_logits = self.model[model_id]._apply_norm_and_lm_head(user_hidden) tt_tokens, tt_log_probs = self.model[model_id].sampling.sample( batched_logits, enable_trace=False, ) ttnn.synchronize_device(self.model[model_id].mesh_device) tokens_host = ttnn.to_torch(ttnn.get_device_tensors(tt_tokens)[0]).reshape(-1) # tt_log_probs may be a LogProbsResult (top-k logprobs mode) or a plain [B] # tensor (scalar logprobs); mirror the single-user handling so # reformat_logprobs receives per-slot LogProbsResult / scalar entries. plain_log_probs_host = ( ttnn.to_torch(ttnn.get_device_tensors(tt_log_probs)[0]).reshape(-1) if tt_log_probs is not None and not isinstance(tt_log_probs, LogProbsResult) else None ) gather_batched_prefill_samples( empty_slots, tokens_host, tt_log_probs, plain_log_probs_host, output_tokens, output_log_probs, ) else: if return_hidden_states: # Embedding models: trace returns hidden states; extract last-token hidden per slot slot_hidden_list = [] for local_idx, slot in enumerate(empty_slots): user_hidden = logits[slot : slot + 1, :, :, :] slot_hidden = self.model[model_id].process_hidden_states_after_prefill_trace( user_hidden, last_token_idx[slot] ) slot_hidden_list.append((local_idx, slot_hidden, last_token_idx[slot])) ttnn.synchronize_device(self.model[model_id].mesh_device) dim = self.model[model_id].args.dim for local_idx, slot_hidden, lt_idx in slot_hidden_list: slot_hidden_torch = ttnn.to_torch(ttnn.get_device_tensors(slot_hidden)[0]).float() pos = int(lt_idx % 32) out = slot_hidden_torch[0, 0, pos, :dim].clone() if out.device.type != "cpu": out = out.cpu() output_tensor[local_idx] = out else: for local_idx, slot in enumerate(empty_slots): user_logits = logits[slot : slot + 1, :, :, :] _logits = self.model[model_id].process_logits_after_prefill_trace( user_logits, last_token_idx[slot] ) _logits = ttnn.to_layout( _logits, ttnn.ROW_MAJOR_LAYOUT, memory_config=ttnn.DRAM_MEMORY_CONFIG ) output_tensor[local_idx] = self.model[model_id].process_output_prefill( _logits.cpu(), last_token_idx=(last_token_idx[slot] % 32) ) break # Non-batched prefill path if enable_trace_current_prompt: last_token_idx_for_trace = last_token_idx if not use_batched_prefill and num_cached_tokens > 0: last_token_idx_for_trace = last_token_idx - num_cached_tokens if return_hidden_states: hidden_states = self.model[model_id].process_hidden_states_after_prefill_trace( logits, last_token_idx_for_trace ) prefill_results.append( { "idx": idx, "model_id": model_id, "last_token_idx": last_token_idx, "hidden_states": hidden_states.cpu(blocking=False), } ) continue else: logits = self.model[model_id].process_logits_after_prefill_trace(logits, last_token_idx_for_trace) else: if return_hidden_states: raise NotImplementedError("return_hidden_states=True requires enable_trace=True") self._append_prefill_result( prefill_results, idx, model_id, last_token_idx, self._post_prefill_tail(model_id, logits, sampling_enabled), sampling_enabled, ) # Only host results are queued above. Release this user's temporary # logits before the next user's prefill trace can overwrite them. del logits if len(prefill_results) > 0: for elem_idx, res in enumerate(prefill_results): idx = res["idx"] last_token_idx = res["last_token_idx"] model_id = res["model_id"] num_cached_tokens = int(start_pos[idx]) if start_pos is not None else 0 last_token_idx_relative = last_token_idx - num_cached_tokens ttnn.synchronize_device(self.model[model_id].mesh_device) if "hidden_states" in res: output_tensor[idx] = self.model[model_id].process_output_prefill_hidden_states( res["hidden_states"], last_token_idx=(last_token_idx_relative % 32) ) elif res["sampling"]: tt_tokens = res["logits"][0] tt_log_probs = res["logits"][1] tokens_host = ttnn.to_torch(ttnn.get_device_tensors(tt_tokens)[0]).reshape(-1)[ last_token_idx_relative % 32 ] if isinstance(tt_log_probs, LogProbsResult): log_probs_host = tt_log_probs.extract_user(last_token_idx_relative % 32) elif tt_log_probs is not None: log_probs_host = ttnn.to_torch(ttnn.get_device_tensors(tt_log_probs)[0]).reshape(-1)[ last_token_idx_relative % 32 ] else: log_probs_host = None output_tokens[idx] = tokens_host if log_probs_host is not None: output_log_probs[idx] = log_probs_host else: output_tensor[idx] = self.model[model_id].process_output_prefill( res["logits"], last_token_idx=(last_token_idx_relative % 32) ) logger.info(f"Finished prefill for all users up to {batch_seq_len} tokens, Starting decode...") if sampling_executed: return output_tokens, reformat_logprobs(output_log_probs, batch_size) else: return output_tensor def _traced_sdpa_q_chunk_size(self, prefill_seq_len, model_id=0): """q_chunk_size a traced prefill of this length is captured with, or None. The traced path hands the op a ``chunk_start_idx`` device tensor that is refreshed per replay, so the captured program config cannot be derived from the offset. It is built with ``chunk_start_idx=0`` instead, which is what this reproduces. """ get_config = getattr(self.model_args[model_id], "get_attn_sdpa_program_config", None) if get_config is None: return None # Deliberately unguarded: a model whose program config this signature does # not describe must say so by not exposing the method. Swallowing the error # here would drop back to block-only alignment and reinstate the wrong # prefix reads this exists to prevent. return get_config(Mode.PREFILL, prefill_seq_len, 0, None).q_chunk_size def _resume_offset_alignment(self, prefill_seq_len, block_size, model_id=0): """Multiple a resume offset must land on for this padded suffix length. The paged ops need ``block_size``; the traced SDPA needs the q_chunk_size its program was captured with. Their least common multiple satisfies both. The q_chunk_size is taken from the model's own program config where it exposes one, so a short suffix keeps its smaller alignment instead of being rounded away. A model without one must declare ``resumed_prefill_token_alignment``: there is no safe default, and guessing block_size is what produces the silent wrong prefix. """ q_chunk = self._traced_sdpa_q_chunk_size(prefill_seq_len, model_id) if q_chunk is None: q_chunk = self.model_capabilities.get("resumed_prefill_token_alignment") if q_chunk is None: raise ValueError( f"{type(self).__name__} resumes a prefill but neither exposes " "`get_attn_sdpa_program_config` on its model_args nor declares " "`model_capabilities['resumed_prefill_token_alignment']`, so the " "alignment its chunked-SDPA program requires cannot be determined." ) return math.lcm(block_size, int(q_chunk)) def _assert_uniform_resume_alignment(self, prefill_seq_len, block_size, expected): """Replica 0's alignment stands for every replica. Fail if it stops doing so. ``create_submeshes`` splits the mesh into submeshes of one shape, and ``initialize_vllm_model`` builds every replica's ``ModelArgs`` from the same arguments, so they all pin the same q_chunk_size. The per-user offsets are therefore aligned once against replica 0 rather than per replica. Nothing in the type system holds that, so check it instead of trusting it. """ for model_id in range(1, self.data_parallel): other = self._resume_offset_alignment(prefill_seq_len, block_size, model_id) assert other == expected, ( f"replica {model_id} needs resume alignment {other} where replica 0 needs " f"{expected} at padded length {prefill_seq_len}; offsets are aligned once " "against replica 0 and would be wrong for this replica." ) def _align_resume_offsets(self, num_cached_per_user, prompt_lens, kv_cache): """Floor each resume offset to what the paged ops and the traced SDPA need. ``block_size`` covers the page-table slice the chunk's K/V is written through: an offset off that multiple shifts every write by ``chunk_start % block_size`` positions. ``q_chunk_size`` covers the SDPA op, which requires ``chunk_start_idx`` to be a multiple of the value its program was built with and answers from the wrong prefix rather than raising when it is not. Under tracing that value is pinned at capture, so it has to be satisfied by the offset. Flooring recomputes at most ``alignment - 1`` tokens whose K/V is rewritten identically into the same blocks, so it is semantically a no-op. Mirrors the ``SDPA_CHUNK_ALIGN`` round-down in ``models/demos/llama3_70b_galaxy/tt/generator.py``. """ if kv_cache is None or kv_cache[0] is None: # Non-paged prefill: there is no page table to slice. return num_cached_per_user block_size = self._paged_prefill_block_size(kv_cache[0]) aligned = [] for i, (num_cached, seq_len) in enumerate(zip(num_cached_per_user, prompt_lens)): if int(num_cached) == 0: # Not a resume. Callers pass a zero-filled start_pos for an # ordinary prefill, so this is the common path and must not # require the model to describe an alignment it never uses. aligned.append(0) continue floored = (int(num_cached) // block_size) * block_size # The alignment depends on the padded suffix length, which depends on # the offset, so it has to settle: flooring lengthens the suffix, a # longer suffix can pin a larger q_chunk_size, and that can demand a # smaller offset again. # # The bound is exact, not generous. A pass that does not break must # lower the offset, which lengthens the suffix, which moves it to a # strictly higher padded bucket: an equal bucket would give an equal # alignment and the pass would have broken. So the bucket count bounds # the passes. It is a loose ceiling in practice, because # get_attn_sdpa_prefill_program_config pins only 64 or 256 and two # passes always suffice. for _ in range(len(get_all_padded_prefill_lengths(int(seq_len))) + 1): padded_suffix = get_padded_prefill_len(int(seq_len) - floored) alignment = self._resume_offset_alignment(padded_suffix, block_size) self._assert_uniform_resume_alignment(padded_suffix, block_size, alignment) settled = (int(num_cached) // alignment) * alignment if settled == floored: break floored = settled else: raise RuntimeError( f"user {i}: resume offset alignment did not settle for " f"start_pos={num_cached}, seq_len={seq_len}, block_size={block_size}" ) if floored != num_cached: logger.debug(f"Resume offset alignment: user {i} start_pos {num_cached} -> {floored}") aligned.append(floored) return aligned @staticmethod def _resumed_warmup_prompt_len(prefill_seq_len, num_cached, capped_warmup_seq_len): """Prompt length whose suffix after ``num_cached`` pads back to this bucket. Only the suffix reaches the kernel, so the prompt has to clear the offset by a full bucket. Spanning the bucket alone leaves no suffix at all once ``block_size`` reaches the smallest traced length. ``capped_warmup_seq_len`` is the ceiling the rest of warmup uses: past it the call is split into chunks and captures a different trace. """ return min(num_cached + prefill_seq_len, capped_warmup_seq_len) def _paged_prefill_block_size(self, kv_cache): """Block size for chunked-prefill page-table padding/slicing. Defaults to the cache's declared block_size. Models whose paged ops address an HMA-shared K/V buffer through a smaller per-layer effective block_size (e.g. gemma4 hybrid kv-cache groups: full-attention head_dim=512 viewing a buffer declared for a head_dim=256 sliding layer) override this so the page table math matches ``paged_fill_cache`` / the chunked SDPA. Non-overriding models are unaffected. """ return get_block_size(kv_cache) def _chunk_prefill_get_last_token(self, *, is_last_chunk, last_token_idx_in_chunk, chunk_size): """``get_last_token`` for one generator-level prefill chunk. Default (legacy): always the last-chunk's relative index. Correct for lm_head on the final chunk, but intermediate chunks then inherit a short index and under-fill their KV — fatal for models that treat ``get_last_token+1`` as the real fill length (Gemma4 bounded sliding). Those models override this. """ del is_last_chunk, chunk_size return (last_token_idx_in_chunk // 32) * 32 def _chunk_prefill_page_table(self, page_table, *, user_id, model_id=-1, kv_cache=None): """Page table + block_size for multi-chunk ``chunk_page_table`` slices. Full-attention ``paged_fill_cache`` writes via ``chunk_page_table`` (absolute block offsets for the current chunk). Returns ``(page_table, block_size)``. Default: the legacy ``page_table`` and ``_paged_prefill_block_size``. Hybrid kv-cache-group models override this to return a full-attention layer's per-layer table and that group's block_size — the legacy table is often group 0 (sliding), whose block IDs / column stride must not be used for full-layer fill. """ del user_id, model_id return page_table, self._paged_prefill_block_size(kv_cache) def prefill_forward_single_user_text( self, tokens, # New tokens to prefill (without the cached tokens), padded by get_padded_prefill_len() page_table, # Cached and new pages user_id, last_token_idx, # Last token index of the full prompt, including the cached tokens kv_cache=None, model_id=-1, num_cached_tokens: int = 0, batch_size=1, **kwargs, ): seq_len = tokens.shape[-1] use_chunked_prefill = seq_len > self.model_args[model_id].max_prefill_chunk_size use_prefix_caching = num_cached_tokens > 0 if use_chunked_prefill or use_prefix_caching: """ Chunked prefill requires paged attention. There are some strange constraints which we must meet: - page_table, which is used in SDPA, must match batch size of inputs, which is 1. This is because SDPA checks that page table batch dim matches input batch dim. Therefore we must slice the page table for the current user. - page_table must also have enough entries in each chunk, so it will be padded with zeros if necessary. - chunked_page_table is the slice of the page table for the current chunk. This is used by paged_fill_cache to keep it otherwise unaware that it is operating on a chunk. - due to the above point, we must always set user_id to 0 for chunked prefill. """ assert page_table is not None, "page_table must be provided for chunked prefill" assert kv_cache is not None, "kv_cache must be provided for chunked prefill" assert last_token_idx is not None and last_token_idx < seq_len + num_cached_tokens, ( f"last_token_idx must be provided and less than seq_len + num_cached_tokens: " f"last_token_idx={last_token_idx}, seq_len={seq_len}, num_cached_tokens={num_cached_tokens}" ) if use_chunked_prefill: # If chunked prefill (more than one chunk is needed), we want to use the maximum chunk size. chunk_size = get_max_prefill_chunk_size(seq_len, self.model_args[model_id].max_prefill_chunk_size) else: # Otherwise we only have one chunk. chunk_size = seq_len last_token_idx_in_seq = last_token_idx - num_cached_tokens # Excluding the cached tokens last_token_idx_in_chunk = last_token_idx_in_seq % chunk_size # Calculate which chunk contains the last_token_idx last_chunk_start = (last_token_idx_in_seq // chunk_size) * chunk_size # Hybrid models may substitute a full-attention per-layer table here # so ``chunk_page_table`` carries the block IDs (and column stride) that # full-layer fill actually writes (legacy ``page_table`` is often # sliding group 0 with a different unified block_size). chunk_source_page_table, block_size = self._chunk_prefill_page_table( page_table, user_id=user_id, model_id=model_id, kv_cache=kv_cache ) page_table_user = chunk_source_page_table[user_id : user_id + 1, :] # Trim over-wide tables (vLLM hybrid pads per-layer tables to # max_num_blocks_per_req) so the pad width below stays non-negative. needed_blocks = num_blocks_in_seq(seq_len + num_cached_tokens, block_size) if page_table_user.shape[1] > needed_blocks: page_table_user = page_table_user[:, :needed_blocks] num_padding_blocks = needed_blocks - page_table_user.shape[1] page_table_user_padded = torch.cat( [page_table_user, torch.zeros(1, num_padding_blocks, dtype=torch.int32)], dim=-1 ) CHUNK_USER_ID = 0 for chunk_start in range(num_cached_tokens, num_cached_tokens + seq_len, chunk_size): # These are absolute, i.e. including the cached tokens chunk_end = chunk_start + chunk_size # These are relative, i.e. excluding the cached tokens chunk_start_relative = chunk_start - num_cached_tokens chunk_end_relative = chunk_end - num_cached_tokens assert chunk_end <= num_cached_tokens + seq_len, ( f"chunk_end should be less or equal to " f"num_cached_tokens + seq_len. " f"Got: chunk_end={chunk_end}, " f"num_cached_tokens={num_cached_tokens}, seq_len={seq_len}" ) # Select tokens for the current chunk. # Cached tokens were already excluded (not part of the input), # so using relative indexes. chunk_tokens = tokens[:, chunk_start_relative:chunk_end_relative] # Select pages for the current chunk. # Cached pages must be skipped as well, # so using absolute indexes. chunk_page_table = page_table_user_padded[:, chunk_start // block_size : chunk_end // block_size] is_last_chunk = chunk_start_relative == last_chunk_start chunk_inputs = self.model[model_id].prepare_inputs_prefill( chunk_tokens, start_pos=chunk_start, page_table=page_table_user_padded, chunk_page_table=chunk_page_table, batch_size=batch_size, user_id=CHUNK_USER_ID, **kwargs, ) ( chunk_prefill_input, chunk_rot_mats_global_prefill, chunk_rot_mats_local_prefill, page_table_tt, chunk_page_table_tt, _chunk_start_idx_tt, ) = chunk_inputs tt_logits = self.model[model_id].ttnn_prefill_forward( chunk_prefill_input, rot_mats_global=chunk_rot_mats_global_prefill, rot_mats_local=chunk_rot_mats_local_prefill, user_id=CHUNK_USER_ID, page_table=page_table_tt, chunk_page_table=chunk_page_table_tt, chunk_start_idx=chunk_start, get_last_token=self._chunk_prefill_get_last_token( is_last_chunk=is_last_chunk, last_token_idx_in_chunk=last_token_idx_in_chunk, chunk_size=chunk_size, ), kv_cache=kv_cache, batch_size=batch_size, **kwargs, ) if is_last_chunk: return tt_logits else: del tt_logits else: inputs = self.model[model_id].prepare_inputs_prefill( tokens, page_table=page_table, batch_size=batch_size, user_id=user_id, **kwargs, ) prefill_input, rot_mats_global_prefill, rot_mats_local_prefill, page_table_tt, *_ = inputs tt_logits = self.model[model_id].ttnn_prefill_forward( prefill_input, rot_mats_global=rot_mats_global_prefill, rot_mats_local=rot_mats_local_prefill, user_id=user_id, page_table=page_table_tt, get_last_token=-1 if batch_size > 1 else (last_token_idx // 32) * 32, kv_cache=kv_cache, batch_size=batch_size, ) return tt_logits # Note: This function is called by vLLM def decode_forward( self, tokens, start_pos, page_table=None, kv_cache=None, enable_trace=True, read_from_device=True, sampling_params: SamplingParams = None, # Should be None if not greedy decoding / sampling on device. prompt_tokens: torch.Tensor | None = None, output_tokens: torch.Tensor | None = None, slot_remap=None, defer_device_sampling: bool = False, *, reload_inputs: bool, reload_page_table: bool, reload_sampling_params: bool, reset_sampling_state: bool, skip_trace_precompile: bool = False, prepare_trace: bool = False, **kwargs, ): if self.mode != Mode.DECODE: self.mode = Mode.DECODE # Switch to decode mode for prefetcher to reintialize sub devices for i in range(len(self.model)): self.model[i].switch_mode(Mode.DECODE) on_device_sampling = (sampling_params is not None) or defer_device_sampling if not enable_trace and not reload_inputs: raise ValueError("Non-traced decode rebuilds all forward inputs and requires reload_inputs=True") # Deferred sampling calls sample_decode_on_device() out of band; stash the # caller's explicit command so that call need not thread it separately. self._decode_reload_inputs = reload_inputs tokens = torch.chunk(tokens, self.data_parallel, 0) start_pos = torch.chunk(start_pos, self.data_parallel, 0) page_table = torch.chunk(page_table, self.data_parallel, 0) if page_table is not None else None decode_kwargs = { "current_pos": start_pos, "tokens": tokens, "page_table": page_table, "kv_cache": kv_cache, "on_device_sampling": on_device_sampling, } if enable_trace: tt_decode_output = self._decode_forward_trace_text( **decode_kwargs, reload_inputs=reload_inputs, reload_page_table=reload_page_table, skip_precompile=skip_trace_precompile, ) elif prepare_trace: tt_decode_output = self._prepare_decode_trace_variant(**decode_kwargs) else: tt_decode_output = self._decode_forward_no_trace_text( **decode_kwargs, ) # Device deferred if defer_device_sampling and on_device_sampling: return tt_decode_output # Device immediate if sampling_params is not None: tt_decode_output = self.sample_decode_on_device( tt_decode_output, sampling_params=sampling_params, start_pos=start_pos, prompt_tokens=prompt_tokens, output_tokens=output_tokens, slot_remap=slot_remap, enable_trace=enable_trace, reload_sampling_params=reload_sampling_params, reset_sampling_state=reset_sampling_state, skip_precompile=skip_trace_precompile, reload_inputs=reload_inputs, ) # Host sampling if read_from_device: to_host = self.read_decode_output(tt_decode_output) output = self.process_decode_output_host(to_host, is_tokens=(sampling_params is not None)) if sampling_params is None: # Host sampling does not invoke the device sampler, but its # dormant per-slot state must still follow the new layout. # Apply only after decode/readback succeeds so a failed call # can be retried with the still-pending, non-idempotent remap. self._apply_sampling_slot_remap(slot_remap) return output if sampling_params is None: self._apply_sampling_slot_remap(slot_remap) return tt_decode_output def _decode_forward_no_trace_text( self, tokens, current_pos, page_table=None, kv_cache=None, on_device_sampling=False, ): """ Performs text decode step. Returns tt_logits on device """ tt_output = [] tt_tokens = [] tt_current_pos = [] tt_rot_mat_idxs = [] tt_page_table = [] for i in range(self.data_parallel): user_page_table = page_table[i] if page_table is not None else None model_i = self.model[i] decode_inputs = model_i.prepare_inputs_decode(tokens[i], current_pos[i], user_page_table) # Compatibility with newer TT model adapters such as Gemma4: decode # input preparation may return auxiliary tensors after the common # four outputs, but the shared generator only consumes those four. ( tt_tokens_i, tt_current_pos_i, tt_rot_mat_idxs_i, tt_page_table_i, *_, ) = decode_inputs tt_tokens.append(tt_tokens_i) tt_current_pos.append(tt_current_pos_i) tt_rot_mat_idxs.append(tt_rot_mat_idxs_i) tt_page_table.append(tt_page_table_i) for i in range(self.data_parallel): user_kv_cache = kv_cache[i] if kv_cache is not None else None decode_out = self.model[i].ttnn_decode_forward( tt_tokens[i], tt_current_pos[i], rot_mat_idxs=tt_rot_mat_idxs[i], page_table=tt_page_table[i], kv_cache=user_kv_cache, on_device_logits=on_device_sampling, ) if isinstance(decode_out, tuple): tt_logits_i, tt_log_probs_i = decode_out else: tt_logits_i, tt_log_probs_i = decode_out, None tt_output.append((tt_logits_i, tt_log_probs_i)) return tt_output def _decode_trace_key(self, on_device_sampling, tokens): return on_device_sampling def _prepare_decode_trace_variant( self, tokens, current_pos, page_table=None, kv_cache=None, on_device_sampling=False ): """Stage a decode variant during an eager warmup call, inside the model's input routing.""" if self._uses_prefetcher(): return self._decode_forward_no_trace_text( tokens, current_pos, page_table=page_table, kv_cache=kv_cache, on_device_sampling=on_device_sampling ) key = self._decode_trace_key(on_device_sampling, tokens) if not hasattr(self, "_prepared_decode_traces"): self._prepared_decode_traces = {} if key not in self._prepared_decode_traces and not self.trace_ids_decode[key]: prepared = self._prepare_decode_trace_text( tokens, current_pos, page_table=page_table, kv_cache=kv_cache, on_device_sampling=on_device_sampling, return_compile_output=True, ) self._prepared_decode_traces[key] = prepared return prepared.pop("compile_output") return self._decode_forward_no_trace_text( tokens, current_pos, page_table=page_table, kv_cache=kv_cache, on_device_sampling=on_device_sampling ) def _prepare_decode_trace_text( self, tokens, current_pos, page_table=None, kv_cache=None, on_device_sampling=False, skip_precompile=False, return_compile_output=False, ): """Phase 1 of decode trace setup: run the compile pass and stage the persistent trace inputs. The trace inputs are long-lived -- refreshed in place for the whole decode loop -- so allocating them after the prefill traces were captured could place them inside a live trace's scratch and let every later prefill replay corrupt them. The first traced prefill call runs this before any capture (via _prepare_decode_trace_once); finalize_deferred_traces only records afterwards. """ # Compile run. Skipped when the caller already has a warmed program cache for this variant # (decode bucketing threads skip_precompile through from main). compile_output = None if not skip_precompile: compile_output = self._decode_forward_no_trace_text( tokens, current_pos, page_table=page_table, kv_cache=kv_cache, on_device_sampling=on_device_sampling, ) logger.info("Done Compiling Model") # Get inputs ready for trace run. device_inputs = [] for i in range(self.data_parallel): user_page_table = page_table[i] if page_table is not None else None host_inputs = self.model[i].prepare_decode_inputs_host( tokens[i], current_pos[i], page_table=user_page_table ) device_inputs_i = copy_host_to_device(host_inputs, mesh_device=self.model_args[i].mesh_device) _maybe_acknowledge_trace_buffers_corruptible(self, device_inputs_i) device_inputs.append(device_inputs_i) # Eager warmup stages this variant before the first capture, including # all sampling programs. Recording then consumes the staged inputs. all_sampling_configs = not self._any_trace_captured() for i in range(self.data_parallel): sampling_module = getattr(self.model[i], "sampling", None) if not on_device_sampling or sampling_module is None or compile_output is None: continue sampling_module.precompile( logits=compile_output[i][0], tt_out_tok=self._decode_token_feedback_buffer(self.model[i], device_inputs[i]), all_configs=all_sampling_configs, ) prepared = { "device_inputs": device_inputs, "kv_cache": kv_cache, "on_device_sampling": on_device_sampling, } if return_compile_output: prepared["compile_output"] = compile_output return prepared def _record_decode_trace_text(self, prepared): """Phase 2 of decode trace setup: capture the trace. Allocation-free outside the capture window -- it binds only to buffers allocated by :meth:`_prepare_decode_trace_text`. """ device_inputs = prepared["device_inputs"] prepared.pop("compile_output", None) kv_cache = prepared["kv_cache"] on_device_sampling = prepared["on_device_sampling"] tt_out_trace = [] trace_ids = {} for i in range(self.data_parallel): sampling_module = getattr(self.model[i], "sampling", None) sampling_trace_enabled = on_device_sampling and sampling_module is not None # Same reasoning as _record_trace_prefill: whatever the model allocates inside the capture # window belongs to the trace being recorded, and recording lane/variant N necessarily runs # while 1..N-1 are live. Acknowledge the window rather than flag it. with trace_allocation_tracker.corruptible_allocation_scope(self.model_args[i].mesh_device): trace_id = ttnn.begin_trace_capture(self.model_args[i].mesh_device, cq_id=0) trace_ids[i] = trace_id user_kv_cache = kv_cache[i] if kv_cache is not None else None model_inputs = device_inputs[i][:4] if len(device_inputs[i]) > 4 else device_inputs[i] # Models that produce extra device inputs beyond the first # four (e.g. Gemma4's host-precomputed per-layer-input at # index 4) feed them into ``ttnn_decode_forward`` via a # model-side stash rather than through the call signature. # Give the model a chance to bind that stash to the # *trace-input* device tensors here, before the trace is # captured — otherwise traced ops stay pointed at whatever # device buffer the compile run produced, and trace replay # reads stale data because ``copy_host_to_device`` only # refreshes ``trace_inputs_decode``. bind_trace_inputs = getattr(self.model[i], "bind_decode_trace_inputs", None) if bind_trace_inputs is not None: bind_trace_inputs(device_inputs[i]) tt_out_trace.append( self.model[i].ttnn_decode_forward( *model_inputs, kv_cache=user_kv_cache, on_device_logits=on_device_sampling, ) ) ttnn.end_trace_capture(self.model_args[i].mesh_device, trace_id, cq_id=0) _maybe_acknowledge_trace_buffers_corruptible(self, tt_out_trace[-1]) if sampling_trace_enabled: # NOTE: sampling trace can be keyed depending on sampling params, # this traces only for the current ones. # tt_out_tok feeds the sampled token back into the decode token # buffer (device_inputs[0]) for the next traced step. Only do this # for models that rely on on-device token feedback. Some token # input buffers are not shaped as sampling outputs (gemma4's is # rank-2; ttnn.sampling requires a rank-4 preallocated output), # so those models opt out and sampling allocates its own output. tt_out_tok = self._decode_token_feedback_buffer(self.model[i], device_inputs[i]) # skip_precompile=True in both cases: either _prepare_decode_trace_text pre-compiled the # sampling pipeline (before any trace was live), or the caller passed skip_precompile and # is asserting the program cache is already warm for this variant. sampling_module.capture_trace(logits=tt_out_trace[i], tt_out_tok=tt_out_tok, skip_precompile=True) logger.info("Done Capturing Decode Trace") return trace_ids, tt_out_trace, *device_inputs def precapture_decode_trace_variants(self, sampling_params, tokens, start_pos, page_table, kv_cache): """Record every decode trace variant a warmup sweep will need, preparing all of them first. Decode traces are keyed by on-device sampling on/off. A sweep that captures the second variant lazily does so with the first variant's traces already live, so its compile pass, its persistent trace inputs and its sampling pre-compile all allocate behind a live trace (measured on Gemma-3-27B DP-4: 58 stranded buffers). Prepare each missing variant before recording any of them, then record in one go. Returns False when nothing could be pre-captured (caller keeps its lazy path). """ variants = [] for param in sampling_params: variant = param is not None if variant not in variants: variants.append(variant) if not variants or page_table is None or self._uses_prefetcher(): return False if self.mode != Mode.DECODE: self.mode = Mode.DECODE for i in range(len(self.model)): self.model[i].switch_mode(Mode.DECODE) tokens = torch.chunk(tokens, self.data_parallel, 0) start_pos = torch.chunk(start_pos, self.data_parallel, 0) page_table = torch.chunk(page_table, self.data_parallel, 0) variants = [(v, self._decode_trace_key(v, tokens)) for v in variants] variants = [(v, key) for v, key in variants if not self.trace_ids_decode[key]] if not variants: return False staged = getattr(self, "_prepared_decode_traces", {}) prepared = [] for variant, key in variants: if key in staged: prep = staged.pop(key) else: prep = self._prepare_decode_trace_text( tokens, start_pos, page_table=page_table, kv_cache=kv_cache, on_device_sampling=variant ) prepared.append((key, prep)) for key, prep in prepared: trace_ids, tt_out_trace, *device_inputs = self._record_decode_trace_text(prep) self.trace_ids_decode[key] = trace_ids self.trace_inputs_decode[key] = device_inputs self.trace_output_decode[key] = tt_out_trace return True def _capture_decode_trace_text( self, tokens, current_pos, page_table=None, kv_cache=None, on_device_sampling=False, skip_precompile=False, ): """Prepare and immediately capture the decode trace. Only safe when no other trace is live on the device. The demo path pre-captures via :meth:`warmup_model_decode` during warmup; this single-shot form is the fallback for callers that reach decode without having warmed up. ``skip_precompile`` is forwarded to the prepare phase, where main's decode-bucketing callers use it to state that the program cache is already warm for this variant. """ key = self._decode_trace_key(on_device_sampling, tokens) staged = getattr(self, "_prepared_decode_traces", {}) prepared = ( staged.pop(key) if key in staged else self._prepare_decode_trace_text( tokens, current_pos, page_table=page_table, kv_cache=kv_cache, on_device_sampling=on_device_sampling, skip_precompile=skip_precompile, ) ) return self._record_decode_trace_text(prepared) def _decode_forward_trace_text( self, tokens, current_pos, page_table=None, kv_cache=None, on_device_sampling=False, *, reload_inputs: bool, reload_page_table: bool, skip_precompile: bool = False, ): """ Run decode forward text with tracing ``reload_inputs`` (from decode_forward): host token/position inputs are authoritative this step and must overwrite every device-resident input. ``reload_page_table`` refreshes only the page table while preserving device-produced token and position state. """ # The trace is different depending on whether we are doing device sampling or not if not self.trace_ids_decode[on_device_sampling]: trace_ids, tt_out_trace, *device_inputs = self._capture_decode_trace_text( tokens, current_pos, page_table=page_table, kv_cache=kv_cache, on_device_sampling=on_device_sampling, skip_precompile=skip_precompile, ) self.trace_ids_decode[on_device_sampling] = trace_ids self.trace_inputs_decode[on_device_sampling] = device_inputs self.trace_output_decode[on_device_sampling] = tt_out_trace for i in range(self.data_parallel): user_page_table = page_table[i] if page_table is not None else None if reload_inputs: # Full resets are required when host token/position inputs are # authoritative again, or for models that explicitly opt out of # partial decode trace input refreshes. host_inputs_i = self.model[i].prepare_decode_inputs_host(tokens[i], current_pos[i], user_page_table) copy_host_to_device( host_tensors=host_inputs_i, device_tensors=self.trace_inputs_decode[on_device_sampling][i], ) elif reload_page_table: # With async device sampling, token/position inputs may # intentionally be stale on host: the previous decode updates # them on device. Page tables still need refreshing when new KV # blocks are allocated, so copy only that trace input and # preserve device-produced tokens. host_inputs_i = self.model[i].prepare_decode_inputs_host(tokens[i], current_pos[i], user_page_table) host_page_table = host_inputs_i[DECODE_PAGE_TABLE_INPUT_IDX] device_page_table = self.trace_inputs_decode[on_device_sampling][i][DECODE_PAGE_TABLE_INPUT_IDX] if host_page_table is not None: ttnn.copy_host_to_device_tensor(host_page_table, device_page_table) for i, trace_id in self.trace_ids_decode[on_device_sampling].items(): ttnn.execute_trace(self.model_args[i].mesh_device, trace_id, cq_id=0, blocking=False) return self.trace_output_decode[on_device_sampling] def _apply_sampling_slot_remap(self, slot_remap) -> None: if slot_remap is None: return global_remap = torch.as_tensor(slot_remap, dtype=torch.long).reshape(-1) if global_remap.numel() % self.data_parallel != 0: raise ValueError( f"slot_remap has {global_remap.numel()} entries, which cannot be " f"split across {self.data_parallel} data-parallel lanes" ) lane_stride = global_remap.numel() // self.data_parallel for i in range(self.data_parallel): sampling_module = getattr(self.model[i], "sampling", None) if sampling_module is None: continue sm_bs = sampling_module.seed_manager.max_batch_size if lane_stride > sm_bs: raise ValueError( f"slot_remap lane width {lane_stride} exceeds sampling state " f"width {sm_bs} for lane {i}" ) lane_base = i * lane_stride lane_remap = global_remap[lane_base : lane_base + lane_stride] - lane_base if torch.any(lane_remap < 0) or torch.any(lane_remap >= lane_stride): raise ValueError( f"slot_remap lane {i} references a slot outside its global " f"range [{lane_base}, {lane_base + lane_stride})" ) rank_remap = torch.arange(sm_bs, dtype=torch.long) rank_remap[:lane_stride] = lane_remap sampling_module.apply_slot_remap(rank_remap) def sample_decode_on_device( self, tt_logits, sampling_params, start_pos=None, prompt_tokens: torch.Tensor | None = None, output_tokens: torch.Tensor | None = None, slot_remap=None, enable_trace=False, *, reload_sampling_params: bool, reset_sampling_state: bool, reload_inputs: bool | None = None, skip_precompile: bool = False, ): """Sample this decode step's tokens on device. ``reload_inputs`` identifies authoritative host positions for seed counter alignment. Deferred callers may omit it after ``decode_forward``; the explicit command from that call is retained for this purpose. """ if reload_inputs is None: reload_inputs = getattr(self, "_decode_reload_inputs", True) # Keep this entry point independently usable by immediate and # separated-sampling callers. self._apply_sampling_slot_remap(slot_remap) # sampling_dp may differ from data_parallel for models that internally # shard users across mesh rows (users_row_sharded) — each row samples # 32 users independently, so sampling params must be chunked by the # number of rows even though data_parallel=1 for the forward pass. sampling_dp_values = [getattr(self.model[i], "sampling_dp", 1) for i in range(self.data_parallel)] assert ( len(set(sampling_dp_values)) == 1 ), f"All model instances must have the same sampling_dp, got {sampling_dp_values}" # NOTE: This assumes data_parallel and sampling_dp are mutually exclusive # (one is always 1). If a future model needs both DP>1 and row-sharded # sampling, this should become data_parallel * sampling_dp_values[0]. sampling_dp = max(self.data_parallel, sampling_dp_values[0]) sampling_params_list = chunk_sampling_params(sampling_params, sampling_dp) prompt_chunks = ( torch.chunk(prompt_tokens, sampling_dp, 0) if prompt_tokens is not None else [None] * sampling_dp ) output_chunks = ( torch.chunk(output_tokens, sampling_dp, 0) if output_tokens is not None else [None] * sampling_dp ) for i in range(self.data_parallel): sampling_module = getattr(self.model[i], "sampling", None) assert sampling_module is not None, "Sampling module not found in model for sampling on device." assert ( sampling_dp % self.data_parallel == 0 ), f"sampling_dp ({sampling_dp}) must be divisible by data_parallel ({self.data_parallel})" cpm = sampling_dp // self.data_parallel start = i * cpm model_chunks = sampling_params_list[start : start + cpm] model_prompt = ( torch.cat([c for c in prompt_chunks[start : start + cpm] if c is not None], 0) if prompt_tokens is not None else None ) model_output = ( torch.cat([c for c in output_chunks[start : start + cpm] if c is not None], 0) if output_tokens is not None else None ) sampling_module.apply_decode_state( model_chunks, reload_sampling_params=reload_sampling_params, reset_sampling_state=reset_sampling_state, prompt_tokens=model_prompt, output_tokens=model_output, ) active_seed_slots = None if start_pos is not None and start_pos[i] is not None: max_seed_slots = sampling_module.seed_manager.max_batch_size start_values = torch.as_tensor(start_pos[i]).reshape(-1).tolist() active_seed_slots = [idx for idx, pos in enumerate(start_values[:max_seed_slots]) if int(pos) >= 0] # A request finishing at the batch tail produces no non-identity # remap, so retire seed state that no longer belongs to a live row. if active_seed_slots is not None: sampling_module.seed_manager.deactivate_slots_except(active_seed_slots) # Register each request's explicit seed into the seed manager and # tie its RNG counter to the absolute decode position before # advancing. Without registration the per-request seed never reaches # the device (the seed manager stays unseeded), so sampling falls # back to per-slot boot RNG and two requests sharing a seed diverge # (this regressed when #45166 dropped these calls from the decode # flow). Position alignment then keeps the stream reproducible even # when vLLM evicts a running request and re-admits it in a different # slot under async scheduling. Mirrors the llama3_70b_galaxy decode path. # # Align only from a trustworthy position (#51981): the counter # self-advances per token, so re-anchoring to a lagging host start_pos # is what breaks reproducibility. Trustworthy means the caller # commanded reload_inputs, or the slot was explicitly reset/reseeded. if active_seed_slots is not None and (reload_inputs or reload_sampling_params or reset_sampling_state): seed_bs = sampling_module.tt_sampling.max_batch_size if len(model_chunks) == 1: seed_values = format_sampling_params(model_chunks[0], seed_bs).seed else: seed_values = [] for chunk in model_chunks: s = format_sampling_params(chunk, seed_bs).seed seed_values += s if isinstance(s, list) else [s] * seed_bs if reset_sampling_state: # A state reset is unconditional even for seed=None: # reset_if_needed would see None == None and skip the # fresh device-seed upload in decode-only sampling mode. sampling_module.seed_manager.reset_seed_from_slots(seed_values, active_seed_slots) reseeded_slots = list(active_seed_slots) elif reload_sampling_params: reseeded_slots = sampling_module.seed_manager.reset_seed_from_slots_if_needed( seed_values, active_seed_slots ) else: reseeded_slots = [] align_slots = active_seed_slots if reload_inputs else reseeded_slots if align_slots: sampling_module.seed_manager.align_seed_counters_to_positions( seed_values, align_slots, start_values ) sampling_module.seed_manager.get_new_values(active_seed_slots) sampled_outputs = [] for i in range(self.data_parallel): sampling_module = getattr(self.model[i], "sampling", None) if sampling_module is None: sampled_outputs.append(tt_logits[i]) continue logits_i = tt_logits[i] if isinstance(logits_i, tuple): logits_i = logits_i[0] # Some models must run the on-device sampling op eagerly rather than from its # own captured trace: the force-argmax path does an all_gather_async whose # multi_device_global_semaphore is taken from get_and_cycle_*() at capture time # and frozen into the trace, so replaying the sampling trace reuses a stale # semaphore and the gather corrupts from the 2nd decode step (#48037). Running # sampling eagerly re-acquires a fresh semaphore each step. sampling_enable_trace = enable_trace and not getattr(self.model[i], "_tt_disable_sampling_trace", False) # Must match the capture-time decision in _capture_decode_trace_text: # only feed the sampled token back into device_inputs[0] for models # that use on-device token feedback (see _decode_token_feedback_buffer). tt_out_tok = ( self._decode_token_feedback_buffer(self.model[i], self.trace_inputs_decode[True][i]) if sampling_enable_trace and self.trace_inputs_decode[True] else None ) sampled_outputs.append( sampling_module.sample( logits=logits_i, tt_out_tok=tt_out_tok, enable_trace=sampling_enable_trace, skip_precompile=skip_precompile, ) ) return sampled_outputs @staticmethod def _decode_token_feedback_buffer(model, device_inputs): """Return the device token buffer to feed the sampled token back into for the next traced decode step, or None if the model doesn't use on-device token feedback. Some models' token input buffer is not a valid sampling output (gemma4's is rank-2; ``ttnn.sampling`` requires a rank-4 preallocated output). Returning None makes sampling allocate its own output instead of writing into ``device_inputs[0]``. """ if not getattr(model, "_tt_supports_decode_token_feedback", True): return None return device_inputs[0] def _prefill_forward_single_user( self, vision_images, vision_mask, tokens, xattn_caches, user_id, total_len, prefill_len, page_table=None, kv_cache=None, cross_page_table=None, model_id=-1, ): """ Performs vision encode step then text prefill. Returns (xattn_caches, cross_attention_masks, full_text_row_masked_out_mask, logits) """ B = tokens.shape[0] last_token_idx = prefill_len - 1 text_only_inference = vision_images is None if not text_only_inference: ( vision_tokens, prefill_cross_attention_masks, prefill_full_text_row_masked_out_mask, decode_cross_attention_masks, decode_full_text_row_masked_out_mask, ) = self.model[model_id].compute_vision_tokens_masks( batch_images=[vision_images], batch_masks=[vision_mask], total_len=total_len, prefill_len=prefill_len, ) if cross_page_table is not None: num_vision_tokens = vision_tokens.shape[2] cross_page_table = self._get_prefill_user_page_table(cross_page_table, kv_cache, num_vision_tokens) else: ( vision_tokens, prefill_cross_attention_masks, prefill_full_text_row_masked_out_mask, decode_cross_attention_masks, decode_full_text_row_masked_out_mask, ) = (None, None, None, None, None) if page_table is not None: page_table = self._get_prefill_user_page_table(page_table, kv_cache, prefill_len) ( tt_h, tt_xattn_mask, tt_full_text_mask_expand_1NSH, tt_full_text_mask_expand_11SD, rot_mats, tt_page_table, tt_cross_page_table, ) = self.model[model_id].prepare_inputs_prefill( tokens, prefill_cross_attention_masks, prefill_full_text_row_masked_out_mask, prefill_len=prefill_len, page_table=page_table, cross_page_table=cross_page_table, text_only_inference=text_only_inference, ) tt_logits = self.model[model_id].ttnn_prefill_forward( tt_h, tt_xattn_mask, tt_full_text_mask_expand_1NSH, tt_full_text_mask_expand_11SD, xattn_caches, rot_mats, user_id, vision_tokens, page_table=tt_page_table, kv_cache=kv_cache, get_last_token=(last_token_idx // 32) * 32, cross_page_table=tt_cross_page_table, text_only_inference=text_only_inference, ) del tt_page_table del tt_cross_page_table return ( xattn_caches, prefill_cross_attention_masks, prefill_full_text_row_masked_out_mask, decode_cross_attention_masks, decode_full_text_row_masked_out_mask, tt_logits, ) # Note: This function is called by vLLM def prefill_forward( self, vision_images, vision_masks, tokens, xattn_caches, total_lens, prompt_lens, page_table=None, kv_cache=None, cross_page_table=None, empty_slots=None, **kwargs, ): if not self.model_args[0].is_llama_vision(): logits = self.prefill_forward_text( tokens, page_table=page_table, kv_cache=kv_cache, prompt_lens=prompt_lens, pixel_values=vision_images, **kwargs, ) return logits, None, None, None, None else: ( output_logits, prefill_output_xattn_masks, prefill_output_full_text_row_masked_out_masks, decode_output_xattn_masks, decode_output_full_text_row_masked_out_masks, ) = self.prefill_forward_llama_vision( vision_images, vision_masks, tokens, xattn_caches, total_lens, prompt_lens, page_table=page_table, kv_cache=kv_cache, cross_page_table=cross_page_table, empty_slots=empty_slots, ) return ( output_logits, prefill_output_xattn_masks, prefill_output_full_text_row_masked_out_masks, decode_output_xattn_masks, decode_output_full_text_row_masked_out_masks, ) # Note: This function is called by vLLM def warmup_vision_encoder(self): """Run the vision encoder once per supported image canvas so its programs exist before any trace. The encoder's programs are keyed on the image geometry: 1, 2 or 4 chunks through the image blocks, and the tile aspect ratio through the tile position embeddings. A request with a new geometry otherwise compiles the whole encoder path (image blocks, positional embeddings, pad and concat, the vision projection) while the decode trace is live - measured on Llama-3.2-11B-Vision T3K batch-1 as 60+ buffers alive across trace replays, arriving with the second and third distinct images. One blank image per supported canvas covers every geometry the transform can produce. """ from PIL import Image as PIL_Image if getattr(self, "already_warmed_up_vision", False): return for model in self.model: transform = getattr(model, "image_transform", None) if transform is None or not hasattr(model, "compute_vision_tokens_masks"): continue base_transform = getattr(transform, "func", transform) resolutions = base_transform.find_supported_resolutions( max_num_chunks=model.max_num_chunks, patch_size=base_transform.size ) for height, width in dict.fromkeys(tuple(r) for r in resolutions): logger.info(f"Warming up vision encoder for a {width}x{height} image") model.compute_vision_tokens_masks( batch_images=[[PIL_Image.new("RGB", (width, height))]], batch_masks=[[[0, -1]]], total_len=128, prefill_len=128, ) self.already_warmed_up_vision = True def prefill_forward_llama_vision( self, vision_images, vision_masks, tokens: torch.Tensor, xattn_caches, total_lens, prompt_lens, page_table=None, kv_cache=None, cross_page_table=None, empty_slots=None, ): """ Batched version of _prefill_forward_single_user for vision model. """ self.warmup_vision_encoder() if page_table is not None: assert isinstance(page_table, torch.Tensor), "page_table mush be torch.Tensor" if cross_page_table is not None: assert isinstance(cross_page_table, torch.Tensor), "cross_page_table mush be torch.Tensor" batch_size, batch_seq_len = tokens.shape max_batch_size_per_model = self.model_args[0].max_batch_size output_logits = torch.zeros(batch_size, 1, self.model_args[0].vocab_size) out_list = [] prefill_output_xattn_masks = [] prefill_output_full_text_row_masked_out_masks = [] decode_output_xattn_masks = [] decode_output_full_text_row_masked_out_masks = [] if empty_slots is None: empty_slots = list(range(batch_size)) for idx, user_id in enumerate(empty_slots): model_id = user_id // max_batch_size_per_model group_user_id = user_id % max_batch_size_per_model if page_table is None else 0 seq_len = int(prompt_lens[idx]) logger.info(f"Prefilling User {user_id + 1} up to {seq_len} tokens") user_page_table = page_table[idx : idx + 1] if page_table is not None else None user_cross_page_table = cross_page_table[idx : idx + 1] if kv_cache is not None else None model_kv_cache = kv_cache[model_id] if kv_cache is not None else None model_xattn_cache = xattn_caches[model_id] if xattn_caches is not None else None ( model_xattn_cache, prefill_cross_attention_masks, prefill_full_text_row_masked_out_mask, decode_cross_attention_masks, decode_full_text_row_masked_out_mask, logits, ) = self._prefill_forward_single_user( vision_images=vision_images[idx], vision_mask=vision_masks[idx], tokens=tokens[idx : idx + 1, :seq_len], # Keep batch dimension xattn_caches=model_xattn_cache, user_id=group_user_id, total_len=total_lens[idx], prefill_len=seq_len, page_table=user_page_table, kv_cache=model_kv_cache, cross_page_table=user_cross_page_table, model_id=model_id, ) if xattn_caches is not None: xattn_caches[model_id] = model_xattn_cache out_list.append(logits) prefill_output_xattn_masks.append(prefill_cross_attention_masks) prefill_output_full_text_row_masked_out_masks.append(prefill_full_text_row_masked_out_mask) decode_output_xattn_masks.append(decode_cross_attention_masks) decode_output_full_text_row_masked_out_masks.append(decode_full_text_row_masked_out_mask) # We gather prefill output at the end of prefill to reduce unnecessary device sync for idx, user_id in enumerate(empty_slots): model_id = user_id // max_batch_size_per_model last_token_idx = prompt_lens[idx] - 1 output_logits[idx] = self.model[model_id].process_output_prefill( out_list[idx].cpu(), 1, last_token_idx=(last_token_idx % 32) ) logger.info(f"Finished prefill for all users up to {batch_seq_len} tokens, Starting decode...") return ( output_logits, prefill_output_xattn_masks, prefill_output_full_text_row_masked_out_masks, decode_output_xattn_masks, decode_output_full_text_row_masked_out_masks, ) # Note: This function is called by vLLM def decode_forward_llama_vision( self, start_pos, tokens, prefill_cross_attention_masks, prefill_full_text_row_masked_out_mask, decode_cross_attention_masks, decode_full_text_row_masked_out_mask, xattn_caches=None, page_table=None, kv_cache=None, cross_page_table=None, enable_trace=True, read_from_device=True, ): # vLLM may warm decode before the first image request. Compile every # canvas before that path captures a trace as well. self.warmup_vision_encoder() B = tokens.shape[0] data_parallel = min(B, self.data_parallel) batch_per_device = B // data_parallel tokens = torch.chunk(tokens, self.data_parallel, 0) start_pos = torch.chunk(start_pos, self.data_parallel, 0) prefill_cross_attention_masks = [ prefill_cross_attention_masks[i * batch_per_device : (i + 1) * batch_per_device] for i in range(data_parallel) ] prefill_full_text_row_masked_out_mask = [ prefill_full_text_row_masked_out_mask[i * batch_per_device : (i + 1) * batch_per_device] for i in range(data_parallel) ] decode_cross_attention_masks = [ decode_cross_attention_masks[i * batch_per_device : (i + 1) * batch_per_device] for i in range(data_parallel) ] decode_full_text_row_masked_out_mask = [ decode_full_text_row_masked_out_mask[i * batch_per_device : (i + 1) * batch_per_device] for i in range(data_parallel) ] page_table = torch.chunk(page_table, self.data_parallel, 0) if page_table is not None else None cross_page_table = ( torch.chunk(cross_page_table, self.data_parallel, 0) if cross_page_table is not None else None ) decode_kwargs = { "position_id": start_pos, "tokens": tokens, "prefill_cross_attention_masks": prefill_cross_attention_masks, "prefill_full_text_row_masked_out_mask": prefill_full_text_row_masked_out_mask, "decode_cross_attention_masks": decode_cross_attention_masks, "decode_full_text_row_masked_out_mask": decode_full_text_row_masked_out_mask, "xattn_caches": xattn_caches, "page_table": page_table, "kv_cache": kv_cache, "cross_page_table": cross_page_table, } if enable_trace: tt_logits = self._easy_trace(**decode_kwargs) else: tt_logits = self._decode_forward_no_trace(**decode_kwargs) if read_from_device: to_host = self.read_decode_output(tt_logits) return self.process_decode_output_host(to_host) else: return tt_logits # Note: This function is called by vLLM def read_decode_output(self, tt_out, async_read=False): """ Input tt_out is list of tuples of (tt_out_tok, tt_log_probs) tt_log_probs can be: ttnn.Tensor (old path), LogProbsResult (new path), or None. """ def _read_logprobs(lp, blocking: bool = True): if lp is None: return None return lp.cpu(blocking=blocking) if not async_read: if isinstance(tt_out[0], tuple): return [(out[0].cpu(), _read_logprobs(out[1])) for out in tt_out] elif isinstance(tt_out[0], ttnn.Tensor): return [out.cpu() for out in tt_out] host_outputs = [] read_events = [] for i in range(self.data_parallel): if isinstance(tt_out[i], tuple): outputs = ( tt_out[i][0].cpu(blocking=False), _read_logprobs(tt_out[i][1], blocking=False), ) host_outputs.append(outputs) elif isinstance(tt_out[i], ttnn.Tensor): outputs = tt_out[i].cpu(blocking=False) host_outputs.append(outputs) read_events.append(ttnn.record_event(self.model[i].mesh_device, 0)) return host_outputs, read_events # Note: This function is called by vLLM def process_decode_output_host(self, tt_out, is_tokens=False): """ Converts the input ttnn host tensors to torch tensors. The input can be logits (if is_tokens=False) or tokens (if is_tokens=True). When the decode output includes logprobs: * Old path: a single logprobs tensor is converted to a torch tensor. * New path (LogProbsResult): the LogProbsResult is converted into a tuple of torch tensors (topk_lp, topk_idx), where each has shape [batch, top_k]. Returns: * If using the old path: (logits, log_probs) where both are torch tensors concatenated across data-parallel ranks. * If any rank uses the new path: (logits, (topk_lp, topk_idx)), where logits, topk_lp, and topk_idx are torch tensors concatenated across data-parallel ranks. """ from models.common.sampling.tt_log_probs import LogProbsResult max_batch_size_per_model = self.model_args[0].max_batch_size logits = [] log_probs = [] for i in range(self.data_parallel): if isinstance(tt_out[i], tuple): logits_i = self.model[i].process_output_decode( tt_out[i][0], max_batch_size_per_model, S=1, is_tokens=is_tokens ) lp = tt_out[i][1] if isinstance(lp, LogProbsResult): # New path: convert LogProbsResult to torch (topk_lp, topk_idx) tuple. # # LogProbsResult contains device tensors of shape (1,1,32,32) — 32 users # × 32 top-k logprobs — replicated across all devices in the mesh. # However, for row-sharded sampling (sampling_dp > 1), each mesh row # independently computes logprobs for its own 32 users, so the content # differs per row even though the tensor is "replicated." # # We cannot use a mesh composer (ConcatMesh2dToTensor) because it would # concatenate all 32 devices including 8 column replicas per row, giving # 8× duplicated data. Instead: # - Row-sharded (sampling_dp > 1): pick one device per row (first in # each row), read its [32, 32] tensor, concatenate rows → [128, 32]. # - Non-row-sharded (sampling_dp == 1): read from a single device. lp_tensor = lp.topk_logprobs_host if lp.topk_logprobs_host is not None else lp.topk_logprobs idx_tensor = lp.topk_indices_host if lp.topk_indices_host is not None else lp.topk_indices sampling_dp = getattr(self.model[i], "sampling_dp", 1) if sampling_dp > 1: # Row-sharded: read one device per row and concatenate rows, cols = self.mesh_device.shape device_tensors_lp = ttnn.get_device_tensors(lp_tensor) device_tensors_idx = ttnn.get_device_tensors(idx_tensor) row_lps = [] row_idxs = [] for row in range(rows): dev_idx = row * cols # first device in this row row_lp = ttnn.to_torch(device_tensors_lp[dev_idx]) row_lps.append(row_lp.reshape(-1, row_lp.shape[-1])[:max_batch_size_per_model]) row_idx = ttnn.to_torch(device_tensors_idx[dev_idx]) row_idxs.append(row_idx.reshape(-1, row_idx.shape[-1])[:max_batch_size_per_model]) topk_lp = torch.cat(row_lps, dim=0).float() topk_idx = torch.cat(row_idxs, dim=0).to(torch.int32) else: # Non-row-sharded: read from first device only device_tensors_lp = ttnn.get_device_tensors(lp_tensor) device_tensors_idx = ttnn.get_device_tensors(idx_tensor) topk_lp = ( ttnn.to_torch(device_tensors_lp[0]) .reshape(-1, device_tensors_lp[0].shape[-1])[:max_batch_size_per_model] .float() ) topk_idx = ( ttnn.to_torch(device_tensors_idx[0]) .reshape(-1, device_tensors_idx[0].shape[-1])[:max_batch_size_per_model] .to(torch.int32) ) logits.append(logits_i) log_probs.append((topk_lp, topk_idx)) elif lp is not None: # Old path: single logprob tensor log_probs_i = self.model[i].process_output_decode( lp, max_batch_size_per_model, S=1, is_tokens=is_tokens, is_log_probs=True ) logits.append(logits_i) log_probs.append(log_probs_i) else: logits.append(logits_i) log_probs.append(torch.ones(logits_i.shape)) elif isinstance(tt_out[i], ttnn.Tensor): logits_i = self.model[i].process_output_decode( tt_out[i], max_batch_size_per_model, S=1, is_tokens=is_tokens ) logits.append(logits_i) log_probs.append(torch.ones(logits_i.shape)) else: raise ValueError(f"Invalid type of tt_out: {type(tt_out[i])}") # Check if any DP rank returned new-path tuples (topk_lp, topk_idx) has_topk = any(isinstance(lp, tuple) for lp in log_probs) if has_topk: # New path: all DP ranks should have tuples. For ranks that # returned a dummy tensor (e.g. sz=0), create matching dummy tuples. normalized = [] for lp in log_probs: if isinstance(lp, tuple): normalized.append(lp) else: # Dummy: shape [B, 32] zeros to match tuple format B = lp.shape[0] normalized.append((torch.zeros(B, 32, dtype=torch.float32), torch.zeros(B, 32, dtype=torch.int32))) all_lp = torch.cat([lp[0] for lp in normalized], 0) all_idx = torch.cat([lp[1] for lp in normalized], 0) return (torch.cat(logits, 0), (all_lp, all_idx)) return (torch.cat(logits, 0), torch.cat(log_probs, 0)) def _decode_forward_no_trace( self, position_id, tokens, prefill_cross_attention_masks, prefill_full_text_row_masked_out_mask, decode_cross_attention_masks, decode_full_text_row_masked_out_mask, xattn_caches=None, page_table=None, kv_cache=None, cross_page_table=None, ): """ Performs text decode step. Returns tt_logits on device """ # forward_decode should be traced callable # decorator does compilation, capture, execute tt_h = [] tt_xattn_mask = [] tt_full_text_mask_expand_1NSH = [] tt_full_text_mask_expand_11SD = [] tt_position_id = [] tt_rot_mats = [] tt_page_table = [] tt_cross_page_table = [] for i in range(self.data_parallel): B, S = tokens[i].shape assert S == 1 user_page_table = page_table[i] if page_table is not None else None user_cross_page_table = cross_page_table[i] if cross_page_table is not None else None ( tt_h_i, tt_xattn_mask_i, tt_full_text_mask_expand_1NSH_i, tt_full_text_mask_expand_11SD_i, tt_position_id_i, tt_rot_mats_i, tt_page_table_i, tt_cross_page_table_i, ) = self.model[i].prepare_inputs_decode( tokens[i], prefill_cross_attention_masks[i], prefill_full_text_row_masked_out_mask[i], decode_cross_attention_masks[i], decode_full_text_row_masked_out_mask[i], position_id=position_id[i], page_table=user_page_table, cross_page_table=user_cross_page_table, ) tt_h.append(tt_h_i) tt_xattn_mask.append(tt_xattn_mask_i) tt_full_text_mask_expand_1NSH.append(tt_full_text_mask_expand_1NSH_i) tt_full_text_mask_expand_11SD.append(tt_full_text_mask_expand_11SD_i) tt_position_id.append(tt_position_id_i) tt_rot_mats.append(tt_rot_mats_i) tt_page_table.append(tt_page_table_i) tt_cross_page_table.append(tt_cross_page_table_i) tt_logits = [] tt_log_probs = [] for i in range(self.data_parallel): user_kv_cache = kv_cache[i] if kv_cache is not None else None xattn_cache = xattn_caches[i] if xattn_caches is not None else None tt_logits_i, tt_log_probs_i = self.model[i].ttnn_decode_forward( tt_h[i], tt_xattn_mask[i], tt_full_text_mask_expand_1NSH[i], tt_full_text_mask_expand_11SD[i], xattn_cache, tt_position_id[i], tt_rot_mats[i], page_table=tt_page_table[i], kv_cache=user_kv_cache, cross_page_table=tt_cross_page_table[i], ) tt_logits.append(tt_logits_i) tt_log_probs.append(tt_log_probs_i) return tt_logits, tt_log_probs def _capture_trace( self, position_id, tokens, prefill_cross_attention_masks, prefill_full_text_row_masked_out_mask, decode_cross_attention_masks, decode_full_text_row_masked_out_mask, xattn_caches, page_table=None, kv_cache=None, cross_page_table=None, ): """ Captures a trace for the decode_forward method. """ tt_h = [] tt_xattn_mask = [] tt_full_text_mask_expand_1NSH = [] tt_full_text_mask_expand_11SD = [] tt_position_id = [] tt_rot_mats = [] tt_page_table = [] tt_cross_page_table = [] for i in range(self.data_parallel): user_page_table = page_table[i] if page_table is not None else None user_cross_page_table = cross_page_table[i] if cross_page_table is not None else None ( tt_h_i, tt_xattn_mask_i, tt_full_text_mask_expand_1NSH_i, tt_full_text_mask_expand_11SD_i, tt_position_id_i, tt_rot_mats_i, tt_page_table_i, tt_cross_page_table_i, ) = self.model[i].prepare_inputs_decode( tokens[i], prefill_cross_attention_masks[i], prefill_full_text_row_masked_out_mask[i], decode_cross_attention_masks[i], decode_full_text_row_masked_out_mask[i], position_id=position_id[i], page_table=user_page_table, cross_page_table=user_cross_page_table, ) tt_h.append(tt_h_i) tt_xattn_mask.append(tt_xattn_mask_i) tt_full_text_mask_expand_1NSH.append(tt_full_text_mask_expand_1NSH_i) tt_full_text_mask_expand_11SD.append(tt_full_text_mask_expand_11SD_i) tt_position_id.append(tt_position_id_i) tt_rot_mats.append(tt_rot_mats_i) tt_page_table.append(tt_page_table_i) tt_cross_page_table.append(tt_cross_page_table_i) # Compile run for i in range(self.data_parallel): user_kv_cache = kv_cache[i] if kv_cache is not None else None xattn_cache = xattn_caches[i] if xattn_caches is not None else None # tt_logits_rm and tt_log_probs_rm unused later, no need to make a list tt_logits_rm, tt_log_probs_rm = self.model[i].ttnn_decode_forward( tt_h[i], tt_xattn_mask[i], tt_full_text_mask_expand_1NSH[i], tt_full_text_mask_expand_11SD[i], xattn_cache, tt_position_id[i], tt_rot_mats[i], page_table=tt_page_table[i], kv_cache=user_kv_cache, cross_page_table=tt_cross_page_table[i], ) logger.info("Done Compiling Model") # Get inputs ready for trace run tt_h = [] tt_xattn_mask = [] tt_full_text_mask_expand_1NSH = [] tt_full_text_mask_expand_11SD = [] tt_position_id = [] tt_rope_id = [] tt_page_table = [] tt_cross_page_table = [] for i in range(self.data_parallel): user_page_table = page_table[i] if page_table is not None else None user_cross_page_table = cross_page_table[i] if cross_page_table is not None else None ( tt_h_i, tt_xattn_mask_i, tt_full_text_mask_expand_1NSH_i, tt_full_text_mask_expand_11SD_i, tt_position_id_i, tt_rope_id_i, tt_page_table_i, tt_cross_page_table_i, ) = self.model[i].prepare_decode_inputs_host( tokens[i], prefill_cross_attention_masks[i], prefill_full_text_row_masked_out_mask[i], decode_cross_attention_masks[i], decode_full_text_row_masked_out_mask[i], position_id[i], page_table=user_page_table, cross_page_table=user_cross_page_table, ) ( tt_h_i, tt_xattn_mask_i, tt_full_text_mask_expand_1NSH_i, tt_full_text_mask_expand_11SD_i, tt_position_id_i, tt_rope_id_i, tt_page_table_i, tt_cross_page_table_i, ) = copy_host_to_device( ( tt_h_i, tt_xattn_mask_i, tt_full_text_mask_expand_1NSH_i, tt_full_text_mask_expand_11SD_i, tt_position_id_i, tt_rope_id_i, tt_page_table_i, tt_cross_page_table_i, ), mesh_device=self.model_args[i].mesh_device, ) tt_h.append(tt_h_i) tt_xattn_mask.append(tt_xattn_mask_i) tt_full_text_mask_expand_1NSH.append(tt_full_text_mask_expand_1NSH_i) tt_full_text_mask_expand_11SD.append(tt_full_text_mask_expand_11SD_i) tt_position_id.append(tt_position_id_i) tt_rope_id.append(tt_rope_id_i) tt_page_table.append(tt_page_table_i) tt_cross_page_table.append(tt_cross_page_table_i) tt_h_trace_input = tt_h tt_logits_rm = [] tt_log_probs_rm = [] trace_ids = {} # Do on-device transformations of inputs before forward for i in range(self.data_parallel): trace_id = ttnn.begin_trace_capture(self.model_args[i].mesh_device, cq_id=0) trace_ids[i] = trace_id B = tokens[i].shape[0] user_kv_cache = kv_cache[i] if kv_cache is not None else None xattn_cache = xattn_caches[i] if xattn_caches is not None else None ( tt_h_transform, tt_rot_mats, tt_xattn_mask_transform, tt_full_text_mask_expand_1NSH_transform, tt_full_text_mask_expand_11SD_transform, ) = self.model[i].transform_decode_inputs_device( tt_h[i], tt_rope_id[i], tt_xattn_mask[i], tt_full_text_mask_expand_1NSH[i], tt_full_text_mask_expand_11SD[i], B=B, ) tt_logits_rm_i, tt_log_probs_rm_i = self.model[i].ttnn_decode_forward( tt_h_transform, tt_xattn_mask_transform, tt_full_text_mask_expand_1NSH_transform, tt_full_text_mask_expand_11SD_transform, xattn_cache, tt_position_id[i], tt_rot_mats, page_table=tt_page_table[i], kv_cache=user_kv_cache, cross_page_table=tt_cross_page_table[i], ) tt_logits_rm.append(tt_logits_rm_i) tt_log_probs_rm.append(tt_log_probs_rm_i) ttnn.end_trace_capture(self.model_args[i].mesh_device, trace_id, cq_id=0) logger.info("Done Capturing Decode Trace") return ( trace_ids, tt_logits_rm, tt_log_probs_rm, tt_h, tt_xattn_mask, tt_full_text_mask_expand_1NSH, tt_full_text_mask_expand_11SD, tt_position_id, tt_rope_id, tt_page_table, tt_cross_page_table, ) def _decode_forward_trace( self, position_id, tokens, prefill_cross_attention_masks, prefill_full_text_row_masked_out_mask, decode_cross_attention_masks, decode_full_text_row_masked_out_mask, page_table, cross_page_table, trace_ids, trace_logits_rm, trace_h, trace_xattn_mask, trace_full_text_mask_expand_1NSH, trace_full_text_mask_expand_11SD, trace_position_id, trace_rope_id, trace_page_table, trace_cross_page_table, ): """ Executes the trace for the decode_forward method but does not read back outputs. """ for i in range(self.data_parallel): user_page_table = page_table[i] if page_table is not None else None user_cross_page_table = cross_page_table[i] if cross_page_table is not None else None ( tt_h, tt_xattn_mask, tt_full_text_mask_expand_1NSH, tt_full_text_mask_expand_11SD, tt_position_id, tt_rope_id, tt_page_table, tt_cross_page_table, ) = self.model[i].prepare_decode_inputs_host( tokens[i], prefill_cross_attention_masks[i], prefill_full_text_row_masked_out_mask[i], decode_cross_attention_masks[i], decode_full_text_row_masked_out_mask[i], position_id=position_id[i], page_table=user_page_table, cross_page_table=user_cross_page_table, ) copy_host_to_device( host_tensors=( tt_h, tt_xattn_mask, tt_full_text_mask_expand_1NSH, tt_full_text_mask_expand_11SD, tt_position_id, tt_rope_id, tt_page_table, tt_cross_page_table, ), device_tensors=( trace_h[i], trace_xattn_mask[i], trace_full_text_mask_expand_1NSH[i], trace_full_text_mask_expand_11SD[i], trace_position_id[i], trace_rope_id[i], trace_page_table[i], trace_cross_page_table[i], ), ) for i, trace_id in trace_ids.items(): ttnn.execute_trace(self.mesh_device, trace_id, cq_id=0, blocking=False) return trace_logits_rm def _easy_trace( self, position_id, tokens, prefill_cross_attention_masks, prefill_full_text_row_masked_out_mask, decode_cross_attention_masks, decode_full_text_row_masked_out_mask, xattn_caches=None, page_table=None, kv_cache=None, cross_page_table=None, ): """ Tracing is easy! Just call this method and we'll handle tracing for you. """ if not hasattr(self, "trace_ids"): ( trace_ids, tt_logits_rm, tt_log_probs_rm, tt_h, tt_xattn_mask, tt_full_text_mask_expand_1NSH, tt_full_text_mask_expand_11SD, tt_position_id, tt_rope_id, tt_page_table, tt_cross_page_table, ) = self._capture_trace( position_id, tokens, prefill_cross_attention_masks, prefill_full_text_row_masked_out_mask, decode_cross_attention_masks, decode_full_text_row_masked_out_mask, xattn_caches, page_table=page_table, kv_cache=kv_cache, cross_page_table=cross_page_table, ) self.trace_ids = trace_ids self.trace_inputs = { "tt_h": tt_h, "tt_xattn_mask": tt_xattn_mask, "tt_full_text_mask_expand_1NSH": tt_full_text_mask_expand_1NSH, "tt_full_text_mask_expand_11SD": tt_full_text_mask_expand_11SD, "tt_position_id": tt_position_id, "tt_rope_id": tt_rope_id, "tt_page_table": tt_page_table, "tt_cross_page_table": tt_cross_page_table, } self.trace_outputs = { "tt_logits_rm": tt_logits_rm, } trace_logits_rm = self._decode_forward_trace( position_id, tokens, prefill_cross_attention_masks, prefill_full_text_row_masked_out_mask, decode_cross_attention_masks, decode_full_text_row_masked_out_mask, page_table, cross_page_table, self.trace_ids, self.trace_outputs["tt_logits_rm"], self.trace_inputs["tt_h"], self.trace_inputs["tt_xattn_mask"], self.trace_inputs["tt_full_text_mask_expand_1NSH"], self.trace_inputs["tt_full_text_mask_expand_11SD"], self.trace_inputs["tt_position_id"], self.trace_inputs["tt_rope_id"], self.trace_inputs["tt_page_table"], self.trace_inputs["tt_cross_page_table"], ) return trace_logits_rm def generate( self, vision_images, vision_mask, prompt_tokens, max_gen_len: int, temperature: float = 0.6, top_p: float = 0.9, ): # Do initial prefill prefill_len = len(prompt_tokens) total_len = prefill_len + max_gen_len # Prepares mask for full length of output prompt_tokens_tensor = torch.tensor(prompt_tokens, dtype=torch.long).reshape(1, -1) # B, S # Suboptimal to allocate caches every time model_id = 0 xattn_caches = self.model[model_id].setup_cache(self.model_args[model_id].max_batch_size) ( xattn_caches, prefill_cross_attention_masks, prefill_full_text_row_masked_out_mask, decode_cross_attention_masks, decode_full_text_row_masked_out_mask, logits, ) = self._prefill_forward_single_user( vision_images, vision_mask, prompt_tokens_tensor, xattn_caches, user_id=0, total_len=total_len, prefill_len=prefill_len, model_id=model_id, ) last_token_idx = prefill_len - 1 logits = self.model[model_id].process_output_prefill(logits.cpu(), 1, last_token_idx=(last_token_idx % 32)) logits = logits.view(1, 1, self.model_args[model_id].vocab_size) prefill_output_xattn_masks = [[] for _ in range(self.data_parallel)] prefill_output_full_text_row_masked_out_masks = [[] for _ in range(self.data_parallel)] decode_output_xattn_masks = [[] for _ in range(self.data_parallel)] decode_output_full_text_row_masked_out_masks = [[] for _ in range(self.data_parallel)] prefill_output_xattn_masks[model_id].append(prefill_cross_attention_masks) prefill_output_full_text_row_masked_out_masks[model_id].append(prefill_full_text_row_masked_out_mask) decode_output_xattn_masks[model_id].append(decode_cross_attention_masks) decode_output_full_text_row_masked_out_masks[model_id].append(decode_full_text_row_masked_out_mask) def sample(logits): if temperature > 0: probs = torch.softmax(logits[:, -1] / temperature, dim=-1) next_token = sample_top_p(probs, top_p) else: next_token = torch.argmax(logits[:, -1], dim=-1) next_token = next_token.reshape(-1) decoder = self.tokenizer or self.processor return next_token, decoder.decode(next_token.tolist()) next_token, text = sample(logits) yield TokenResult( token=next_token[0].item(), text=text, ) for gen_idx in range(max_gen_len - 1): position_id = torch.tensor([prefill_len + gen_idx]) next_token_tensor = next_token.reshape(1, 1) # B, S logits = self.decode_forward_llama_vision( position_id, next_token_tensor, prefill_output_xattn_masks, prefill_output_full_text_row_masked_out_masks, decode_output_xattn_masks, decode_output_full_text_row_masked_out_masks, [xattn_caches], enable_trace=False, ) if isinstance(logits, tuple): logits = logits[0] next_token, text = sample(logits) yield TokenResult( token=next_token[0].item(), text=text, ) def chat_completion( self, messages, temperature=0.6, top_p: float = 0.9, max_gen_len=None, ): model_id = 0 if max_gen_len is None or max_gen_len == 0 or max_gen_len >= self.model[model_id].configuration.max_seq_len: max_gen_len = self.model[model_id].configuration.max_seq_len - 1 encoder = self.processor or self.tokenizer model_input = encoder.apply_chat_template(messages, add_generation_prompt=True, tokenize=True, return_dict=True) vision_images = extract_images_from_messages(messages) or None vision_mask = None if vision_images is not None: vision_mask = create_vision_mask(model_input["input_ids"][0], encoder.image_token_id) or None tokens = [] stop_reason = None for result in self.generate( vision_images=vision_images, vision_mask=vision_mask, prompt_tokens=model_input["input_ids"][0], max_gen_len=max_gen_len, temperature=temperature, top_p=top_p, ): tokens.append(result.token) if result.text == "<|eot_id|>": stop_reason = StopReason.end_of_turn elif result.text == "<|eom_id|>": stop_reason = StopReason.end_of_message if stop_reason is None: stop_reason = StopReason.out_of_tokens decoder = self.tokenizer or self.processor message = decoder.decode(tokens, skip_special_tokens=True) return CompletionMessage(message) def text_completion( self, content, temperature: float = 0.6, top_p: float = 0.9, max_gen_len=None, ): """Supports only vision models at the moment""" model_id = 0 if max_gen_len is None or max_gen_len == 0 or max_gen_len >= self.model[model_id].configuration.max_seq_len: max_gen_len = self.model[model_id].configuration.max_seq_len - 1 vision_images = [] image_token = getattr(self.processor, "image_token", None) or getattr(self.tokenizer, "image_token", None) text = encode_content(content, vision_images, image_token) vision_images = vision_images or None model_input = self.processor(text=text, images=vision_images, add_special_tokens=False) vision_mask = None if vision_images is not None: vision_mask = create_vision_mask(model_input["input_ids"][0], self.processor.image_token_id) or None tokens = [] for result in self.generate( vision_images=vision_images, vision_mask=vision_mask, prompt_tokens=model_input["input_ids"], max_gen_len=max_gen_len, temperature=temperature, top_p=top_p, ): tokens.append(result.token) decoder = self.tokenizer or self.processor generation = decoder.decode(tokens, skip_special_tokens=True) return generation def _get_prefill_user_page_table( self, page_table, kv_cache, prefill_len, trace_enabled=False, prefill_seq_len=None, use_batched_prefill=False, user_id=None, padded_batch_size=None, use_full_prompt_len=False, ): block_size = get_block_size(kv_cache) if use_batched_prefill: batch_dim = padded_batch_size if padded_batch_size is not None else self.model_args[0].max_batch_size num_blocks = num_blocks_in_seq(prefill_seq_len, block_size) page_table = page_table[:, :num_blocks] if trace_enabled: if page_table.shape[1] < num_blocks: padding = torch.ones(page_table.shape[0], num_blocks - page_table.shape[1], dtype=torch.int32) * -1 page_table = torch.cat([page_table, padding], dim=1) padded_page_table = torch.ones(batch_dim, page_table.shape[1], dtype=torch.int32) * -1 assert user_id is not None for i, user in enumerate(user_id): padded_page_table[user, :] = page_table[i, :] return padded_page_table else: # Compatibility with VLLM warmup: prefill kernels run on the padded # prefill length (for example 32-token prompts become 128-token # kernels), so the page table must expose blocks for that padded # length even on the non-traced compile path. if use_full_prompt_len: target_prefill_len = prefill_len else: target_prefill_len = prefill_seq_len if prefill_seq_len is not None else prefill_len num_blocks = num_blocks_in_seq(target_prefill_len, block_size) if page_table.shape[1] < num_blocks: padding = torch.ones(1, num_blocks - page_table.shape[1], dtype=torch.int32) * -1 page_table = torch.cat([page_table, padding], dim=1) return page_table[:, :num_blocks] def release_persistent_capture(self) -> None: """Release model-lifetime traces once, while the mesh is still open. The plugin calls this on the model before it closes the mesh. An override must chain to ``super()``: the destructor below runs only this base method, because a destructor may fire after the mesh closed. """ if getattr(self, "_generator_capture_released", False): return self._generator_capture_released = True try: # Release prefill traces if hasattr(self, "trace_id_prefill"): for trace_key, trace_id in self.trace_id_prefill.items(): if trace_id is not None: # Extract model_id from trace_key (format: "{prefill_seq_len}_{model_id}" or "{prefill_seq_len}_{model_id}_{batch_size}") parts = trace_key.split("_") model_id = int(parts[1]) if len(parts) >= 2 else 0 try: ttnn.release_trace(self.model_args[model_id].mesh_device, trace_id) except Exception: pass # Ignore errors during cleanup # Release prefill sampling traces if hasattr(self, "trace_id_prefill_sampling"): for trace_key, trace_id in self.trace_id_prefill_sampling.items(): if trace_id is not None: parts = trace_key.split("_") if parts and parts[0] == "sampling" and len(parts) >= 3: m_id = int(parts[2]) else: m_id = int(parts[-1]) if len(parts) >= 2 else 0 try: ttnn.release_trace(self.model_args[m_id].mesh_device, trace_id) except Exception: pass # Release all sampling traces, including every decode-bucket namespace. for model in getattr(self, "model", []): sampling_module = getattr(model, "sampling", None) if sampling_module is not None and hasattr(sampling_module, "reset_trace"): try: sampling_module.reset_trace() except Exception: pass # Release decode traces decode_trace_stores = [] for bucket_store in getattr(self, "_bucket_trace_store", {}).values(): if bucket_store is not None: decode_trace_stores.append(bucket_store[0]) if hasattr(self, "trace_ids_decode"): decode_trace_stores.append(self.trace_ids_decode) released_decode_traces = set() for trace_store in decode_trace_stores: for sampling_key, trace_ids_dict in trace_store.items(): if trace_ids_dict is not None: for model_id, trace_id in trace_ids_dict.items(): trace_key = (model_id, trace_id) if trace_id is not None and trace_key not in released_decode_traces: released_decode_traces.add(trace_key) try: ttnn.release_trace(self.model_args[model_id].mesh_device, trace_id) except Exception: pass # Ignore errors during cleanup # Release vision traces if present if hasattr(self, "trace_ids"): for model_id, trace_id in self.trace_ids.items(): if trace_id is not None: try: ttnn.release_trace(self.mesh_device, trace_id) except Exception: pass # Ignore errors during cleanup except Exception: pass # Ignore any errors during trace cleanup def __del__(self): # Base traces only. Subclass releases need an open mesh and run through # the plugin's release_persistent_capture call at shutdown. Generator.release_persistent_capture(self) # Workaround for issue #19052 if self.data_parallel > 1: for m in self.model: ttnn.close_mesh_device(m.mesh_device) if hasattr(super(Generator, self), "__del__"): super().__del__() def _mesh_shape_tuple(mesh_shape): return tuple(int(dim) for dim in mesh_shape) def _galaxy_data_parallel_submesh_shape(devices_per_group): # Galaxy DP groups should follow the 4x8 row-oriented view recommended by # the runtime, so DP=4 maps to four routeable 1x8 T3K-like submeshes. if devices_per_group >= 8 and devices_per_group % 8 == 0: return ttnn.MeshShape(devices_per_group // 8, 8) # Smaller DP groups still use contiguous 1D row submeshes; callers select # linear CCL when these groups are too small for ring topology. return ttnn.MeshShape(1, devices_per_group) def create_submeshes(mesh_device, data_parallel): mesh_device_type = getattr(ttnn, "MeshDevice", None) if mesh_device_type is None: mesh_device_type = getattr(getattr(ttnn, "device", None), "Device", None) if mesh_device_type is None or not isinstance(mesh_device, mesh_device_type) or data_parallel == 1: return [mesh_device] num_rows, num_cols = _mesh_shape_tuple(mesh_device.shape) num_devices = num_rows * num_cols assert num_devices % data_parallel == 0, f"Unsupported device split: {num_devices} devices, {data_parallel} groups" if num_devices == 32: if (num_rows, num_cols) != (4, 8): logger.info(f"Reshaping 32-device mesh from {(num_rows, num_cols)} to (4, 8) for DP submeshes") mesh_device.reshape(ttnn.MeshShape(4, 8)) return mesh_device.create_submeshes(_galaxy_data_parallel_submesh_shape(num_devices // data_parallel)) return mesh_device.create_submeshes(ttnn.MeshShape(1, num_devices // data_parallel))