File size: 6,465 Bytes
2415c4c | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 | # 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")
|