Download code/models/common/sampling/README.md from tt-hous/clef: direct link, hf CLI and curl.
- Browser
- Download file 14.5 kB
-
https://huggingface.co/tt-hous/clef/resolve/main/code/models/common/sampling/README.md
- Command line
-
hf download hf://tt-hous/clef/code/models/common/sampling/README.md
-
curl -L -o README.md https://huggingface.co/tt-hous/clef/resolve/main/code/models/common/sampling/README.md
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 bothTTSamplingandTTPenalties, 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 byTTSampling.LogProbsCalculator: computes per-token log-probabilities across a sharded vocabulary using numerically stable log-softmax (global max / sum-exp reduction across devices).
Quick Start
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
vocab_size: int # actual vocabulary size (unpadded)
cluster_shape: tuple # (rows, cols) of the device mesh, e.g. (4, 8)
Optional (with defaults):
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_parallellives above this package. It means multiple TT model instances / submeshes process different requests in parallel.sampling_dplives 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.
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:
- Sending
slot_remaponly on device-sampling decodes leaves model-owned recurrent/conv/RoPE state in the old slot when host sampling changes the layout. - 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)andread_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.
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(inmodels/demos/deepseek_v3/tt/generator.py) to use the common one