Download code/models/tt_transformers/tt/generator.py from tt-hous/clef: direct link, hf CLI and curl.
- Browser
- Download file 191 kB
-
https://huggingface.co/tt-hous/clef/resolve/main/code/models/tt_transformers/tt/generator.py
- Command line
-
hf download hf://tt-hous/clef/code/models/tt_transformers/tt/generator.py
-
curl -L -o generator.py https://huggingface.co/tt-hous/clef/resolve/main/code/models/tt_transformers/tt/generator.py
191 kB
| # 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 | |
| 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 | |
| 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)) | |