Download code/models/common/tests/test_sampling.py from tt-hous/clef: direct link, hf CLI and curl.
- Browser
- Download file 81 kB
-
https://huggingface.co/tt-hous/clef/resolve/main/code/models/common/tests/test_sampling.py
- Command line
-
hf download hf://tt-hous/clef/code/models/common/tests/test_sampling.py
-
curl -L -o test_sampling.py https://huggingface.co/tt-hous/clef/resolve/main/code/models/common/tests/test_sampling.py
81 kB
| # SPDX-FileCopyrightText: © 2025 Tenstorrent USA, Inc. | |
| # SPDX-License-Identifier: Apache-2.0 | |
| from types import SimpleNamespace | |
| import pytest | |
| import torch | |
| import torch.nn.functional as F | |
| from ttnn.tools import trace_allocation_tracker | |
| import ttnn | |
| from models.common.sampling import ( | |
| LogProbsCalculator, | |
| SamplingGenerator, | |
| SamplingParams, | |
| SeedManager, | |
| TTSampling, | |
| broadcast_sampling_params, | |
| format_sampling_params, | |
| scatter_sampling_params_to_slots, | |
| ) | |
| from models.common.sampling._utils import topk_would_route_to_large_indices | |
| from models.common.sampling.generator import ( | |
| MAX_UINT32, | |
| _acknowledge_trace_buffers_corruptible, | |
| _hash_request_seed_to_device_seed, | |
| ) | |
| from models.common.sampling.tt_log_probs import MAX_TOP_LOGPROBS, LogProbsResult | |
| from models.common.utility_functions import comp_pcc, is_blackhole | |
| def test_sampling_precompile_preserves_logits_and_request_state(monkeypatch, all_configs): | |
| """Compiling an in-place penalty path must not penalize the next real replay.""" | |
| logits = torch.tensor([1.0, 2.0, 3.0]) | |
| original = logits.clone() | |
| monkeypatch.setattr(ttnn, "clone", torch.clone) | |
| log_probs = SimpleNamespace(logprobs_enabled=[False], num_logprobs=[0], enable_log_probs=False) | |
| def set_log_probs_mode(enabled, num_logprobs): | |
| log_probs.logprobs_enabled = enabled if isinstance(enabled, list) else [enabled] | |
| log_probs.num_logprobs = num_logprobs if isinstance(num_logprobs, list) else [num_logprobs] | |
| log_probs.enable_log_probs = any(log_probs.logprobs_enabled) | |
| log_probs.set_log_probs_mode = set_log_probs_mode | |
| sampling = SamplingGenerator.__new__(SamplingGenerator) | |
| sampling.sub_core_grids = None | |
| sampling._penalties_active = True | |
| sampling._trace_states = {} | |
| sampling.tt_sampling = SimpleNamespace( | |
| log_probs_calculator=log_probs, _force_argmax_sampling=True, _allow_force_argmax_sampling=True | |
| ) | |
| compiled = [] | |
| def run_sampling(scratch, *, penalties_on, tt_out_tok, count_tokens): | |
| assert not count_tokens, "Warmup must not add dummy samples to request history" | |
| if penalties_on: | |
| scratch.sub_(2.0) | |
| compiled.append(penalties_on) | |
| sampling._run_sampling = run_sampling | |
| sampling.precompile(logits, all_configs=all_configs) | |
| assert compiled | |
| torch.testing.assert_close(logits, original, rtol=0, atol=0) | |
| assert sampling._penalties_active is True | |
| assert sampling.tt_sampling._force_argmax_sampling is True | |
| assert log_probs.logprobs_enabled == [False] | |
| assert log_probs.num_logprobs == [0] | |
| def test_sampling_trace_buffer_reuse_is_bucket_only(monkeypatch): | |
| marked = [] | |
| monkeypatch.setattr(trace_allocation_tracker, "acknowledge_corruptible", marked.append) | |
| _acknowledge_trace_buffers_corruptible(None, ["default"]) | |
| _acknowledge_trace_buffers_corruptible(1, ["input", None, ("output",)]) | |
| assert marked == ["input", "output"] | |
| def test_sampling_trace_bucket_isolation(): | |
| """Default users keep one flat namespace; Qwen bucket widths get distinct slots.""" | |
| sampling = SamplingGenerator.__new__(SamplingGenerator) | |
| sampling._trace_states = {} | |
| sampling._active_trace_bucket = None | |
| default_key, default_slot = sampling._trace_slot(False, False, True) | |
| assert default_key.bucket is None | |
| assert sampling._trace_slot(False, False, True)[1] is default_slot | |
| sampling.set_trace_bucket(1) | |
| width1_key, width1_slot = sampling._trace_slot(False, False, True) | |
| sampling.set_trace_bucket(8) | |
| width8_key, width8_slot = sampling._trace_slot(False, False, True) | |
| assert width1_key.bucket == 1 and width8_key.bucket == 8 | |
| assert width1_slot is not default_slot | |
| assert width8_slot is not default_slot | |
| assert width8_slot is not width1_slot | |
| sampling.set_trace_bucket(None) | |
| assert sampling._trace_slot(False, False, True)[1] is default_slot | |
| assert len(sampling._trace_states) == 3 | |
| # --------------------------------------------------------------------------- | |
| # Helper: simulate per-device top-k gather (mirrors TTSampling behaviour) | |
| # --------------------------------------------------------------------------- | |
| def _simulate_gathered_topk(torch_logits, num_devices, top_k=32): | |
| """Simulate the per-device top-k + all-gather that TTSampling performs. | |
| Args: | |
| torch_logits: Full logits tensor, shape (1, 1, B, V). | |
| num_devices: Number of TP devices. | |
| top_k: Per-device top-k count. | |
| Returns: | |
| gathered_values: (1, 1, B, num_devices * top_k) raw logit values. | |
| gathered_indices: (1, 1, B, num_devices * top_k) global vocab indices. | |
| """ | |
| V = torch_logits.shape[-1] | |
| shard_size = V // num_devices | |
| all_values = [] | |
| all_indices = [] | |
| for d in range(num_devices): | |
| shard = torch_logits[:, :, :, d * shard_size : (d + 1) * shard_size] | |
| vals, local_idx = torch.topk(shard, top_k, dim=-1) | |
| global_idx = local_idx + d * shard_size | |
| all_values.append(vals) | |
| all_indices.append(global_idx) | |
| gathered_values = torch.cat(all_values, dim=-1) | |
| gathered_indices = torch.cat(all_indices, dim=-1) | |
| return gathered_values, gathered_indices | |
| # =========================================================================== | |
| # Top-K logprobs tests (TG Galaxy only) | |
| # =========================================================================== | |
| # Common TG Galaxy device parametrization for all new tests | |
| TG_SHAPE = [1, 1, 32, 8 * 16032] # Llama on TG with 8-chip TP sharded vocab | |
| TG_DEVICE_PARAMS = { | |
| "fabric_config": ttnn.FabricConfig.FABRIC_1D_RING, | |
| "dispatch_core_axis": ttnn.DispatchCoreAxis.COL, | |
| } | |
| TG_MESH_SHAPE = (8, 4) | |
| TG_SUB_CORE_GRIDS = ttnn.CoreRangeSet( | |
| [ | |
| ttnn.CoreRange(ttnn.CoreCoord(1, 0), ttnn.CoreCoord(3, 9)), | |
| ttnn.CoreRange(ttnn.CoreCoord(5, 0), ttnn.CoreCoord(6, 9)), | |
| ] | |
| ) | |
| TG_NUM_TP_DEVICES = 8 # TP dimension for Galaxy | |
| def _make_host_only_seed_manager(max_batch_size=4): | |
| return SeedManager(SimpleNamespace(_sampling_dp=1), max_batch_size=max_batch_size) | |
| def test_seed_manager_seed_params_do_not_fallback_to_slot_zero(): | |
| seed_manager = _make_host_only_seed_manager() | |
| assert seed_manager._seed_from_slot_params([11], 0) == 11 | |
| assert seed_manager._seed_from_slot_params([11], 1) is None | |
| assert seed_manager._seed_from_slot_params(torch.tensor([22]), 1) is None | |
| assert seed_manager._seed_from_slot_params((44,), 0) == 44 | |
| assert seed_manager._seed_from_slot_params(33, 3) == 33 | |
| def test_seed_manager_updates_lazy_buffer_with_request_position_hash_and_preserves_default_source(): | |
| class Buffer: | |
| def __init__(self): | |
| self.source = torch.arange(4, dtype=torch.int64) | |
| self.updates = [] | |
| def update(self, source): | |
| self.source = source | |
| self.updates.append(source.clone()) | |
| buffer = Buffer() | |
| defaults = buffer.source.clone() | |
| seed_manager = SeedManager(max_batch_size=4, seed_buffer=buffer) | |
| seed_manager.reset_seed_from_slots((707, None, None, None), range(4)) | |
| seed_manager.align_seed_counters_to_positions((707, None, None, None), [0], [13], offset=1) | |
| values = seed_manager.get_new_values([0]) | |
| assert values == (_hash_request_seed_to_device_seed(707, 14), MAX_UINT32, MAX_UINT32, MAX_UINT32) | |
| assert torch.equal(buffer.updates[-1], torch.tensor(values)) | |
| assert torch.equal(buffer.source, defaults) | |
| seed_manager.restore_default_device_values() | |
| assert torch.equal(buffer.updates[-1], defaults) | |
| def test_seed_counter_position_alignment_skips_out_of_bounds_slots(): | |
| seed_manager = _make_host_only_seed_manager() | |
| seed_manager.align_seed_counters_to_positions([101, None, 303], [0, 2], [5], offset=1) | |
| assert seed_manager.seed_counters == [6, 0, 0, 0] | |
| def test_slot_remap_condense_relabels_destination_and_vacates_source(): | |
| """A condense map moves the source slot's RNG state to its new slot and leaves | |
| the vacated source unseeded. | |
| vLLM's condense moves the highest live request down into the lowest empty slot | |
| (``InputBatch.condense``: ``_slot_remap[empty_index] = _slot_remap[last_req_index]``), | |
| so a source that is not itself a destination has genuinely been vacated. | |
| """ | |
| seed_manager = _make_host_only_seed_manager(max_batch_size=4) | |
| seed_manager.reset_seed([42, 99], [0, 3]) # slot0=42, slot3=99 | |
| assert seed_manager.seeds == [42, None, None, 99] | |
| # Condense: the request in slot3 moves into empty slot1. remap[1]=3; indices | |
| # 0/2/3 keep their identity values (the map does not mark slot3 as empty). | |
| seed_manager.apply_slot_remap(torch.tensor([0, 3, 2, 3], dtype=torch.int32)) | |
| assert seed_manager.seeds[1] == 99 # relabelled into its new slot | |
| assert seed_manager.seeds[3] is None # source vacated | |
| assert seed_manager.seed_counters[3] == 0 | |
| assert seed_manager.seeds[0] == 42 # untouched slot keeps its seed | |
| assert seed_manager._seed_active is True | |
| def test_slot_remap_identity_is_a_noop(): | |
| """The steady state is an identity map (vLLM pops the remap and resets it to | |
| identity every step, and lane-DP never condenses at all), which must not touch | |
| any slot's RNG state.""" | |
| seed_manager = _make_host_only_seed_manager(max_batch_size=4) | |
| seed_manager.reset_seed([42, 43], [0, 1]) | |
| identity = torch.tensor([0, 1, 2, 3], dtype=torch.int32) | |
| for _ in range(50): | |
| seed_manager.apply_slot_remap(identity) | |
| assert seed_manager.seeds == [42, 43, None, None] | |
| assert seed_manager._seed_active is True | |
| def test_duplicate_request_seeds_get_distinct_device_streams(): | |
| """Concurrent slots sharing one request seed must not draw identical streams. | |
| Regression test for #53077: n>1 completions of one prompt with a fixed seed | |
| land in different slots with the same request seed, and the device seed was | |
| derived from (seed, position) only, so every completion came out identical. | |
| """ | |
| seed_manager = _make_host_only_seed_manager(max_batch_size=4) | |
| seed_manager.reset_seed([1234, 1234, 1234, 1234], [0, 1, 2, 3]) | |
| assert sorted(seed_manager.seed_salts) == [0, 1, 2, 3] | |
| first_draws = [seed_manager._next_device_seed_for_slot(slot) for slot in range(4)] | |
| assert len(set(first_draws)) == 4, f"duplicate-seed slots drew identical device seeds: {first_draws}" | |
| def test_duplicate_request_seeds_can_share_one_stream_when_salting_is_disabled(): | |
| """Independent vLLM requests with the same seed remain bit-identical.""" | |
| seed_manager = SeedManager( | |
| SimpleNamespace(_sampling_dp=1), | |
| max_batch_size=4, | |
| salt_duplicate_seeds=False, | |
| ) | |
| seed_manager.reset_seed([1234, 1234, 1234, 1234], [0, 1, 2, 3]) | |
| assert seed_manager.seed_salts == [0, 0, 0, 0] | |
| first_draws = [seed_manager._next_device_seed_for_slot(slot) for slot in range(4)] | |
| assert len(set(first_draws)) == 1 | |
| def test_unique_seed_stream_is_unchanged_and_slot_independent(): | |
| """A request whose seed is unique among active slots keeps salt 0, so its | |
| stream matches the pre-salt derivation and does not depend on the slot.""" | |
| manager_a = _make_host_only_seed_manager(max_batch_size=4) | |
| manager_a.reset_seed([777], [0]) | |
| manager_b = _make_host_only_seed_manager(max_batch_size=4) | |
| manager_b.reset_seed([777], [3]) | |
| draws_a = [manager_a._next_device_seed_for_slot(0) for _ in range(4)] | |
| draws_b = [manager_b._next_device_seed_for_slot(3) for _ in range(4)] | |
| assert manager_a.seed_salts[0] == 0 | |
| assert draws_a == draws_b | |
| def test_seed_salt_travels_with_slot_remap(): | |
| seed_manager = _make_host_only_seed_manager(max_batch_size=4) | |
| seed_manager.reset_seed([55, 55], [0, 3]) # duplicates: slot0 salt 0, slot3 salt 1 | |
| assert seed_manager.seed_salts[3] == 1 | |
| # Condense: slot3's request moves into empty slot1. | |
| seed_manager.apply_slot_remap(torch.tensor([0, 3, 2, 3], dtype=torch.int32)) | |
| assert seed_manager.seed_salts[1] == 1 # stream identity survives the move | |
| assert seed_manager.seed_salts[3] == 0 # vacated slot cleared | |
| def test_seed_salt_does_not_recollide_with_surviving_duplicate(): | |
| seed_manager = _make_host_only_seed_manager(max_batch_size=4) | |
| seed_manager.reset_seed([55, 55], [0, 1]) # slot0 salt 0, slot1 salt 1 | |
| # Slot 0 finishes and is vacated; a new request with the same seed arrives. | |
| seed_manager.apply_slot_remap(torch.tensor([1, 1, 2, 3], dtype=torch.int32)) | |
| assert seed_manager.seeds == [55, None, None, None] | |
| seed_manager.reset_seed([55], [2]) | |
| # The newcomer must not reuse the surviving request's salt. | |
| survivor_slot = 0 | |
| assert seed_manager.seed_salts[2] != seed_manager.seed_salts[survivor_slot] | |
| def test_seed_salt_survives_decode_re_registration_after_sibling_finishes(): | |
| """The unconditional decode-path re-registration (first decode after any | |
| admission) must not recompute a running request's salt: after its same-seed | |
| sibling finishes, the recomputed salt would drop to the sibling's and the | |
| survivor's remaining tokens would replay the finished stream.""" | |
| seed_manager = _make_host_only_seed_manager(max_batch_size=4) | |
| seed_manager.reset_seed([42, 42], [0, 1]) # A salt 0, B salt 1 | |
| assert seed_manager.seed_salts[1] == 1 | |
| # A finishes; condense moves B down into slot 0 (salt travels). | |
| seed_manager.apply_slot_remap(torch.tensor([1, 1, 2, 3], dtype=torch.int32)) | |
| assert seed_manager.seed_salts[0] == 1 | |
| # An unrelated admission triggers reset_batch: every active slot re-registers. | |
| seed_manager.reset_seed([7], [2]) | |
| seed_manager.reset_seed_from_slots([42, None, 7, None], [0, 2]) | |
| assert seed_manager.seed_salts[0] == 1 # B keeps its stream mid-generation | |
| def test_finished_tail_request_ghost_seed_is_cleared(): | |
| """A request finishing at the batch tail is never vacated by condense; the | |
| decode-path deactivate must drop it so a later unique-seed request still | |
| gets salt 0 (seeded reproducibility).""" | |
| seed_manager = _make_host_only_seed_manager(max_batch_size=4) | |
| seed_manager.reset_seed([42], [3]) | |
| # Tail completion: identity remap makes no moves, the ghost stays. | |
| seed_manager.apply_slot_remap(torch.tensor([0, 1, 2, 3], dtype=torch.int32)) | |
| assert seed_manager.seeds[3] == 42 | |
| seed_manager.deactivate_slots_except([0]) | |
| assert seed_manager.seeds[3] is None | |
| # A fresh unique seed-42 request must land on salt 0. | |
| seed_manager.reset_seed([42], [1]) | |
| assert seed_manager.seed_salts[1] == 0 | |
| def test_deactivating_last_seeded_slot_rearms_the_unseeded_push(): | |
| """When the last seeded request finishes, the device still holds non-SKIP | |
| reinit values; unless _reseted is set, get_new_values early-returns forever | |
| and every surviving user's PRNG reinitializes to the same stale seed each | |
| token (frozen sampling).""" | |
| seed_manager = _make_host_only_seed_manager(max_batch_size=4) | |
| seed_manager.reset_seed([42], [3]) | |
| seed_manager._reseted = False # simulate the post-push steady state | |
| seed_manager.deactivate_slots_except([0]) | |
| assert seed_manager._seed_active is False | |
| assert seed_manager._reseted is True | |
| def test_remap_overwriting_last_seeded_slot_rearms_the_unseeded_push(): | |
| seed_manager = _make_host_only_seed_manager(max_batch_size=4) | |
| seed_manager.reset_seed([42], [0]) | |
| seed_manager._reseted = False | |
| # Condense moves the (unseeded) request from slot 1 over the seeded slot 0. | |
| seed_manager.apply_slot_remap(torch.tensor([1, 1, 2, 3], dtype=torch.int32)) | |
| assert seed_manager._seed_active is False | |
| assert seed_manager._reseted is True | |
| def test_prefill_admission_into_same_slot_fully_resets_seed_state(): | |
| """reset_seed registers a NEW request: even when the slot already holds the | |
| same seed value, the counter must restart and the salt must be recomputed | |
| (its old salt may reflect siblings that no longer exist).""" | |
| seed_manager = _make_host_only_seed_manager(max_batch_size=4) | |
| seed_manager.reset_seed([42, 42], [0, 1]) | |
| assert seed_manager.seed_salts[1] == 1 | |
| seed_manager._next_device_seed_for_slot(1) | |
| assert seed_manager.seed_counters[1] == 1 | |
| seed_manager.deactivate_slots_except([1]) # slot 0's request finished | |
| seed_manager.reset_seed([42], [1]) # new same-seed request admitted into slot 1 | |
| assert seed_manager.seed_salts[1] == 0 | |
| assert seed_manager.seed_counters[1] == 0 | |
| def test_finished_requests_release_seeds_before_permuted_prefill(): | |
| """A unique seed must replay from prefill, before decode can reconcile slots.""" | |
| manager = _make_host_only_seed_manager(max_batch_size=32) | |
| slots = [0, 10, 20, 31] | |
| seeds = [100, 101, 102, 103] | |
| first = {} | |
| for slot, seed in zip(slots, seeds): | |
| manager.reset_seed([seed], [slot]) | |
| first[seed] = [manager._next_device_seed_for_slot(slot) for _ in range(4)] | |
| for slot in slots: | |
| manager.release_slot(slot) | |
| for slot, seed in zip(reversed(slots), seeds): | |
| manager.reset_seed([seed], [slot]) | |
| assert manager.seed_salts[slot] == 0 | |
| assert [manager._next_device_seed_for_slot(slot) for _ in range(4)] == first[seed] | |
| def test_release_keeps_surviving_duplicate_stream_and_frees_only_finished_salt(): | |
| manager = _make_host_only_seed_manager() | |
| manager.reset_seed([55, 55], [0, 1]) | |
| manager._next_device_seed_for_slot(1) | |
| survivor_rng = manager.rngs[1].getstate() | |
| manager.release_slot(0) | |
| manager.release_slot(0) # Idempotent; a live sibling still owns salt 1. | |
| assert manager.seed_salts[1] == 1 | |
| assert manager.seed_counters[1] == 1 | |
| assert manager.rngs[1].getstate() == survivor_rng | |
| assert manager._next_device_seed_for_slot(1) == _hash_request_seed_to_device_seed(55, 1, 1) | |
| manager.reset_seed([55], [3]) | |
| assert manager.seed_salts[3] == 0 | |
| assert manager.seed_salts[1] == 1 | |
| def test_releasing_last_seeded_request_rearms_unseeded_sampling(): | |
| manager = _RecordingSeedManager(4) | |
| manager.reset_seed([42], [3]) | |
| manager.get_new_values([3]) | |
| manager.release_slot(3) | |
| assert not manager._seed_active | |
| assert manager._reseted | |
| first = manager.get_new_values([0]) | |
| assert all(0 < seed < MAX_UINT32 for seed in first) | |
| assert manager.get_new_values([0]) == (MAX_UINT32,) * 4 | |
| assert manager.get_new_values([0]) is None | |
| def test_generator_release_request_routes_global_state_slot(capacity, replicas, slot): | |
| from models.tt_transformers.tt.generator import Generator | |
| managers = [_make_host_only_seed_manager(capacity) for _ in range(replicas)] | |
| for manager in managers: | |
| manager.reset_seed([17, 17], [0, capacity - 1]) | |
| generator = SimpleNamespace( | |
| model_args=[SimpleNamespace(max_batch_size=capacity)], | |
| data_parallel=replicas, | |
| model=[SimpleNamespace(sampling=SimpleNamespace(seed_manager=m)) for m in managers], | |
| _slots_prefilled_since_decode={0, slot}, | |
| ) | |
| Generator.release_request(generator, slot) | |
| rank, local_slot = divmod(slot, capacity) | |
| assert managers[rank].seeds[local_slot] is None | |
| assert managers[rank].seeds[0] == 17 | |
| assert generator._slots_prefilled_since_decode == {0} | |
| for other_rank, manager in enumerate(managers): | |
| if other_rank != rank: | |
| assert manager.seeds[-1] == 17 | |
| for invalid_slot in (-1, capacity * replicas): | |
| with pytest.raises(ValueError, match="outside"): # allow-pytest.raises: host-only bounds regression | |
| Generator.release_request(generator, invalid_slot) | |
| def test_broadcast_sampling_params_preserves_none_list_fields(): | |
| params = SamplingParams(temperature=[1.0, 1.0], top_k=[1, 1], top_p=[1.0, 1.0], seed=[None, 42]) | |
| broadcast = broadcast_sampling_params(params, 0, slot_len=4) | |
| assert broadcast.seed == [None, None, None, None] | |
| def test_format_sampling_params_uses_device_argmax_sentinel_for_greedy_rows(): | |
| params = format_sampling_params( | |
| SamplingParams(temperature=0.0, top_k=32, top_p=0.95), | |
| max_batch_size=32, | |
| ) | |
| assert params.temperature[0] == 1.0 | |
| assert params.top_k[0] == 1 | |
| assert params.top_p[0] == 0.0 | |
| # --------------------------------------------------------------------------- | |
| # Seeded decode reproducibility under async scheduling (#51981). | |
| # Host-only: the generator is built with __new__ and driven with stub modules. | |
| # --------------------------------------------------------------------------- | |
| SEED_TEST_BATCH = 32 # format_sampling_params requires a multiple of 32 | |
| class _RecordingSeedManager(SeedManager): | |
| """SeedManager that records each step's device seed vector instead of pushing it.""" | |
| def __init__(self, max_batch_size=SEED_TEST_BATCH): | |
| super().__init__(SimpleNamespace(_sampling_dp=1), max_batch_size=max_batch_size) | |
| self.pushed = [] | |
| def write_device_seed_values(self, seed_values): | |
| self.pushed.append(list(seed_values)) | |
| class _StubSamplingModule: | |
| def __init__(self, max_batch_size=SEED_TEST_BATCH): | |
| self.seed_manager = _RecordingSeedManager(max_batch_size) | |
| self.tt_sampling = SimpleNamespace(max_batch_size=max_batch_size) | |
| def apply_decode_state(self, sampling_params_chunks, **kwargs): | |
| pass | |
| def sample(self, logits=None, **kwargs): | |
| return logits | |
| def _make_stub_generator(max_batch_size=SEED_TEST_BATCH): | |
| from models.tt_transformers.tt.generator import Generator | |
| generator = Generator.__new__(Generator) | |
| generator.data_parallel = 1 | |
| sampling = _StubSamplingModule(max_batch_size) | |
| generator.model = [SimpleNamespace(sampling=sampling, sampling_dp=1)] | |
| return generator, sampling | |
| def _decode_sampling_step(generator, seeds, positions, reload_inputs, max_batch_size=SEED_TEST_BATCH): | |
| """Run one device-sampling decode step; return the per-slot device seeds pushed.""" | |
| seeds = list(seeds) + [None] * (max_batch_size - len(seeds)) | |
| positions = list(positions) + [-1] * (max_batch_size - len(positions)) | |
| params = SamplingParams( | |
| temperature=[1.0] * max_batch_size, | |
| top_k=[32] * max_batch_size, | |
| top_p=[1.0] * max_batch_size, | |
| seed=seeds, | |
| ) | |
| generator.sample_decode_on_device( | |
| [None], | |
| sampling_params=params, | |
| start_pos=[torch.tensor(positions, dtype=torch.int32)], | |
| reload_sampling_params=True, | |
| reset_sampling_state=False, | |
| reload_inputs=reload_inputs, | |
| ) | |
| return generator.model[0].sampling.seed_manager.pushed[-1] | |
| def _expected_seed_stream(request_seed, first_position, num_steps): | |
| """Device seeds for `num_steps` consecutive tokens starting at `first_position`.""" | |
| positions = range(first_position, first_position + num_steps) | |
| return [_hash_request_seed_to_device_seed(request_seed, pos + 1) for pos in positions] | |
| def test_seed_stream_is_independent_of_host_position_lag(): | |
| """The counter must self-advance from the last authoritative anchor; re-anchoring | |
| to a lagging host position replays device seeds (#51981).""" | |
| generator, sampling = _make_stub_generator() | |
| pushed = [_decode_sampling_step(generator, [7], [100], reload_inputs=True)[0]] | |
| # Device is at 101, 102, 103; the host reports the previous position and | |
| # stalls entirely when a readback is late. | |
| for lagging_host_pos in (100, 101, 101): | |
| pushed.append(_decode_sampling_step(generator, [7], [lagging_host_pos], reload_inputs=False)[0]) | |
| assert pushed == _expected_seed_stream(7, 100, 4) | |
| assert len(set(pushed)) == 4 # no replayed seed | |
| assert sampling.seed_manager.seed_counters[0] == 105 # anchored at 101, one per token | |
| def test_seed_stream_matches_across_different_host_lags(): | |
| """Same seed and true positions, but one run has async overlap engaged (host | |
| lags) and the other does not. Equal streams is what `seed=` promises.""" | |
| lagged, _ = _make_stub_generator() | |
| exact, _ = _make_stub_generator() | |
| lagged_stream = [_decode_sampling_step(lagged, [7], [100], reload_inputs=True)[0]] | |
| exact_stream = [_decode_sampling_step(exact, [7], [100], reload_inputs=True)[0]] | |
| for true_pos in (101, 102, 103): | |
| lagged_stream.append(_decode_sampling_step(lagged, [7], [true_pos - 1], reload_inputs=False)[0]) | |
| exact_stream.append(_decode_sampling_step(exact, [7], [true_pos], reload_inputs=False)[0]) | |
| assert lagged_stream == exact_stream | |
| def test_seed_counters_realign_when_host_inputs_are_authoritative(): | |
| """A batch reset re-anchors every active slot: vLLM may have evicted and | |
| re-admitted the request elsewhere, so the resident counter is untrustworthy.""" | |
| generator, sampling = _make_stub_generator() | |
| _decode_sampling_step(generator, [7], [100], reload_inputs=True) | |
| sampling.seed_manager.seed_counters[0] = 0 # state moved behind our back | |
| pushed = _decode_sampling_step(generator, [7], [200], reload_inputs=True) | |
| assert pushed[0] == _hash_request_seed_to_device_seed(7, 201) | |
| def test_newly_seeded_slot_is_aligned_even_on_a_non_authoritative_step(): | |
| """A freshly admitted slot's position comes from its prefill, so it is | |
| authoritative even when the rest of the batch's host inputs are stale; | |
| otherwise its reset-to-zero counter starts the stream at the wrong offset.""" | |
| generator, _ = _make_stub_generator() | |
| _decode_sampling_step(generator, [7], [100], reload_inputs=True) | |
| # Slot 1 admitted mid-flight at position 5; slot 0 keeps decoding with a lag. | |
| pushed = _decode_sampling_step(generator, [7, 11], [100, 5], reload_inputs=False) | |
| assert pushed[0] == _hash_request_seed_to_device_seed(7, 102) # unmoved, self-advanced | |
| assert pushed[1] == _hash_request_seed_to_device_seed(11, 6) # anchored to its prefill position | |
| def test_reset_seed_from_slots_if_needed_reports_the_slots_it_reset(): | |
| seed_manager = _make_host_only_seed_manager() | |
| seed_manager.reset_seed_from_slots([42, 43, None, None], [0, 1, 2, 3]) | |
| assert seed_manager.reset_seed_from_slots_if_needed([42, 43, None, None], [0, 1, 2, 3]) == [] | |
| assert seed_manager.reset_seed_from_slots_if_needed([42, 99, None, None], [0, 1, 2, 3]) == [1] | |
| def test_scatter_sampling_params_to_slots_moves_params_to_their_slot_row(): | |
| """A batched prefill samples slot row s with the params of the request there.""" | |
| params = SamplingParams(temperature=[0.1, 0.2, 0.3], top_k=[1, 2, 3], top_p=[0.5, 0.6, 0.7], seed=[7, 8, 9]) | |
| scattered = scatter_sampling_params_to_slots(params, [2, 0, 5], slot_len=8) | |
| assert scattered.temperature[2] == 0.1 and scattered.temperature[0] == 0.2 | |
| assert scattered.temperature[5] == 0.3 | |
| assert scattered.top_k[2] == 1 and scattered.top_k[0] == 2 and scattered.top_k[5] == 3 | |
| assert scattered.top_p[2] == 0.5 and scattered.top_p[0] == 0.6 and scattered.top_p[5] == 0.7 | |
| # Unoccupied rows carry the last request's values, so they stay valid instead of | |
| # sampling from a formatter default. | |
| assert scattered.temperature[1] == 0.3 | |
| # SeedManager.reset_seed is given the slot list separately and maps seeds itself. | |
| assert scattered.seed == [7, 8, 9] | |
| # The input is never mutated. | |
| assert params.temperature == [0.1, 0.2, 0.3] | |
| def test_scatter_sampling_params_to_slots_is_identity_for_dense_slots(): | |
| params = format_sampling_params(SamplingParams(temperature=[0.5, 0.5], top_k=[4, 4], top_p=[0.9, 0.9]), 32) | |
| scattered = scatter_sampling_params_to_slots(params, list(range(2)), slot_len=32) | |
| assert scattered.temperature[:2] == params.temperature[:2] | |
| assert scattered.top_k[:2] == params.top_k[:2] | |
| def _skip_if_not_galaxy(mesh_device): | |
| """Skip test if not running on TG Galaxy (32 devices).""" | |
| if mesh_device.get_num_devices() != 32: | |
| pytest.skip(f"Test requires TG Galaxy (32 devices), got {mesh_device.get_num_devices()}") | |
| def _push_topk_test_tensors_to_tg(torch_tensor, gathered_values, gathered_indices, mesh_device): | |
| """Push logits, topk values, and topk indices to a TG Galaxy mesh device.""" | |
| logits_tt = ttnn.from_torch( | |
| torch_tensor, | |
| device=mesh_device, | |
| dtype=ttnn.bfloat16, | |
| layout=ttnn.TILE_LAYOUT, | |
| mesh_mapper=ttnn.ShardTensor2dMesh(mesh_device, dims=(-1, None), mesh_shape=list(mesh_device.shape)), | |
| memory_config=ttnn.DRAM_MEMORY_CONFIG, | |
| ) | |
| topk_values_tt = ttnn.from_torch( | |
| gathered_values, | |
| device=mesh_device, | |
| dtype=ttnn.bfloat16, | |
| layout=ttnn.TILE_LAYOUT, | |
| mesh_mapper=ttnn.ReplicateTensorToMesh(mesh_device), | |
| memory_config=ttnn.DRAM_MEMORY_CONFIG, | |
| ) | |
| topk_indices_tt = ttnn.from_torch( | |
| gathered_indices.to(torch.int32), | |
| device=mesh_device, | |
| dtype=ttnn.int32, | |
| layout=ttnn.TILE_LAYOUT, | |
| mesh_mapper=ttnn.ReplicateTensorToMesh(mesh_device), | |
| memory_config=ttnn.DRAM_MEMORY_CONFIG, | |
| ) | |
| return logits_tt, topk_values_tt, topk_indices_tt | |
| def test_log_probs_calculation(shape, mesh_device): | |
| seed = 1234 | |
| torch.manual_seed(seed) | |
| log_probs_calculator = LogProbsCalculator(mesh_device) | |
| torch_tensor = torch.randn(shape) | |
| # shuffle the tensor in last 2 dimensions | |
| for i in range(shape[-2]): | |
| torch_tensor[:, :, i, :] = torch_tensor[:, :, i, torch.randperm(shape[-1])] | |
| argmax_tensor = torch.argmax(torch_tensor, dim=-1, keepdim=True) | |
| indices_tensor = argmax_tensor.reshape( | |
| argmax_tensor.shape[0], argmax_tensor.shape[1], argmax_tensor.shape[-1], argmax_tensor.shape[-2] | |
| ) | |
| # Push inputs to device | |
| logits_tensor = ttnn.from_torch( | |
| torch_tensor, | |
| device=mesh_device, | |
| dtype=ttnn.bfloat16, | |
| layout=ttnn.TILE_LAYOUT, | |
| mesh_mapper=ttnn.ShardTensorToMesh(mesh_device, dim=-1), | |
| memory_config=ttnn.DRAM_MEMORY_CONFIG, | |
| ) | |
| ttnn_indices_tensor = ttnn.from_torch( | |
| indices_tensor, | |
| device=mesh_device, | |
| dtype=ttnn.int32, | |
| layout=ttnn.TILE_LAYOUT, | |
| mesh_mapper=ttnn.ReplicateTensorToMesh(mesh_device), | |
| memory_config=ttnn.DRAM_MEMORY_CONFIG, | |
| ) | |
| log_probs_calculator.set_log_probs_mode(True) | |
| tt_log_probs = log_probs_calculator.calculate_log_probs(logits_tensor, ttnn_indices_tensor) | |
| log_probs_tt_host = ttnn.to_torch(tt_log_probs, mesh_composer=ttnn.ConcatMeshToTensor(mesh_device, dim=3)) | |
| log_probs_tt_host = log_probs_tt_host[:, :, :1, :32] | |
| # Calculate log-probs for each user on each chip using torch | |
| log_probs_torch = F.log_softmax(torch_tensor.float(), dim=-1) | |
| log_probs_torch_argmax = torch.gather(log_probs_torch, dim=-1, index=argmax_tensor) | |
| log_probs_torch_argmax = torch.reshape(log_probs_torch_argmax, (1, 1, 1, 32)) | |
| passing, pcc = comp_pcc(log_probs_torch_argmax, log_probs_tt_host, pcc=0.99) | |
| print(f"pcc={pcc}") | |
| assert passing, f"Assertion failed, PCC={pcc}" | |
| def _shard_logits_2d_mesh(logits_host, mesh_device): | |
| """Shard vocab along mesh TP axis (matches test_sampling_1d._make_logits_tt).""" | |
| cluster_shape = tuple(mesh_device.shape) | |
| if cluster_shape[-1] >= cluster_shape[-2]: | |
| shard_dims = (None, -1) | |
| else: | |
| shard_dims = (-1, None) | |
| return ttnn.from_torch( | |
| logits_host, | |
| device=mesh_device, | |
| dtype=ttnn.bfloat16, | |
| layout=ttnn.TILE_LAYOUT, | |
| mesh_mapper=ttnn.ShardTensor2dMesh(mesh_device, dims=shard_dims, mesh_shape=cluster_shape), | |
| memory_config=ttnn.DRAM_MEMORY_CONFIG, | |
| ) | |
| def test_log_probs_calculation_shard_tensor_2d_mesh_1x8(mesh_device): | |
| """LogProbsCalculator with ShardTensor2dMesh on 1×8 — the path Sampling1D uses. | |
| test_log_probs_calculation shards via ShardTensorToMesh(dim=-1), which does not | |
| exercise the 1×N _all_gather_cluster_axis bug fixed in tt_log_probs.py. | |
| """ | |
| if mesh_device.get_num_devices() != 8: | |
| pytest.skip(f"Test targets 1×8 mesh, got {mesh_device.get_num_devices()} devices") | |
| batch_size = 32 | |
| vocab_size = 32768 | |
| shape = [1, 1, batch_size, vocab_size] | |
| torch.manual_seed(42) | |
| log_probs_calculator = LogProbsCalculator(mesh_device) | |
| torch_tensor = torch.randn(shape, dtype=torch.bfloat16) | |
| for i in range(batch_size): | |
| torch_tensor[:, :, i, :] = torch_tensor[:, :, i, torch.randperm(vocab_size)] | |
| # Pin a few batch slots to tokens on different chips (4096 tokens/chip on 1×8). | |
| pinned_tokens = [(0, 100), (1, 20000), (2, 30000), (3, 5000)] # chips 0, 4, 7, 1 | |
| for batch_idx, token_id in pinned_tokens: | |
| torch_tensor[:, :, batch_idx, token_id] = 10.0 | |
| argmax_tensor = torch.argmax(torch_tensor.float(), dim=-1, keepdim=True) | |
| for batch_idx, token_id in pinned_tokens: | |
| assert argmax_tensor[0, 0, batch_idx, 0].item() == token_id | |
| indices_tensor = argmax_tensor.reshape( | |
| argmax_tensor.shape[0], argmax_tensor.shape[1], argmax_tensor.shape[-1], argmax_tensor.shape[-2] | |
| ) | |
| logits_tensor = _shard_logits_2d_mesh(torch_tensor, mesh_device) | |
| ttnn_indices_tensor = ttnn.from_torch( | |
| indices_tensor, | |
| device=mesh_device, | |
| dtype=ttnn.int32, | |
| layout=ttnn.TILE_LAYOUT, | |
| mesh_mapper=ttnn.ReplicateTensorToMesh(mesh_device), | |
| memory_config=ttnn.DRAM_MEMORY_CONFIG, | |
| ) | |
| log_probs_calculator.set_log_probs_mode(True) | |
| tt_log_probs = log_probs_calculator.calculate_log_probs(logits_tensor, ttnn_indices_tensor) | |
| assert tt_log_probs is not None | |
| log_probs_tt_host = ttnn.to_torch(tt_log_probs, mesh_composer=ttnn.ConcatMeshToTensor(mesh_device, dim=3)) | |
| log_probs_tt_host = log_probs_tt_host[:, :, :1, :batch_size] | |
| log_probs_torch = F.log_softmax(torch_tensor.float(), dim=-1) | |
| log_probs_torch_argmax = torch.gather(log_probs_torch, dim=-1, index=argmax_tensor) | |
| log_probs_torch_argmax = log_probs_torch_argmax.reshape(1, 1, 1, batch_size) | |
| passing, pcc = comp_pcc(log_probs_torch_argmax, log_probs_tt_host, pcc=0.99) | |
| assert passing, f"logprobs PCC below threshold: {pcc}" | |
| def test_log_probs_returns_none_when_disabled(shape, mesh_device): | |
| """Test that calculate_log_probs returns None when enable_log_probs is False.""" | |
| log_probs_calculator = LogProbsCalculator(mesh_device) | |
| torch_tensor = torch.randn(shape) | |
| argmax_tensor = torch.argmax(torch_tensor, dim=-1, keepdim=True) | |
| indices_tensor = argmax_tensor.reshape( | |
| argmax_tensor.shape[0], argmax_tensor.shape[1], argmax_tensor.shape[-1], argmax_tensor.shape[-2] | |
| ) | |
| logits_tensor = ttnn.from_torch( | |
| torch_tensor, | |
| device=mesh_device, | |
| dtype=ttnn.bfloat16, | |
| layout=ttnn.TILE_LAYOUT, | |
| mesh_mapper=ttnn.ShardTensorToMesh(mesh_device, dim=-1), | |
| memory_config=ttnn.DRAM_MEMORY_CONFIG, | |
| ) | |
| ttnn_indices_tensor = ttnn.from_torch( | |
| indices_tensor, | |
| device=mesh_device, | |
| dtype=ttnn.int32, | |
| layout=ttnn.TILE_LAYOUT, | |
| mesh_mapper=ttnn.ReplicateTensorToMesh(mesh_device), | |
| memory_config=ttnn.DRAM_MEMORY_CONFIG, | |
| ) | |
| # Log probs disabled (default) - should return None | |
| log_probs_calculator.set_log_probs_mode(False) | |
| result = log_probs_calculator.calculate_log_probs(logits_tensor, ttnn_indices_tensor) | |
| assert result is None, f"Expected None when log_probs disabled, got {type(result)}" | |
| # Log probs enabled - should return a tensor (not None) | |
| log_probs_calculator.set_log_probs_mode(True) | |
| num_devices = mesh_device.get_num_devices() | |
| result = log_probs_calculator.calculate_log_probs(logits_tensor, ttnn_indices_tensor) | |
| if num_devices in (8, 32) and log_probs_calculator.num_devices_for_sharding >= 2: | |
| assert result is not None, "Expected tensor when log_probs enabled on supported device" | |
| else: | |
| assert result is None, "Expected None on unsupported device count" | |
| def test_log_probs_with_sub_core_grids_on_galaxy(shape, mesh_device): | |
| seed = 1234 | |
| torch.manual_seed(seed) | |
| sub_core_grids = ttnn.CoreRangeSet( | |
| [ | |
| ttnn.CoreRange(ttnn.CoreCoord(1, 0), ttnn.CoreCoord(3, 9)), | |
| ttnn.CoreRange(ttnn.CoreCoord(5, 0), ttnn.CoreCoord(6, 9)), | |
| ] | |
| ) | |
| log_probs_calculator = LogProbsCalculator(mesh_device, sub_core_grids) | |
| torch_tensor = torch.randn(shape) | |
| # shuffle the tensor in last 2 dimensions | |
| for i in range(shape[-2]): | |
| torch_tensor[:, :, i, :] = torch_tensor[:, :, i, torch.randperm(shape[-1])] | |
| argmax_tensor = torch.argmax(torch_tensor, dim=-1, keepdim=True) | |
| indices_tensor = argmax_tensor.reshape( | |
| argmax_tensor.shape[0], argmax_tensor.shape[1], argmax_tensor.shape[-1], argmax_tensor.shape[-2] | |
| ) | |
| if mesh_device.get_num_devices() == 8: | |
| mesh_mapper = ttnn.ShardTensorToMesh(mesh_device, dim=-1) | |
| elif mesh_device.get_num_devices() == 32: | |
| mesh_mapper = ttnn.ShardTensor2dMesh(mesh_device, dims=(-1, None), mesh_shape=list(mesh_device.shape)) | |
| else: | |
| raise ValueError(f"Unsupported number of devices: {mesh_device.get_num_devices()}") | |
| logits_tensor = ttnn.from_torch( | |
| torch_tensor, | |
| device=mesh_device, | |
| dtype=ttnn.bfloat16, | |
| layout=ttnn.TILE_LAYOUT, | |
| mesh_mapper=mesh_mapper, | |
| memory_config=ttnn.DRAM_MEMORY_CONFIG, | |
| ) | |
| ttnn_indices_tensor = ttnn.from_torch( | |
| indices_tensor, | |
| device=mesh_device, | |
| dtype=ttnn.int32, | |
| layout=ttnn.TILE_LAYOUT, | |
| mesh_mapper=ttnn.ReplicateTensorToMesh(mesh_device), | |
| memory_config=ttnn.DRAM_MEMORY_CONFIG, | |
| ) | |
| log_probs_calculator.set_log_probs_mode(True) | |
| tt_log_probs = log_probs_calculator.calculate_log_probs(logits_tensor, ttnn_indices_tensor) | |
| log_probs_tt_host = ttnn.to_torch(tt_log_probs, mesh_composer=ttnn.ConcatMeshToTensor(mesh_device, dim=3)) | |
| # slice from (1,1,32,256) -> (1,1,1,32) | |
| log_probs_tt_host = log_probs_tt_host[:, :, :1, :32] | |
| log_probs_torch = F.log_softmax(torch_tensor.float(), dim=-1) | |
| log_probs_torch_argmax = torch.gather(log_probs_torch, dim=-1, index=argmax_tensor) | |
| log_probs_torch_argmax = torch.reshape(log_probs_torch_argmax, (1, 1, 1, 32)) | |
| passing, pcc = comp_pcc(log_probs_torch_argmax, log_probs_tt_host, pcc=0.99) | |
| print(f"pcc={pcc}") | |
| assert passing, f"Assertion failed, PCC={pcc}" | |
| # =========================================================================== | |
| # New top-K logprobs tests (TG Galaxy only) | |
| # =========================================================================== | |
| def test_top_k_log_probs_on_galaxy(shape, mesh_device): | |
| """Top-K logprobs PCC check on TG Galaxy (32-device 2D mesh).""" | |
| _skip_if_not_galaxy(mesh_device) | |
| torch.manual_seed(1234) | |
| batch_size = shape[2] | |
| calc = LogProbsCalculator(mesh_device, TG_SUB_CORE_GRIDS, batch_size=batch_size, use_topk_logprobs=True) | |
| torch_tensor = torch.randn(shape) | |
| for i in range(batch_size): | |
| torch_tensor[:, :, i, :] = torch_tensor[:, :, i, torch.randperm(shape[-1])] | |
| log_probs_torch = F.log_softmax(torch_tensor.to(torch.float16), dim=-1) | |
| gathered_values, gathered_indices = _simulate_gathered_topk(torch_tensor, TG_NUM_TP_DEVICES) | |
| argmax_tensor = torch.argmax(torch_tensor, dim=-1, keepdim=True) | |
| logits_tt, topk_values_tt, topk_indices_tt = _push_topk_test_tensors_to_tg( | |
| torch_tensor, gathered_values, gathered_indices, mesh_device | |
| ) | |
| calc.set_log_probs_mode([True] * batch_size, num_logprobs=[5] * batch_size) | |
| result = calc.calculate_topk_log_probs(logits_tt, topk_values_tt, topk_indices_tt) | |
| assert result is not None, "Expected LogProbsResult, got None" | |
| assert isinstance(result, LogProbsResult) | |
| host_results = calc.transfer_logprobs_to_host(result, argmax_tensor.squeeze()) | |
| composer = calc._build_mesh_composer() | |
| topk_logprobs_host = ttnn.to_torch(result.topk_logprobs, mesh_composer=composer) | |
| topk_logprobs_host = topk_logprobs_host[0, 0, ...] | |
| topk_indices_host = ttnn.to_torch(result.topk_indices, mesh_composer=composer) | |
| topk_indices_host = topk_indices_host[0, 0, ...].long() | |
| expected_logprobs = torch.gather( | |
| log_probs_torch.squeeze(0).squeeze(0), | |
| dim=-1, | |
| index=topk_indices_host, | |
| ) | |
| passing, pcc = comp_pcc(expected_logprobs, topk_logprobs_host, pcc=0.99) | |
| print(f"Galaxy top-K logprobs PCC={pcc}") | |
| assert passing, f"Galaxy top-K logprobs PCC failed: {pcc}" | |
| for user_idx in range(batch_size): | |
| r = host_results[user_idx] | |
| assert r is not None | |
| sampled_id = argmax_tensor[0, 0, user_idx, 0].item() | |
| torch_lp = log_probs_torch[0, 0, user_idx, sampled_id].item() | |
| assert abs(r["returned_token"]["logprob"] - torch_lp) < 0.05 | |
| def test_top_k_log_probs_returns_none_when_not_needed(shape, mesh_device): | |
| """calculate_topk_log_probs returns None when disabled.""" | |
| _skip_if_not_galaxy(mesh_device) | |
| batch_size = shape[2] | |
| calc = LogProbsCalculator(mesh_device, TG_SUB_CORE_GRIDS, batch_size=batch_size, use_topk_logprobs=True) | |
| torch_tensor = torch.randn(shape) | |
| gathered_values, gathered_indices = _simulate_gathered_topk(torch_tensor, TG_NUM_TP_DEVICES) | |
| argmax_tensor = torch.argmax(torch_tensor, dim=-1, keepdim=True) | |
| logits_tt, topk_values_tt, topk_indices_tt = _push_topk_test_tensors_to_tg( | |
| torch_tensor, gathered_values, gathered_indices, mesh_device | |
| ) | |
| calc.set_log_probs_mode(False, num_logprobs=0) | |
| result = calc.calculate_topk_log_probs(logits_tt, topk_values_tt, topk_indices_tt) | |
| assert result is None, "Expected None when logprobs disabled" | |
| calc.set_log_probs_mode(True, num_logprobs=0) | |
| assert calc.topk_logprobs_needed # needed for sampled token logprob | |
| result = calc.calculate_topk_log_probs(logits_tt, topk_values_tt, topk_indices_tt) | |
| assert result is not None, "Expected LogProbsResult when logprobs enabled" | |
| sampled_ids = argmax_tensor.squeeze() | |
| host_results = calc.transfer_logprobs_to_host(result, sampled_ids) | |
| assert len(host_results) == batch_size | |
| for i in range(batch_size): | |
| r = host_results[i] | |
| assert r is not None | |
| assert r["returned_token"]["token_idx"] == int(sampled_ids[i].item()) | |
| assert len(r["top_logprobs"]["token_indices"]) == 0 | |
| def test_per_user_logprobs_enabled(shape, mesh_device): | |
| """Mixed per-user logprobs: only even users enabled.""" | |
| _skip_if_not_galaxy(mesh_device) | |
| torch.manual_seed(42) | |
| batch_size = shape[2] | |
| calc = LogProbsCalculator(mesh_device, TG_SUB_CORE_GRIDS, batch_size=batch_size, use_topk_logprobs=True) | |
| torch_tensor = torch.randn(shape) | |
| for i in range(batch_size): | |
| torch_tensor[:, :, i, :] = torch_tensor[:, :, i, torch.randperm(shape[-1])] | |
| log_probs_torch = F.log_softmax(torch_tensor.to(torch.float16), dim=-1) | |
| gathered_values, gathered_indices = _simulate_gathered_topk(torch_tensor, TG_NUM_TP_DEVICES) | |
| argmax_tensor = torch.argmax(torch_tensor, dim=-1, keepdim=True) | |
| logits_tt, topk_values_tt, topk_indices_tt = _push_topk_test_tensors_to_tg( | |
| torch_tensor, gathered_values, gathered_indices, mesh_device | |
| ) | |
| enable_log_probs = [i % 2 == 0 for i in range(batch_size)] | |
| num_logprobs_list = [5 if i % 2 == 0 else 0 for i in range(batch_size)] | |
| calc.set_log_probs_mode(enable_log_probs, num_logprobs=num_logprobs_list) | |
| result = calc.calculate_topk_log_probs(logits_tt, topk_values_tt, topk_indices_tt) | |
| assert result is not None | |
| sampled_ids = argmax_tensor.squeeze() | |
| host_results = calc.transfer_logprobs_to_host(result, sampled_ids) | |
| for i in range(batch_size): | |
| if enable_log_probs[i]: | |
| assert host_results[i] is not None | |
| sampled_id = int(sampled_ids[i].item()) | |
| torch_lp = log_probs_torch[0, 0, i, sampled_id].item() | |
| assert abs(host_results[i]["returned_token"]["logprob"] - torch_lp) < 0.05 | |
| else: | |
| assert host_results[i] is None | |
| def test_set_log_probs_mode_validation(shape, mesh_device): | |
| """Verify set_log_probs_mode internal state.""" | |
| _skip_if_not_galaxy(mesh_device) | |
| batch_size = shape[2] | |
| calc = LogProbsCalculator(mesh_device, TG_SUB_CORE_GRIDS, batch_size=batch_size, use_topk_logprobs=True) | |
| calc.set_log_probs_mode(True) | |
| assert calc.enable_log_probs is True | |
| assert all(calc.logprobs_enabled) | |
| assert calc.topk_logprobs_needed # needed for sampled token logprob | |
| calc.set_log_probs_mode(True, num_logprobs=5) | |
| assert calc.topk_logprobs_needed is True | |
| assert all(n == 5 for n in calc.num_logprobs) | |
| enable_list = [True, False, True] + [False] * (batch_size - 3) | |
| num_lp_list = [10, 0, 3] + [0] * (batch_size - 3) | |
| calc.set_log_probs_mode(enable_list, num_logprobs=num_lp_list) | |
| assert calc.enable_log_probs is True | |
| assert calc.topk_logprobs_needed is True | |
| assert calc.logprobs_enabled == enable_list | |
| assert calc.num_logprobs == num_lp_list | |
| calc.set_log_probs_mode(False, num_logprobs=0) | |
| assert calc.enable_log_probs is False | |
| calc.set_log_probs_mode(True, num_logprobs=0) | |
| assert calc.enable_log_probs is True | |
| assert calc.topk_logprobs_needed # needed for sampled token logprob | |
| calc.set_log_probs_mode(False, num_logprobs=0) | |
| calc.set_log_probs_mode([True, True], num_logprobs=[10, 15], empty_slots=[2, 5]) | |
| assert calc.logprobs_enabled[2] is True | |
| assert calc.logprobs_enabled[5] is True | |
| assert calc.logprobs_enabled[0] is False | |
| assert calc.num_logprobs[2] == 10 | |
| assert calc.num_logprobs[5] == 15 | |
| calc.set_log_probs_mode(False, num_logprobs=0) | |
| calc.set_log_probs_mode(True, num_logprobs=7, empty_slots=[0, 3, 4]) | |
| assert all(calc.logprobs_enabled[i] for i in [0, 3, 4]) | |
| assert calc.logprobs_enabled[1] is False | |
| calc.set_log_probs_mode([True], num_logprobs=[20], empty_slots=[1]) | |
| assert calc.logprobs_enabled[1] is True | |
| assert calc.num_logprobs[0] == 7 | |
| assert calc.num_logprobs[1] == 20 | |
| def test_top_k_logprobs_pcc_torch_vs_tt(shape, mesh_device): | |
| """Compare host (PyTorch bfloat16) vs device (bfloat16) logprobs for full batch.""" | |
| _skip_if_not_galaxy(mesh_device) | |
| torch.manual_seed(9999) | |
| batch_size = shape[2] | |
| requested_logprobs = MAX_TOP_LOGPROBS | |
| calc = LogProbsCalculator(mesh_device, TG_SUB_CORE_GRIDS, batch_size=batch_size, use_topk_logprobs=True) | |
| torch_tensor = torch.randn(shape).to(torch.bfloat16) | |
| for i in range(batch_size): | |
| torch_tensor[:, :, i, :] = torch_tensor[:, :, i, torch.randperm(shape[-1])] | |
| log_probs_torch = F.log_softmax(torch_tensor, dim=-1, dtype=torch.bfloat16) | |
| gathered_values, gathered_indices = _simulate_gathered_topk(torch_tensor, TG_NUM_TP_DEVICES) | |
| argmax_tensor = torch.argmax(torch_tensor, dim=-1, keepdim=True) | |
| logits_tt, topk_values_tt, topk_indices_tt = _push_topk_test_tensors_to_tg( | |
| torch_tensor, gathered_values, gathered_indices, mesh_device | |
| ) | |
| calc.set_log_probs_mode([True] * batch_size, num_logprobs=[requested_logprobs] * batch_size) | |
| result = calc.calculate_topk_log_probs(logits_tt, topk_values_tt, topk_indices_tt) | |
| assert result is not None | |
| sampled_ids = argmax_tensor.squeeze() | |
| host_results = calc.transfer_logprobs_to_host(result, sampled_ids) | |
| for user in range(batch_size): | |
| r = host_results[user] | |
| assert r is not None | |
| device_sampled_lp = r["returned_token"]["logprob"] | |
| token_idx = r["returned_token"]["token_idx"] | |
| torch_sampled_lp = log_probs_torch[0, 0, user, token_idx].item() | |
| assert abs(device_sampled_lp - torch_sampled_lp) < 0.05 | |
| top_indices = r["top_logprobs"]["token_indices"] | |
| top_lps_device = torch.tensor(r["top_logprobs"]["logprobs"], dtype=torch.float32) | |
| assert len(top_indices) == requested_logprobs | |
| top_lps_torch = log_probs_torch[0, 0, user, top_indices].float() | |
| passing, pcc = comp_pcc(top_lps_torch.unsqueeze(0), top_lps_device.unsqueeze(0), pcc=0.98) | |
| assert passing, ( | |
| f"User {user} top-{requested_logprobs} logprobs PCC failed: {pcc}\n" | |
| f" device: {top_lps_device[:5].tolist()}...\n" | |
| f" torch: {top_lps_torch[:5].tolist()}..." | |
| ) | |
| # =========================================================================== | |
| # TTSampling top-k path on a single device | |
| # =========================================================================== | |
| def test_num_single_device_vocab_splits(padded_vocab_size, expected_splits): | |
| assert TTSampling.num_single_device_vocab_splits(padded_vocab_size) == expected_splits | |
| def test_untilize_chunk_count(width, expected): | |
| assert TTSampling._untilize_chunk_count(width) == expected | |
| def test_untilize_chunk_width(width, num_chunks, expected_split, expected_last): | |
| split = TTSampling._untilize_chunk_width(width, num_chunks) | |
| assert split == expected_split | |
| assert split % 32 == 0 | |
| assert -(-width // split) == num_chunks | |
| assert width - split * (num_chunks - 1) == expected_last | |
| def test_ttsampling_force_argmax_matches_row_max_on_wide_vocab(vocab_size, mesh_device): | |
| """Greedy params through the force-argmax fast path must pick the row maximum.""" | |
| torch.manual_seed(42) | |
| batch_size = 32 | |
| args = _single_device_sampling_args(mesh_device, vocab_size) | |
| args.model_config = {"SAMPLING_AG_CONFIG": {"allow_force_argmax": True, "num_links": 1, "topology": None}} | |
| sampler = TTSampling( | |
| args=args, | |
| mesh_device=mesh_device, | |
| tt_ccl=None, | |
| k=torch.ones(batch_size), | |
| p=torch.zeros(batch_size), | |
| temp=torch.ones(batch_size), | |
| ) | |
| assert sampler.force_argmax_sampling, "greedy params must take the argmax fast path" | |
| logits_host = torch.randn(1, 1, batch_size, vocab_size) | |
| # Exercise both sides of every chunk boundary, including the last element. | |
| # Negative logits make accidental zero padding observable as a wrong argmax. | |
| logits_host = -logits_host.abs() - 2 | |
| split = TTSampling._untilize_chunk_width(vocab_size, TTSampling._untilize_chunk_count(vocab_size)) | |
| boundary_indices = [0, vocab_size - 1] | |
| for boundary in range(split, vocab_size, split): | |
| boundary_indices.extend((boundary - 1, boundary)) | |
| for user in range(batch_size): | |
| logits_host[0, 0, user, boundary_indices[user % len(boundary_indices)]] = -1 | |
| logits_tt = ttnn.from_torch(logits_host, dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT, device=mesh_device) | |
| logits_bf16 = ttnn.to_torch(logits_tt).float().reshape(batch_size, vocab_size) | |
| tokens_tt, _log_probs = sampler(logits_tt) | |
| tokens = ttnn.to_torch(tokens_tt).flatten()[:batch_size].long() | |
| row_max = logits_bf16.max(dim=-1).values | |
| for user in range(batch_size): | |
| token = int(tokens[user]) | |
| assert 0 <= token < vocab_size, f"user {user}: token {token} outside [0, {vocab_size})" | |
| assert logits_bf16[user, token].item() == row_max[user].item(), ( | |
| f"user {user}: token {token} has logit {logits_bf16[user, token].item():.6f}, " | |
| f"but the row maximum is {row_max[user].item():.6f}" | |
| ) | |
| def test_uneven_untilize_preserves_logits_and_argmax_under_trace(mesh_device): | |
| # Isolate the argmax conversion: the single-device constructor also validates | |
| # a top-k split, which deliberately does not support this padded width. | |
| sampler = TTSampling.__new__(TTSampling) | |
| sampler._force_argmax_sub_core_grids = None | |
| width = 262208 | |
| split = sampler._untilize_chunk_width(width, sampler._untilize_chunk_count(width)) | |
| boundaries = [0, width - 1] | |
| for boundary in range(split, width, split): | |
| boundaries.extend((boundary - 1, boundary)) | |
| expected = torch.tensor([boundaries[row % len(boundaries)] for row in range(32)]) | |
| logits = torch.full((1, 1, 32, width), -2.0, dtype=torch.bfloat16) | |
| logits[0, 0, torch.arange(32), expected] = -1.0 | |
| device_logits = ttnn.from_torch(logits, device=mesh_device, dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT) | |
| untilized = sampler._untilize_for_argmax(device_logits) | |
| assert torch.equal(ttnn.to_torch(untilized), logits) | |
| tokens = ttnn.argmax(untilized, dim=-1, keepdim=False) | |
| assert torch.equal(ttnn.to_torch(tokens).flatten().long(), expected) | |
| ttnn.deallocate(tokens) | |
| ttnn.deallocate(untilized) | |
| trace_id = ttnn.begin_trace_capture(mesh_device, cq_id=0) | |
| untilized = sampler._untilize_for_argmax(device_logits) | |
| tokens = ttnn.argmax(untilized, dim=-1, keepdim=False) | |
| ttnn.end_trace_capture(mesh_device, trace_id, cq_id=0) | |
| try: | |
| # Reuse the capture with different maxima to detect stale inputs/outputs. | |
| for expected in (expected.flip(0), expected): | |
| logits.fill_(-2) | |
| logits[0, 0, torch.arange(32), expected] = -1 | |
| host_logits = ttnn.from_torch(logits, dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT) | |
| ttnn.copy_host_to_device_tensor(host_logits, device_logits) | |
| ttnn.execute_trace(mesh_device, trace_id, cq_id=0, blocking=True) | |
| assert torch.equal(ttnn.to_torch(tokens).flatten().long(), expected) | |
| finally: | |
| ttnn.release_trace(mesh_device, trace_id) | |
| def _single_device_sampling_args(mesh_device, vocab_size, max_top_k=32, max_batch_size=32): | |
| """Minimal args for TTSampling on a 1x1 mesh: no vocab padding, no force-argmax.""" | |
| grid = mesh_device.compute_with_storage_grid_size() | |
| sub_core_grids = ttnn.CoreRangeSet([ttnn.CoreRange(ttnn.CoreCoord(0, 0), ttnn.CoreCoord(grid.x - 1, grid.y - 1))]) | |
| return SimpleNamespace( | |
| vocab_size=vocab_size, | |
| padded_vocab_size=vocab_size, | |
| max_batch_size=max_batch_size, | |
| max_top_k=max_top_k, | |
| cluster_shape=(1, 1), | |
| sub_core_grids=sub_core_grids, | |
| sub_core_grid_topk=sub_core_grids, | |
| start_core=ttnn.CoreCoord(0, 0), | |
| ) | |
| def test_ttsampling_topk_matches_argmax_on_single_device(vocab_size, mesh_device): | |
| """top-k=1 through TTSampling must select the row maximum on a 1x1 mesh. | |
| On a single device TTSampling splits the logits in half and runs ttnn.topk on each half | |
| (``multi_step_reduction``), so this covers the split path end to end for both top-k program | |
| factories. Regression test for the half-width local indices buffer, which made the multi-core | |
| factory page its index tiles past the end of that buffer and return indices that did not | |
| belong to the values it returned. | |
| """ | |
| torch.manual_seed(42) | |
| batch_size = 32 | |
| sampler = TTSampling( | |
| args=_single_device_sampling_args(mesh_device, vocab_size), | |
| mesh_device=mesh_device, | |
| tt_ccl=None, | |
| k=torch.ones(batch_size), # top-1 | |
| p=torch.zeros(batch_size), | |
| temp=torch.ones(batch_size), | |
| ) | |
| assert sampler.multi_step_reduction, "a 1x1 mesh is expected to take the split top-k path" | |
| assert not sampler.force_argmax_sampling, "this test must exercise the top-k path, not argmax" | |
| logits_host = torch.randn(1, 1, batch_size, vocab_size) | |
| logits_tt = ttnn.from_torch(logits_host, dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT, device=mesh_device) | |
| # Compare against what the device actually holds: bfloat16 rounding creates ties, so the | |
| # sampled token need not be torch.argmax of the fp32 logits, but its value must be the maximum. | |
| logits_bf16 = ttnn.to_torch(logits_tt).float().reshape(batch_size, vocab_size) | |
| tokens_tt, _log_probs = sampler(logits_tt) | |
| tokens = ttnn.to_torch(tokens_tt).flatten()[:batch_size].long() | |
| row_max = logits_bf16.max(dim=-1).values | |
| failures = [] | |
| for user in range(batch_size): | |
| token = int(tokens[user]) | |
| if not 0 <= token < vocab_size: | |
| failures.append(f" user {user}: token {token} outside [0, {vocab_size})") | |
| continue | |
| value = logits_bf16[user, token].item() | |
| if value < row_max[user].item(): | |
| failures.append( | |
| f" user {user}: token {token} has logit {value:.6f}, but the row maximum is " | |
| f"{row_max[user].item():.6f} at index {int(logits_bf16[user].argmax())}" | |
| ) | |
| header = f"{len(failures)}/{batch_size} users did not sample the row maximum (vocab_size={vocab_size})" | |
| assert not failures, header + ":\n" + "\n".join(failures) | |
| def test_ttsampling_duplicate_request_seeds_sample_diverse_tokens(mesh_device): | |
| """A batch of users sharing one request seed must not all sample the same token. | |
| Device-level regression test for #53077: every user gets identical logits (a flat-ish | |
| top-32 so the multinomial draw is what differentiates them) and the identical request | |
| seed. Before per-slot seed salting, every slot derived the same device seed at the same | |
| position and the whole batch sampled the same token. The draw must also be reproducible: | |
| a fresh manager with the same seed produces the same tokens. | |
| """ | |
| torch.manual_seed(7) | |
| batch_size = 32 | |
| vocab_size = 32768 | |
| def _sample_once(): | |
| sampler = TTSampling( | |
| args=_single_device_sampling_args(mesh_device, vocab_size), | |
| mesh_device=mesh_device, | |
| tt_ccl=None, | |
| k=torch.full((batch_size,), 32), | |
| p=torch.ones(batch_size), | |
| temp=torch.ones(batch_size), | |
| ) | |
| assert not sampler.force_argmax_sampling | |
| seed_manager = SeedManager(sampler, max_batch_size=batch_size) | |
| seed_manager.reset_seed([1234] * batch_size, list(range(batch_size))) | |
| seed_manager.get_new_values(list(range(batch_size))) | |
| # One logits row replicated across the batch: only the RNG stream can differ. | |
| row = torch.zeros(1, 1, 1, vocab_size) | |
| row[..., :32] = 5.0 # 32 equally-likely candidates, everything else improbable | |
| logits_host = row.expand(1, 1, batch_size, vocab_size).contiguous() | |
| logits_tt = ttnn.from_torch(logits_host, dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT, device=mesh_device) | |
| tokens_tt, _ = sampler(logits_tt) | |
| return ttnn.to_torch(tokens_tt).flatten()[:batch_size].long().tolist() | |
| tokens_first = _sample_once() | |
| tokens_second = _sample_once() | |
| assert all(0 <= t < vocab_size for t in tokens_first) | |
| assert len(set(tokens_first)) > 1, ( | |
| f"all {batch_size} users with the same request seed sampled token {tokens_first[0]} -- " | |
| "duplicate-seed slots are drawing from identical RNG streams (#53077)" | |
| ) | |
| assert tokens_first == tokens_second, "same request seed must reproduce the same tokens across runs" | |
| def test_topk_authoritative_route_decision(mesh_device, shape, k, expected): | |
| """Exercise the C++ route policy used by both top-k and sampling. | |
| The cells pin important merged-policy boundaries (#53464: k_multiple=16, | |
| max_width=1<<19, no MoE-gate arm) without duplicating that policy in Python. | |
| """ | |
| x = ttnn.from_torch( | |
| torch.zeros(shape, dtype=torch.bfloat16), dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT, device=mesh_device | |
| ) | |
| # Routing is Blackhole-only: off-BH the predicate short-circuits to False, | |
| # so every cell's expectation collapses to False there (still asserted -- | |
| # this doubles as off-BH never-routes coverage). | |
| expected_here = expected and ttnn.device.is_blackhole(mesh_device) | |
| assert topk_would_route_to_large_indices(x, k) is expected_here | |
| ttnn.deallocate(x) | |
| def test_ttsampling_greedy_tied_max_picks_lowest_index(vocab_size, mesh_device): | |
| """Greedy rows must resolve an exact-value tie at the maximum to the LOWEST global index. | |
| This is the determinism contract that _adjust_values_for_tiebreak provides: among the | |
| GATHERED candidates, the lowest-global-index tied maximum wins. The tie plateau here | |
| (4 tokens) fits inside every per-shard top-k, so the lowest tied index is always | |
| gathered and the sampled token must be TIE_START exactly, for every user, on every | |
| path (routed full-row BH, chunked stock split, WH). | |
| A plateau WIDER than a shard's k additionally requires the top-k itself to keep the | |
| lowest indices (stable=True gather completeness) -- see the KNOWN LIMITATION note in | |
| _adjust_values_for_tiebreak. | |
| """ | |
| batch_size = 32 | |
| TIE_START = 777 # deliberately not 0: index 0 could win by zero-initialized accident | |
| TIE_LEN = 4 | |
| sampler = TTSampling( | |
| args=SimpleNamespace( | |
| vocab_size=vocab_size, | |
| padded_vocab_size=vocab_size, | |
| max_batch_size=batch_size, | |
| max_top_k=32, | |
| cluster_shape=(1, 1), | |
| ), | |
| mesh_device=mesh_device, | |
| tt_ccl=None, | |
| k=torch.ones(batch_size), # greedy: top-1 | |
| p=torch.zeros(batch_size), | |
| temp=torch.ones(batch_size), | |
| ) | |
| assert not sampler.force_argmax_sampling | |
| torch.manual_seed(7) | |
| # Tail strictly below the tie plateau (all values in [-2, -1)); plateau tied at 0.0. | |
| logits_host = torch.rand(1, 1, batch_size, vocab_size) - 2.0 | |
| logits_host[..., TIE_START : TIE_START + TIE_LEN] = 0.0 | |
| logits_tt = ttnn.from_torch(logits_host, dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT, device=mesh_device) | |
| tokens_tt, _log_probs = sampler(logits_tt) | |
| tokens = ttnn.to_torch(tokens_tt).flatten()[:batch_size].long() | |
| mismatched = [(u, int(tokens[u])) for u in range(batch_size) if int(tokens[u]) != TIE_START] | |
| assert not mismatched, ( | |
| f"{len(mismatched)}/{batch_size} greedy users did not pick the lowest tied index " | |
| f"{TIE_START}: {mismatched[:8]}" | |
| ) | |
| def test_ttsampling_routed_full_row_topk_end_to_end(vocab_size, mesh_device): | |
| """End-to-end TTSampling over the ROUTED full-row ttnn.topk path (Blackhole only). | |
| With sub_core_grid_topk=None at a routing-eligible width, TTSampling replaces the | |
| single-device vocab split with one full-row ttnn.topk that takes the Blackhole | |
| topk_large_indices composite. Greedy (top-1) users must still sample a row maximum | |
| of the bf16 logits, and every token must be a valid global vocab position -- i.e. | |
| the routed indices are global, not per-chunk. | |
| """ | |
| torch.manual_seed(42) | |
| batch_size = 32 | |
| sampler = TTSampling( | |
| args=SimpleNamespace( | |
| vocab_size=vocab_size, | |
| padded_vocab_size=vocab_size, | |
| max_batch_size=batch_size, | |
| max_top_k=32, | |
| cluster_shape=(1, 1), | |
| ), | |
| mesh_device=mesh_device, | |
| tt_ccl=None, | |
| k=torch.ones(batch_size), # top-1 | |
| p=torch.zeros(batch_size), | |
| temp=torch.ones(batch_size), | |
| ) | |
| assert sampler.multi_step_reduction | |
| assert sampler.sub_core_grid_topk is None | |
| assert not sampler.force_argmax_sampling | |
| logits_host = torch.randn(1, 1, batch_size, vocab_size) | |
| logits_tt = ttnn.from_torch(logits_host, dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT, device=mesh_device) | |
| # The exact gate the forward pass evaluates: this parametrization must route. | |
| assert topk_would_route_to_large_indices( | |
| logits_tt, sampler._num_vocab_splits * sampler.max_top_k | |
| ), "test premise broken: this cell no longer routes -- update the parametrization" | |
| logits_bf16 = ttnn.to_torch(logits_tt).float().reshape(batch_size, vocab_size) | |
| tokens_tt, _log_probs = sampler(logits_tt) | |
| tokens = ttnn.to_torch(tokens_tt).flatten()[:batch_size].long() | |
| row_max = logits_bf16.max(dim=-1).values | |
| failures = [] | |
| for user in range(batch_size): | |
| token = int(tokens[user]) | |
| if not 0 <= token < vocab_size: | |
| failures.append(f" user {user}: token {token} outside [0, {vocab_size})") | |
| continue | |
| value = logits_bf16[user, token].item() | |
| if value < row_max[user].item(): | |
| failures.append( | |
| f" user {user}: token {token} has logit {value:.6f}, but the row maximum is " | |
| f"{row_max[user].item():.6f} at index {int(logits_bf16[user].argmax())}" | |
| ) | |
| header = f"{len(failures)}/{batch_size} users did not sample the row maximum (vocab_size={vocab_size})" | |
| assert not failures, header + ":\n" + "\n".join(failures) | |
| def _routed_single_device_args(mesh_device, vocab_size, max_top_k=32, max_batch_size=32): | |
| """_single_device_sampling_args with the top-k sub-grid released. | |
| TTSampling only relaxes its call shape when sub_core_grid_topk is None, so this is the | |
| one knob that separates the routed full-row path from the chunked split. Everything else | |
| (including sub_core_grids, which only places programs) stays identical to the chunked | |
| helper, so a routed-vs-chunked comparison varies exactly one thing. | |
| """ | |
| args = _single_device_sampling_args(mesh_device, vocab_size, max_top_k, max_batch_size) | |
| args.sub_core_grid_topk = None | |
| return args | |
| def test_ttsampling_routed_full_row_random_sampling_stays_in_top_k(vocab_size, mesh_device): | |
| """The ROUTED path must respect per-user top-k for RANDOM users (k > 1), not just greedy ones. | |
| Every other routed-path test runs k=1, where the routed and chunked candidate sets are | |
| provably equivalent (both contain the global maximum), so they cannot observe the thing | |
| this path actually changes: the candidate row is now the global top-(num_splits*k) instead | |
| of the union of the per-chunk top-ks. ttnn.sampling sorts and masks that row itself, so a | |
| k=32 draw must still land inside the true global top-32. | |
| Asserted against the k-th largest VALUE rather than the top-k index set: bfloat16 rounding | |
| creates ties at the boundary, so a token outside torch's top-k indices can still be a | |
| legitimate draw if its value ties the k-th largest. | |
| """ | |
| torch.manual_seed(1234) | |
| batch_size = 32 | |
| top_k = 32 | |
| sampler = TTSampling( | |
| args=_routed_single_device_args(mesh_device, vocab_size, max_top_k=top_k), | |
| mesh_device=mesh_device, | |
| tt_ccl=None, | |
| k=torch.full((batch_size,), top_k), # random sampling, not greedy | |
| p=torch.ones(batch_size), # no nucleus filtering | |
| temp=torch.ones(batch_size), | |
| ) | |
| assert sampler.multi_step_reduction | |
| assert sampler.sub_core_grid_topk is None | |
| assert not sampler.force_argmax_sampling, "this test must exercise the sampling path, not argmax" | |
| logits_host = torch.randn(1, 1, batch_size, vocab_size) | |
| logits_tt = ttnn.from_torch(logits_host, dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT, device=mesh_device) | |
| assert topk_would_route_to_large_indices( | |
| logits_tt, sampler._num_vocab_splits * sampler.max_top_k | |
| ), "test premise broken: this cell no longer routes -- update the parametrization" | |
| logits_bf16 = ttnn.to_torch(logits_tt).float().reshape(batch_size, vocab_size) | |
| tokens = ttnn.to_torch(sampler(logits_tt)[0]).flatten()[:batch_size].long() | |
| # Value of the k-th largest entry per row: any legitimate top-k draw is >= this. | |
| kth_value = logits_bf16.topk(top_k, dim=-1).values[:, -1] | |
| failures = [] | |
| for user in range(batch_size): | |
| token = int(tokens[user]) | |
| if not 0 <= token < vocab_size: | |
| failures.append(f" user {user}: token {token} outside [0, {vocab_size})") | |
| continue | |
| value = logits_bf16[user, token].item() | |
| if value < kth_value[user].item(): | |
| failures.append( | |
| f" user {user}: token {token} has logit {value:.6f}, below the top-{top_k} " | |
| f"cutoff {kth_value[user].item():.6f} -- drawn from outside the candidate set" | |
| ) | |
| assert not failures, f"{len(failures)}/{batch_size} random users drew outside top-{top_k}:\n" + "\n".join(failures) | |
| # A k=32 draw with per-user seeds must not collapse to one token for the whole batch; | |
| # that would mean the candidate row or the RNG stream degenerated. | |
| assert len(set(tokens.tolist())) > 1, f"all {batch_size} random users drew the same token {int(tokens[0])}" | |
| def test_ttsampling_routed_full_row_matches_chunked_path(vocab_size, mesh_device): | |
| """The routed full-row top-k must produce the SAME tokens as the chunked split it replaces. | |
| The other routed tests assert only that the sampled token is a row maximum -- an invariant | |
| the chunked path already satisfied before de-chunking, so they cannot catch the two paths | |
| diverging. This runs both on identical logits with only sub_core_grid_topk differing | |
| (None -> routed full row, set -> chunked split) and requires bit-identical tokens. | |
| Greedy (k=1) on purpose: it makes the comparison deterministic without depending on the two | |
| samplers drawing identical RNG streams. bfloat16 rounding of randn produces ties at the | |
| maximum for some rows, so this also pins that _adjust_values_for_tiebreak resolves a tie to | |
| the same global index on both paths. | |
| """ | |
| torch.manual_seed(7) | |
| batch_size = 32 | |
| def sample(args): | |
| sampler = TTSampling( | |
| args=args, | |
| mesh_device=mesh_device, | |
| tt_ccl=None, | |
| k=torch.ones(batch_size), # greedy: deterministic, no RNG dependence | |
| p=torch.zeros(batch_size), | |
| temp=torch.ones(batch_size), | |
| ) | |
| assert sampler.multi_step_reduction | |
| assert not sampler.force_argmax_sampling | |
| logits_tt = ttnn.from_torch(logits_host, dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT, device=mesh_device) | |
| routes = sampler.sub_core_grid_topk is None and topk_would_route_to_large_indices( | |
| logits_tt, sampler._num_vocab_splits * sampler.max_top_k | |
| ) | |
| tokens = ttnn.to_torch(sampler(logits_tt)[0]).flatten()[:batch_size].long().tolist() | |
| return tokens, routes | |
| logits_host = torch.randn(1, 1, batch_size, vocab_size) | |
| routed_tokens, routed = sample(_routed_single_device_args(mesh_device, vocab_size)) | |
| chunked_tokens, chunked_routes = sample(_single_device_sampling_args(mesh_device, vocab_size)) | |
| assert routed, "test premise broken: the sub_core_grid_topk=None arm did not take the routed path" | |
| assert not chunked_routes, "test premise broken: the sub-grid arm must stay on the chunked path" | |
| mismatched = [ | |
| (user, routed_tokens[user], chunked_tokens[user]) | |
| for user in range(batch_size) | |
| if routed_tokens[user] != chunked_tokens[user] | |
| ] | |
| assert not mismatched, ( | |
| f"{len(mismatched)}/{batch_size} users got a different token from the routed full-row " | |
| f"top-k than from the chunked split (vocab_size={vocab_size}); " | |
| f"(user, routed, chunked): {mismatched[:8]}" | |
| ) | |
| def test_ttsampling_routed_full_row_resolves_tie_plateau_wider_than_shard_k(vocab_size, mesh_device): | |
| """A tie plateau WIDER than one shard's top-k still resolves to the lowest global index. | |
| This is the case test_ttsampling_greedy_tied_max_picks_lowest_index documents but does not | |
| cover: its 4-token plateau fits inside every per-chunk top-k, so both paths gather all of it. | |
| With TIE_LEN=48 > max_top_k=32 the chunked path's first chunk must drop 16 of the tied maxima | |
| through the unstable top-k network, and the lowest tied index can be among the dropped ones -- | |
| the KNOWN LIMITATION on _adjust_values_for_tiebreak. | |
| The routed full row takes top-(num_splits * max_top_k) = 64 in ONE call, so all 48 tied | |
| maxima reach the gathered candidate set and the lowest-index tie-break is exact. That is the | |
| concrete determinism win of de-chunking, so it is asserted on the routed path only. | |
| """ | |
| batch_size = 32 | |
| TIE_START = 777 # deliberately not 0: index 0 could win by zero-initialized accident | |
| TIE_LEN = 48 # > max_top_k (32), <= num_vocab_splits * max_top_k (64) | |
| sampler = TTSampling( | |
| args=_routed_single_device_args(mesh_device, vocab_size), | |
| mesh_device=mesh_device, | |
| tt_ccl=None, | |
| k=torch.ones(batch_size), # greedy: top-1 | |
| p=torch.zeros(batch_size), | |
| temp=torch.ones(batch_size), | |
| ) | |
| assert sampler.multi_step_reduction | |
| assert sampler.sub_core_grid_topk is None | |
| assert not sampler.force_argmax_sampling | |
| assert TIE_LEN > sampler.max_top_k, "plateau must exceed one shard's k or this adds no coverage" | |
| assert TIE_LEN <= sampler._num_vocab_splits * sampler.max_top_k, "plateau must fit the routed candidate row" | |
| torch.manual_seed(7) | |
| # Tail strictly below the plateau (all values in [-2, -1)); the plateau is tied at exactly 0.0. | |
| logits_host = torch.rand(1, 1, batch_size, vocab_size) - 2.0 | |
| logits_host[..., TIE_START : TIE_START + TIE_LEN] = 0.0 | |
| logits_tt = ttnn.from_torch(logits_host, dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT, device=mesh_device) | |
| assert topk_would_route_to_large_indices( | |
| logits_tt, sampler._num_vocab_splits * sampler.max_top_k | |
| ), "test premise broken: this cell no longer routes -- update the parametrization" | |
| tokens = ttnn.to_torch(sampler(logits_tt)[0]).flatten()[:batch_size].long() | |
| mismatched = [(u, int(tokens[u])) for u in range(batch_size) if int(tokens[u]) != TIE_START] | |
| assert not mismatched, ( | |
| f"{len(mismatched)}/{batch_size} greedy users did not pick the lowest index {TIE_START} of a " | |
| f"{TIE_LEN}-wide tie plateau: {mismatched[:8]}" | |
| ) | |