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")