Download code/models/common/tests/modules/sampling/test_sampling_1d.py from tt-hous/clef: direct link, hf CLI and curl.
- Browser
- Download file 71.4 kB
-
https://huggingface.co/tt-hous/clef/resolve/main/code/models/common/tests/modules/sampling/test_sampling_1d.py
- Command line
-
hf download hf://tt-hous/clef/code/models/common/tests/modules/sampling/test_sampling_1d.py
-
curl -L -o test_sampling_1d.py https://huggingface.co/tt-hous/clef/resolve/main/code/models/common/tests/modules/sampling/test_sampling_1d.py
71.4 kB
| # SPDX-FileCopyrightText: © 2025 Tenstorrent USA, Inc. | |
| # SPDX-License-Identifier: Apache-2.0 | |
| """Tests for Sampling1D module.""" | |
| import pytest | |
| import torch | |
| import ttnn | |
| from models.common.auto_compose import to_torch_auto_compose | |
| from models.common.modules.sampling.sampling_1d import Sampling1D, Sampling1DConfig, _resolve_sampling1d_config | |
| # 1D module suites target the T3K; skip when the host system is a Galaxy. | |
| pytestmark = pytest.mark.usefixtures("skip_on_galaxy_system") | |
| # --------------------------------------------------------------------------- | |
| # Model name constants (match test_mlp_1d.py naming convention) | |
| # --------------------------------------------------------------------------- | |
| LLAMA_1B = "meta-llama/Llama-3.2-1B-Instruct" | |
| LLAMA_3B = "meta-llama/Llama-3.2-3B-Instruct" | |
| LLAMA_8B = "meta-llama/Llama-3.1-8B-Instruct" | |
| LLAMA_11B = "meta-llama/Llama-3.2-11B-Vision-Instruct" | |
| LLAMA_70B = "meta-llama/Llama-3.3-70B-Instruct" | |
| MISTRAL_7B = "mistralai/Mistral-7B-Instruct-v0.3" | |
| MIXTRAL_8X7B = "mistralai/Mixtral-8x7B-v0.1" | |
| QWEN25_72B = "Qwen/Qwen2.5-72B-Instruct" | |
| QWEN3_32B = "Qwen/Qwen3-32B" | |
| _slow = pytest.mark.slow | |
| def _sub_core_grids_for_32_users(): | |
| return ttnn.CoreRangeSet({ttnn.CoreRange(ttnn.CoreCoord(0, 0), ttnn.CoreCoord(7, 3))}) | |
| def _list_collected_sampling_cases() -> list[pytest.param]: | |
| """ | |
| Collected from TTTv1 demo runs (Phase B of test_case_collection.md). | |
| Each entry is: | |
| (mesh_shape, vocab_size, k, p, temp, force_argmax, hf_model_name) | |
| Source CSVs: sampling_generator_config_collected.csv, | |
| sampling_generator_params_collected.csv | |
| Deduplicated by (topology, vocab, k, p, temp, force_argmax). | |
| """ | |
| # fmt: off | |
| return [ | |
| # --- (1,1) Mistral7B v32768 --- | |
| pytest.param((1, 1), 32768, 1, 1.0, 1.0, False, MISTRAL_7B, id="1x1-Mistral7B-v32768-k1-p1.0-t1.0"), | |
| pytest.param((1, 1), 32768, 10, 0.9, 1.0, False, MISTRAL_7B, id="1x1-Mistral7B-v32768-k10-p0.9-t1.0", marks=_slow), | |
| # --- (1,1) Llama8B v128256 --- | |
| pytest.param((1, 1), 128256, 1, 1.0, 1.0, True, LLAMA_8B, id="1x1-Llama8B-v128256-k1-p1.0-t1.0-argmax"), | |
| # --- (1,1) Llama1B v128256 --- | |
| pytest.param((1, 1), 128256, 1, 0.08, 1.0, False, LLAMA_1B, id="1x1-Llama1B-v128256-k1-p0.08-t1.0", marks=_slow), | |
| pytest.param((1, 1), 128256, 1, 1.0, 1.0, False, LLAMA_1B, id="1x1-Llama1B-v128256-k1-p1.0-t1.0", marks=_slow), | |
| pytest.param((1, 1), 128256, 10, 0.9, 1.0, False, LLAMA_1B, id="1x1-Llama1B-v128256-k10-p0.9-t1.0", marks=_slow), | |
| # --- (1,2) Mistral7B v32768 --- | |
| pytest.param((1, 2), 32768, 1, 1.0, 1.0, False, MISTRAL_7B, id="1x2-Mistral7B-v32768-k1-p1.0-t1.0"), | |
| pytest.param((1, 2), 32768, 10, 0.9, 1.0, False, MISTRAL_7B, id="1x2-Mistral7B-v32768-k10-p0.9-t1.0", marks=_slow), | |
| # --- (1,2) Llama8B v128256 --- | |
| pytest.param((1, 2), 128256, 1, 1.0, 1.0, True, LLAMA_8B, id="1x2-Llama8B-v128256-k1-p1.0-t1.0-argmax"), | |
| # --- (1,2) Llama1B v128256 --- | |
| pytest.param((1, 2), 128256, 1, 0.08, 1.0, False, LLAMA_1B, id="1x2-Llama1B-v128256-k1-p0.08-t1.0", marks=_slow), | |
| pytest.param((1, 2), 128256, 1, 1.0, 1.0, False, LLAMA_1B, id="1x2-Llama1B-v128256-k1-p1.0-t1.0", marks=_slow), | |
| pytest.param((1, 2), 128256, 10, 0.9, 1.0, False, LLAMA_1B, id="1x2-Llama1B-v128256-k10-p0.9-t1.0", marks=_slow), | |
| # --- (1,8) Mixtral8x7B v32000 --- | |
| pytest.param((1, 8), 32000, 1, 0.08, 1.0, False, MIXTRAL_8X7B, id="1x8-Mixtral8x7B-v32000-k1-p0.08-t1.0"), | |
| pytest.param((1, 8), 32000, 1, 1.0, 1.0, False, MIXTRAL_8X7B, id="1x8-Mixtral8x7B-v32000-k1-p1.0-t1.0", marks=_slow), | |
| pytest.param((1, 8), 32000, 10, 0.9, 1.0, False, MIXTRAL_8X7B, id="1x8-Mixtral8x7B-v32000-k10-p0.9-t1.0", marks=_slow), | |
| # --- (1,8) Mistral7B v32768 --- | |
| pytest.param((1, 8), 32768, 1, 1.0, 1.0, False, MISTRAL_7B, id="1x8-Mistral7B-v32768-k1-p1.0-t1.0"), | |
| pytest.param((1, 8), 32768, 10, 0.9, 1.0, False, MISTRAL_7B, id="1x8-Mistral7B-v32768-k10-p0.9-t1.0", marks=_slow), | |
| # --- (1,8) Llama8B v128256 --- | |
| pytest.param((1, 8), 128256, 1, 1.0, 1.0, True, LLAMA_8B, id="1x8-Llama8B-v128256-k1-p1.0-t1.0-argmax"), | |
| # --- (1,8) Llama1B v128256 --- | |
| pytest.param((1, 8), 128256, 1, 0.08, 1.0, False, LLAMA_1B, id="1x8-Llama1B-v128256-k1-p0.08-t1.0", marks=_slow), | |
| pytest.param((1, 8), 128256, 1, 1.0, 1.0, False, LLAMA_1B, id="1x8-Llama1B-v128256-k1-p1.0-t1.0", marks=_slow), | |
| pytest.param((1, 8), 128256, 10, 0.9, 1.0, False, LLAMA_1B, id="1x8-Llama1B-v128256-k10-p0.9-t1.0", marks=_slow), | |
| # --- (1,8) Qwen3-32B v151936 --- | |
| pytest.param((1, 8), 151936, 10, 0.9, 1.0, False, QWEN3_32B, id="1x8-Qwen3-32B-v151936-k10-p0.9-t1.0"), | |
| # --- (1,8) Qwen2.5-72B v152064 --- | |
| pytest.param((1, 8), 152064, 1, 0.08, 1.0, False, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-k1-p0.08-t1.0"), | |
| pytest.param((1, 8), 152064, 1, 1.0, 1.0, False, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-k1-p1.0-t1.0", marks=_slow), | |
| pytest.param((1, 8), 152064, 10, 0.9, 1.0, False, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-k10-p0.9-t1.0", marks=_slow), | |
| ] | |
| # fmt: on | |
| # ============================================================================== | |
| # Unit tests: Config (no device) | |
| # ============================================================================== | |
| class TestConfigUnit: | |
| def test_config_defaults(self): | |
| cfg = Sampling1DConfig(vocab_size=1024) | |
| assert cfg.max_batch_size == 32 | |
| assert cfg.max_top_k == 32 | |
| assert cfg.allow_force_argmax is False | |
| assert cfg.num_gather_links == 1 | |
| assert cfg.mesh_device is None | |
| assert cfg.index_offsets is None | |
| assert cfg.seeds is None | |
| def test_config_custom(self): | |
| cfg = Sampling1DConfig(vocab_size=128256, max_top_k=64, allow_force_argmax=True) | |
| assert cfg.vocab_size == 128256 | |
| assert cfg.max_top_k == 64 | |
| assert cfg.allow_force_argmax is True | |
| def test_config_not_resolved_without_device(self): | |
| cfg = Sampling1DConfig(vocab_size=1024) | |
| assert not cfg.is_resolved() | |
| def test_config_not_resolved_multi_device_no_ccl(self): | |
| """is_resolved() returns False when multi-device mesh but tt_ccl is None (line 65).""" | |
| from unittest.mock import MagicMock | |
| mock_device = MagicMock() | |
| mock_device.get_num_devices.return_value = 2 | |
| cfg = Sampling1DConfig(vocab_size=1024, mesh_device=mock_device, tt_ccl=None) | |
| assert not cfg.is_resolved() | |
| # ============================================================================== | |
| # Device tests | |
| # ============================================================================== | |
| class TestSampling1DDevice: | |
| def test_resolve_config(self, ttnn_mesh_device, vocab_size): | |
| cfg = Sampling1DConfig(vocab_size=vocab_size, mesh_device=ttnn_mesh_device) | |
| resolved = _resolve_sampling1d_config(cfg) | |
| assert resolved.is_resolved() | |
| assert resolved.start_core is not None | |
| assert resolved.sampling_memory_config is not None | |
| assert resolved.index_offsets is not None | |
| assert resolved.seeds is not None | |
| def test_load_device_buffers(self, ttnn_mesh_device, vocab_size): | |
| sampler = Sampling1D(vocab_size=vocab_size, mesh_device=ttnn_mesh_device) | |
| sampler.load_device_buffers() | |
| assert sampler._device_buffers_loaded | |
| assert isinstance(sampler._index_offsets, ttnn.Tensor) | |
| assert isinstance(sampler._seeds, ttnn.Tensor) | |
| assert isinstance(sampler._user_ids, ttnn.Tensor) | |
| def test_force_argmax(self, ttnn_mesh_device, vocab_size): | |
| """allow_force_argmax=True, k/p/temp=None → matches torch.argmax.""" | |
| sampler = Sampling1D(vocab_size=vocab_size, mesh_device=ttnn_mesh_device, allow_force_argmax=True) | |
| sampler.load_device_buffers() | |
| B = sampler.config.max_batch_size | |
| torch.manual_seed(42) | |
| logits_host = torch.randn(1, 1, B, vocab_size, dtype=torch.bfloat16) | |
| logits_tt = _make_logits_tt(logits_host, ttnn_mesh_device) | |
| tokens_tt, log_probs = sampler.decode_forward(logits_tt) | |
| tokens_host = to_torch_auto_compose(tokens_tt) | |
| expected_argmax = logits_host.float().argmax(dim=-1) | |
| # .long(): ttnn.argmax → uint32, torch.argmax → int64; normalize before compare | |
| tokens_flat = tokens_host.flatten()[:B].long() | |
| expected_flat = expected_argmax.flatten()[:B].long() | |
| assert torch.equal( | |
| tokens_flat, expected_flat | |
| ), f"Argmax mismatch: got {tokens_flat[:5]} vs expected {expected_flat[:5]}" | |
| def test_error_on_partial_params(self, ttnn_mesh_device, vocab_size, expect_error): | |
| """k provided but not p/temp → ValueError.""" | |
| sampler = Sampling1D(vocab_size=vocab_size, mesh_device=ttnn_mesh_device) | |
| logits_host = torch.randn(1, 1, 32, vocab_size, dtype=torch.bfloat16) | |
| logits_tt = _make_logits_tt(logits_host, ttnn_mesh_device) | |
| k_tt = ttnn.from_torch(torch.ones(32), device=ttnn_mesh_device, dtype=ttnn.uint32, layout=ttnn.ROW_MAJOR_LAYOUT) | |
| with expect_error(ValueError, "k, p, temp must all be provided"): | |
| sampler.decode_forward(logits_tt, k=k_tt) | |
| def test_from_model_args(self, ttnn_mesh_device, vocab_size): | |
| """from_model_args backward compat factory.""" | |
| class MockArgs: | |
| padded_vocab_size = vocab_size | |
| sub_core_grids = None | |
| sub_core_grid_topk = None | |
| start_core = ttnn.CoreCoord(0, 0) | |
| max_top_k = 32 | |
| sampler = Sampling1D.from_model_args(ttnn_mesh_device, None, MockArgs()) | |
| assert sampler.config.vocab_size == vocab_size | |
| assert sampler.config.mesh_device is ttnn_mesh_device | |
| # ------------------------------------------------------------------ | |
| # CCL introspection (_bind_strategy lines 116-126) | |
| # ------------------------------------------------------------------ | |
| def test_bind_strategy_ccl_introspection_with_kwargs(self, ttnn_mesh_device): | |
| """_bind_strategy correctly detects buffer_key support on line_all_gather.""" | |
| from dataclasses import replace | |
| sampler = Sampling1D(vocab_size=1024, mesh_device=ttnn_mesh_device) | |
| class MockCCL: | |
| def line_all_gather(self, tensor, dim, cluster_axis, memory_config, num_links, buffer_key=None): | |
| return tensor | |
| sampler.config = replace(sampler.config, tt_ccl=MockCCL()) | |
| sampler._bind_strategy() | |
| assert sampler._line_all_gather_supports_buffer_key | |
| def test_bind_strategy_ccl_introspection_no_kwargs(self, ttnn_mesh_device): | |
| """_bind_strategy detects when line_all_gather does NOT support buffer_key.""" | |
| from dataclasses import replace | |
| sampler = Sampling1D(vocab_size=1024, mesh_device=ttnn_mesh_device) | |
| class MockCCL: | |
| def line_all_gather(self, tensor, dim, cluster_axis, memory_config, num_links): | |
| return tensor | |
| sampler.config = replace(sampler.config, tt_ccl=MockCCL()) | |
| sampler._bind_strategy() | |
| assert not sampler._line_all_gather_supports_buffer_key | |
| def test_bind_strategy_ccl_introspection_exception(self, ttnn_mesh_device): | |
| """_bind_strategy handles TypeError from inspect.signature gracefully (lines 125-126).""" | |
| from dataclasses import replace | |
| from unittest.mock import patch | |
| sampler = Sampling1D(vocab_size=1024, mesh_device=ttnn_mesh_device) | |
| class MockCCL: | |
| def line_all_gather(self, *args, **kwargs): | |
| return args[0] | |
| sampler.config = replace(sampler.config, tt_ccl=MockCCL()) | |
| with patch( | |
| "models.common.modules.sampling.sampling_1d.inspect.signature", side_effect=TypeError("Cannot inspect") | |
| ): | |
| sampler._bind_strategy() | |
| assert not sampler._line_all_gather_supports_buffer_key | |
| # ------------------------------------------------------------------ | |
| # Error paths (lines 178, 186) | |
| # ------------------------------------------------------------------ | |
| def test_error_all_none_no_force_argmax(self, ttnn_mesh_device, vocab_size, expect_error): | |
| """decode_forward with all-None k/p/temp when allow_force_argmax=False → ValueError (line 178).""" | |
| sampler = Sampling1D(vocab_size=vocab_size, mesh_device=ttnn_mesh_device) | |
| logits_host = torch.randn(1, 1, 32, vocab_size, dtype=torch.bfloat16) | |
| logits_tt = _make_logits_tt(logits_host, ttnn_mesh_device) | |
| with expect_error(ValueError, "allow_force_argmax is False"): | |
| sampler.decode_forward(logits_tt) | |
| def test_forward_dispatches_to_decode_forward(self, ttnn_mesh_device, vocab_size): | |
| """forward() delegates to decode_forward() (line 186).""" | |
| sampler = Sampling1D(vocab_size=vocab_size, mesh_device=ttnn_mesh_device, allow_force_argmax=True) | |
| logits_host = torch.randn(1, 1, 32, vocab_size, dtype=torch.bfloat16) | |
| logits_tt = _make_logits_tt(logits_host, ttnn_mesh_device) | |
| result = sampler.forward(logits_tt) | |
| assert result is not None | |
| assert len(result) == 2 # (token_ids, log_probs) | |
| # ------------------------------------------------------------------ | |
| # _perform_all_gather with line_all_gather (lines 367-377) | |
| # ------------------------------------------------------------------ | |
| def test_perform_all_gather_with_mock_ccl(self, ttnn_mesh_device, vocab_size): | |
| """_perform_all_gather passes the buffer_key kwarg when line_all_gather supports it.""" | |
| B, K = 32, 32 | |
| sampler = Sampling1D(vocab_size=vocab_size, mesh_device=ttnn_mesh_device) | |
| sampler.load_device_buffers() | |
| captured_kwargs = {} | |
| def mock_line_ag(tensor, **kwargs): | |
| captured_kwargs.update(kwargs) | |
| return tensor | |
| sampler._line_all_gather = mock_line_ag | |
| sampler._line_all_gather_supports_buffer_key = True | |
| test_tensor = ttnn.from_torch( | |
| torch.zeros(1, 1, B, K, dtype=torch.bfloat16), | |
| device=ttnn_mesh_device, | |
| dtype=ttnn.bfloat16, | |
| layout=ttnn.TILE_LAYOUT, | |
| memory_config=ttnn.DRAM_MEMORY_CONFIG, | |
| ) | |
| result = sampler._perform_all_gather( | |
| test_tensor, | |
| dim=3, | |
| cluster_axis=None, | |
| memory_config=ttnn.DRAM_MEMORY_CONFIG, | |
| num_links=1, | |
| buffer_key="TEST_KEY", | |
| ) | |
| assert result is test_tensor | |
| assert captured_kwargs.get("buffer_key") == "TEST_KEY" | |
| # ------------------------------------------------------------------ | |
| # from_model_args model_config branches (lines 406-408, 416-419) | |
| # ------------------------------------------------------------------ | |
| def test_from_model_args_with_galaxy_num_links(self, ttnn_mesh_device): | |
| """from_model_args reads num_gather_links from GALAXY_NUM_LINKS in model_config (lines 406-408).""" | |
| class MockArgs: | |
| padded_vocab_size = 1024 | |
| sub_core_grids = None | |
| sub_core_grid_topk = None | |
| start_core = ttnn.CoreCoord(0, 0) | |
| max_top_k = 32 | |
| model_config = {"GALAXY_NUM_LINKS": 4} | |
| sampler = Sampling1D.from_model_args(ttnn_mesh_device, None, MockArgs(), model_config=model_config) | |
| # max_top_k=32 → 32//32=1, max_links=4 → min(1, 4) = 1 | |
| assert sampler.config.num_gather_links == 1 | |
| def test_from_model_args_with_sampling_ag_config(self, ttnn_mesh_device): | |
| """from_model_args reads allow_force_argmax/num_links/topology from SAMPLING_AG_CONFIG (lines 416-419).""" | |
| class MockArgs: | |
| padded_vocab_size = 1024 | |
| sub_core_grids = None | |
| sub_core_grid_topk = None | |
| start_core = ttnn.CoreCoord(0, 0) | |
| max_top_k = 32 | |
| model_config = { | |
| "SAMPLING_AG_CONFIG": { | |
| "allow_force_argmax": True, | |
| "num_links": 3, | |
| "topology": ttnn.Topology.Linear, | |
| } | |
| } | |
| sampler = Sampling1D.from_model_args(ttnn_mesh_device, None, MockArgs(), model_config=model_config) | |
| assert sampler.config.allow_force_argmax is True | |
| assert sampler.config.num_argmax_gather_links == 3 | |
| assert sampler.config.ag_topology == ttnn.Topology.Linear | |
| # ------------------------------------------------------------------ | |
| # Buffer passthrough: _resolve_buf ttnn.Tensor path (lines 493-494) | |
| # and _materialize ttnn.Tensor path (line 554) | |
| # ------------------------------------------------------------------ | |
| def test_resolve_buf_tensor_passthrough_and_materialize(self, ttnn_mesh_device, vocab_size): | |
| """Pre-existing ttnn.Tensor passes through _resolve_buf (493-494) and _materialize (554).""" | |
| cluster_shape = tuple(ttnn_mesh_device.shape) | |
| num_devices_in_mesh = 2 if list(cluster_shape) == [1, 1] else max(cluster_shape) | |
| B, K = 32, 32 | |
| replicate_mapper = ttnn.ShardTensor2dMesh(ttnn_mesh_device, dims=(None, None), mesh_shape=cluster_shape) | |
| offsets_host = torch.zeros(1, 1, B, K * num_devices_in_mesh, dtype=torch.int64) | |
| pre_tensor = ttnn.from_torch( | |
| offsets_host, | |
| device=ttnn_mesh_device, | |
| dtype=ttnn.int32, | |
| layout=ttnn.TILE_LAYOUT, | |
| mesh_mapper=replicate_mapper, | |
| memory_config=ttnn.DRAM_MEMORY_CONFIG, | |
| ) | |
| cfg = Sampling1DConfig(vocab_size=vocab_size, mesh_device=ttnn_mesh_device, index_offsets=pre_tensor) | |
| resolved = _resolve_sampling1d_config(cfg) | |
| assert resolved.index_offsets is pre_tensor # ttnn.Tensor passthrough in _resolve_buf | |
| sampler = Sampling1D.from_config(cfg) | |
| sampler.load_device_buffers() | |
| assert sampler._index_offsets is pre_tensor # ttnn.Tensor passthrough in _materialize | |
| # ------------------------------------------------------------------ | |
| # Buffer passthrough: _resolve_buf LazyBuffer path (line 495) | |
| # ------------------------------------------------------------------ | |
| def test_resolve_buf_lazy_buffer_passthrough(self, ttnn_mesh_device, vocab_size): | |
| """Pre-existing LazyBuffer with device=None → resolve_lazy_buffer fills in device (line 495).""" | |
| from models.common.modules.lazy_buffer import LazyBuffer | |
| cluster_shape = tuple(ttnn_mesh_device.shape) | |
| num_devices_in_mesh = 2 if list(cluster_shape) == [1, 1] else max(cluster_shape) | |
| B, K = 32, 32 | |
| replicate_mapper = ttnn.ShardTensor2dMesh(ttnn_mesh_device, dims=(None, None), mesh_shape=cluster_shape) | |
| partial_lb = LazyBuffer( | |
| source=torch.zeros(1, 1, B, K * num_devices_in_mesh, dtype=torch.int32), | |
| dtype=ttnn.int32, | |
| layout=ttnn.TILE_LAYOUT, | |
| device=None, # device not set — resolve_lazy_buffer fills it in | |
| mesh_mapper=replicate_mapper, | |
| memory_config=ttnn.DRAM_MEMORY_CONFIG, | |
| ) | |
| cfg = Sampling1DConfig(vocab_size=vocab_size, mesh_device=ttnn_mesh_device, index_offsets=partial_lb) | |
| resolved = _resolve_sampling1d_config(cfg) | |
| assert isinstance(resolved.index_offsets, LazyBuffer) | |
| assert resolved.index_offsets.device is ttnn_mesh_device # filled in by resolve_lazy_buffer | |
| def test_rejects_galaxy(self, ttnn_mesh_device, expect_error): | |
| """from_model_args should reject 2D (Galaxy) topologies.""" | |
| class FakeMesh: | |
| shape = (2, 4) | |
| def get_num_devices(self): | |
| return 8 | |
| class MockArgs: | |
| padded_vocab_size = 1024 | |
| sub_core_grids = None | |
| sub_core_grid_topk = None | |
| start_core = ttnn.CoreCoord(0, 0) | |
| max_top_k = 32 | |
| with expect_error(ValueError, "1D mesh topologies"): | |
| Sampling1D.from_model_args(FakeMesh(), None, MockArgs()) | |
| # ============================================================================== | |
| # Shared helper: tie-break-aware mismatch assertion | |
| # ============================================================================== | |
| def _assert_no_true_mismatches( | |
| device_tokens: "torch.Tensor", | |
| ref_tokens: "torch.Tensor", | |
| logits_2d: "torch.Tensor", | |
| *, | |
| test_label: str, | |
| quant_tolerance: float = 0.0, | |
| ): | |
| """Classify index mismatches as tie-breaks vs true mismatches. Assert zero true mismatches. | |
| In low-precision formats (bfloat16, bfloat8_b), multiple elements can share the same | |
| representable value. When ttnn.topk picks a different index among tied elements, that's | |
| a TIE-BREAK (acceptable). Only a TRUE-MISMATCH (device picked a genuinely lower value) | |
| indicates a kernel bug. | |
| Args: | |
| device_tokens: [B] int tensor of device-chosen indices. | |
| ref_tokens: [B] int tensor of reference indices. | |
| logits_2d: [B, V] tensor for value lookups (same precision as device input). | |
| test_label: Printed header, e.g. "top-k=1 vs argmax (V=32768, mesh=(1,1))". | |
| quant_tolerance: Max acceptable value delta for quantization-boundary effects | |
| (0.0 for bfloat16, ~0.032 for bfloat8_b block-float). | |
| """ | |
| B = len(device_tokens) | |
| num_mismatches = (device_tokens != ref_tokens).sum().item() | |
| true_mismatches = 0 | |
| tie_breaks = 0 | |
| true_mismatch_details = [] | |
| for b in range(B): | |
| dev_idx = int(device_tokens[b].item()) | |
| ref_idx = int(ref_tokens[b].item()) | |
| if dev_idx == ref_idx: | |
| continue | |
| dev_val = logits_2d[b, dev_idx].float().item() | |
| ref_val = logits_2d[b, ref_idx].float().item() | |
| delta = ref_val - dev_val | |
| if delta > quant_tolerance: | |
| true_mismatches += 1 | |
| true_mismatch_details.append( | |
| f" batch {b}: device idx={dev_idx} (val={dev_val:.6f}) " | |
| f"< ref idx={ref_idx} (val={ref_val:.6f}), delta={delta:.6f}" | |
| ) | |
| else: | |
| tie_breaks += 1 | |
| # Report | |
| print(f"\n--- {test_label} ---") | |
| print(f" index mismatches: {num_mismatches}/{B} (tie-breaks: {tie_breaks}, true: {true_mismatches})") | |
| if num_mismatches > 0: | |
| for b in range(B): | |
| dev_idx = int(device_tokens[b].item()) | |
| ref_idx = int(ref_tokens[b].item()) | |
| if dev_idx == ref_idx: | |
| continue | |
| dev_val = logits_2d[b, dev_idx].float().item() | |
| ref_val = logits_2d[b, ref_idx].float().item() | |
| delta = ref_val - dev_val | |
| if delta > quant_tolerance: | |
| label = "TRUE-MISMATCH" | |
| elif delta > 0: | |
| label = "QUANT-BOUNDARY" | |
| else: | |
| label = "TIE-BREAK" | |
| print( | |
| f" batch {b} [{label}]: device idx={dev_idx} (val={dev_val:.6f}), " | |
| f"ref idx={ref_idx} (val={ref_val:.6f}), delta={delta:.6f}" | |
| ) | |
| # Assert | |
| assert true_mismatches == 0, ( | |
| f"{test_label}: {true_mismatches}/{B} TRUE mismatches " | |
| f"(device picked a lower value, not a tie-break)\n" | |
| f" ({num_mismatches} total index disagreements, {tie_breaks} are tie-breaks)\n" | |
| + "\n".join(true_mismatch_details) | |
| ) | |
| return num_mismatches, tie_breaks, true_mismatches | |
| # ============================================================================== | |
| # VS Reference tests — sampling correctness against torch golden | |
| # ============================================================================== | |
| def _make_logits_tt(logits_host, ttnn_mesh_device, *, shard_vocab=False): | |
| """Create logits on device. shard_vocab=True shards the last dim across devices (for top-k path).""" | |
| cluster_shape = tuple(ttnn_mesh_device.shape) | |
| if not shard_vocab or max(cluster_shape) == 1: | |
| shard_dims = (None, None) | |
| elif cluster_shape[-1] >= cluster_shape[-2]: | |
| shard_dims = (None, -1) | |
| else: | |
| shard_dims = (-1, None) | |
| return ttnn.from_torch( | |
| logits_host, | |
| device=ttnn_mesh_device, | |
| dtype=ttnn.bfloat16, | |
| layout=ttnn.TILE_LAYOUT, | |
| mesh_mapper=ttnn.ShardTensor2dMesh(ttnn_mesh_device, dims=shard_dims, mesh_shape=cluster_shape), | |
| ) | |
| def _make_sampling_params(ttnn_mesh_device, B, *, k_val=1, p_val=0.0, temp_val=1.0, cluster_shape=(1, 1)): | |
| """Helper: create k/p/temp device tensors for Sampling1D.decode_forward().""" | |
| k = ttnn.from_torch( | |
| torch.full((B,), k_val, dtype=torch.int32), | |
| device=ttnn_mesh_device, | |
| dtype=ttnn.uint32, | |
| layout=ttnn.ROW_MAJOR_LAYOUT, | |
| mesh_mapper=ttnn.ShardTensor2dMesh(ttnn_mesh_device, dims=(None, None), mesh_shape=cluster_shape), | |
| ) | |
| p = ttnn.from_torch( | |
| torch.full((B,), p_val), | |
| device=ttnn_mesh_device, | |
| dtype=ttnn.bfloat16, | |
| layout=ttnn.ROW_MAJOR_LAYOUT, | |
| mesh_mapper=ttnn.ShardTensor2dMesh(ttnn_mesh_device, dims=(None, None), mesh_shape=cluster_shape), | |
| ) | |
| temp = ttnn.from_torch( | |
| torch.full((B,), temp_val), | |
| device=ttnn_mesh_device, | |
| dtype=ttnn.bfloat16, | |
| layout=ttnn.ROW_MAJOR_LAYOUT, | |
| mesh_mapper=ttnn.ShardTensor2dMesh(ttnn_mesh_device, dims=(None, None), mesh_shape=cluster_shape), | |
| ) | |
| return k, p, temp | |
| def test_sampling1d_topk1_vs_argmax( | |
| ttnn_mesh_device, mesh_shape, vocab_size, k_val, p_val, temp_val, force_argmax, hf_model_name | |
| ): | |
| """ | |
| Top-k=1, p=0.0, temp=1.0 should produce the same result as torch.argmax. | |
| This is the primary correctness test: with k=1 the sampling degenerates to argmax, | |
| giving us an exact reference to compare against. | |
| """ | |
| torch.manual_seed(42) | |
| B = 32 | |
| sampler = Sampling1D(vocab_size=vocab_size, mesh_device=ttnn_mesh_device) | |
| logits_host = torch.randn(1, 1, B, vocab_size, dtype=torch.bfloat16) | |
| logits_tt = _make_logits_tt(logits_host, ttnn_mesh_device, shard_vocab=True) | |
| cluster_shape = tuple(ttnn_mesh_device.shape) | |
| k, p, temp = _make_sampling_params( | |
| ttnn_mesh_device, B, k_val=1, p_val=0.0, temp_val=1.0, cluster_shape=cluster_shape | |
| ) | |
| tokens_tt, _ = sampler.decode_forward(logits_tt, k=k, p=p, temp=temp) | |
| tokens_host = to_torch_auto_compose(tokens_tt).flatten()[:B] | |
| # Naive fp32 reference — shows how many index disagreements arise from bfloat16 precision. | |
| # These are NOT correctness failures; they demonstrate why the bf16-sharded reference below | |
| # is necessary. All fp32 mismatches should be tie-breaks (same bf16 value, different index). | |
| fp32_expected = logits_host.float().argmax(dim=-1).flatten()[:B] | |
| fp32_mismatches = (tokens_host.long() != fp32_expected.long()).sum().item() | |
| mesh_label = tuple(ttnn_mesh_device.shape) | |
| print(f"\n fp32 argmax mismatches: {fp32_mismatches}/{B} (V={vocab_size}, mesh={mesh_label})") | |
| # Bfloat16-aware sharded reference: shard the vocab the same way the device does, | |
| # find top-1 per shard, then pick the global winner. This accounts for bfloat16 precision | |
| # loss at shard boundaries that torch.argmax on float32 doesn't see. | |
| num_devices = max(ttnn_mesh_device.shape) | |
| if num_devices == 1: | |
| num_shards = 2 # single device splits vocab in half internally | |
| else: | |
| num_shards = num_devices | |
| logits_bf16 = logits_host.squeeze().bfloat16() # [B, V] in bfloat16 | |
| shard_size = vocab_size // num_shards | |
| # For each batch element, find the global argmax by comparing shard-local argmaxes | |
| bf16_expected = torch.zeros(B, dtype=torch.long) | |
| for b in range(B): | |
| best_val = float("-inf") | |
| best_idx = 0 | |
| for s in range(num_shards): | |
| shard = logits_bf16[b, s * shard_size : (s + 1) * shard_size] | |
| local_idx = shard.float().argmax().item() | |
| local_val = shard[local_idx].float().item() | |
| if local_val > best_val: | |
| best_val = local_val | |
| best_idx = s * shard_size + local_idx | |
| bf16_expected[b] = best_idx | |
| _assert_no_true_mismatches( | |
| tokens_host.long(), | |
| bf16_expected, | |
| logits_bf16, | |
| test_label=f"top-k=1 vs bf16-sharded argmax (V={vocab_size}, mesh={mesh_label})", | |
| ) | |
| def test_sampling1d_argmax_vs_reference(ttnn_mesh_device): | |
| """ | |
| Force-argmax path (k/p/temp=None, allow_force_argmax=True) vs torch.argmax. | |
| Tests the all-gather-free argmax path on single device. | |
| """ | |
| torch.manual_seed(99) | |
| B = 32 | |
| vocab_size = 1024 | |
| sampler = Sampling1D( | |
| vocab_size=vocab_size, | |
| mesh_device=ttnn_mesh_device, | |
| allow_force_argmax=True, | |
| ) | |
| logits_host = torch.randn(1, 1, B, vocab_size, dtype=torch.bfloat16) | |
| logits_tt = _make_logits_tt(logits_host, ttnn_mesh_device) | |
| tokens_tt, _ = sampler.decode_forward(logits_tt) | |
| # .long(): ttnn.argmax → uint32, torch.argmax → int64; normalize before compare | |
| tokens_host = to_torch_auto_compose(tokens_tt).flatten()[:B].long() | |
| expected = logits_host.float().argmax(dim=-1).flatten()[:B].long() | |
| assert torch.equal( | |
| tokens_host, expected | |
| ), f"argmax path mismatch:\n got: {tokens_host[:8]}\n expected: {expected[:8]}" | |
| def test_sampling1d_qwen3_32b_uses_compact_tail_mask(ttnn_mesh_device): | |
| sampler = Sampling1D(vocab_size=152064, valid_vocab_size=151936, mesh_device=ttnn_mesh_device) | |
| sampler.load_device_buffers() | |
| assert sampler._invalid_vocab_mask is None | |
| assert isinstance(sampler._invalid_vocab_tail_mask, ttnn.Tensor) | |
| assert sampler._invalid_vocab_tail_width == 128 | |
| def test_sampling1d_topk_masks_qwen3_32b_padded_tail(ttnn_mesh_device): | |
| B = 32 | |
| valid_vocab_size = 151936 | |
| padded_vocab_size = 152064 | |
| sampler = Sampling1D(vocab_size=padded_vocab_size, valid_vocab_size=valid_vocab_size, mesh_device=ttnn_mesh_device) | |
| logits_host = torch.full((1, 1, B, padded_vocab_size), -1.0, dtype=torch.bfloat16) | |
| logits_host[..., valid_vocab_size:] = 0.0 | |
| logits_tt = _make_logits_tt(logits_host, ttnn_mesh_device, shard_vocab=True) | |
| cluster_shape = tuple(ttnn_mesh_device.shape) | |
| k, p, temp = _make_sampling_params( | |
| ttnn_mesh_device, B, k_val=1, p_val=0.0, temp_val=1.0, cluster_shape=cluster_shape | |
| ) | |
| tokens_tt, _ = sampler.decode_forward(logits_tt, k=k, p=p, temp=temp) | |
| tokens_host = to_torch_auto_compose(tokens_tt).flatten()[:B].long() | |
| assert torch.all(tokens_host < valid_vocab_size) | |
| def test_sampling1d_padded_tail_with_sub_core_grids_runs(ttnn_mesh_device): | |
| B = 32 | |
| valid_vocab_size = 151936 | |
| padded_vocab_size = 152064 | |
| sub_core_grids = _sub_core_grids_for_32_users() | |
| sampler = Sampling1D( | |
| vocab_size=padded_vocab_size, | |
| valid_vocab_size=valid_vocab_size, | |
| mesh_device=ttnn_mesh_device, | |
| sub_core_grids=sub_core_grids, | |
| sub_core_grid_topk=sub_core_grids, | |
| start_core=ttnn.CoreCoord(0, 0), | |
| ) | |
| sampler.load_device_buffers() | |
| logits_host = torch.full((1, 1, B, padded_vocab_size), -1.0, dtype=torch.bfloat16) | |
| logits_host[..., valid_vocab_size:] = 0.0 | |
| logits_tt = _make_logits_tt(logits_host, ttnn_mesh_device, shard_vocab=True) | |
| cluster_shape = tuple(ttnn_mesh_device.shape) | |
| k, p, temp = _make_sampling_params( | |
| ttnn_mesh_device, B, k_val=1, p_val=0.0, temp_val=1.0, cluster_shape=cluster_shape | |
| ) | |
| worker_sub_device_id = ttnn.SubDeviceId(0) | |
| worker_sub_device = ttnn.SubDevice([sub_core_grids]) | |
| sub_device_manager = ttnn_mesh_device.create_sub_device_manager([worker_sub_device], 0) | |
| stall_group_set = False | |
| manager_loaded = False | |
| try: | |
| ttnn_mesh_device.load_sub_device_manager(sub_device_manager) | |
| manager_loaded = True | |
| ttnn_mesh_device.set_sub_device_stall_group([worker_sub_device_id]) | |
| stall_group_set = True | |
| tokens_tt, _ = sampler.decode_forward(logits_tt, k=k, p=p, temp=temp) | |
| tokens_host = to_torch_auto_compose(tokens_tt).flatten()[:B].long() | |
| assert torch.all(tokens_host < valid_vocab_size) | |
| finally: | |
| if stall_group_set: | |
| ttnn_mesh_device.reset_sub_device_stall_group() | |
| if manager_loaded: | |
| ttnn_mesh_device.clear_loaded_sub_device_manager() | |
| ttnn_mesh_device.remove_sub_device_manager(sub_device_manager) | |
| def test_sampling1d_argmax_slices_qwen3_32b_padded_tail(ttnn_mesh_device): | |
| B = 32 | |
| valid_vocab_size = 151936 | |
| padded_vocab_size = 152064 | |
| sampler = Sampling1D( | |
| vocab_size=padded_vocab_size, | |
| valid_vocab_size=valid_vocab_size, | |
| mesh_device=ttnn_mesh_device, | |
| allow_force_argmax=True, | |
| ) | |
| logits_host = torch.full((1, 1, B, padded_vocab_size), -1.0, dtype=torch.bfloat16) | |
| logits_host[..., valid_vocab_size:] = 0.0 | |
| logits_tt = _make_logits_tt(logits_host, ttnn_mesh_device, shard_vocab=True) | |
| tokens_tt, _ = sampler.decode_forward(logits_tt) | |
| tokens_host = to_torch_auto_compose(tokens_tt).flatten()[:B].long() | |
| assert torch.all(tokens_host < valid_vocab_size) | |
| def test_sampling1d_topk32_in_range(ttnn_mesh_device): | |
| """ | |
| Top-k=32, p=1.0 → sampled token must be within the top-32 set for every batch element. | |
| This is a statistical correctness test: we don't know which token will be sampled | |
| (it's stochastic), but it MUST be one of the top-32 tokens by logit value. | |
| """ | |
| torch.manual_seed(77) | |
| B = 32 | |
| vocab_size = 1024 | |
| sampler = Sampling1D(vocab_size=vocab_size, mesh_device=ttnn_mesh_device) | |
| logits_host = torch.randn(1, 1, B, vocab_size, dtype=torch.bfloat16) | |
| logits_tt = _make_logits_tt(logits_host, ttnn_mesh_device, shard_vocab=True) | |
| cluster_shape = tuple(ttnn_mesh_device.shape) | |
| k, p, temp = _make_sampling_params( | |
| ttnn_mesh_device, B, k_val=32, p_val=1.0, temp_val=1.0, cluster_shape=cluster_shape | |
| ) | |
| tokens_tt, _ = sampler.decode_forward(logits_tt, k=k, p=p, temp=temp) | |
| tokens_host = to_torch_auto_compose(tokens_tt).flatten()[:B] | |
| # Compute the top-32 token set per batch element | |
| _, top32_indices = logits_host.float().squeeze().topk(32, dim=-1) # [B, 32] | |
| for b in range(B): | |
| sampled_token = tokens_host[b].item() | |
| top32_set = set(top32_indices[b].tolist()) | |
| assert sampled_token in top32_set, f"Batch {b}: sampled token {sampled_token} not in top-32 set" | |
| def _hf_valid_token_set(logits_row: "torch.Tensor", k: int, p: float, temp: float) -> set: | |
| """Compute the set of tokens eligible under top-k / top-p / temperature filtering. | |
| Mirrors the pipeline inside ttnn.sampling: | |
| 1. Temperature: divide logits by temp (skipped if temp == 1.0) | |
| 2. Top-k: zero out all but top-k tokens | |
| 3. Top-p: zero out tokens outside the cumulative-probability nucleus | |
| Uses HuggingFace's LogitsWarper classes so this reference is auditable against | |
| the transformers library rather than a hand-rolled implementation. | |
| Returns the set of token ids that have finite logit after filtering — any | |
| sampled token MUST come from this set. | |
| """ | |
| from transformers.generation.logits_process import TemperatureLogitsWarper, TopKLogitsWarper, TopPLogitsWarper | |
| # Warpers expect input_ids (unused here, pass None) and a [1, V] float32 scores tensor. | |
| scores = logits_row.float().unsqueeze(0) # [1, V] | |
| if temp != 1.0: | |
| scores = TemperatureLogitsWarper(temperature=temp)(None, scores) | |
| if k > 0: | |
| scores = TopKLogitsWarper(top_k=k)(None, scores) | |
| if 0.0 < p < 1.0: | |
| scores = TopPLogitsWarper(top_p=p)(None, scores) | |
| # Tokens with -inf logit are filtered out; all others are valid candidates. | |
| return set(scores[0].isfinite().nonzero(as_tuple=False).squeeze(-1).tolist()) | |
| def test_sampling1d_token_in_valid_set(ttnn_mesh_device, k, p, temp, max_boundary_violations): | |
| """Sampled token must lie within the HF-derived valid candidate set (up to bf16 boundary). | |
| For each (k, p, temp), the HuggingFace pipeline | |
| TemperatureLogitsWarper → TopKLogitsWarper → TopPLogitsWarper | |
| defines which tokens are eligible. Any sampled token MUST come from this set. | |
| Precision note: ttnn.sampling runs its softmax/cumsum in bfloat16, while the HF | |
| reference uses float32. Tokens near the nucleus cutoff may fall on different sides | |
| of the cumulative-probability threshold. max_boundary_violations allows for this; | |
| it is zero when p ∈ {0.0, 1.0} (no nucleus threshold exists) and small-but-nonzero | |
| otherwise. Violations significantly above max indicate a real correctness regression. | |
| """ | |
| torch.manual_seed(42) | |
| B = 32 | |
| vocab_size = 1024 | |
| sampler = Sampling1D(vocab_size=vocab_size, mesh_device=ttnn_mesh_device) | |
| logits_host = torch.randn(1, 1, B, vocab_size, dtype=torch.bfloat16) | |
| logits_tt = _make_logits_tt(logits_host, ttnn_mesh_device, shard_vocab=True) | |
| cluster_shape = tuple(ttnn_mesh_device.shape) | |
| k_tt, p_tt, temp_tt = _make_sampling_params( | |
| ttnn_mesh_device, B, k_val=k, p_val=p, temp_val=temp, cluster_shape=cluster_shape | |
| ) | |
| tokens_tt, _ = sampler.decode_forward(logits_tt, k=k_tt, p=p_tt, temp=temp_tt) | |
| tokens_host = to_torch_auto_compose(tokens_tt).flatten()[:B] | |
| # Build per-batch-element valid sets from bf16 logits (same precision as device input) | |
| logits_2d = logits_host.squeeze().bfloat16() # [B, V] | |
| violations = [] | |
| for b in range(B): | |
| valid = _hf_valid_token_set(logits_2d[b], k=k, p=p, temp=temp) | |
| token = tokens_host[b].item() | |
| if token not in valid: | |
| violations.append((b, token, len(valid))) | |
| assert len(violations) <= max_boundary_violations, ( | |
| f"k={k} p={p} temp={temp}: {len(violations)}/{B} tokens outside valid set " | |
| f"(max allowed={max_boundary_violations} for bf16 boundary):\n" | |
| + "\n".join(f" batch {b}: token {tok} not in {n}-token valid set" for b, tok, n in violations[:5]) | |
| ) | |
| def test_sampling1d_deterministic_with_same_seed(ttnn_mesh_device): | |
| """ | |
| Two decode_forward calls with the same seed tensor should produce the same tokens. | |
| """ | |
| torch.manual_seed(42) | |
| B = 32 | |
| vocab_size = 1024 | |
| sampler = Sampling1D(vocab_size=vocab_size, mesh_device=ttnn_mesh_device) | |
| logits_host = torch.randn(1, 1, B, vocab_size, dtype=torch.bfloat16) | |
| cluster_shape = tuple(ttnn_mesh_device.shape) | |
| k, p, temp = _make_sampling_params( | |
| ttnn_mesh_device, B, k_val=32, p_val=0.9, temp_val=0.8, cluster_shape=cluster_shape | |
| ) | |
| # Use explicit seed tensor | |
| seed_tensor = ttnn.from_torch( | |
| torch.arange(B, dtype=torch.int64).to(torch.int32), | |
| device=ttnn_mesh_device, | |
| dtype=ttnn.uint32, | |
| layout=ttnn.ROW_MAJOR_LAYOUT, | |
| memory_config=ttnn.DRAM_MEMORY_CONFIG, | |
| ) | |
| # First call | |
| logits_tt1 = _make_logits_tt(logits_host, ttnn_mesh_device, shard_vocab=True) | |
| tokens1, _ = sampler.decode_forward(logits_tt1, k=k, p=p, temp=temp, seeds=seed_tensor) | |
| tokens1_host = to_torch_auto_compose(tokens1).flatten()[:B] | |
| # Second call with same seed | |
| logits_tt2 = _make_logits_tt(logits_host, ttnn_mesh_device, shard_vocab=True) | |
| tokens2, _ = sampler.decode_forward(logits_tt2, k=k, p=p, temp=temp, seeds=seed_tensor) | |
| tokens2_host = to_torch_auto_compose(tokens2).flatten()[:B] | |
| assert torch.equal( | |
| tokens1_host, tokens2_host | |
| ), f"Same seed produced different tokens:\n call1: {tokens1_host[:8]}\n call2: {tokens2_host[:8]}" | |
| # ============================================================================== | |
| # Isolation tests — ttnn.topk + ttnn.all_gather without Sampling1D | |
| # ============================================================================== | |
| def test_topk_allgather_isolation(ttnn_mesh_device, vocab_size): | |
| """ | |
| Minimal reproducer: ttnn.topk + ttnn.all_gather on multi-device, bypassing Sampling1D. | |
| Runs the raw op pipeline that Sampling1D._topk_multi_device performs: | |
| 1. Shard logits across devices along the vocab dim | |
| 2. ttnn.topk per device (local top-K) | |
| 3. ttnn.all_gather values and indices across devices | |
| 4. Add index offsets for global vocab indices | |
| 5. Pick global top-1 from gathered results | |
| Compare against a bfloat16-sharded torch reference. This isolates whether | |
| mismatches come from topk+all_gather or from downstream ops (sampling, typecast, etc.). | |
| """ | |
| torch.manual_seed(42) | |
| B = 32 | |
| K = 32 # max_top_k, matches Sampling1D default | |
| cluster_shape = tuple(ttnn_mesh_device.shape) | |
| num_devices = max(cluster_shape) | |
| per_device_vocab = vocab_size // num_devices | |
| # -- 1. Shard logits across devices (same as _make_logits_tt with shard_vocab=True) -- | |
| logits_host = torch.randn(1, 1, B, vocab_size, dtype=torch.bfloat16) | |
| shard_dims = (None, -1) if cluster_shape[-1] >= cluster_shape[-2] else (-1, None) | |
| logits_tt = ttnn.from_torch( | |
| logits_host, | |
| device=ttnn_mesh_device, | |
| dtype=ttnn.bfloat16, | |
| layout=ttnn.TILE_LAYOUT, | |
| mesh_mapper=ttnn.ShardTensor2dMesh(ttnn_mesh_device, dims=shard_dims, mesh_shape=cluster_shape), | |
| ) | |
| # -- 2. Build local_indices buffer replicated on all devices -- | |
| replicate_mapper = ttnn.ShardTensor2dMesh(ttnn_mesh_device, dims=(None, None), mesh_shape=cluster_shape) | |
| local_indices_host = torch.zeros(1, 1, B, per_device_vocab, dtype=torch.int32) | |
| for i in range(per_device_vocab): | |
| local_indices_host[:, :, :, i] = i | |
| local_indices_tt = ttnn.from_torch( | |
| local_indices_host, | |
| device=ttnn_mesh_device, | |
| dtype=ttnn.uint16, | |
| layout=ttnn.TILE_LAYOUT, | |
| mesh_mapper=replicate_mapper, | |
| memory_config=ttnn.DRAM_MEMORY_CONFIG, | |
| ) | |
| # -- 3. ttnn.topk per device -- | |
| topk_values, topk_indices = ttnn.topk( | |
| logits_tt, | |
| k=K, | |
| dim=-1, | |
| indices_tensor=local_indices_tt, | |
| ) | |
| # -- 4. all_gather values and indices along the vocab dim -- | |
| sampling_cluster_axis = None if 1 in cluster_shape else 0 | |
| gathered_values = ttnn.all_gather( | |
| topk_values, | |
| dim=3, | |
| num_links=1, | |
| memory_config=ttnn.DRAM_MEMORY_CONFIG, | |
| cluster_axis=sampling_cluster_axis, | |
| topology=ttnn.Topology.Linear, | |
| ) | |
| ttnn.deallocate(topk_values) | |
| gathered_indices = ttnn.all_gather( | |
| topk_indices, | |
| dim=3, | |
| num_links=1, | |
| memory_config=ttnn.DRAM_MEMORY_CONFIG, | |
| cluster_axis=sampling_cluster_axis, | |
| topology=ttnn.Topology.Linear, | |
| ) | |
| ttnn.deallocate(topk_indices) | |
| # -- 5. Add per-device offsets to convert local → global vocab indices -- | |
| offsets_host = torch.zeros(1, 1, B, K * num_devices, dtype=torch.int64) | |
| for d in range(num_devices): | |
| offsets_host[:, :, :, d * K : (d + 1) * K] = d * per_device_vocab | |
| index_offsets_tt = ttnn.from_torch( | |
| offsets_host, | |
| device=ttnn_mesh_device, | |
| dtype=ttnn.int32, | |
| layout=ttnn.TILE_LAYOUT, | |
| mesh_mapper=replicate_mapper, | |
| memory_config=ttnn.DRAM_MEMORY_CONFIG, | |
| ) | |
| gathered_indices_int32 = ttnn.typecast(gathered_indices, dtype=ttnn.int32) | |
| global_indices = ttnn.add(index_offsets_tt, gathered_indices_int32, dtype=ttnn.int32) | |
| # -- 6. Read back to host -- | |
| values_host = to_torch_auto_compose(gathered_values).squeeze()[:B] # [B, K*num_devices] | |
| indices_host = to_torch_auto_compose(global_indices).squeeze()[:B] # [B, K*num_devices] | |
| # -- 7. Pick global top-1: find max value position then look up its global index -- | |
| top1_pos = values_host.float().argmax(dim=-1) | |
| device_top1 = torch.tensor([indices_host[b, top1_pos[b]].item() for b in range(B)], dtype=torch.long) | |
| # -- 8. Bfloat16-sharded torch reference (same method as test_sampling1d_topk1_vs_argmax) -- | |
| logits_bf16 = logits_host.squeeze().bfloat16() # [B, V] | |
| bf16_expected = torch.zeros(B, dtype=torch.long) | |
| for b in range(B): | |
| best_val = float("-inf") | |
| best_idx = 0 | |
| for s in range(num_devices): | |
| shard = logits_bf16[b, s * per_device_vocab : (s + 1) * per_device_vocab] | |
| local_idx = shard.float().argmax().item() | |
| local_val = shard[local_idx].float().item() | |
| if local_val > best_val: | |
| best_val = local_val | |
| best_idx = s * per_device_vocab + local_idx | |
| bf16_expected[b] = best_idx | |
| # -- 9. Report and assert (tie-break-aware) -- | |
| _assert_no_true_mismatches( | |
| device_top1, | |
| bf16_expected, | |
| logits_bf16, | |
| test_label=f"topk+all_gather isolation (V={vocab_size}, mesh={cluster_shape})", | |
| ) | |
| def test_ttnn_sampling_isolation(ttnn_mesh_device, vocab_size): | |
| """ | |
| Hypothesis 2: does ttnn.sampling introduce mismatches at k=1, p=0.0, temp=1.0? | |
| Builds correct gathered_values + global_indices via the topk+all_gather pipeline | |
| (confirmed 0 mismatches in test_topk_allgather_isolation), then runs the remaining | |
| steps from Sampling1D._sample_topk verbatim: | |
| - ttnn.typecast (uint16 → int32) | |
| - ttnn.add (index offsets) | |
| - ttnn.untilize (TILE → ROW_MAJOR, required by ttnn.sampling) | |
| - ttnn.manual_seed + ttnn.sampling(k=1, p=0.0, temp=1.0) | |
| Compares against the bfloat16-sharded torch argmax reference. | |
| Any mismatches here can be attributed to ttnn.sampling itself. | |
| Observed results (seed=42, B=32): | |
| v1024-1x2: 1/32 mismatches ← ttnn.sampling | |
| v1024-1x8: 1/32 mismatches ← same batch as 1x2, topology-independent | |
| v32000-1x2: 4/32 mismatches ← ttnn.sampling | |
| v32000-1x8: 7/32 mismatches ← 4 shared with 1x2 + 3 additional | |
| The growing mismatch count with more devices (1x2→1x8) is not caused by | |
| all_gather (test_topk_allgather_isolation confirms 0/32 there). The extra | |
| mismatches at 1x8 come from ttnn.sampling seeing a wider candidate buffer | |
| (K*8=256 entries vs K*2=64), causing its internal softmax reduction to | |
| diverge from argmax on more batches. | |
| """ | |
| torch.manual_seed(42) | |
| B = 32 | |
| K = 32 # max_top_k | |
| cluster_shape = tuple(ttnn_mesh_device.shape) | |
| num_devices = max(cluster_shape) | |
| per_device_vocab = vocab_size // num_devices | |
| replicate_mapper = ttnn.ShardTensor2dMesh(ttnn_mesh_device, dims=(None, None), mesh_shape=cluster_shape) | |
| # ---- Step A: topk + all_gather (confirmed correct, 0 mismatches) -------- | |
| logits_host = torch.randn(1, 1, B, vocab_size, dtype=torch.bfloat16) | |
| shard_dims = (None, -1) if cluster_shape[-1] >= cluster_shape[-2] else (-1, None) | |
| logits_tt = ttnn.from_torch( | |
| logits_host, | |
| device=ttnn_mesh_device, | |
| dtype=ttnn.bfloat16, | |
| layout=ttnn.TILE_LAYOUT, | |
| mesh_mapper=ttnn.ShardTensor2dMesh(ttnn_mesh_device, dims=shard_dims, mesh_shape=cluster_shape), | |
| ) | |
| local_indices_host = torch.zeros(1, 1, B, per_device_vocab, dtype=torch.int32) | |
| for i in range(per_device_vocab): | |
| local_indices_host[:, :, :, i] = i | |
| local_indices_tt = ttnn.from_torch( | |
| local_indices_host, | |
| device=ttnn_mesh_device, | |
| dtype=ttnn.uint16, | |
| layout=ttnn.TILE_LAYOUT, | |
| mesh_mapper=replicate_mapper, | |
| memory_config=ttnn.DRAM_MEMORY_CONFIG, | |
| ) | |
| topk_values, topk_indices = ttnn.topk(logits_tt, k=K, dim=-1, indices_tensor=local_indices_tt) | |
| sampling_cluster_axis = None if 1 in cluster_shape else 0 | |
| gathered_values = ttnn.all_gather( | |
| topk_values, | |
| dim=3, | |
| num_links=1, | |
| memory_config=ttnn.DRAM_MEMORY_CONFIG, | |
| cluster_axis=sampling_cluster_axis, | |
| topology=ttnn.Topology.Linear, | |
| ) | |
| ttnn.deallocate(topk_values) | |
| gathered_indices = ttnn.all_gather( | |
| topk_indices, | |
| dim=3, | |
| num_links=1, | |
| memory_config=ttnn.DRAM_MEMORY_CONFIG, | |
| cluster_axis=sampling_cluster_axis, | |
| topology=ttnn.Topology.Linear, | |
| ) | |
| ttnn.deallocate(topk_indices) | |
| # ---- Step B: index offset addition (same as _sample_topk lines 233-253) - | |
| gathered_indices_int32 = ttnn.typecast(gathered_indices, dtype=ttnn.int32) | |
| offsets_host = torch.zeros(1, 1, B, K * num_devices, dtype=torch.int64) | |
| for d in range(num_devices): | |
| offsets_host[:, :, :, d * K : (d + 1) * K] = d * per_device_vocab | |
| index_offsets_tt = ttnn.from_torch( | |
| offsets_host, | |
| device=ttnn_mesh_device, | |
| dtype=ttnn.int32, | |
| layout=ttnn.TILE_LAYOUT, | |
| mesh_mapper=replicate_mapper, | |
| memory_config=ttnn.DRAM_MEMORY_CONFIG, | |
| ) | |
| global_indices_tiled = ttnn.add( | |
| index_offsets_tt, gathered_indices_int32, dtype=ttnn.int32, memory_config=ttnn.DRAM_MEMORY_CONFIG | |
| ) | |
| ttnn.deallocate(gathered_indices_int32) | |
| global_indices_rm = ttnn.untilize(global_indices_tiled, use_multicore=True) | |
| ttnn.deallocate(global_indices_tiled) | |
| # ---- Step C: seed + ttnn.sampling(k=1, p=0.0, temp=1.0) ----------------- | |
| seeds_host = torch.arange(B, dtype=torch.int32) | |
| seeds_tt = ttnn.from_torch( | |
| seeds_host, | |
| device=ttnn_mesh_device, | |
| dtype=ttnn.uint32, | |
| layout=ttnn.ROW_MAJOR_LAYOUT, | |
| mesh_mapper=replicate_mapper, | |
| memory_config=ttnn.DRAM_MEMORY_CONFIG, | |
| ) | |
| user_ids_tt = ttnn.from_torch( | |
| seeds_host, | |
| device=ttnn_mesh_device, | |
| dtype=ttnn.uint32, | |
| layout=ttnn.ROW_MAJOR_LAYOUT, | |
| mesh_mapper=replicate_mapper, | |
| memory_config=ttnn.DRAM_MEMORY_CONFIG, | |
| ) | |
| ttnn.manual_seed(seeds=seeds_tt, user_ids=user_ids_tt) | |
| k_tt = ttnn.from_torch( | |
| torch.ones(B, dtype=torch.int32), | |
| device=ttnn_mesh_device, | |
| dtype=ttnn.uint32, | |
| layout=ttnn.ROW_MAJOR_LAYOUT, | |
| mesh_mapper=replicate_mapper, | |
| memory_config=ttnn.DRAM_MEMORY_CONFIG, | |
| ) | |
| p_tt = ttnn.from_torch( | |
| torch.zeros(B, dtype=torch.float32), | |
| device=ttnn_mesh_device, | |
| dtype=ttnn.bfloat16, | |
| layout=ttnn.ROW_MAJOR_LAYOUT, | |
| mesh_mapper=replicate_mapper, | |
| memory_config=ttnn.DRAM_MEMORY_CONFIG, | |
| ) | |
| temp_tt = ttnn.from_torch( | |
| torch.ones(B, dtype=torch.float32), | |
| device=ttnn_mesh_device, | |
| dtype=ttnn.bfloat16, | |
| layout=ttnn.ROW_MAJOR_LAYOUT, | |
| mesh_mapper=replicate_mapper, | |
| memory_config=ttnn.DRAM_MEMORY_CONFIG, | |
| ) | |
| sampled_tokens = ttnn.sampling(gathered_values, global_indices_rm, k=k_tt, p=p_tt, temp=temp_tt) | |
| ttnn.deallocate(gathered_values) | |
| ttnn.deallocate(global_indices_rm) | |
| # ---- Step D: compare against bfloat16-sharded torch reference ----------- | |
| tokens_host = to_torch_auto_compose(sampled_tokens).flatten()[:B] | |
| # Bfloat16-sharded reference (same as other tests in this file) | |
| logits_bf16 = logits_host.squeeze().bfloat16() | |
| bf16_expected = torch.zeros(B, dtype=torch.long) | |
| for b in range(B): | |
| best_val, best_idx = float("-inf"), 0 | |
| for s in range(num_devices): | |
| shard = logits_bf16[b, s * per_device_vocab : (s + 1) * per_device_vocab] | |
| local_idx = shard.float().argmax().item() | |
| local_val = shard[local_idx].float().item() | |
| if local_val > best_val: | |
| best_val, best_idx = local_val, s * per_device_vocab + local_idx | |
| bf16_expected[b] = best_idx | |
| _assert_no_true_mismatches( | |
| tokens_host.long(), | |
| bf16_expected, | |
| logits_bf16, | |
| test_label=f"ttnn.sampling isolation (V={vocab_size}, mesh={cluster_shape})", | |
| ) | |
| # ============================================================================== | |
| # Minimal reproducer: ttnn.topk correctness at various input widths | |
| # ============================================================================== | |
| def test_ttnn_topk_correctness(ttnn_mesh_device, input_width, dtype): | |
| """ | |
| Minimal reproducer for ttnn.topk index disagreements at various input widths and dtypes. | |
| Isolates ttnn.topk on a SINGLE device with NO Sampling1D, NO all_gather, | |
| NO ttnn.sampling — just the raw topk op. We compare: | |
| 1. Top-1 (argmax): device top-1 index vs torch.topk top-1 | |
| 2. Top-K set: whether the device's top-32 index SET matches torch's top-32 set | |
| 3. Top-K values: whether the returned values match the expected values | |
| Note: ttnn.topk only supports BFLOAT16 and BFLOAT8_B inputs (enforced by the kernel | |
| at topk_device_operation.cpp:146). Float32 is not a valid input dtype. | |
| Observed results (seed=42, B=32, K=32, single device 1x1): | |
| bfloat16 (~128 distinct values in the typical randn range): | |
| Width Top-1 mismatches Top-K set mismatches Top-1 value mismatches | |
| ----- ---------------- -------------------- ---------------------- | |
| 512 0/32 3/32 0/32 | |
| 1024 1/32 3/32 0/32 | |
| 2048 1/32 14/32 0/32 | |
| 4096 3/32 7/32 0/32 | |
| 8192 6/32 12/32 0/32 | |
| 16384 8/32 13/32 0/32 | |
| 32768 17/32 14/32 0/32 | |
| All top-1 mismatches are TIE-BREAKS (delta=0.0000). Zero true mismatches. | |
| bfloat8_b (even fewer distinct values → more ties expected): | |
| Serves as a confirmation of the tie-breaking hypothesis: with coarser | |
| quantization, ties are more frequent, so index disagreements should | |
| increase compared to bfloat16 at the same width. | |
| Key findings: | |
| - NOT a comparison bug. Every top-1 mismatch has delta=0.0000: the device picks a | |
| different index that has the SAME value as the reference's top-1. | |
| - Values are always correct (0/32 value mismatches at every width). ttnn.topk finds | |
| the right maximum value; it just returns a different index among tied elements. | |
| - Root cause: low-precision tie-breaking non-determinism. With few distinct | |
| representable values, ties become very common as width grows, explaining | |
| why mismatches scale with width. | |
| - Implication: the "10/32 mismatches" in Sampling1D with vocab=32768 are NOT a | |
| ttnn.topk kernel bug — they are precision tie-breaks. Any element with the max | |
| value is a valid argmax. Tests should only assert on true mismatches (device | |
| picked a genuinely lower value), not on tie-breaks. | |
| """ | |
| # Both supported dtypes use bfloat16 as the host torch dtype; bfloat8_b is a | |
| # device-side block-float format that ttnn converts from bfloat16 on the host. | |
| torch_dtype = torch.bfloat16 | |
| dtype_tag = "bf16" if dtype == ttnn.bfloat16 else "bf8b" | |
| torch.manual_seed(42) | |
| B = 32 | |
| K = 32 | |
| # -- 1. Create input tensor [1, 1, B, W] in the target dtype -- | |
| input_host = torch.randn(1, 1, B, input_width, dtype=torch_dtype) | |
| input_tt = ttnn.from_torch( | |
| input_host, | |
| device=ttnn_mesh_device, | |
| dtype=dtype, | |
| layout=ttnn.TILE_LAYOUT, | |
| memory_config=ttnn.DRAM_MEMORY_CONFIG, | |
| ) | |
| # -- 2. Create local_indices [1, 1, B, W] with 0-based range -- | |
| # This matches the production code: indices_tensor[..., i] = i | |
| local_indices_host = torch.zeros(1, 1, B, input_width, dtype=torch.int32) | |
| for i in range(input_width): | |
| local_indices_host[:, :, :, i] = i | |
| local_indices_tt = ttnn.from_torch( | |
| local_indices_host, | |
| device=ttnn_mesh_device, | |
| dtype=ttnn.uint16, | |
| layout=ttnn.TILE_LAYOUT, | |
| memory_config=ttnn.DRAM_MEMORY_CONFIG, | |
| ) | |
| # -- 3. Run ttnn.topk -- | |
| topk_values_tt, topk_indices_tt = ttnn.topk(input_tt, k=K, dim=-1, indices_tensor=local_indices_tt) | |
| # -- 4. Read back to host -- | |
| values_host = to_torch_auto_compose(topk_values_tt).squeeze()[:B] # [B, K] | |
| indices_host = to_torch_auto_compose(topk_indices_tt).squeeze()[:B] # [B, K] | |
| ttnn.deallocate(topk_values_tt) | |
| ttnn.deallocate(topk_indices_tt) | |
| ttnn.deallocate(input_tt) | |
| ttnn.deallocate(local_indices_tt) | |
| # -- 5. Torch reference on the SAME input (compare in native precision) -- | |
| input_2d = input_host.squeeze() # [B, W] | |
| ref_values, ref_indices = torch.topk(input_2d.float(), k=K, dim=-1) | |
| ref_values = ref_values.to(torch_dtype) # compare in the same dtype | |
| # -- 6a. Check top-1 (argmax) correctness -- | |
| device_top1 = indices_host[:, 0].long() # first column = top-1 index | |
| ref_top1 = ref_indices[:, 0].long() | |
| top1_mismatches = (device_top1 != ref_top1).sum().item() | |
| # -- 6b. Check top-K set correctness (order doesn't matter) -- | |
| set_mismatches = 0 | |
| missing_details = [] | |
| for b in range(B): | |
| device_set = set(indices_host[b].long().tolist()) | |
| ref_set = set(ref_indices[b].tolist()) | |
| missing_from_device = ref_set - device_set | |
| if missing_from_device: | |
| set_mismatches += 1 | |
| if len(missing_details) < 5: | |
| extra_in_device = device_set - ref_set | |
| missing_details.append( | |
| f" batch {b}: missing {len(missing_from_device)} ref indices, " | |
| f"has {len(extra_in_device)} wrong indices\n" | |
| f" missing (first 5): {sorted(missing_from_device)[:5]}\n" | |
| f" extra (first 5): {sorted(extra_in_device)[:5]}" | |
| ) | |
| # -- 6c. Check top-1 value correctness (did it at least get the max value right?) -- | |
| device_top1_vals = values_host[:, 0].float() | |
| ref_top1_vals = ref_values[:, 0].float() | |
| val_mismatches = (device_top1_vals != ref_top1_vals).sum().item() | |
| # -- 6d. Tie-break-aware assertion on top-1 indices -- | |
| # bfloat8_b uses block-float quantization (shared exponent per block of 32 elements). | |
| # Elements within a block can lose up to 2 ULPs (~0.03125 at typical magnitudes) | |
| # relative to the bfloat16 host view. | |
| quant_tolerance = 0.032 if dtype == ttnn.bfloat8_b else 0.0 | |
| # Print additional top-K diagnostics before the shared assertion prints its report | |
| print(f"\n top-1 index mismatches: {top1_mismatches}/{B}") | |
| print(f" top-K set mismatches: {set_mismatches}/{B}") | |
| print(f" top-1 value mismatches: {val_mismatches}/{B}") | |
| _assert_no_true_mismatches( | |
| device_top1, | |
| ref_top1, | |
| input_2d, | |
| test_label=f"ttnn.topk isolation (width={input_width}, dtype={dtype_tag}, B={B}, K={K})", | |
| quant_tolerance=quant_tolerance, | |
| ) | |
| # ============================================================================== | |
| # Logprobs plumbing (enable_log_probs per-call arg) | |
| # ============================================================================== | |
| def test_sampling1d_logprobs_disabled_returns_none(ttnn_mesh_device): | |
| """Default (enable_log_probs=False) → log_probs is None on every mesh, both paths.""" | |
| torch.manual_seed(42) | |
| B = 32 | |
| vocab_size = 1024 | |
| sampler = Sampling1D(vocab_size=vocab_size, mesh_device=ttnn_mesh_device, allow_force_argmax=True) | |
| logits_host = torch.randn(1, 1, B, vocab_size, dtype=torch.bfloat16) | |
| cluster_shape = tuple(ttnn_mesh_device.shape) | |
| # Top-k path | |
| logits_tt = _make_logits_tt(logits_host, ttnn_mesh_device, shard_vocab=True) | |
| k, p, temp = _make_sampling_params( | |
| ttnn_mesh_device, B, k_val=1, p_val=0.0, temp_val=1.0, cluster_shape=cluster_shape | |
| ) | |
| _, log_probs = sampler.decode_forward(logits_tt, k=k, p=p, temp=temp) | |
| assert log_probs is None, "top-k path must return None when enable_log_probs=False" | |
| # Argmax path | |
| logits_tt2 = _make_logits_tt(logits_host, ttnn_mesh_device) | |
| _, log_probs_argmax = sampler.decode_forward(logits_tt2) | |
| assert log_probs_argmax is None, "argmax path must return None when enable_log_probs=False" | |
| def test_sampling1d_argmax_never_emits_logprobs(ttnn_mesh_device): | |
| """Argmax contract (P0): argmax path returns None even when enable_log_probs=True.""" | |
| torch.manual_seed(42) | |
| B = 32 | |
| vocab_size = 1024 | |
| sampler = Sampling1D(vocab_size=vocab_size, mesh_device=ttnn_mesh_device, allow_force_argmax=True) | |
| logits_host = torch.randn(1, 1, B, vocab_size, dtype=torch.bfloat16) | |
| logits_tt = _make_logits_tt(logits_host, ttnn_mesh_device) | |
| _, log_probs = sampler.decode_forward(logits_tt, enable_log_probs=True) | |
| assert log_probs is None, "argmax path must never compute logprobs (force-argmax ⇒ no logprobs)" | |
| def test_sampling1d_logprobs_topk(ttnn_mesh_device): | |
| """enable_log_probs=True on the top-k path. | |
| The old single-token logprob path only computes on multi-device shards with | |
| num_devices ∈ {8, 32} (T3K 1×8). On 1×1/1×2 the calculator returns None even when enabled. | |
| On 1×8, the returned logprob must match torch.log_softmax(logits)[sampled_token] within | |
| bf16 reduction tolerance. PCC is intentionally not used here because the k=1 random-bf16 case | |
| is near-constant and can degenerate to zero variance on device. | |
| """ | |
| torch.manual_seed(42) | |
| B = 32 | |
| vocab_size = 32768 # divisible by 8 | |
| sampler = Sampling1D(vocab_size=vocab_size, mesh_device=ttnn_mesh_device) | |
| logits_host = torch.randn(1, 1, B, vocab_size, dtype=torch.bfloat16) | |
| logits_tt = _make_logits_tt(logits_host, ttnn_mesh_device, shard_vocab=True) | |
| cluster_shape = tuple(ttnn_mesh_device.shape) | |
| k, p, temp = _make_sampling_params( | |
| ttnn_mesh_device, B, k_val=1, p_val=0.0, temp_val=1.0, cluster_shape=cluster_shape | |
| ) | |
| tokens_tt, log_probs = sampler.decode_forward(logits_tt, k=k, p=p, temp=temp, enable_log_probs=True) | |
| num_devices = max(cluster_shape) | |
| if num_devices not in (8, 32): | |
| assert log_probs is None, f"logprobs unsupported on {num_devices} devices → expected None" | |
| return | |
| assert log_probs is not None, "logprobs must be computed on a 1×8 mesh when enabled" | |
| # output_tensor shape (1,1,1,B), replicated across devices — match test_sampling.py read path | |
| mesh_composer = ttnn.ConcatMeshToTensor(ttnn_mesh_device, dim=3) | |
| lp_host = ttnn.to_torch(log_probs, mesh_composer=mesh_composer)[:, :, 0, :B].reshape(-1).float() | |
| tokens_host = to_torch_auto_compose(tokens_tt).flatten()[:B].long() | |
| # Reference: log_softmax over the full vocab (fp32), indexed at the sampled token. | |
| ref_log_softmax = torch.log_softmax(logits_host.float().squeeze(), dim=-1) # [B, V] | |
| ref_lp = ref_log_softmax[torch.arange(B), tokens_host] | |
| max_abs_error = torch.max(torch.abs(ref_lp - lp_host)).item() | |
| assert max_abs_error <= 5e-2, f"logprobs max abs error {max_abs_error:.6f} exceeds bf16 tolerance" | |
| # ============================================================================== | |
| # Trace capture — on-device sampling under begin/end_trace_capture (N150/N300) | |
| # ============================================================================== | |
| # | |
| # The tests above prove the sampler is correct *eagerly* on every mesh. They do NOT exercise the | |
| # thing that actually gates on-device sampling in a model: whether the sampling op graph can be | |
| # captured inside ttnn.begin_trace_capture / ttnn.end_trace_capture. TracedLLMExecutor runs | |
| # model.sampling.decode_forward *inside* trace capture (executor.py:_capture_decode_trace). | |
| # | |
| # The TTTv2 perf-recovery handoff gated models' supports_on_device_sampling on num_devices >= 8, | |
| # assuming the sub-8-device Linear+barrier all-gather could not be trace-captured. The cases below | |
| # DISPROVE that at the op level: argmax and top-k both capture AND replay correctly on 1x1 (N150, | |
| # no CCL) and 1x2 (N300, Linear+barrier all-gather). See the handoff doc's "wrong assumption" note. | |
| # | |
| # Note the dict-form ttnn_mesh_device param: it carries trace_region_size so the device is opened | |
| # with a trace region (the fixture does not set one by default). | |
| _TRACE_REGION_SIZE = 32 << 20 | |
| def test_sampling1d_trace_capture(ttnn_mesh_device, mode): | |
| """Sampling1D.decode_forward must trace-capture and replay correctly on N150/N300. | |
| Mirrors TracedLLMExecutor._capture_decode_trace: warmup-compile + load_device_buffers OUTSIDE | |
| capture (those issue device writes that are illegal mid-capture), then run decode_forward | |
| inside begin/end_trace_capture and replay, asserting replayed tokens == eager tokens. | |
| All four (mesh x mode) combos pass — including the 1x2 Linear+barrier all-gather the | |
| ``num_devices >= 8`` model gate assumes is not capturable. | |
| """ | |
| mesh_device = ttnn_mesh_device | |
| num_devices = mesh_device.get_num_devices() | |
| cluster_shape = tuple(mesh_device.shape) | |
| torch.manual_seed(0) | |
| B = 32 | |
| vocab_size = 128256 # Llama-class vocab; per-device shard on 1x2 (64128) >> max_top_k=32 | |
| logits_host = torch.randn(1, 1, B, vocab_size, dtype=torch.bfloat16) | |
| sampler = Sampling1D( | |
| vocab_size=vocab_size, | |
| mesh_device=mesh_device, | |
| max_batch_size=B, | |
| allow_force_argmax=True, | |
| pad_to_power_of_2=(mode == "topk" and num_devices > 1), | |
| ) | |
| # k/p/temp built ONCE, OUTSIDE capture, and the same persistent tensors are referenced inside | |
| # it — exactly as TracedLLMExecutor caches them (executor.py:_get_decode_sampling_kpt). | |
| # Building them inside capture would itself be an illegal in-capture host->device write. | |
| kpt = ( | |
| _make_sampling_params(mesh_device, B, k_val=1, p_val=0.0, temp_val=1.0, cluster_shape=cluster_shape) | |
| if mode == "topk" | |
| else None | |
| ) | |
| def _forward(logits_tt): | |
| if mode == "argmax": | |
| return sampler.decode_forward(logits_tt) | |
| k, p, temp = kpt | |
| return sampler.decode_forward(logits_tt, k=k, p=p, temp=temp) | |
| # Persistent input buffer, reused across warmup / capture / replay (as the executor does). | |
| logits_tt = _make_logits_tt(logits_host, mesh_device, shard_vocab=True) | |
| # --- Warmup: load buffers + JIT-compile the sampling program OUTSIDE capture --- | |
| sampler.load_device_buffers() | |
| eager_tok, _ = _forward(logits_tt) | |
| ttnn.synchronize_device(mesh_device) | |
| eager_host = to_torch_auto_compose(eager_tok).flatten()[:B].long() | |
| # --- Capture (close+release in finally so a failed capture can't wedge the next case) --- | |
| trace_id = ttnn.begin_trace_capture(mesh_device, cq_id=0) | |
| captured_tok = None | |
| try: | |
| captured_tok, _ = _forward(logits_tt) | |
| finally: | |
| ttnn.end_trace_capture(mesh_device, trace_id, cq_id=0) | |
| ttnn.synchronize_device(mesh_device) | |
| # --- Replay (same logits -> same tokens) --- | |
| ttnn.execute_trace(mesh_device, trace_id, cq_id=0, blocking=True) | |
| replay_host = to_torch_auto_compose(captured_tok).flatten()[:B].long() | |
| ttnn.release_trace(mesh_device, trace_id) | |
| assert torch.equal(replay_host, eager_host), ( | |
| f"[{mode} mesh={cluster_shape}] traced tokens diverged from eager:\n" | |
| f" eager: {eager_host[:8]}\n replay: {replay_host[:8]}" | |
| ) | |