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