File size: 14,506 Bytes
b025706
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
# Sampling Module Overview

The `models.common.sampling` package bundles everything needed to run on-device
sampling (top-k / top-p / temperature/ seed) plus presence/frequency/repetition
penalties with optional trace capture.

## Key Components
- `SamplingGenerator`: high-level class that owns both `TTSampling` and
  `TTPenalties`, exposes helper methods to reset sampling parameters, penalties,
  prompt/output state, and to run sampling with or without trace capture.
- `format_sampling_params`: utility that pads/clamps sampling parameters to the
  hardware-friendly layout expected by `TTSampling`.
- `LogProbsCalculator`: computes per-token log-probabilities across a sharded
  vocabulary using numerically stable log-softmax (global max / sum-exp
  reduction across devices).

## Quick Start
```python
from models.common.sampling import SamplingGenerator, format_sampling_params

sampling = SamplingGenerator(args=args, mesh_device=mesh_device, tt_ccl=tt_ccl)

params = format_sampling_params(user_params, max_batch_size=32)
sampling.reset_sampling_params(params)

sampling.reset_seed(seed)

sampling.reset_prompt_tokens(prompt_tokens)   # torch tensor shaped [B, S]
sampling.reset_output_state(output_tokens)

tt_tokens = sampling.sample(
    tt_logits,
    tt_out_tok=tt_out_buffer,
)
```

`SamplingGenerator.sample()` accepts `enable_trace=True` to record/replay
sampling traces.

## File Map

| File | Purpose |
|---|---|
| `generator.py` | `SamplingGenerator` orchestrator; `SamplingParams`; `format_sampling_params`; `broadcast_sampling_params`; `chunk_sampling_params`; `SeedManager` |
| `tt_sampling.py` | `TTSampling` — on-device top-k/top-p/temp with multi-device all-gather |
| `tt_penalties.py` | `TTPenalties` — presence / frequency / repetition penalties |
| `tt_log_probs.py` | `LogProbsCalculator` — log-softmax across sharded vocabulary |
| `_utils.py` | Shared helpers: `clamp`, `is_default_value`, `filter_none`, `split_list` |

## Required `args` Attributes

```python
vocab_size: int           # actual vocabulary size (unpadded)
cluster_shape: tuple      # (rows, cols) of the device mesh, e.g. (4, 8)
```

Optional (with defaults):

```python
padded_vocab_size: int    # tile-aligned total vocab; defaults to vocab_size
max_batch_size: int       # per sampling row; default 32
max_top_k: int            # default 32
sampling_dp: int          # >1 for multi-row DP; default 1
sub_core_grids            # CoreRangeSet or None
model_config: dict        # keys: GALAXY_NUM_LINKS, DECODE_SAMPLING_INPUT_MEMCFG, SAMPLING_AG_CONFIG
```

## `data_parallel` vs `sampling_dp`

These are different concepts and should not be mixed:

- **`data_parallel`** lives above this package. It means multiple TT model
  instances / submeshes process different requests in parallel.
- **`sampling_dp`** lives inside this package. It means one TT model instance
  has multiple independent sampling groups, usually one per mesh row.

For `sampling_dp > 1`:
- logits are still computed per sampling group
- but sampling params, seeds, and penalty state are flattened to
  `max_batch_size * sampling_dp`
- those flattened host tensors are then row-sharded onto the device

Decode already follows this contract by using `chunk_sampling_params(...)`
plus `apply_decode_state(...)`.

## Param Distribution API

**`SamplingParams`**: Canonical dataclass for sampling parameters (temp, top_k, top_p, penalties, seed, log_probs). Import from `models.common.sampling`. vLLM has its own duck-type-compatible `TTSamplingParams`.

**`broadcast_sampling_params(params, idx, slot_len=32)`**: Expand a single user's params to fill `slot_len` slots. Used during prefill.

**`chunk_sampling_params(params, sampling_dp)`**: Split a SamplingParams into `sampling_dp` pieces. List fields split evenly; scalars replicated. Works with duck-typed objects (vLLM).

**`SamplingGenerator.apply_prefill_state(...)`**: Reset params, seeds, prompt tokens, and output state for a prefill request.

**`SamplingGenerator.apply_decode_state(chunks, ...)`**: Execute the sampling
half of vLLM's explicit decode update contract. `reload_sampling_params=True`
formats/merges and uploads parameters; `reset_sampling_state=True` rebuilds
prompt/output penalty state. The flags are independent. The method does NOT
advance seeds — callers apply slot remaps first, reset/align seeds when state is
reset, and call `seed_manager.get_new_values()` exactly once per sampled token.
Both command flags are required at every decode call.

This contract includes the unconditional first-decode reseed for `seed=None`
also addressed by
[tt-metal#51556](https://github.com/tenstorrent/tt-metal/pull/51556).
It does not depend on that PR.
`reset_sampling_state=True` calls `reset_seed_from_slots(...)` rather than the
conditional helper, ensuring decode-only sampling actually initializes and
uploads fresh device seeds when both the requested and cached seed are `None`.

## vLLM Decode Update Contract

Refactored vLLM model adapters advertise
`decode_input_update_contract = 1`. The vLLM TT plugin sends these adapters
four boolean commands on every decode:

- `reload_inputs`: copy every forward trace input.
- `reload_page_table`: copy only page-table inputs while preserving
  device-produced token/position state.
- `reload_sampling_params`: upload sampling configuration.
- `reset_sampling_state`: rebuild mutable penalty/RNG state for the layout.

When a version-1 layout transition produces a non-identity `slot_remap`, the
plugin sends it in either sampling mode, including host-sampling steps.
`slot_remap[i] = j` means every persistent state owned by new slot `i` must take
the continuing request state from old slot `j` before the forward reads it.
This is broader than sampler state: recurrent or convolution state indexed by
decode slot must be remapped too. Stateless adapters accept and may ignore the
value.

There are two distinct ways to implement this incompletely:

1. Sending `slot_remap` only on device-sampling decodes leaves model-owned
   recurrent/conv/RoPE state in the old slot when host sampling changes the
   layout.
2. Sending it on every decode but neither remapping nor invalidating dormant
   sampler state leaves seed/RNG/parameter/penalty state in the old slot during
   host sampling. A later switch back to device sampling can then resume the
   wrong request's state.

Every slot-owning subsystem must therefore consume each supplied remap exactly
once on the accepted version-1 decode that carries it. State read by the
forward is remapped before that read. A dormant sampler consumes it after
successful decode/readback, which preserves retry safety because slot remaps
are non-idempotent. Its slot-addressable host RNG state is remapped immediately.
The sharded device parameter and penalty buffers are instead marked invalid:
their next activation must command both `reload_sampling_params` and
`reset_sampling_state`, rebuilding them from authoritative host state before
they are read. A slot-scoped prefill may rebuild only its newly admitted rows,
but it does not clear the whole-device invalidation; a later decode still needs
a whole-device penalty reset together with the parameter upload. The plugin
issues both commands on every layout or sampling-mode transition, and the
sampling generator rejects a later direct caller that omits that full rebuild.
This authoritative rebuild replaces a physical remap for those buffers;
inactivity alone does not silently accept stale state.
Version-0 adapters retain their historical remap behavior unchanged.

For a merged lane-DP call the remap uses global lane-major slots, while each
model replica's seed manager owns a rank-local padded array. The shared
generator splits by the actual scheduler lane stride, subtracts the lane base,
and pads the untouched tail with a local identity mapping. Splitting by the
sampler's padded width instead is incorrect whenever scheduler capacity is
smaller than that width, and passing absolute lane-1+ indices to a local seed
manager is out of bounds.

State that is not addressable by vLLM slot cannot be remapped. Unseeded
on-device RNG is the known exception: its state lives in per-core hardware PRNG
registers with no slot-to-slot move primitive. A commanded
`reset_sampling_state` therefore reinitializes it instead. Explicitly seeded,
slot-addressable counters still follow the request through `slot_remap`. Any
additional exception must be documented beside the code that skips it and must
be physically unmovable, not merely inconvenient to move.

Generators execute these commands without adding page-table comparisons,
sampling-mode checks, or model-specific forced reloads. The corresponding
vLLM plugin falls back to the legacy `reset_batch` interface for adapters that
do not advertise the contract, preserving their existing reload and overlap
behavior. vLLM warns that correctness is not guaranteed on that compatibility
path. This lets vLLM land first and adapters opt in as they are refactored. The
marker is negotiation metadata on vLLM-facing adapters only; all refactored
generator APIs require direct callers, including demos and warmup code, to
provide all four commands. No model-side fallback heuristics are restored. Any
demo-side decision to retain traced inputs is made at the call site. Warmup
reloads each device-sampling parameter configuration but does not request a
sampling-state reset because it has no request-owned prompt/output history.
The SGLang bridge explicitly uses host sampling and reloads its authoritative
token, position, and page table every step.

Trace selection does not authorize a reload. Direct callers switching sampling
mode or a Gemma4 decode batch bucket must explicitly reload inputs for the
selected trace. The plugin commands full reloads on sampling-mode and request
layout transitions; its decode bucket is derived from that request layout.
Galaxy also carries `reload_inputs` from forward to separated sampling so it
can realign seeded RNG counters on any authoritative full reload, even without
a penalty-state reset. It never realigns from stale steady-decode positions.

Stable-slot adapters must also preserve unscheduled rows when admitting a
prefill. DeepSeek starts from its cached full-batch sampling parameters and
scatters only the incoming requests into their assigned slots; filling every
other row from an incoming request silently changes live decodes. Its partial
prefill also resets prompt/output penalty history and seed state only for the
incoming slots. Those slots are not the full live set, so continuing seeds are
not retired. Reset-only sampling commands still need slot-indexed seeds:
Galaxy formats/pads the request-ordered parameters whenever either
`reload_sampling_params` or `reset_sampling_state` needs their seed values.
When Galaxy condenses slots during host sampling, it remaps its cached
per-slot parameter vectors with the same snapshot as its seed state so a later
partial prefill starts from the current layout.

`model_capabilities["supports_async_decode"]` is separate from contract
versioning. It certifies that a vLLM wrapper supports split async readback and
device-resident sampled-token feedback; wrappers without it receive explicit
full-input reload commands instead.

### Requirements for `supports_async_decode=True`

A model wrapper may opt in only if all of the following hold:

- `decode_forward(..., read_from_device=False)` and
  `read_decode_output(..., async_read=True)` split submission from observational
  readback.
- Device sampling writes the selected token into the persistent token input
  consumed by the next decode.
- Decode forward advances the persistent position exactly once; sampling and
  readback never advance it.
- Page tables can be refreshed without copying or rebinding token, position, or
  RoPE trace inputs.
- All four reload commands are honored independently, without model-local
  heuristics escalating page-table-only refresh into a full reload.
- Slot remap applies before the forward to every persistent slot-indexed state
  that the forward reads, in both sampling modes. After a successful
  host-sampling decode, dormant RNG state is remapped and device sampler state
  is invalidated exactly once. Before device sampling reads that state, an
  authoritative parameter upload and penalty reset rebuild it. Seed advancement
  follows, once per sampled token.
- Persistent buffers remain valid through deferred readback and until an
  explicit reload replaces them.

If any item is unsupported, leave the capability absent or `False`. vLLM will
disable async scheduling and issue a full input reload for every decode. The
authoritative mode definitions, transition matrix, and correctness invariant
live in the paired standalone plugin document
[`docs/DECODE_RELOAD_CONTRACT.md`](https://github.com/tenstorrent/vllm-tt-plugin/blob/main/docs/DECODE_RELOAD_CONTRACT.md).

## Pitfalls

**`padded_vocab_size` vs `vocab_size`**: TTSampling device offsets for global token IDs must use the padded vocab size to match how the LM head shards logits across devices. Using unpadded `vocab_size` for offsets shifts token IDs from devices 1+ and produces garbled output.

**Padded vocab logits**: If the LM head pads output weights beyond the real tokenizer vocabulary, the sampler must mask those padded token IDs before force-argmax or local top-k. Zero-padded LM-head weights are useful for legal sharded matmul shapes, but they are not a sampling mask.

**`sampling_dp`**: When >1, k/p/temp tensors must have length `max_batch_size * sampling_dp` and are row-sharded via `ShardTensor2dMesh(dims=(0, None))`. Use `chunk_sampling_params` + `apply_decode_state` to distribute params across mesh rows.

**Batched prefill + on-device sampling**: This path is only valid when the
runtime prefill compute layout matches the sampling-group layout. If a model
uses `sampling_dp > 1` but does not expose a row-sharded batched-prefill input
contract, batched prefill must fall back to sequential prefill for correctness.

**Trace invalidation**: Changing `force_argmax_sampling` state invalidates captured traces. Force-argmax is triggered when callers pass k=1, p=1.0, temp=1.0 (note: p=1.0 means "no top-p filtering", distinct from the internal initialization default of p=0). `SamplingGenerator.reset_sampling_params` handles this.

## Future Work

- Consolidate DeepSeek's minimal `SamplingParams` (in `models/demos/deepseek_v3/tt/generator.py`) to use the common one