# SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc. # SPDX-License-Identifier: Apache-2.0 import os from itertools import product import torch from loguru import logger from models.common.sampling.sampling_params import SamplingParams class WarmupForwardMixin: """ This class is used by vLLM. Mixin class that provides decode warmup functionality for generator classes. This class should be inherited by any generator class that needs to warm up the decode forward pass. It requires the following to be defined in the inheriting class: - self.decode_forward(): method to perform decode forward pass """ def _create_sampling_params(self, can_sample_on_device, batch_size, greedy_only: bool = False): """ greedy_only: when True, warmup only covers greedy decoding on device (temperature=0.0, top_k=1, top_p=1.0). When False (the default), warmup also exercises non-greedy variants — temperature/top_k/top_p, presence/frequency/repetition penalties, and log_probs. """ if not can_sample_on_device: return [None] sampling_configs = [] if not greedy_only: # Full warmup pre-captures every penalties × log_probs permutation so # no on-device-sampling request ever pays a one-time trace-capture # cost on first use. Each permutation is a *separate resident trace*, # though, and for large MoE models the combined trace region can run # into the gigabytes (e.g. Gemma4-26B-A4B). ``TT_LEAN_DECODE_WARMUP`` # restricts the sweep to the plain (no-penalty, no-logprob) sampling # config — what a throughput benchmark with default sampling actually # exercises — trading a one-time runtime capture for the rarer # penalty/logprob request shapes in exchange for a much smaller trace # region. Greedy and ``None`` are still captured below. if os.environ.get("TT_LEAN_DECODE_WARMUP"): penalty_logprob_combos = [(False, False)] else: penalty_logprob_combos = list(product([True, False], repeat=2)) for penalties, log_probs in penalty_logprob_combos: presence_penalty, frequency_penalty, repetition_penalty = None, None, None if penalties: presence_penalty = [1.2] * batch_size frequency_penalty = [1.2] * batch_size repetition_penalty = [1.5] * batch_size enable_log_probs = [log_probs] * batch_size temperature = [1.0] * batch_size top_k = [10] * batch_size top_p = [0.9] * batch_size sampling_configs.append( SamplingParams( temperature=temperature, top_k=top_k, top_p=top_p, presence_penalty=presence_penalty, frequency_penalty=frequency_penalty, repetition_penalty=repetition_penalty, enable_log_probs=enable_log_probs, ) ) sampling_configs.append( SamplingParams( temperature=[0.0] * batch_size, top_k=[1] * batch_size, top_p=[1.0] * batch_size, ) ) sampling_configs.append(None) return sampling_configs def _create_decode_warmup_inputs(self, max_batch_size, num_blocks): tokens = torch.zeros(max_batch_size, 1, dtype=torch.int32) start_pos = torch.zeros(max_batch_size, dtype=torch.int32) page_table = torch.zeros(max_batch_size, num_blocks, dtype=torch.int32) return tokens, start_pos, page_table def warmup_model_decode( self, kv_cache, enable_trace, max_batch_size, num_blocks, can_sample_on_device, read_from_device=True, greedy_only: bool = False, skip_trace_precompile: bool = False, ): """ This function is called by vLLM """ sampling_params = self._create_sampling_params(can_sample_on_device, max_batch_size, greedy_only=greedy_only) tokens, start_pos, page_table = self._create_decode_warmup_inputs(max_batch_size, num_blocks) logger.info("Starting decode warmup") logger.info(f"Tokens shape: {tokens.shape}") logger.info(f"Start pos shape: {start_pos.shape}") logger.info(f"Page table shape: {page_table.shape}") # Record every trace variant this sweep needs before any of them is live (see # Generator.precapture_decode_trace_variants); the per-param passes below then only replay. precapture = getattr(self, "precapture_decode_trace_variants", None) if enable_trace and not skip_trace_precompile and precapture is not None: if precapture(sampling_params, tokens, start_pos, page_table, kv_cache): logger.info("Pre-captured decode trace variants before the sampling sweep") for param in sampling_params: logger.info(f"Warming up decode for sampling params: {param}") decode_kwargs = dict( tokens=tokens, start_pos=start_pos, page_table=page_table, kv_cache=kv_cache, enable_trace=enable_trace, read_from_device=read_from_device, sampling_params=param, reload_inputs=True, reload_page_table=False, reload_sampling_params=param is not None, # Warmup has no request-owned prompt/output history. The old # reset_batch=False path compiled each sampling configuration # without rebuilding penalty state; preserve that behavior. reset_sampling_state=False, ) if skip_trace_precompile: decode_kwargs["skip_trace_precompile"] = True if not enable_trace and hasattr(self, "_prepare_decode_trace_variant"): # Run through decode_forward so model-specific page-table routing # is active while staging. decode_kwargs["prepare_trace"] = True self.decode_forward(**decode_kwargs) logger.info("Decode warmup completed")