Download code/models/common/tests/test_tt_sampling.py from tt-hous/clef: direct link, hf CLI and curl.
- Browser
- Download file 96.5 kB
-
https://huggingface.co/tt-hous/clef/resolve/main/code/models/common/tests/test_tt_sampling.py
- Command line
-
hf download hf://tt-hous/clef/code/models/common/tests/test_tt_sampling.py
-
curl -L -o test_tt_sampling.py https://huggingface.co/tt-hous/clef/resolve/main/code/models/common/tests/test_tt_sampling.py
96.5 kB
| # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC | |
| # SPDX-License-Identifier: Apache-2.0 | |
| """ | |
| Pure TT-metal sampling tests inspired by vLLM sampling tests. | |
| It should keep the high-level organization from vLLM request-level tests, | |
| but validate behavior directly at token-ID level using synthetic logits and | |
| on-device sampling primitives ('TTSampling' + 'SamplingGenerator'). | |
| """ | |
| from __future__ import annotations | |
| from collections import Counter | |
| from dataclasses import dataclass | |
| import pytest | |
| import torch | |
| import ttnn | |
| from models.common.sampling.generator import SamplingGenerator, SamplingParams, format_sampling_params | |
| from models.common.sampling.tt_sampling import TTSampling | |
| # TEST NOTES: | |
| # - 'mesh_device' and 'device_params' come from repo-root 'conftest.py'. | |
| # - 'device_params' defaults to '{}' (no explicit fabric config). | |
| # - 'models/common/tests/conftest.py' provides additional common-test fixtures | |
| # (e.g. 'ttnn_mesh_device'), but this file uses 'mesh_device'. | |
| # - For all multi-device paths in this file, we pass an explicit ring fabric | |
| # config so tests do not rely on implicit environment defaults. | |
| # - Token ID bands are intentionally disjoint across suites for clearer failure triage. | |
| # --- Constants & helpers --- | |
| BATCH_SIZE = 32 | |
| MAX_TOP_K = 32 | |
| VOCAB_SIZE = 32000 | |
| FAST_NUM_TRIES = 6 | |
| FAST_NUM_STEPS = 5 | |
| MULTI_DEVICE_MESHES = [1, (4, 8)] | |
| # Fabric needs at least two participating chips; since #56669 opening a 1-chip mesh with a fabric | |
| # config is a fatal, so on N150 / P150 the same tests run without one (the ops fall back to the | |
| # single-device path, which is what these tests exercise there anyway). | |
| RING_FABRIC_DEVICE_PARAMS = ( | |
| [{"fabric_config": ttnn.FabricConfig.FABRIC_1D_RING}] if len(ttnn.get_device_ids()) > 1 else [{}] | |
| ) | |
| # Lane positions used to sweep the "odd lane out" tests below. The sampling | |
| # writer kernel reads each user's candidates out of a 32x32 tile where users | |
| # 0-15 live in tile faces 0/1 and users 16-31 in faces 2/3, so cover both sides | |
| # of that boundary as well as the first and last lane. | |
| ODD_LANE_INDICES = [0, 15, 16, BATCH_SIZE - 1] | |
| # The face-boundary lanes exercise per-core writer arithmetic that does not | |
| # depend on the mesh, so multi-device runs only sweep the endpoints. | |
| ODD_LANE_INDICES_MULTI_DEVICE = [0, BATCH_SIZE - 1] | |
| BAND_TOKENS_PER_USER = 8 | |
| # Lightweight args container expected by TTSampling/SamplingGenerator in tests. | |
| class _SamplingArgs: | |
| vocab_size: int | |
| padded_vocab_size: int | |
| max_batch_size: int | |
| max_top_k: int | |
| cluster_shape: tuple[int, int] | |
| sampling_all_gather_axis: int | |
| sampling_dp: int | |
| sub_core_grids: ttnn.CoreRangeSet | None | |
| sub_core_grid_topk: ttnn.CoreRangeSet | None | |
| start_core: ttnn.CoreCoord | |
| model_config: dict | |
| def compute_per_device_vocab(vocab_size: int, num_tp: int) -> int: | |
| per_device = (((vocab_size + num_tp - 1) // num_tp + 31) // 32) * 32 | |
| return 1 << (per_device - 1).bit_length() | |
| def per_lane_params(temperature, top_k, top_p, *, size: int = BATCH_SIZE) -> SamplingParams: | |
| """Full-length per-lane sampling params. | |
| A **scalar** temperature makes format_sampling_params treat only lane 0 as active: every | |
| other lane takes the temperature=0.0 default and comes back GREEDY, not sampled (see | |
| TestFormatSamplingParamsLanes::test_all_scalar_input_is_lane_zero_plus_greedy_padding). | |
| That is correct, documented behaviour -- but it means any assertion about more than one | |
| lane sampling has to pass full-length lists, or 31 of the 32 lanes just return the argmax | |
| and the assertion is either vacuous or wrong. Use this whenever the whole batch is | |
| expected to sample; pass a scalar directly only when lane-0-only is the point. (#38316) | |
| """ | |
| return SamplingParams( | |
| temperature=[temperature] * size, | |
| top_k=[top_k] * size, | |
| top_p=[top_p] * size, | |
| ) | |
| def broadcast(value, *, size: int = BATCH_SIZE): | |
| if isinstance(value, list): | |
| assert len(value) == size, f"Expected list of length {size}, got {len(value)}" | |
| return list(value) | |
| return [value] * size | |
| def safe_sync(mesh_device): | |
| try: | |
| ttnn.synchronize_device(mesh_device) | |
| except Exception: | |
| # Cleanup best-effort only; sync failures should not mask test assertions. | |
| pass | |
| def default_sub_core_grids(mesh_device) -> ttnn.CoreRangeSet: | |
| """The full Tensix compute grid, as the pool TTSampling carves its lane cores from. | |
| This must be set. ttnn.manual_seed maps seed[i] -> the core whose index in the | |
| enumerated core grid equals user_ids[i] (arange(max_batch_size)), and ttnn.sampling | |
| assigns lanes to cores over its own grid. Per-lane seeding is only correct when both | |
| ops see the SAME core set in the SAME order, which TTSampling guarantees by carving | |
| exactly max_batch_size cores out of `sub_core_grids` and passing that one set to both | |
| (see tt_sampling.py `_sampling_sub_core_grids`). Leaving this None makes both ops fall | |
| back to the unrestricted grid independently, so the mapping is no longer pinned and | |
| per-lane seeds land on the wrong lanes. Mirrors the construction in | |
| models/common/tests/test_sampling.py::_single_device_sampling_args. (#38316) | |
| """ | |
| grid = mesh_device.compute_with_storage_grid_size() | |
| return ttnn.CoreRangeSet([ttnn.CoreRange(ttnn.CoreCoord(0, 0), ttnn.CoreCoord(grid.x - 1, grid.y - 1))]) | |
| def make_sampling_args(mesh_device, sampling_dp: int = 1) -> _SamplingArgs: | |
| """Build sampling args for synthetic tests on the current mesh shape.""" | |
| cluster_shape = tuple(mesh_device.shape) | |
| # For 1D meshes, gather across the non-singleton axis. | |
| if cluster_shape[0] == 1 and cluster_shape[1] > 1: | |
| sampling_all_gather_axis = 1 | |
| elif cluster_shape[1] == 1 and cluster_shape[0] > 1: | |
| sampling_all_gather_axis = 0 | |
| else: | |
| sampling_all_gather_axis = 1 if cluster_shape[0] > 1 and cluster_shape[1] > 1 else 0 | |
| num_tp = cluster_shape[sampling_all_gather_axis] if cluster_shape[sampling_all_gather_axis] > 0 else 1 | |
| per_device_vocab = compute_per_device_vocab(VOCAB_SIZE, num_tp) | |
| padded_vocab_size = per_device_vocab * num_tp | |
| sub_core_grids = default_sub_core_grids(mesh_device) | |
| return _SamplingArgs( | |
| vocab_size=VOCAB_SIZE, | |
| padded_vocab_size=padded_vocab_size, | |
| max_batch_size=BATCH_SIZE, | |
| max_top_k=MAX_TOP_K, | |
| cluster_shape=cluster_shape, | |
| sampling_all_gather_axis=sampling_all_gather_axis, | |
| sampling_dp=sampling_dp, | |
| sub_core_grids=sub_core_grids, | |
| sub_core_grid_topk=sub_core_grids, | |
| start_core=ttnn.CoreCoord(0, 0), | |
| model_config={}, | |
| ) | |
| def make_sharded_logits(torch_logits: torch.Tensor, mesh_device, args: _SamplingArgs): | |
| """Create device logits with vocab sharded along the sampling all-gather axis.""" | |
| if mesh_device.get_num_devices() == 1: | |
| mesh_mapper = None | |
| elif args.cluster_shape[0] > 1 and args.cluster_shape[1] > 1: | |
| dims = (None, 3) if args.sampling_all_gather_axis == 1 else (3, None) | |
| mesh_mapper = ttnn.ShardTensor2dMesh(mesh_device, dims=dims, mesh_shape=args.cluster_shape) | |
| else: | |
| mesh_mapper = ttnn.ShardTensorToMesh(mesh_device, dim=3) | |
| return ttnn.from_torch( | |
| torch_logits, | |
| device=mesh_device, | |
| dtype=ttnn.bfloat16, | |
| layout=ttnn.TILE_LAYOUT, | |
| memory_config=ttnn.DRAM_MEMORY_CONFIG, | |
| mesh_mapper=mesh_mapper, | |
| ) | |
| def infer_effective_batch_size( | |
| torch_logits: torch.Tensor, | |
| batch_size: int | None, | |
| *, | |
| max_batch_size: int = BATCH_SIZE, | |
| ) -> int: | |
| if torch_logits.ndim != 4: | |
| raise ValueError( | |
| f"Expected torch_logits with rank 4 [1, 1, batch, vocab], got shape {tuple(torch_logits.shape)}" | |
| ) | |
| inferred_batch_size = int(torch_logits.shape[2]) | |
| effective_batch_size = inferred_batch_size if batch_size is None else int(batch_size) | |
| if effective_batch_size < 1: | |
| raise ValueError(f"batch_size must be >= 1, got {effective_batch_size}") | |
| if effective_batch_size > inferred_batch_size: | |
| raise ValueError( | |
| f"batch_size ({effective_batch_size}) cannot exceed logits batch dimension ({inferred_batch_size})" | |
| ) | |
| if effective_batch_size > max_batch_size: | |
| raise ValueError(f"batch_size ({effective_batch_size}) cannot exceed max test batch size ({max_batch_size})") | |
| return effective_batch_size | |
| def pad_logits_to_max_batch(torch_logits: torch.Tensor, *, max_batch_size: int = BATCH_SIZE) -> torch.Tensor: | |
| """Pad active-batch logits to max batch expected by TT sampling kernels.""" | |
| current_batch = int(torch_logits.shape[2]) | |
| if current_batch == max_batch_size: | |
| return torch_logits | |
| padded = torch.empty( | |
| (torch_logits.shape[0], torch_logits.shape[1], max_batch_size, torch_logits.shape[3]), | |
| dtype=torch_logits.dtype, | |
| device=torch_logits.device, | |
| ) | |
| padded[:, :, :current_batch, :] = torch_logits | |
| # Inactive lanes are filled with a valid copy to avoid all -inf rows. | |
| padded[:, :, current_batch:, :] = torch_logits[:, :, :1, :] | |
| return padded | |
| def validate_token_id(token: int, vocab_size: int, *, field: str): | |
| if not (0 <= token < vocab_size): | |
| raise ValueError(f"{field} token id {token} out of range [0, {vocab_size - 1}]") | |
| def extract_tokens(tt_out_tok, batch_size: int = BATCH_SIZE, device_idx: int = 0) -> list[int]: | |
| """Extract token IDs from one mesh device (default device 0). | |
| Using a single device view matches existing TT sampling test conventions and | |
| keeps checks lightweight; pass device_idx to debug alternate device views. | |
| """ | |
| device_tensors = ttnn.get_device_tensors(tt_out_tok) | |
| if not (0 <= device_idx < len(device_tensors)): | |
| raise ValueError(f"device_idx {device_idx} out of range for {len(device_tensors)} device tensors") | |
| out_torch = ttnn.to_torch(device_tensors[device_idx]).reshape(-1).to(torch.int64) | |
| return out_torch[:batch_size].tolist() | |
| def extract_tokens_all_devices(tt_out_tok, batch_size: int = BATCH_SIZE) -> list[list[int]]: | |
| """Extract token IDs for all device views from a mesh output tensor.""" | |
| device_tensors = ttnn.get_device_tensors(tt_out_tok) | |
| return [extract_tokens(tt_out_tok, batch_size=batch_size, device_idx=i) for i in range(len(device_tensors))] | |
| def representative_device_indices(mesh_device) -> list[int]: | |
| """Choose a small set of device indices for cross-device consistency checks.""" | |
| num_devices = mesh_device.get_num_devices() | |
| if num_devices <= 1: | |
| return [0] | |
| shape = tuple(mesh_device.shape) | |
| if len(shape) == 2 and shape[0] > 1 and shape[1] > 1: | |
| cols = shape[1] | |
| return [row * cols for row in range(shape[0])] | |
| return [0, num_devices - 1] | |
| def build_hot_logits( | |
| args: _SamplingArgs, | |
| *, | |
| batch_size: int = BATCH_SIZE, | |
| hot_tokens: list[int] | None = None, | |
| per_user_hot_tokens: list[list[int]] | None = None, | |
| base_logit: float = -10.0, | |
| top_logit: float = 10.0, | |
| step: float = 0.25, | |
| ) -> torch.Tensor: | |
| """Create logits with a small hot token set that dominates sampling.""" | |
| if hot_tokens is None and per_user_hot_tokens is None: | |
| raise ValueError("Either hot_tokens or per_user_hot_tokens must be provided") | |
| if hot_tokens is not None: | |
| for tok in hot_tokens: | |
| validate_token_id(tok, args.vocab_size, field="hot_tokens") | |
| if per_user_hot_tokens is None: | |
| per_user_hot_tokens = [list(hot_tokens)] * batch_size | |
| assert len(per_user_hot_tokens) == batch_size | |
| for user_idx, user_hot in enumerate(per_user_hot_tokens): | |
| for tok in user_hot: | |
| validate_token_id(tok, args.vocab_size, field=f"per_user_hot_tokens[{user_idx}]") | |
| logits = torch.full((1, 1, batch_size, args.padded_vocab_size), base_logit, dtype=torch.float32) | |
| for user_idx, user_hot in enumerate(per_user_hot_tokens): | |
| for rank, tok in enumerate(user_hot): | |
| logits[0, 0, user_idx, tok] = top_logit - rank * step | |
| logits[:, :, :, args.vocab_size :] = -float("inf") | |
| return logits | |
| def build_disjoint_band_logits( | |
| args: _SamplingArgs, | |
| *, | |
| base_token: int, | |
| tokens_per_user: int = BAND_TOKENS_PER_USER, | |
| batch_size: int = BATCH_SIZE, | |
| **kwargs, | |
| ) -> tuple[torch.Tensor, list[list[int]]]: | |
| """Give every user its own contiguous, non-overlapping hot-token band. | |
| Disjoint bands make cross-user leakage directly observable: a token returned | |
| for user i can only have come from user i's own logits row, so a single | |
| equality check localises the failure to one lane. | |
| Returns (logits, bands) where bands[i][0] is user i's max-logit token. | |
| """ | |
| bands = [ | |
| [base_token + user_idx * tokens_per_user + rank for rank in range(tokens_per_user)] | |
| for user_idx in range(batch_size) | |
| ] | |
| logits = build_hot_logits(args, per_user_hot_tokens=bands, batch_size=batch_size, **kwargs) | |
| return logits, bands | |
| def build_penalty_logits( | |
| args: _SamplingArgs, | |
| *, | |
| target_token: int, | |
| batch_size: int = BATCH_SIZE, | |
| base_logit: float = 4.0, | |
| target_logit: float = 5.0, | |
| ) -> torch.Tensor: | |
| """Create logits where one token is the greedy target before penalties.""" | |
| validate_token_id(target_token, args.vocab_size, field="target_token") | |
| logits = torch.full((1, 1, batch_size, args.padded_vocab_size), base_logit, dtype=torch.float32) | |
| logits[:, :, :, target_token] = target_logit | |
| logits[:, :, :, args.vocab_size :] = -float("inf") | |
| return logits | |
| def run_ttsampling_once( | |
| mesh_device, | |
| args: _SamplingArgs, | |
| torch_logits: torch.Tensor, | |
| *, | |
| top_k, | |
| top_p, | |
| temperature, | |
| batch_size: int | None = None, | |
| ) -> list[int]: | |
| """Run one direct TTSampling forward and return host token IDs. | |
| NOTE: this bypasses 'format_sampling_params', so 'top_k'/'top_p'/'temperature' | |
| are the raw *device-side* values. In particular 'temperature' is the inverse | |
| temperature (1/T) that the kernel multiplies the logits by — it is not the | |
| user-facing temperature, and 0.0 is not "greedy" here. For greedy, pass | |
| 'top_k=1' with 'temperature=1.0'. | |
| """ | |
| assert all( | |
| t > 0 for t in broadcast(temperature) | |
| ), f"temperature is the device-side inverse temperature (1/T) and must be > 0, got {temperature}" | |
| effective_batch_size = infer_effective_batch_size(torch_logits, batch_size, max_batch_size=BATCH_SIZE) | |
| padded_logits = pad_logits_to_max_batch(torch_logits, max_batch_size=BATCH_SIZE) | |
| tt_sampling = None | |
| tt_input = None | |
| tt_logits = None | |
| tt_tokens = None | |
| tt_log_probs = None | |
| try: | |
| tt_sampling = TTSampling( | |
| mesh_device=mesh_device, | |
| tt_ccl=None, | |
| args=args, | |
| k=torch.tensor(broadcast(top_k), dtype=torch.int64), | |
| p=torch.tensor(broadcast(top_p), dtype=torch.float32), | |
| temp=torch.tensor(broadcast(temperature), dtype=torch.float32), | |
| ) | |
| tt_logits = padded_logits | |
| tt_input = make_sharded_logits(tt_logits, mesh_device, args) | |
| tt_tokens, tt_log_probs = tt_sampling(tt_input) | |
| return extract_tokens(tt_tokens, effective_batch_size) | |
| finally: | |
| if tt_log_probs is not None: | |
| del tt_log_probs | |
| if tt_tokens is not None: | |
| del tt_tokens | |
| if tt_input is not None: | |
| del tt_input | |
| if tt_logits is not None: | |
| del tt_logits | |
| if tt_sampling is not None: | |
| del tt_sampling | |
| safe_sync(mesh_device) | |
| def run_sampling_generator( | |
| mesh_device, | |
| args: _SamplingArgs, | |
| torch_logits: torch.Tensor, | |
| sampling_params: SamplingParams, | |
| *, | |
| num_steps: int = 1, | |
| advance_seeds: bool = True, | |
| seed_values: list[int] | None = None, | |
| write_seed_values_to_device: bool = False, | |
| batch_size: int | None = None, | |
| device_idx: int = 0, | |
| state_setup=None, | |
| enable_trace: bool = False, | |
| ) -> list[list[int]]: | |
| """Run SamplingGenerator for num_steps and return per-step token lists.""" | |
| effective_batch_size = infer_effective_batch_size(torch_logits, batch_size, max_batch_size=BATCH_SIZE) | |
| padded_logits = pad_logits_to_max_batch(torch_logits, max_batch_size=BATCH_SIZE) | |
| sg = None | |
| tt_input = None | |
| tt_tokens = None | |
| tt_log_probs = None | |
| outputs = [] | |
| try: | |
| sg = SamplingGenerator(args=args, mesh_device=mesh_device, tt_ccl=None) | |
| formatted = format_sampling_params(sampling_params, BATCH_SIZE) | |
| sg.reset_sampling_params(formatted) | |
| if seed_values is not None: | |
| if len(seed_values) > BATCH_SIZE: | |
| raise ValueError(f"seed_values length ({len(seed_values)}) cannot exceed BATCH_SIZE ({BATCH_SIZE})") | |
| if len(seed_values) > effective_batch_size and not write_seed_values_to_device: | |
| raise ValueError( | |
| f"seed_values length ({len(seed_values)}) cannot exceed active batch size ({effective_batch_size})" | |
| ) | |
| user_ids = list(range(len(seed_values))) | |
| sg.seed_manager.reset_seed(seed_values, user_ids) | |
| else: | |
| # Mirror production's apply_prefill_state: reset_sampling_params does | |
| # not touch the SeedManager, so an unseeded request must still call | |
| # reset_seed (seeds=None) to move it out of the fresh _reseted=False | |
| # state. Without this the SeedManager stays in the steady "SKIP" state | |
| # and get_new_values() never pushes entropy seeds, making unseeded | |
| # sampling deterministic. | |
| sg.seed_manager.reset_seed(None, list(range(BATCH_SIZE))) | |
| if state_setup is not None: | |
| state_setup(sg) | |
| if enable_trace and sg._penalties_active: | |
| raise ValueError( | |
| "enable_trace=True is incompatible with penalties in this harness: TTPenalties.apply() " | |
| "rewrites the traced input tensor in place, and trace replay cannot re-upload it." | |
| ) | |
| if write_seed_values_to_device: | |
| if seed_values is None: | |
| raise ValueError("write_seed_values_to_device=True requires seed_values") | |
| sg.seed_manager.write_device_seed_values(seed_values) | |
| # Trace replay binds the captured input address, so `_validate_trace_inputs` | |
| # requires the SAME logits tensor object on every step after capture. Upload it | |
| # once up front for traced runs; penalties are rejected below because | |
| # TTPenalties.apply() would rewrite that one tensor in place. | |
| if enable_trace: | |
| tt_input = make_sharded_logits(padded_logits, mesh_device, args) | |
| for _ in range(num_steps): | |
| # SamplingGenerator keeps per-user RNG state in SeedManager. | |
| if advance_seeds: | |
| sg.seed_manager.get_new_values() | |
| if not enable_trace: | |
| # TTPenalties.apply() rewrites the logits tensor in place, so upload a | |
| # pristine copy every step. Reusing one device tensor would compound | |
| # each step's penalties onto the previous step's already-penalised | |
| # logits, whereas real decode gets fresh logits from the LM head. | |
| if tt_input is not None: | |
| ttnn.deallocate(tt_input) | |
| tt_input = make_sharded_logits(padded_logits, mesh_device, args) | |
| tt_tokens, tt_log_probs = sg.sample(tt_input, enable_trace=enable_trace) | |
| if enable_trace: | |
| # ttnn.execute_trace replays with blocking=False, so the output buffer is | |
| # only valid once the device has caught up. | |
| ttnn.synchronize_device(mesh_device) | |
| outputs.append(extract_tokens(tt_tokens, effective_batch_size, device_idx=device_idx)) | |
| return outputs | |
| finally: | |
| # Release the captured trace before the generator goes away. mesh_device is | |
| # function-scoped so the region would be reclaimed at device close anyway, but a | |
| # leaked trace breaks any future test that captures twice in one device session. | |
| if enable_trace and sg is not None: | |
| try: | |
| sg.reset_trace() | |
| except Exception: | |
| pass | |
| if tt_log_probs is not None: | |
| del tt_log_probs | |
| if tt_tokens is not None: | |
| del tt_tokens | |
| if tt_input is not None: | |
| del tt_input | |
| if sg is not None: | |
| del sg | |
| safe_sync(mesh_device) | |
| def assert_tokens_in_vocab(tokens: list[int], vocab_size: int = VOCAB_SIZE): | |
| assert all( | |
| 0 <= tok < vocab_size for tok in tokens | |
| ), f"Found out-of-range token(s) for vocab_size={vocab_size}: {tokens}" | |
| def flatten_steps(outputs: list[list[int]]) -> list[int]: | |
| return [tok for step in outputs for tok in step] | |
| # --- Test: prefill parameter behavior --- | |
| class TestPrefillWithDifferentParams: | |
| def test_prefill_temperature_varied_in_batch(self, mesh_device, device_params): | |
| args = make_sampling_args(mesh_device) | |
| hot_tokens = [100, 101, 102, 103, 104, 105, 106, 107] | |
| hot_token_set = set(hot_tokens) | |
| logits = build_hot_logits(args, hot_tokens=hot_tokens) | |
| # Every lane gets a different temperature, including greedy lanes (0.0). | |
| temperature = ([0.0, 0.5, 1.0, 2.0] * (BATCH_SIZE // 4))[:BATCH_SIZE] | |
| greedy_lanes = [i for i, t in enumerate(temperature) if t == 0.0] | |
| params = SamplingParams(temperature=temperature, top_k=8, top_p=1.0) | |
| seeds = [3100 + i for i in range(BATCH_SIZE)] | |
| out1 = run_sampling_generator( | |
| mesh_device, args, logits, params, num_steps=1, advance_seeds=True, seed_values=seeds | |
| )[0] | |
| out2 = run_sampling_generator( | |
| mesh_device, args, logits, params, num_steps=1, advance_seeds=True, seed_values=seeds | |
| )[0] | |
| assert out1 == out2, "Same per-user seeds should replay exactly for stochastic prefill config" | |
| assert_tokens_in_vocab(out1, args.vocab_size) | |
| unexpected = [tok for tok in out1 if tok not in hot_token_set] | |
| assert not unexpected, f"Sampled tokens outside expected hot set: {unexpected}, out={out1}" | |
| for lane in greedy_lanes: | |
| assert out1[lane] == hot_tokens[0], ( | |
| f"temperature=0.0 lane {lane} must be greedy on the max-logit token " | |
| f"{hot_tokens[0]}, got {out1[lane]}. out={out1}" | |
| ) | |
| def test_prefill_temperature_varied_between_batches(self, mesh_device): | |
| args = make_sampling_args(mesh_device) | |
| logits = build_hot_logits(args, hot_tokens=[120, 121, 122, 123, 124, 125, 126, 127]) | |
| params = SamplingParams(temperature=2.0, top_k=8, top_p=1.0) | |
| outputs = run_sampling_generator( | |
| mesh_device, args, logits, params, num_steps=FAST_NUM_TRIES, advance_seeds=True | |
| ) | |
| user0 = [step[0] for step in outputs] | |
| assert len(set(user0)) >= 2, f"Expected variation across runs for stochastic sampling, got {user0}" | |
| def test_prefill_topk_mixed_greedy_and_stochastic(self, mesh_device): | |
| args = make_sampling_args(mesh_device) | |
| logits = build_hot_logits(args, hot_tokens=[140, 141, 142, 143, 144, 145, 146, 147]) | |
| half = BATCH_SIZE // 2 | |
| temperature = [0.0] * half + [1.5] * (BATCH_SIZE - half) # First half is greedy; second half is stochastic. | |
| top_k = [32] * half + [8] * (BATCH_SIZE - half) | |
| top_p = [1.0] * BATCH_SIZE | |
| params = SamplingParams(temperature=temperature, top_k=top_k, top_p=top_p) | |
| seed_values = [1000 + i for i in range(BATCH_SIZE)] | |
| out1 = run_sampling_generator( | |
| mesh_device, args, logits, params, num_steps=1, advance_seeds=True, seed_values=seed_values | |
| )[0] | |
| out2 = run_sampling_generator( | |
| mesh_device, args, logits, params, num_steps=1, advance_seeds=True, seed_values=seed_values | |
| )[0] | |
| assert out1 == out2, "Same per-user seeds must replay exactly across batches" | |
| greedy = out1[:half] | |
| stochastic = out1[half:] | |
| assert len(set(greedy)) == 1, f"Greedy half should be deterministic and identical, got {greedy}" | |
| assert len(set(stochastic)) >= 2, f"Stochastic half should vary across slots, got {stochastic}" | |
| def test_prefill_seeding(self, mesh_device): | |
| args = make_sampling_args(mesh_device) | |
| logits = build_hot_logits(args, hot_tokens=[160, 161, 162, 163, 164, 165, 166, 167]) | |
| thirds = BATCH_SIZE // 3 | |
| greedy_count = BATCH_SIZE - 2 * thirds | |
| temperature = [0.0] * greedy_count + [1.5] * (BATCH_SIZE - greedy_count) | |
| top_k = [32] * greedy_count + [8] * (BATCH_SIZE - greedy_count) | |
| top_p = [1.0] * BATCH_SIZE | |
| params = SamplingParams(temperature=temperature, top_k=top_k, top_p=top_p) | |
| seed_values = [2000 + i for i in range(BATCH_SIZE)] | |
| out1 = run_sampling_generator( | |
| mesh_device, args, logits, params, num_steps=1, advance_seeds=True, seed_values=seed_values | |
| )[0] | |
| out2 = run_sampling_generator( | |
| mesh_device, args, logits, params, num_steps=1, advance_seeds=True, seed_values=seed_values | |
| )[0] | |
| assert out1 == out2, "Replay with same seeds must be deterministic for every slot" | |
| assert len(set(out1[greedy_count:])) >= 2, "Different stochastic seeds should yield diverse tokens" | |
| def test_prefill_topk_1_is_greedy(self, mesh_device): | |
| args = make_sampling_args(mesh_device) | |
| hot_tokens = [180, 181, 182, 183] | |
| logits = build_hot_logits(args, hot_tokens=hot_tokens) | |
| # top_k=1 collapses to argmax whatever the temperature scaling is, so the | |
| # max-logit token must come back for every inverse temperature. | |
| # top_p=1.0 (disabled) keeps this portable across top-p conventions. | |
| for inverse_temperature in (0.25, 1.0, 4.0): | |
| tokens = run_ttsampling_once(mesh_device, args, logits, top_k=1, top_p=1.0, temperature=inverse_temperature) | |
| assert all(tok == hot_tokens[0] for tok in tokens), ( | |
| f"top_k=1 should be greedy at inverse temperature {inverse_temperature}, " | |
| f"expected all {hot_tokens[0]}, got {tokens}" | |
| ) | |
| def test_greedy_picks_max_logit(self, mesh_device): | |
| args = make_sampling_args(mesh_device) | |
| logits = build_hot_logits(args, hot_tokens=[42, 43, 44]) # 42 has highest logit | |
| tokens = run_ttsampling_once(mesh_device, args, logits, top_k=1, top_p=1.0, temperature=1.0) | |
| assert all(tok == 42 for tok in tokens), "Greedy should always pick the max logit token" | |
| def test_run_ttsampling_once_respects_logits_batch_size(self, mesh_device): | |
| args = make_sampling_args(mesh_device) | |
| logits = build_hot_logits(args, batch_size=2, hot_tokens=[60, 61, 62]) | |
| tokens = run_ttsampling_once(mesh_device, args, logits, top_k=1, top_p=1.0, temperature=1.0) | |
| assert len(tokens) == 2, f"Expected 2 tokens for batch_size=2 logits, got {len(tokens)}" | |
| assert_tokens_in_vocab(tokens, args.vocab_size) | |
| assert all(tok == 60 for tok in tokens), f"Greedy should pick the max-logit token 60, got {tokens}" | |
| def test_run_sampling_generator_respects_logits_batch_size(self, mesh_device): | |
| args = make_sampling_args(mesh_device) | |
| logits = build_hot_logits(args, batch_size=2, hot_tokens=[80, 81, 82]) | |
| params = SamplingParams(temperature=[0.0, 1.0], top_k=[1, 3], top_p=[1.0, 1.0]) | |
| outputs = run_sampling_generator(mesh_device, args, logits, params, num_steps=1, advance_seeds=True) | |
| assert len(outputs) == 1 | |
| assert len(outputs[0]) == 2, f"Expected 2 tokens for batch_size=2 logits, got {len(outputs[0])}" | |
| assert_tokens_in_vocab(outputs[0], args.vocab_size) | |
| def test_extract_tokens_invalid_device_idx_raises(self, mesh_device, expect_error): | |
| args = make_sampling_args(mesh_device) | |
| logits = build_hot_logits(args, hot_tokens=[84, 85, 86]) | |
| params = SamplingParams(temperature=1.0, top_k=3, top_p=1.0) | |
| invalid_device_idx = mesh_device.get_num_devices() | |
| with expect_error(ValueError, r"device_idx .* out of range"): | |
| run_sampling_generator( | |
| mesh_device, | |
| args, | |
| logits, | |
| params, | |
| num_steps=1, | |
| advance_seeds=False, | |
| device_idx=invalid_device_idx, | |
| ) | |
| def test_build_hot_logits_rejects_out_of_vocab_token(self, mesh_device, expect_error): | |
| args = make_sampling_args(mesh_device) | |
| bad_token = args.vocab_size | |
| with expect_error(ValueError, r"out of range"): | |
| build_hot_logits(args, hot_tokens=[bad_token]) | |
| def test_build_penalty_logits_rejects_out_of_vocab_target(self, mesh_device, expect_error): | |
| args = make_sampling_args(mesh_device) | |
| bad_token = args.vocab_size | |
| with expect_error(ValueError, r"out of range"): | |
| build_penalty_logits(args, target_token=bad_token) | |
| def test_top_p_restricts_candidate_set(self, mesh_device): | |
| args = make_sampling_args(mesh_device) | |
| hot_tokens = [1900, 1901, 1902] | |
| logits = build_hot_logits( | |
| args, | |
| hot_tokens=hot_tokens, | |
| base_logit=-12.0, | |
| top_logit=2.0, | |
| step=0.2, | |
| ) | |
| seeds = [31000 + i for i in range(BATCH_SIZE)] | |
| # Per-lane lists so all 32 lanes sample. With a scalar temperature only lane 0 would, | |
| # making "1902 in full_set" a 6-draw coin flip (~15% failure) instead of 192 draws. | |
| params_low = per_lane_params(1.0, 3, 0.6) | |
| params_full = per_lane_params(1.0, 3, 1.0) | |
| out_low = run_sampling_generator( | |
| mesh_device, | |
| args, | |
| logits, | |
| params_low, | |
| num_steps=FAST_NUM_TRIES, | |
| advance_seeds=True, | |
| seed_values=seeds, | |
| ) | |
| out_full = run_sampling_generator( | |
| mesh_device, | |
| args, | |
| logits, | |
| params_full, | |
| num_steps=FAST_NUM_TRIES, | |
| advance_seeds=True, | |
| seed_values=seeds, | |
| ) | |
| for step_tokens in out_low + out_full: | |
| assert_tokens_in_vocab(step_tokens, args.vocab_size) | |
| low_flat = flatten_steps(out_low) | |
| full_flat = flatten_steps(out_full) | |
| low_set = set(low_flat) | |
| full_set = set(full_flat) | |
| allowed_low = {1900, 1901} | |
| allowed_full = {1900, 1901, 1902} | |
| assert low_set.issubset(allowed_low), ( | |
| f"top_p=0.6 sampled outside allowed set {sorted(allowed_low)}. " | |
| f"low_set={sorted(low_set)}, low_hist={Counter(low_flat).most_common(6)}" | |
| ) | |
| assert full_set.issubset(allowed_full), ( | |
| f"top_p=1.0 sampled outside expected hot set {sorted(allowed_full)}. " | |
| f"full_set={sorted(full_set)}, full_hist={Counter(full_flat).most_common(6)}" | |
| ) | |
| assert 1902 in full_set, ( | |
| "top_p=1.0 did not sample token 1902 at least once. " | |
| f"low_set={sorted(low_set)}, full_set={sorted(full_set)}, " | |
| f"low_hist={Counter(low_flat).most_common(6)}, full_hist={Counter(full_flat).most_common(6)}" | |
| ) | |
| # --- Test: per-request penalties --- | |
| class TestRepetitionPenaltyPerRequest: | |
| def test_different_repetition_penalties(self, mesh_device, device_params): | |
| args = make_sampling_args(mesh_device) | |
| target_token = 500 | |
| logits = build_penalty_logits(args, target_token=target_token) | |
| penalties = ([1.0, 1.0, 1.2, 1.5, 2.0, 3.0, 4.0, 5.0] * 4)[:BATCH_SIZE] | |
| params = self._get_sampling_params(penalties) | |
| tokens = run_sampling_generator( | |
| mesh_device, args, logits, params, state_setup=lambda sg: self._state_setup(sg, target_token) | |
| )[0] | |
| low_penalty = [ | |
| tok for i, tok in enumerate(tokens) if penalties[i] <= 1.0 | |
| ] # Keep boundary-sensitive values (e.g. 1.2) out of strict assertions. | |
| high_penalty = [tok for i, tok in enumerate(tokens) if penalties[i] >= 1.5] | |
| assert all(tok == target_token for tok in low_penalty), "Low repetition penalty should keep target token" | |
| assert any(tok != target_token for tok in high_penalty), "High repetition penalties should alter output" | |
| def test_repetition_penalty_vs_no_penalty(self, mesh_device, device_params): | |
| args = make_sampling_args(mesh_device) | |
| target_token = 520 | |
| logits = build_penalty_logits(args, target_token=target_token) | |
| penalties = [1.0 if i % 2 == 0 else 2.5 for i in range(BATCH_SIZE)] | |
| params = self._get_sampling_params(penalties) | |
| tokens = run_sampling_generator( | |
| mesh_device, args, logits, params, state_setup=lambda sg: self._state_setup(sg, target_token) | |
| )[0] | |
| no_penalty = [tokens[i] for i in range(0, BATCH_SIZE, 2)] | |
| with_penalty = [tokens[i] for i in range(1, BATCH_SIZE, 2)] | |
| assert all(tok == target_token for tok in no_penalty), "No-penalty lanes should keep target" | |
| assert all(tok != target_token for tok in with_penalty), "Penalty lanes should change token" | |
| assert no_penalty[0] != with_penalty[0], "Penalty and no-penalty outputs should differ" | |
| def test_repetition_penalty_persists_across_steps(self, mesh_device, device_params): | |
| args = make_sampling_args(mesh_device) | |
| target_token = 540 | |
| logits = build_penalty_logits(args, target_token=target_token) | |
| penalties = [1.0 if i % 2 == 0 else 2.5 for i in range(BATCH_SIZE)] | |
| params = self._get_sampling_params(penalties) | |
| outputs = run_sampling_generator( | |
| mesh_device, | |
| args, | |
| logits, | |
| params, | |
| num_steps=3, | |
| advance_seeds=False, | |
| ) | |
| for step_tokens in outputs: | |
| assert_tokens_in_vocab(step_tokens, args.vocab_size) | |
| even_idxs = range(0, BATCH_SIZE, 2) | |
| odd_idxs = range(1, BATCH_SIZE, 2) | |
| assert all(outputs[0][i] == target_token for i in odd_idxs), "Odd lanes should start at target token" | |
| assert all(outputs[0][i] == target_token for i in even_idxs), "Even lanes should start at target token" | |
| assert all(outputs[1][i] != target_token for i in odd_idxs), "Odd lanes should leave target at step 1" | |
| assert all(outputs[2][i] != target_token for i in odd_idxs), "Odd lanes should stay off target at step 2" | |
| assert all(outputs[1][i] == target_token for i in even_idxs), "Even lanes should keep target at step 1" | |
| assert all(outputs[2][i] == target_token for i in even_idxs), "Even lanes should keep target at step 2" | |
| def _state_setup(self, sg, target_token: int): | |
| seen = torch.full((BATCH_SIZE, 1), target_token, dtype=torch.int64) | |
| sg.reset_prompt_tokens(seen) | |
| sg.reset_output_state(tokens=seen) | |
| def _get_sampling_params(self, penalties): | |
| return SamplingParams( | |
| temperature=[0.0] * BATCH_SIZE, | |
| top_k=[32] * BATCH_SIZE, | |
| top_p=[1.0] * BATCH_SIZE, | |
| repetition_penalty=penalties, | |
| ) | |
| class TestPresencePenaltyPerRequest: | |
| def test_different_presence_penalties(self, mesh_device, device_params): | |
| args = make_sampling_args(mesh_device) | |
| target_token = 700 | |
| logits = build_penalty_logits(args, target_token=target_token) | |
| penalties = ([0.0, 0.5, 1.0, 2.0, 3.0, -0.5, -1.0, 4.0] * 4)[:BATCH_SIZE] | |
| params = self._get_sampling_params(penalties) | |
| tokens = run_sampling_generator( | |
| mesh_device, args, logits, params, state_setup=lambda sg: self._state_setup(sg, target_token) | |
| )[0] | |
| assert any(tok == target_token for tok in tokens), "Some lanes should retain target token" | |
| assert any(tok != target_token for tok in tokens), "Some lanes should shift off target with higher penalties" | |
| def test_presence_penalty_mixed_batch(self, mesh_device, device_params): | |
| args = make_sampling_args(mesh_device) | |
| target_token = 720 | |
| logits = build_penalty_logits(args, target_token=target_token) | |
| penalties = [0.0 if i % 2 == 0 else 2.0 for i in range(BATCH_SIZE)] | |
| params = self._get_sampling_params(penalties) | |
| tokens = run_sampling_generator( | |
| mesh_device, args, logits, params, state_setup=lambda sg: self._state_setup(sg, target_token) | |
| )[0] | |
| no_penalty = [tokens[i] for i in range(0, BATCH_SIZE, 2)] | |
| with_penalty = [tokens[i] for i in range(1, BATCH_SIZE, 2)] | |
| assert all(tok == target_token for tok in no_penalty), "No presence-penalty lanes should keep target" | |
| assert all(tok != target_token for tok in with_penalty), "Presence-penalty lanes should move off target" | |
| def test_presence_penalty_persists_across_steps(self, mesh_device, device_params): | |
| args = make_sampling_args(mesh_device) | |
| target_token = 740 | |
| logits = build_penalty_logits(args, target_token=target_token) | |
| penalties = [0.0 if i % 2 == 0 else 2.0 for i in range(BATCH_SIZE)] | |
| params = self._get_sampling_params(penalties) | |
| outputs = run_sampling_generator( | |
| mesh_device, | |
| args, | |
| logits, | |
| params, | |
| num_steps=3, | |
| advance_seeds=False, | |
| ) | |
| for step_tokens in outputs: | |
| assert_tokens_in_vocab(step_tokens, args.vocab_size) | |
| even_idxs = range(0, BATCH_SIZE, 2) | |
| odd_idxs = range(1, BATCH_SIZE, 2) | |
| assert all(outputs[0][i] == target_token for i in odd_idxs), "Odd lanes should start at target token" | |
| assert all(outputs[0][i] == target_token for i in even_idxs), "Even lanes should start at target token" | |
| assert all(outputs[1][i] != target_token for i in odd_idxs), "Odd lanes should leave target at step 1" | |
| assert all(outputs[2][i] != target_token for i in odd_idxs), "Odd lanes should stay off target at step 2" | |
| assert all(outputs[1][i] == target_token for i in even_idxs), "Even lanes should keep target at step 1" | |
| assert all(outputs[2][i] == target_token for i in even_idxs), "Even lanes should keep target at step 2" | |
| def _state_setup(self, sg, target_token: int): | |
| seen = torch.full((BATCH_SIZE, 1), target_token, dtype=torch.int64) | |
| sg.reset_prompt_tokens(seen) | |
| sg.reset_output_state(tokens=seen) | |
| def _get_sampling_params(self, penalties): | |
| return SamplingParams( | |
| temperature=[0.0] * BATCH_SIZE, | |
| top_k=[32] * BATCH_SIZE, | |
| top_p=[1.0] * BATCH_SIZE, | |
| presence_penalty=penalties, | |
| ) | |
| class TestFrequencyPenaltyPerRequest: | |
| def test_different_frequency_penalties(self, mesh_device, device_params): | |
| args = make_sampling_args(mesh_device) | |
| target_token = 900 | |
| logits = build_penalty_logits(args, target_token=target_token) | |
| penalties = ([0.0, 0.5, 1.0, 2.0, 3.0, -0.5, -1.0, 4.0] * 4)[:BATCH_SIZE] | |
| params = self._get_sampling_params(penalties) | |
| tokens = run_sampling_generator( | |
| mesh_device, args, logits, params, state_setup=lambda sg: self._state_setup(sg, target_token) | |
| )[0] | |
| assert any(tok == target_token for tok in tokens), "Some lanes should retain target token" | |
| assert any(tok != target_token for tok in tokens), "Some lanes should shift off target with higher penalties" | |
| def test_frequency_penalty_mixed_batch(self, mesh_device, device_params): | |
| args = make_sampling_args(mesh_device) | |
| target_token = 920 | |
| logits = build_penalty_logits(args, target_token=target_token) | |
| penalties = [0.0 if i % 2 == 0 else 2.0 for i in range(BATCH_SIZE)] | |
| params = self._get_sampling_params(penalties) | |
| tokens = run_sampling_generator( | |
| mesh_device, args, logits, params, state_setup=lambda sg: self._state_setup(sg, target_token) | |
| )[0] | |
| no_penalty = [tokens[i] for i in range(0, BATCH_SIZE, 2)] | |
| with_penalty = [tokens[i] for i in range(1, BATCH_SIZE, 2)] | |
| assert all(tok == target_token for tok in no_penalty), "No frequency-penalty lanes should keep target" | |
| assert all(tok != target_token for tok in with_penalty), "Frequency-penalty lanes should move off target" | |
| def test_frequency_penalty_accumulates_across_steps(self, mesh_device, device_params): | |
| args = make_sampling_args(mesh_device) | |
| target_token = 940 | |
| logits = build_penalty_logits(args, target_token=target_token) | |
| penalties = [0.1 if i % 2 == 0 else 0.6 for i in range(BATCH_SIZE)] | |
| params = self._get_sampling_params(penalties) | |
| outputs = run_sampling_generator( | |
| mesh_device, | |
| args, | |
| logits, | |
| params, | |
| num_steps=4, | |
| advance_seeds=False, | |
| ) | |
| for step_tokens in outputs: | |
| assert_tokens_in_vocab(step_tokens, args.vocab_size) | |
| even_idxs = range(0, BATCH_SIZE, 2) | |
| odd_idxs = range(1, BATCH_SIZE, 2) | |
| assert all(outputs[0][i] == target_token for i in odd_idxs), "Odd lanes should start at target token" | |
| assert all(outputs[1][i] == target_token for i in odd_idxs), "Odd lanes should still hit target at step 1" | |
| assert all(outputs[2][i] != target_token for i in odd_idxs), "Odd lanes should leave target at step 2" | |
| assert all(outputs[3][i] != target_token for i in odd_idxs), "Odd lanes should stay off target at step 3" | |
| assert all(outputs[0][i] == target_token for i in even_idxs), "Even lanes should start at target token" | |
| assert all(outputs[1][i] == target_token for i in even_idxs), "Even lanes should keep target at step 1" | |
| assert all(outputs[2][i] == target_token for i in even_idxs), "Even lanes should keep target at step 2" | |
| assert all(outputs[3][i] == target_token for i in even_idxs), "Even lanes should keep target at step 3" | |
| def _state_setup(self, sg, target_token: int): | |
| seen = torch.full((BATCH_SIZE, 1), target_token, dtype=torch.int64) | |
| sg.reset_prompt_tokens(seen) | |
| sg.reset_output_state(tokens=seen) | |
| def _get_sampling_params(self, penalties): | |
| return SamplingParams( | |
| temperature=[0.0] * BATCH_SIZE, | |
| top_k=[32] * BATCH_SIZE, | |
| top_p=[1.0] * BATCH_SIZE, | |
| frequency_penalty=penalties, | |
| ) | |
| # --- Test: seed behavior --- | |
| class TestSeededSamplingPerRequest: | |
| def test_run_sampling_generator_rejects_seed_vector_longer_than_max_batch(self, mesh_device, expect_error): | |
| args = make_sampling_args(mesh_device) | |
| logits = build_hot_logits(args, hot_tokens=[990, 991, 992]) | |
| params = SamplingParams(temperature=1.0, top_k=3, top_p=1.0) | |
| too_many_seeds = list(range(BATCH_SIZE + 1)) | |
| with expect_error(ValueError, r"cannot exceed BATCH_SIZE"): | |
| run_sampling_generator( | |
| mesh_device, | |
| args, | |
| logits, | |
| params, | |
| num_steps=1, | |
| advance_seeds=True, | |
| seed_values=too_many_seeds, | |
| ) | |
| def test_run_sampling_generator_rejects_seed_vector_longer_than_effective_batch(self, mesh_device, expect_error): | |
| args = make_sampling_args(mesh_device) | |
| logits = build_hot_logits(args, batch_size=2, hot_tokens=[993, 994, 995]) | |
| params = SamplingParams(temperature=[1.0, 1.0], top_k=[3, 3], top_p=[1.0, 1.0]) | |
| too_many_for_active_batch = [11, 12, 13] | |
| with expect_error(ValueError, r"active batch size"): | |
| run_sampling_generator( | |
| mesh_device, | |
| args, | |
| logits, | |
| params, | |
| num_steps=1, | |
| advance_seeds=True, | |
| seed_values=too_many_for_active_batch, | |
| ) | |
| def test_seeded_replay_is_deterministic(self, mesh_device): | |
| """Same seed vector must replay exactly across independent generator instances. | |
| Split out of the per-lane seed investigation below so that these invariants -- | |
| replay determinism, token range, and top-k containment -- keep gating CI while the | |
| per-lane seed question is still open. A regression to argmax-only sampling or an | |
| out-of-range token fails here, not silently as an xfail. (#38316) | |
| """ | |
| args = make_sampling_args(mesh_device) | |
| hot_tokens = [1000, 1001, 1002, 1003, 1004, 1005, 1006, 1007] | |
| hot_token_set = set(hot_tokens) | |
| logits = build_hot_logits(args, hot_tokens=hot_tokens) | |
| params = per_lane_params(1.5, 8, 1.0) | |
| seeds_a = list(range(BATCH_SIZE)) | |
| out_a1 = run_sampling_generator( | |
| mesh_device, args, logits, params, num_steps=1, advance_seeds=True, seed_values=seeds_a | |
| )[0] | |
| out_a2 = run_sampling_generator( | |
| mesh_device, args, logits, params, num_steps=1, advance_seeds=True, seed_values=seeds_a | |
| )[0] | |
| assert out_a1 == out_a2, "Same seed vector should replay exactly across independent runs" | |
| for outputs in (out_a1, out_a2): | |
| assert_tokens_in_vocab(outputs, args.vocab_size) | |
| unexpected = [tok for tok in outputs if tok not in hot_token_set] | |
| assert not unexpected, ( | |
| f"Sampled tokens outside expected hot set {sorted(hot_token_set)}: {unexpected}. " f"outputs={outputs}" | |
| ) | |
| def test_different_seeds_produce_different_outputs(self, mesh_device): | |
| """Distinct per-lane seeds must produce a distinct draw per lane. | |
| Previously xfailed as a suspected per-lane seeding defect (blamed on #50685). It was a | |
| bug in this test: it passed SCALAR sampling params, so format_sampling_params treated | |
| only lane 0 as active and lanes 1..31 came back greedy on the argmax. Run 32289587326 | |
| with the marker removed showed it exactly -- out_a1 = [1002, 1000, 1000, ... 1000] with | |
| hot_tokens[0] == 1000 -- so the "distinct seeds" assertion was comparing 31 deterministic | |
| argmax lanes plus one stochastic lane whose single draw happened to repeat. Fixed by | |
| passing per-lane lists; there is no product defect and #50685 is not implicated. | |
| """ | |
| args = make_sampling_args(mesh_device) | |
| hot_tokens = [1000, 1001, 1002, 1003, 1004, 1005, 1006, 1007] | |
| logits = build_hot_logits(args, hot_tokens=hot_tokens) | |
| # Per-lane lists: a scalar temperature would leave lanes 1..31 greedy and this test | |
| # would be asserting against 31 argmax lanes. See per_lane_params. | |
| params = per_lane_params(1.5, 8, 1.0) | |
| seeds_a = list(range(BATCH_SIZE)) | |
| out_a1 = run_sampling_generator( | |
| mesh_device, args, logits, params, num_steps=1, advance_seeds=True, seed_values=seeds_a | |
| )[0] | |
| # 32 distinct seeds over a flat-ish top-8 distribution: the lanes must not all | |
| # land on the same token. This is the reported symptom -- lanes 1..31 coming out | |
| # identical despite distinct seeds. | |
| assert len(set(out_a1)) > 1, ( | |
| f"32 distinct per-lane seeds produced one token across all lanes: {out_a1}. " | |
| "Per-lane seeding is not taking effect." | |
| ) | |
| # Shifting every seed must move at least one lane; identical output for a different | |
| # seed vector means the RNG path is inert (argmax-only). | |
| seeds_b = [s + 12345 for s in seeds_a] | |
| out_b = run_sampling_generator( | |
| mesh_device, args, logits, params, num_steps=1, advance_seeds=True, seed_values=seeds_b | |
| )[0] | |
| changed_indices = [i for i, (a, b) in enumerate(zip(out_a1, out_b)) if a != b] | |
| assert len(changed_indices) > 0, ( | |
| "Different seed vector produced identical outputs for this stochastic config; " | |
| f"backend appears argmax-only (seed/RNG regression). out_a1={out_a1}, out_b={out_b}" | |
| ) | |
| def test_same_seeds_reproduce_across_batches(self, mesh_device, device_params): | |
| args = make_sampling_args(mesh_device) | |
| logits = build_hot_logits(args, hot_tokens=[1100, 1101, 1102, 1103, 1104, 1105, 1106, 1107]) | |
| params = per_lane_params(1.25, 8, 1.0) | |
| seeds = [500 + i for i in range(BATCH_SIZE)] | |
| out1 = run_sampling_generator( | |
| mesh_device, args, logits, params, num_steps=FAST_NUM_STEPS, advance_seeds=True, seed_values=seeds | |
| ) | |
| out2 = run_sampling_generator( | |
| mesh_device, args, logits, params, num_steps=FAST_NUM_STEPS, advance_seeds=True, seed_values=seeds | |
| ) | |
| assert out1 == out2, "Same seed replay should match exactly for all users and steps" | |
| for device_idx in representative_device_indices(mesh_device)[1:]: | |
| out_device = run_sampling_generator( | |
| mesh_device, | |
| args, | |
| logits, | |
| params, | |
| num_steps=FAST_NUM_STEPS, | |
| advance_seeds=True, | |
| seed_values=seeds, | |
| device_idx=device_idx, | |
| ) | |
| assert out_device == out1, f"Device view mismatch for device_idx={device_idx} in seeded replay" | |
| def test_specific_seed_reproducible(self, mesh_device, seed): | |
| args = make_sampling_args(mesh_device) | |
| logits = build_hot_logits(args, hot_tokens=[1200, 1201, 1202, 1203, 1204, 1205, 1206, 1207]) | |
| params = SamplingParams(temperature=1.25, top_k=8, top_p=1.0) | |
| seeds = [10000 + i for i in range(BATCH_SIZE)] | |
| seeds[0] = seed | |
| out1 = run_sampling_generator( | |
| mesh_device, args, logits, params, num_steps=1, advance_seeds=True, seed_values=seeds | |
| )[0] | |
| out2 = run_sampling_generator( | |
| mesh_device, args, logits, params, num_steps=1, advance_seeds=True, seed_values=seeds | |
| )[0] | |
| assert out1[0] == out2[0], f"Seed {seed} should be reproducible for user-0" | |
| def test_batch1_seed_reproducible(self, mesh_device, seed): | |
| args = make_sampling_args(mesh_device) | |
| logits = build_hot_logits(args, hot_tokens=[1300, 1301, 1302, 1303, 1304, 1305, 1306, 1307]) | |
| params = SamplingParams(temperature=2.0, top_k=8, top_p=1.0) | |
| results = [] | |
| for _ in range(FAST_NUM_TRIES): | |
| seeds = [seed] + [100000 + i for i in range(1, BATCH_SIZE)] | |
| out = run_sampling_generator( | |
| mesh_device, args, logits, params, num_steps=1, advance_seeds=True, seed_values=seeds | |
| )[0] | |
| results.append(out[0]) | |
| assert len(set(results)) == 1, f"Single seeded slot should reproduce, got {results}" | |
| def test_batch1_no_seed_varied(self, mesh_device): | |
| args = make_sampling_args(mesh_device) | |
| logits = build_hot_logits(args, hot_tokens=[1400, 1401, 1402, 1403, 1404, 1405, 1406, 1407]) | |
| params = SamplingParams(temperature=2.0, top_k=8, top_p=1.0) | |
| results = [] | |
| for _ in range(FAST_NUM_TRIES): | |
| out = run_sampling_generator(mesh_device, args, logits, params, num_steps=1, advance_seeds=True)[0] | |
| results.append(out[0]) | |
| assert len(set(results)) >= 2, f"Unseeded single slot should vary across requests, got {results}" | |
| def test_uniform_seed_deterministic(self, mesh_device, seed): | |
| args = make_sampling_args(mesh_device) | |
| logits = build_hot_logits(args, hot_tokens=[1500, 1501, 1502, 1503, 1504, 1505, 1506, 1507]) | |
| params = per_lane_params(1.0, 8, 1.0) | |
| seeds = [seed] * BATCH_SIZE | |
| out1 = run_sampling_generator( | |
| mesh_device, args, logits, params, num_steps=1, advance_seeds=True, seed_values=seeds | |
| )[0] | |
| out2 = run_sampling_generator( | |
| mesh_device, args, logits, params, num_steps=1, advance_seeds=True, seed_values=seeds | |
| )[0] | |
| assert out1 == out2, f"Uniform seed {seed} should be deterministic across runs" | |
| def test_uniform_noseed_varied(self, mesh_device): | |
| args = make_sampling_args(mesh_device) | |
| logits = build_hot_logits(args, hot_tokens=[1600, 1601, 1602, 1603, 1604, 1605, 1606, 1607]) | |
| params = SamplingParams(temperature=2.0, top_k=8, top_p=1.0) | |
| outputs = run_sampling_generator( | |
| mesh_device, args, logits, params, num_steps=FAST_NUM_STEPS, advance_seeds=True | |
| ) | |
| user0 = [step[0] for step in outputs] | |
| assert len(set(user0)) >= 2, f"Expected unseeded variation across steps, got {user0}" | |
| def test_seed_0_produces_deterministic_outputs(self, mesh_device): | |
| args = make_sampling_args(mesh_device) | |
| logits = build_hot_logits(args, hot_tokens=[1700, 1701, 1702, 1703, 1704, 1705, 1706, 1707]) | |
| params = SamplingParams(temperature=1.0, top_k=8, top_p=1.0) | |
| seeds = [0] * BATCH_SIZE | |
| out1 = run_sampling_generator( | |
| mesh_device, args, logits, params, num_steps=1, advance_seeds=True, seed_values=seeds | |
| )[0] | |
| out2 = run_sampling_generator( | |
| mesh_device, args, logits, params, num_steps=1, advance_seeds=True, seed_values=seeds | |
| )[0] | |
| assert out1 == out2, "Seed 0 must be deterministic across requests" | |
| def test_negative_seed_does_not_crash(self, mesh_device): | |
| args = make_sampling_args(mesh_device) | |
| logits = build_hot_logits(args, hot_tokens=[1800, 1801, 1802, 1803, 1804, 1805, 1806, 1807]) | |
| params = SamplingParams(temperature=1.0, top_k=8, top_p=1.0) | |
| seeds = [-1] + [2000 + i for i in range(1, BATCH_SIZE)] | |
| out = run_sampling_generator( | |
| mesh_device, | |
| args, | |
| logits, | |
| params, | |
| num_steps=1, | |
| advance_seeds=False, | |
| seed_values=seeds, | |
| write_seed_values_to_device=True, | |
| )[0] | |
| assert_tokens_in_vocab(out, args.vocab_size) | |
| # --- Test: batch isolation --- | |
| class TestBatchIsolation: | |
| def test_mixed_params_batch(self, mesh_device, device_params): | |
| args = make_sampling_args(mesh_device) | |
| logits = build_hot_logits(args, hot_tokens=[2000, 2001, 2002, 2003, 2004, 2005, 2006, 2007]) | |
| temperature_template = [0.0, 1.5, 0.5, 0.5, 1.0, 0.01, 2.0, 0.7] | |
| top_k_template = [32, 8, 8, 8, 1, 32, 5, 8] | |
| top_p_template = [1.0, 1.0, 1.0, 0.7, 1.0, 1.0, 1.0, 0.9] | |
| repetition_template = [1.0, 1.0, 2.0, 1.0, 1.0, 1.0, 1.0, 1.5] | |
| presence_template = [0.0, 0.0, 0.0, 2.0, 0.0, 0.0, 0.0, 1.0] | |
| frequency_template = [0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 2.0] | |
| temperature = (temperature_template * 4)[:BATCH_SIZE] | |
| top_k = (top_k_template * 4)[:BATCH_SIZE] | |
| top_p = (top_p_template * 4)[:BATCH_SIZE] | |
| repetition = (repetition_template * 4)[:BATCH_SIZE] | |
| presence = (presence_template * 4)[:BATCH_SIZE] | |
| frequency = (frequency_template * 4)[:BATCH_SIZE] | |
| seeds = [9000 + i for i in range(BATCH_SIZE)] | |
| params = SamplingParams( | |
| temperature=temperature, | |
| top_k=top_k, | |
| top_p=top_p, | |
| repetition_penalty=repetition, | |
| presence_penalty=presence, | |
| frequency_penalty=frequency, | |
| ) | |
| out1 = run_sampling_generator( | |
| mesh_device, | |
| args, | |
| logits, | |
| params, | |
| num_steps=1, | |
| advance_seeds=True, | |
| seed_values=seeds, | |
| state_setup=self._state_setup, | |
| )[0] | |
| out2 = run_sampling_generator( | |
| mesh_device, | |
| args, | |
| logits, | |
| params, | |
| num_steps=1, | |
| advance_seeds=True, | |
| seed_values=seeds, | |
| state_setup=self._state_setup, | |
| )[0] | |
| assert out1 == out2, "Deterministic lanes (and seeded lanes) should replay identically across runs" | |
| assert len(set(out1)) >= 2, "Mixed batch should not collapse to a single token" | |
| def test_outputs_not_mixed_different_prompts(self, mesh_device, device_params): | |
| args = make_sampling_args(mesh_device) | |
| tokens_per_user = 4 | |
| per_user_hot = [] | |
| expected_sets = [] | |
| token_cursor = 2200 | |
| for _ in range(BATCH_SIZE): | |
| hot = [token_cursor + i for i in range(tokens_per_user)] | |
| per_user_hot.append(hot) | |
| expected_sets.append(set(hot)) | |
| token_cursor += tokens_per_user | |
| logits = build_hot_logits(args, per_user_hot_tokens=per_user_hot) | |
| params = SamplingParams(temperature=1.0, top_k=tokens_per_user, top_p=1.0) | |
| seeds = [10000 + i for i in range(BATCH_SIZE)] | |
| tokens = run_sampling_generator( | |
| mesh_device, args, logits, params, num_steps=1, advance_seeds=True, seed_values=seeds | |
| )[0] | |
| for i, tok in enumerate(tokens): | |
| assert ( | |
| tok in expected_sets[i] | |
| ), f"User {i} token leaked across users: tok={tok}, expected={expected_sets[i]}" | |
| for device_idx in representative_device_indices(mesh_device)[1:]: | |
| device_tokens = run_sampling_generator( | |
| mesh_device, | |
| args, | |
| logits, | |
| params, | |
| num_steps=1, | |
| advance_seeds=True, | |
| seed_values=seeds, | |
| device_idx=device_idx, | |
| )[0] | |
| assert device_tokens == tokens, f"Device view mismatch for device_idx={device_idx} in batch isolation test" | |
| def test_same_prompt_users_get_identical_logits(self, mesh_device, device_params): | |
| """Identical prompts must agree under greedy, and vary under distinct seeds. | |
| Split out of the uniform-seed sub-case below (#38316): these two invariants have | |
| nothing to do with per-lane seeding and must keep gating CI. | |
| """ | |
| args = make_sampling_args(mesh_device) | |
| # All users share the exact same hot-token distribution. | |
| hot_tokens = [2500, 2501, 2502, 2503, 2504, 2505, 2506, 2507] | |
| logits = build_hot_logits(args, hot_tokens=hot_tokens) | |
| # --- Greedy: every user should pick the top-logit token. --- | |
| greedy_params = SamplingParams( | |
| temperature=[0.0] * BATCH_SIZE, | |
| top_k=[32] * BATCH_SIZE, | |
| top_p=[1.0] * BATCH_SIZE, | |
| ) | |
| greedy_tokens = run_sampling_generator( | |
| mesh_device, args, logits, greedy_params, num_steps=1, advance_seeds=False | |
| )[0] | |
| assert ( | |
| len(set(greedy_tokens)) == 1 | |
| ), f"All users with the same prompt under greedy should pick the same token, got {greedy_tokens}" | |
| assert ( | |
| greedy_tokens[0] == hot_tokens[0] | |
| ), f"Greedy should pick the highest-logit token {hot_tokens[0]}, got {greedy_tokens[0]}" | |
| # --- Stochastic with different seeds: should see variation. --- | |
| # Use full-length lists (like the greedy sub-case above) so temperature applies | |
| # uniformly to every lane. A scalar temperature would only configure lane 0 and | |
| # leave lanes 1..31 on the greedy default. | |
| stochastic_params = SamplingParams( | |
| temperature=[1.5] * BATCH_SIZE, | |
| top_k=[8] * BATCH_SIZE, | |
| top_p=[1.0] * BATCH_SIZE, | |
| ) | |
| diverse_seeds = [5000 + i for i in range(BATCH_SIZE)] | |
| out_diverse = run_sampling_generator( | |
| mesh_device, | |
| args, | |
| logits, | |
| stochastic_params, | |
| num_steps=FAST_NUM_TRIES, | |
| advance_seeds=True, | |
| seed_values=diverse_seeds, | |
| ) | |
| all_tokens = flatten_steps(out_diverse) | |
| assert ( | |
| len(set(all_tokens)) >= 2 | |
| ), f"Different seeds on the same prompt should produce variation, got {set(all_tokens)}" | |
| assert_tokens_in_vocab(all_tokens, args.vocab_size) | |
| unexpected = [tok for tok in all_tokens if tok not in set(hot_tokens)] | |
| assert not unexpected, f"Sampled tokens outside expected hot set: {unexpected}" | |
| def test_uniform_seed_diverges_across_lanes(self, mesh_device, device_params): | |
| """One request seed shared by every lane must still give each lane its own draw. | |
| This asserts the SEED SALT, not agreement. SeedManager._set_slot_seed assigns | |
| seed_salts[slot] = _next_free_salt(slot, seed): the first slot holding a given seed | |
| gets salt 0 and each later duplicate gets the next free value, and the salt is mixed | |
| into _hash_request_seed_to_device_seed. Without it, n>1 completions of one prompt at | |
| a fixed seed came out byte-identical (#53077). So a uniform seed vector is exactly the | |
| case the salt exists to separate, and identical tokens across lanes would mean the | |
| salt is not being applied. | |
| A request whose seed is unique among live slots always lands on salt 0, so | |
| single-seeded reproducibility (test_specific_seed_reproducible, | |
| test_seeded_replay_is_deterministic) is unaffected. | |
| Deliberately paired with | |
| test_sampling.py::test_ttsampling_duplicate_request_seeds_sample_diverse_tokens, which | |
| asserts the same invariant one level down on raw TTSampling + SeedManager. This one | |
| goes through SamplingGenerator and format_sampling_params, so it also covers the | |
| wiring between them. Change one, check the other. | |
| """ | |
| args = make_sampling_args(mesh_device) | |
| hot_tokens = [2500, 2501, 2502, 2503, 2504, 2505, 2506, 2507] | |
| hot_token_set = set(hot_tokens) | |
| logits = build_hot_logits(args, hot_tokens=hot_tokens) | |
| stochastic_params = SamplingParams( | |
| temperature=[1.5] * BATCH_SIZE, | |
| top_k=[8] * BATCH_SIZE, | |
| top_p=[1.0] * BATCH_SIZE, | |
| ) | |
| uniform_seed = [7777] * BATCH_SIZE | |
| out_uniform = run_sampling_generator( | |
| mesh_device, | |
| args, | |
| logits, | |
| stochastic_params, | |
| num_steps=1, | |
| advance_seeds=True, | |
| seed_values=uniform_seed, | |
| )[0] | |
| assert len(set(out_uniform)) > 1, ( | |
| "Every lane sharing seed 7777 produced the same token, so the per-slot seed salt " | |
| f"is not reaching the device: {out_uniform}. This is the #53077 failure mode." | |
| ) | |
| assert_tokens_in_vocab(out_uniform, args.vocab_size) | |
| unexpected = [tok for tok in out_uniform if tok not in hot_token_set] | |
| assert not unexpected, f"Sampled tokens outside expected hot set: {unexpected}" | |
| # Replay must still be exact: the salt is derived from the live slot layout, which is | |
| # identical for a fresh generator given the same seed vector. | |
| out_again = run_sampling_generator( | |
| mesh_device, | |
| args, | |
| logits, | |
| stochastic_params, | |
| num_steps=1, | |
| advance_seeds=True, | |
| seed_values=uniform_seed, | |
| )[0] | |
| assert out_uniform == out_again, ( | |
| "Salted per-slot seeds must still replay exactly for the same seed vector. " | |
| f"first={out_uniform}, second={out_again}" | |
| ) | |
| def _state_setup(self, sg): | |
| seen = torch.full((BATCH_SIZE, 1), 2000, dtype=torch.int64) | |
| sg.reset_prompt_tokens(seen) | |
| sg.reset_output_state(tokens=seen) | |
| # --- Test: batch size variations --- | |
| class TestBatchSizeVariations: | |
| def _make_active_user_logits(self, args: _SamplingArgs, active_users: int): | |
| per_user_hot = [] | |
| cursor = 2800 | |
| for i in range(BATCH_SIZE): | |
| if i < active_users: | |
| hot = [cursor, cursor + 1, cursor + 2] | |
| cursor += 3 | |
| else: | |
| hot = [42] | |
| per_user_hot.append(hot) | |
| return build_hot_logits(args, per_user_hot_tokens=per_user_hot), per_user_hot | |
| def test_small_batch_different_params(self, mesh_device): | |
| args = make_sampling_args(mesh_device) | |
| logits, per_user_hot = self._make_active_user_logits(args, active_users=2) | |
| params = SamplingParams( | |
| temperature=[0.0, 1.0], | |
| top_k=[1, 3], | |
| top_p=[1.0, 1.0], | |
| ) | |
| tokens = run_sampling_generator(mesh_device, args, logits, params, num_steps=1, advance_seeds=True)[0] | |
| assert tokens[0] in set(per_user_hot[0]) | |
| assert tokens[1] in set(per_user_hot[1]) | |
| assert_tokens_in_vocab(tokens, args.vocab_size) | |
| def test_full_batch_different_params(self, mesh_device): | |
| args = make_sampling_args(mesh_device) | |
| logits, per_user_hot = self._make_active_user_logits(args, active_users=BATCH_SIZE) | |
| temperature = [0.5 + (i * 0.02) for i in range(BATCH_SIZE)] | |
| params = SamplingParams(temperature=temperature, top_k=[3] * BATCH_SIZE, top_p=[1.0] * BATCH_SIZE) | |
| tokens = run_sampling_generator(mesh_device, args, logits, params, num_steps=1, advance_seeds=True)[0] | |
| for i, tok in enumerate(tokens): | |
| assert tok in set(per_user_hot[i]), f"User {i} expected token from {per_user_hot[i]}, got {tok}" | |
| def test_partial_batch_different_params(self, mesh_device): | |
| args = make_sampling_args(mesh_device) | |
| active_users = BATCH_SIZE // 2 | |
| logits, per_user_hot = self._make_active_user_logits(args, active_users=active_users) | |
| temperature = [0.0 if i % 2 == 0 else 1.0 for i in range(active_users)] | |
| top_k = [1 if i % 2 == 0 else 3 for i in range(active_users)] | |
| params = SamplingParams(temperature=temperature, top_k=top_k, top_p=[1.0] * active_users) | |
| tokens = run_sampling_generator(mesh_device, args, logits, params, num_steps=1, advance_seeds=True)[0] | |
| for i in range(active_users): | |
| assert tokens[i] in set(per_user_hot[i]), f"Active user {i} expected token from {per_user_hot[i]}" | |
| assert_tokens_in_vocab(tokens, args.vocab_size) | |
| # --- Test: mixed-parameter batches --- | |
| class TestMixedParameterBatches: | |
| def test_all_parameter_types_in_batch(self, mesh_device): | |
| args = make_sampling_args(mesh_device) | |
| logits = build_hot_logits(args, hot_tokens=[3000, 3001, 3002, 3003, 3004, 3005, 3006, 3007]) | |
| temperature_template = [0.0, 1.5, 1.0, 0.5, 0.5, 0.5, 0.5, 1.0] | |
| top_k_template = [32, 8, 10, 8, 8, 8, 8, 8] | |
| top_p_template = [1.0, 1.0, 1.0, 1.0, 1.0, 1.0, 0.7, 0.5] | |
| repetition_template = [1.0, 1.0, 1.0, 3.0, 1.0, 1.0, 1.5, 1.0] | |
| presence_template = [0.0, 0.0, 0.0, 0.0, 2.0, 0.0, 1.0, 0.0] | |
| frequency_template = [0.0, 0.0, 0.0, 0.0, 0.0, 2.0, 1.0, 0.0] | |
| temperature = (temperature_template * 4)[:BATCH_SIZE] | |
| top_k = (top_k_template * 4)[:BATCH_SIZE] | |
| top_p = (top_p_template * 4)[:BATCH_SIZE] | |
| repetition = (repetition_template * 4)[:BATCH_SIZE] | |
| presence = (presence_template * 4)[:BATCH_SIZE] | |
| frequency = (frequency_template * 4)[:BATCH_SIZE] | |
| seeds = [12000 + i for i in range(BATCH_SIZE)] | |
| def _state_setup(sg): | |
| seen = torch.full((BATCH_SIZE, 1), 3000, dtype=torch.int64) | |
| sg.reset_prompt_tokens(seen) | |
| sg.reset_output_state(tokens=seen) | |
| params = SamplingParams( | |
| temperature=temperature, | |
| top_k=top_k, | |
| top_p=top_p, | |
| repetition_penalty=repetition, | |
| presence_penalty=presence, | |
| frequency_penalty=frequency, | |
| ) | |
| out1 = run_sampling_generator( | |
| mesh_device, | |
| args, | |
| logits, | |
| params, | |
| num_steps=1, | |
| advance_seeds=True, | |
| seed_values=seeds, | |
| state_setup=_state_setup, | |
| )[0] | |
| out2 = run_sampling_generator( | |
| mesh_device, | |
| args, | |
| logits, | |
| params, | |
| num_steps=1, | |
| advance_seeds=True, | |
| seed_values=seeds, | |
| state_setup=_state_setup, | |
| )[0] | |
| assert out1 == out2, "Mixed parameter batch should replay deterministically with fixed seeds" | |
| assert_tokens_in_vocab(out1, args.vocab_size) | |
| assert len(set(out1)) >= 2, "Mixed parameter batch should produce non-trivial diversity" | |
| # --- Test: a single odd lane out (1 greedy vs 31 stochastic, and the mirror) --- | |
| class TestSingleGreedyLaneInStochasticBatch: | |
| """One greedy user in an otherwise stochastic batch, and the mirror case. | |
| This is the configuration that stresses per-user sampling hardest, and it is | |
| not covered by uniformly-greedy or uniformly-stochastic batches: | |
| * The force-argmax fast path cannot fire, so the top-k pipeline has to honour | |
| a k=1 lane sitting next to 31 lanes that each draw from their own RNG | |
| stream. Every k/p/temp value is indexed per core, and the k=1 lane takes | |
| the ``k <= FACE_WIDTH`` branch of the writer kernel's candidate walk while | |
| its neighbours do not. | |
| * A batch that is *nearly* uniform is exactly where an "any lane" predicate | |
| passes for a "all lanes" rule -- see | |
| ``test_force_argmax_needs_every_lane_greedy``. | |
| Every user gets its own disjoint hot-token band so cross-lane leakage shows | |
| up as a concrete per-user failure rather than a diversity statistic. | |
| """ | |
| BASE_TOKEN = 4000 | |
| def _band_params(self, odd_lane, *, base_top_k, base_temperature, odd_top_k, odd_temperature, **overrides): | |
| """Uniform params for every lane, overridden on ``odd_lane`` only.""" | |
| temperature = [base_temperature] * BATCH_SIZE | |
| top_k = [base_top_k] * BATCH_SIZE | |
| if odd_lane is not None: | |
| temperature[odd_lane] = odd_temperature | |
| top_k[odd_lane] = odd_top_k | |
| return SamplingParams(temperature=temperature, top_k=top_k, top_p=[1.0] * BATCH_SIZE, **overrides) | |
| def _greedy_lane_params(self, greedy_lane, *, temperature=1.0): | |
| """31 stochastic lanes plus (optionally) one greedy lane.""" | |
| return self._band_params( | |
| greedy_lane, | |
| base_top_k=BAND_TOKENS_PER_USER, | |
| base_temperature=temperature, | |
| odd_top_k=1, | |
| odd_temperature=0.0, | |
| ) | |
| def _assert_no_band_leakage(self, tokens, bands, *, context=""): | |
| for user, tok in enumerate(tokens): | |
| assert tok in set(bands[user]), ( | |
| f"User {user} sampled token {tok} outside its own band " | |
| f"[{bands[user][0]}..{bands[user][-1]}]{context}. tokens={tokens}" | |
| ) | |
| def test_one_greedy_lane_rest_stochastic(self, mesh_device, device_params, greedy_lane): | |
| args = make_sampling_args(mesh_device) | |
| logits, bands = build_disjoint_band_logits(args, base_token=self.BASE_TOKEN) | |
| params = self._greedy_lane_params(greedy_lane) | |
| seeds = [40000 + i for i in range(BATCH_SIZE)] | |
| outputs = run_sampling_generator( | |
| mesh_device, | |
| args, | |
| logits, | |
| params, | |
| num_steps=FAST_NUM_STEPS, | |
| advance_seeds=True, | |
| seed_values=seeds, | |
| ) | |
| expected_greedy = bands[greedy_lane][0] | |
| for step, tokens in enumerate(outputs): | |
| assert_tokens_in_vocab(tokens, args.vocab_size) | |
| self._assert_no_band_leakage(tokens, bands, context=f" at step {step}") | |
| assert tokens[greedy_lane] == expected_greedy, ( | |
| f"Greedy lane {greedy_lane} must pick its max-logit token {expected_greedy} " | |
| f"at every step, got {tokens[greedy_lane]} at step {step}. tokens={tokens}" | |
| ) | |
| # The 31 stochastic lanes must still be sampling. If the k=1 lane dragged | |
| # the batch into argmax, every lane would return its own band rank 0. | |
| off_top_picks = sum( | |
| 1 for tokens in outputs for user, tok in enumerate(tokens) if user != greedy_lane and tok != bands[user][0] | |
| ) | |
| assert off_top_picks > 0, ( | |
| f"Stochastic lanes never left their band's top token; the greedy lane " | |
| f"{greedy_lane} appears to have forced argmax on the whole batch. outputs={outputs}" | |
| ) | |
| # Same seeds must replay bit-exactly, and every device view must agree. | |
| replay = run_sampling_generator( | |
| mesh_device, | |
| args, | |
| logits, | |
| params, | |
| num_steps=FAST_NUM_STEPS, | |
| advance_seeds=True, | |
| seed_values=seeds, | |
| ) | |
| assert replay == outputs, "Mixed greedy/stochastic batch must replay exactly for fixed seeds" | |
| for device_idx in representative_device_indices(mesh_device)[1:]: | |
| out_device = run_sampling_generator( | |
| mesh_device, | |
| args, | |
| logits, | |
| params, | |
| num_steps=FAST_NUM_STEPS, | |
| advance_seeds=True, | |
| seed_values=seeds, | |
| device_idx=device_idx, | |
| ) | |
| assert out_device == outputs, f"Device view mismatch for device_idx={device_idx}" | |
| def test_greedy_lane_ignores_seeds(self, mesh_device, greedy_lane): | |
| """The greedy lane must be seed-independent while its neighbours are not.""" | |
| args = make_sampling_args(mesh_device) | |
| logits, bands = build_disjoint_band_logits(args, base_token=self.BASE_TOKEN) | |
| params = self._greedy_lane_params(greedy_lane) | |
| seeds_a = [41000 + i for i in range(BATCH_SIZE)] | |
| seeds_b = [s + 987654 for s in seeds_a] | |
| out_a = run_sampling_generator( | |
| mesh_device, args, logits, params, num_steps=1, advance_seeds=True, seed_values=seeds_a | |
| )[0] | |
| out_b = run_sampling_generator( | |
| mesh_device, args, logits, params, num_steps=1, advance_seeds=True, seed_values=seeds_b | |
| )[0] | |
| expected_greedy = bands[greedy_lane][0] | |
| for label, tokens in (("seeds_a", out_a), ("seeds_b", out_b)): | |
| self._assert_no_band_leakage(tokens, bands, context=f" ({label})") | |
| assert tokens[greedy_lane] == expected_greedy, ( | |
| f"Greedy lane {greedy_lane} changed with the seed vector ({label}): " | |
| f"expected {expected_greedy}, got {tokens[greedy_lane]}" | |
| ) | |
| changed = [i for i in range(BATCH_SIZE) if i != greedy_lane and out_a[i] != out_b[i]] | |
| assert len(changed) >= 2, ( | |
| "Shifting the seed vector should move several stochastic lanes; only " | |
| f"{changed} changed. out_a={out_a}, out_b={out_b}" | |
| ) | |
| def test_greedy_lane_does_not_perturb_other_lanes(self, mesh_device, greedy_lane): | |
| """Flipping one lane to k=1 must leave every other lane's token untouched. | |
| Each user is sampled on its own core with its own k/p/temp and its own | |
| rand tile, so lane ``greedy_lane`` going greedy is invisible to the other | |
| 31 lanes. Any difference means per-user state is being shared. | |
| """ | |
| args = make_sampling_args(mesh_device) | |
| logits, bands = build_disjoint_band_logits(args, base_token=self.BASE_TOKEN) | |
| seeds = [42000 + i for i in range(BATCH_SIZE)] | |
| # Only k differs between the two runs: temperature=0.0 is normalised to an | |
| # inverse temperature of 1.0, which is what temperature=1.0 maps to too. | |
| all_stochastic = self._greedy_lane_params(None) | |
| one_greedy = self._greedy_lane_params(greedy_lane) | |
| out_reference = run_sampling_generator( | |
| mesh_device, args, logits, all_stochastic, num_steps=1, advance_seeds=True, seed_values=seeds | |
| )[0] | |
| out_mixed = run_sampling_generator( | |
| mesh_device, args, logits, one_greedy, num_steps=1, advance_seeds=True, seed_values=seeds | |
| )[0] | |
| self._assert_no_band_leakage(out_reference, bands, context=" (all stochastic)") | |
| self._assert_no_band_leakage(out_mixed, bands, context=" (one greedy)") | |
| perturbed = [ | |
| (user, out_reference[user], out_mixed[user]) | |
| for user in range(BATCH_SIZE) | |
| if user != greedy_lane and out_reference[user] != out_mixed[user] | |
| ] | |
| assert not perturbed, ( | |
| f"Making lane {greedy_lane} greedy changed other lanes (user, expected, got): {perturbed}. " | |
| f"reference={out_reference}, mixed={out_mixed}" | |
| ) | |
| assert ( | |
| out_mixed[greedy_lane] == bands[greedy_lane][0] | |
| ), f"Greedy lane {greedy_lane} should pick {bands[greedy_lane][0]}, got {out_mixed[greedy_lane]}" | |
| def test_one_stochastic_lane_rest_greedy(self, mesh_device, stochastic_lane): | |
| """Mirror case: 31 greedy lanes must not silence the one sampling lane.""" | |
| args = make_sampling_args(mesh_device) | |
| logits, bands = build_disjoint_band_logits(args, base_token=self.BASE_TOKEN) | |
| params = self._band_params( | |
| stochastic_lane, | |
| base_top_k=1, | |
| base_temperature=0.0, | |
| odd_top_k=BAND_TOKENS_PER_USER, | |
| odd_temperature=1.0, | |
| ) | |
| # Enough steps that a genuinely stochastic lane repeating one token is | |
| # far less likely than any other cause of failure. | |
| num_steps = 2 * FAST_NUM_TRIES | |
| seeds = [43000 + i for i in range(BATCH_SIZE)] | |
| outputs = run_sampling_generator( | |
| mesh_device, args, logits, params, num_steps=num_steps, advance_seeds=True, seed_values=seeds | |
| ) | |
| for step, tokens in enumerate(outputs): | |
| assert_tokens_in_vocab(tokens, args.vocab_size) | |
| self._assert_no_band_leakage(tokens, bands, context=f" at step {step}") | |
| for user in range(BATCH_SIZE): | |
| if user == stochastic_lane: | |
| continue | |
| assert tokens[user] == bands[user][0], ( | |
| f"Greedy lane {user} must pick {bands[user][0]} at every step, " | |
| f"got {tokens[user]} at step {step}" | |
| ) | |
| sampled = [step[stochastic_lane] for step in outputs] | |
| assert len(set(sampled)) >= 2, ( | |
| f"Lane {stochastic_lane} was the only stochastic lane and never varied over " | |
| f"{num_steps} steps; the greedy majority appears to have forced it to argmax. " | |
| f"sampled={sampled}" | |
| ) | |
| def test_one_greedy_lane_with_penalties(self, mesh_device, greedy_lane): | |
| """A penalty on the greedy lane must retarget only that lane. | |
| The greedy lane's top band token is penalised below its runner-up, so the | |
| expected token is exact; the other 31 lanes carry no penalty and must be | |
| untouched by the penalty state written for their neighbour. | |
| """ | |
| args = make_sampling_args(mesh_device) | |
| logits, bands = build_disjoint_band_logits(args, base_token=self.BASE_TOKEN) | |
| # top_logit=10.0 with step=0.25: dividing rank 0 by 4.0 gives 2.5, well | |
| # below rank 1 at 9.75, so the greedy pick moves by exactly one rank. | |
| repetition = [1.0] * BATCH_SIZE | |
| repetition[greedy_lane] = 4.0 | |
| params = self._band_params( | |
| greedy_lane, | |
| base_top_k=BAND_TOKENS_PER_USER, | |
| base_temperature=1.0, | |
| odd_top_k=1, | |
| odd_temperature=0.0, | |
| repetition_penalty=repetition, | |
| ) | |
| def _state_setup(sg): | |
| # -1 marks "no token seen", so only the greedy lane carries history. | |
| seen = torch.full((BATCH_SIZE, 1), -1, dtype=torch.int64) | |
| seen[greedy_lane, 0] = bands[greedy_lane][0] | |
| sg.reset_output_state(tokens=seen) | |
| seeds = [44000 + i for i in range(BATCH_SIZE)] | |
| tokens = run_sampling_generator( | |
| mesh_device, | |
| args, | |
| logits, | |
| params, | |
| num_steps=1, | |
| advance_seeds=True, | |
| seed_values=seeds, | |
| state_setup=_state_setup, | |
| )[0] | |
| assert_tokens_in_vocab(tokens, args.vocab_size) | |
| self._assert_no_band_leakage(tokens, bands) | |
| assert tokens[greedy_lane] == bands[greedy_lane][1], ( | |
| f"Penalised greedy lane {greedy_lane} should fall back to its runner-up " | |
| f"{bands[greedy_lane][1]}, got {tokens[greedy_lane]}. tokens={tokens}" | |
| ) | |
| def test_force_argmax_needs_every_lane_greedy(self, mesh_device, odd_lane): | |
| """One non-greedy lane must keep the whole batch off the argmax fast path. | |
| ``TTSampling._is_force_argmax_sampling`` is an all-lanes predicate. If it | |
| ever degraded to "any lane", a single greedy user would silently force | |
| argmax on the 31 users who asked to sample -- and because the fast path | |
| skips the top-k/top-p/RNG pipeline entirely, nothing downstream would | |
| notice. This asserts the predicate directly, with the fast path enabled | |
| in ``model_config`` (the other tests leave it off). | |
| """ | |
| args = make_sampling_args(mesh_device) | |
| args.model_config = { | |
| "SAMPLING_AG_CONFIG": { | |
| "allow_force_argmax": True, | |
| "num_links": 1, | |
| "topology": ttnn.Topology.Linear, | |
| } | |
| } | |
| logits, _ = build_disjoint_band_logits(args, base_token=self.BASE_TOKEN) | |
| observed = {} | |
| def _probe(label): | |
| def _record(sg): | |
| observed[label] = sg.tt_sampling.force_argmax_sampling | |
| return _record | |
| # num_steps=0 runs the host-side param plumbing without any sampling. | |
| all_greedy = self._band_params(None, base_top_k=1, base_temperature=0.0, odd_top_k=1, odd_temperature=0.0) | |
| run_sampling_generator(mesh_device, args, logits, all_greedy, num_steps=0, state_setup=_probe("all_greedy")) | |
| # temperature=1.0 also normalises to an inverse temperature of 1.0, so | |
| # only top_k distinguishes this lane -- the k branch of the predicate. | |
| one_stochastic = self._band_params( | |
| odd_lane, | |
| base_top_k=1, | |
| base_temperature=0.0, | |
| odd_top_k=BAND_TOKENS_PER_USER, | |
| odd_temperature=1.0, | |
| ) | |
| run_sampling_generator( | |
| mesh_device, args, logits, one_stochastic, num_steps=0, state_setup=_probe("one_stochastic") | |
| ) | |
| assert observed["all_greedy"] is True, "A fully greedy batch should take the force-argmax fast path" | |
| assert observed["one_stochastic"] is False, ( | |
| f"Lane {odd_lane} requested top_k={BAND_TOKENS_PER_USER} but the batch still took the " | |
| "force-argmax fast path, which skips top-k/top-p/RNG entirely" | |
| ) | |
| # --- Test: traced sampling path --- | |
| # Production decode calls SamplingGenerator.sample(..., enable_trace=True) -- that is the | |
| # signature default. Every other test in this file runs enable_trace=False, so | |
| # _trace_slot / capture_trace / _execute_trace / _validate_trace_inputs are otherwise | |
| # never executed. Two things constrain what a traced test can look like: | |
| # * sample() computes `use_internal_trace = enable_trace and not | |
| # seed_manager.has_active_request_seed()`, so a test with explicit request seeds | |
| # silently falls back to the eager path and proves nothing. These use unseeded runs. | |
| # * trace replay binds the captured input address, so the same logits tensor object must | |
| # be reused on every step (run_sampling_generator does this when enable_trace=True) and | |
| # penalties -- which rewrite that tensor in place -- cannot be combined with it. | |
| # A trace region has to be reserved at device open, hence the device_params override; the | |
| # rest of the suite opens with none ("No trace region size for 1" in the CI log). | |
| _TRACE_DEVICE_PARAMS = [{"trace_region_size": 23887872}] | |
| class TestTracedSampling: | |
| def test_traced_greedy_matches_untraced(self, mesh_device, device_params): | |
| """Capture-then-replay must return exactly what the eager path returns. | |
| Greedy (temperature=0.0, top_k=1) makes the expected token independent of RNG state, | |
| so the traced result is directly comparable to the untraced one across replays -- | |
| step 1 captures, steps 2 and 3 go through _execute_trace. | |
| """ | |
| args = make_sampling_args(mesh_device) | |
| hot_tokens = [3300, 3301, 3302, 3303] | |
| logits = build_hot_logits(args, hot_tokens=hot_tokens) | |
| params = SamplingParams( | |
| temperature=[0.0] * BATCH_SIZE, | |
| top_k=[1] * BATCH_SIZE, | |
| top_p=[1.0] * BATCH_SIZE, | |
| ) | |
| untraced = run_sampling_generator( | |
| mesh_device, args, logits, params, num_steps=1, advance_seeds=True, enable_trace=False | |
| )[0] | |
| traced = run_sampling_generator( | |
| mesh_device, args, logits, params, num_steps=3, advance_seeds=True, enable_trace=True | |
| ) | |
| assert all( | |
| tok == hot_tokens[0] for tok in untraced | |
| ), f"Eager greedy should pick the max-logit token {hot_tokens[0]}, got {untraced}" | |
| for step, tokens in enumerate(traced): | |
| assert tokens == untraced, ( | |
| f"Traced step {step} diverged from the eager result. " f"traced={tokens}, untraced={untraced}" | |
| ) | |
| def test_trace_rejects_mismatched_logits_tensor(self, mesh_device, device_params, expect_error): | |
| """_validate_trace_inputs must reject a logits tensor the trace was not captured on. | |
| Replay reuses the captured input buffer, so silently accepting a different tensor | |
| would sample stale logits. Exercised directly because run_sampling_generator always | |
| reuses one tensor by construction. | |
| """ | |
| args = make_sampling_args(mesh_device) | |
| logits = build_hot_logits(args, hot_tokens=[3400, 3401, 3402]) | |
| padded = pad_logits_to_max_batch(logits, max_batch_size=BATCH_SIZE) | |
| params = SamplingParams( | |
| temperature=[0.0] * BATCH_SIZE, | |
| top_k=[1] * BATCH_SIZE, | |
| top_p=[1.0] * BATCH_SIZE, | |
| ) | |
| sg = None | |
| captured_input = None | |
| other_input = None | |
| try: | |
| sg = SamplingGenerator(args=args, mesh_device=mesh_device, tt_ccl=None) | |
| sg.reset_sampling_params(format_sampling_params(params, BATCH_SIZE)) | |
| # Unseeded, or sample() would bypass the trace entirely. | |
| sg.seed_manager.reset_seed(None, list(range(BATCH_SIZE))) | |
| sg.seed_manager.get_new_values() | |
| captured_input = make_sharded_logits(padded, mesh_device, args) | |
| sg.sample(captured_input, enable_trace=True) # captures | |
| ttnn.synchronize_device(mesh_device) | |
| other_input = make_sharded_logits(padded, mesh_device, args) | |
| with expect_error(ValueError, r"does not match the tensor used during trace capture"): | |
| sg.sample(other_input, enable_trace=True) | |
| finally: | |
| if sg is not None: | |
| try: | |
| sg.reset_trace() | |
| except Exception: | |
| pass | |
| for tensor in (captured_input, other_input): | |
| if tensor is not None: | |
| del tensor | |
| if sg is not None: | |
| del sg | |
| safe_sync(mesh_device) | |
| # --- Test: format_sampling_params lane semantics (host-only, no device) --- | |
| # These need no device and no mesh_device fixture. They belong beside the other host-only | |
| # sampling tests in models/common/tests/test_sampling.py, but that file is not wired into | |
| # any pipeline today, so they live here to actually run. Move them when it is. (#38316) | |
| class TestFormatSamplingParamsLanes: | |
| """Pin the per-lane contract of format_sampling_params. | |
| A scalar companion field alongside a per-user temperature must reach every active lane. | |
| Getting this wrong is invisible: the padding defaults are top_k=1 (greedy), top_p=1.0 | |
| and penalty no-ops, so an under-broadcast field silently changes how real lanes sample | |
| while every assertion about token ranges still passes. | |
| """ | |
| def _fmt(**kwargs): | |
| return format_sampling_params(SamplingParams(**kwargs), BATCH_SIZE) | |
| def test_all_scalar_input_is_lane_zero_plus_greedy_padding(self): | |
| """The historical all-scalar shape must be untouched by the broadcast rule.""" | |
| out = self._fmt(temperature=0.5, top_k=8, top_p=0.9) | |
| assert out.temperature[0] == pytest.approx(1 / 0.5) # inverted for the device | |
| assert out.top_k[0] == 8 | |
| assert out.top_p[0] == pytest.approx(0.9) | |
| # Inactive lanes are padded to the greedy representation (temp 1.0 / k 1 / p 0.0). | |
| for lane in range(1, BATCH_SIZE): | |
| assert out.top_k[lane] == 1, f"lane {lane}" | |
| assert out.temperature[lane] == pytest.approx(1.0), f"lane {lane}" | |
| def test_scalar_companion_broadcasts_across_active_lanes(self): | |
| """A scalar top_k/top_p must not leave active lanes on the greedy default.""" | |
| active = 8 | |
| out = self._fmt(temperature=[0.8] * active, top_k=8, top_p=0.9) | |
| for lane in range(active): | |
| assert out.top_k[lane] == 8, f"active lane {lane} fell back to the k=1 default" | |
| assert out.top_p[lane] == pytest.approx(0.9), f"active lane {lane}" | |
| assert out.temperature[lane] == pytest.approx(1 / 0.8), f"active lane {lane}" | |
| for lane in range(active, BATCH_SIZE): | |
| assert out.top_k[lane] == 1, f"inactive lane {lane} should be greedy" | |
| def test_single_element_companion_stays_lane_scoped(self): | |
| """[x] keeps targeting lane 0; reinterpreting it as a broadcast would change | |
| sampling for existing callers' other lanes.""" | |
| out = self._fmt(temperature=[1.0, 1.0], top_k=[8], top_p=1.0) | |
| assert out.top_k[0] == 8 | |
| assert out.top_k[1] == 1, "a 1-element list must not broadcast onto lane 1" | |
| def test_mismatched_companion_length_is_rejected(self): | |
| """Silently padding the gap would turn real lanes greedy.""" | |
| try: | |
| self._fmt(temperature=[1.0] * 8, top_k=[3, 3, 3], top_p=1.0) | |
| except ValueError as exc: | |
| assert "top_k" in str(exc) and "8 active lanes" in str(exc) | |
| else: | |
| raise AssertionError("a top_k list shorter than the active lanes must be rejected") | |
| def test_scalar_penalty_reaches_every_active_lane(self): | |
| """Regression: a scalar penalty used to land on lane 0 only, leaving the rest on the | |
| no-op default with no diagnostic.""" | |
| active = 8 | |
| out = self._fmt( | |
| temperature=[1.0] * active, | |
| top_k=8, | |
| top_p=1.0, | |
| repetition_penalty=1.2, | |
| presence_penalty=0.5, | |
| frequency_penalty=0.25, | |
| ) | |
| for lane in range(active): | |
| assert out.repetition_penalty[lane] == pytest.approx(1.2), f"active lane {lane}" | |
| assert out.presence_penalty[lane] == pytest.approx(0.5), f"active lane {lane}" | |
| assert out.frequency_penalty[lane] == pytest.approx(0.25), f"active lane {lane}" | |
| for lane in range(active, BATCH_SIZE): | |
| assert out.repetition_penalty[lane] == pytest.approx(1.0), f"inactive lane {lane}" | |
| assert out.presence_penalty[lane] == pytest.approx(0.0), f"inactive lane {lane}" | |
| def test_unset_penalties_are_all_no_ops(self): | |
| """The SamplingParams penalty defaults are scalars equal to the padding defaults, so | |
| a caller who never sets them must be unaffected by the broadcast rule.""" | |
| out = self._fmt(temperature=[1.0] * 8, top_k=8, top_p=1.0) | |
| assert all(v == pytest.approx(1.0) for v in out.repetition_penalty) | |
| assert all(v == pytest.approx(0.0) for v in out.presence_penalty) | |
| assert all(v == pytest.approx(0.0) for v in out.frequency_penalty) | |
| def test_scalar_seed_stays_on_lane_zero(self): | |
| """seed is deliberately NOT broadcast: one seed on every lane means every lane draws | |
| the same token, which a caller has to ask for explicitly.""" | |
| out = self._fmt(temperature=[1.0] * 8, top_k=8, top_p=1.0, seed=1234) | |
| assert out.seed[0] == 1234 | |
| assert all(s is None for s in out.seed[1:]), f"a scalar seed must not broadcast: {out.seed[:4]}" | |
| def test_log_probs_scalar_broadcasts_to_all_lanes(self): | |
| """enable_log_probs selects an output format rather than shaping a distribution, so | |
| it broadcasts to max_batch_size, not to the active lane count.""" | |
| out = self._fmt(temperature=[1.0, 1.0], top_k=8, top_p=1.0, enable_log_probs=True) | |
| assert all(out.enable_log_probs), "enable_log_probs should cover every lane" | |