# SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc. # SPDX-License-Identifier: Apache-2.0 import ast import inspect from pathlib import Path from types import SimpleNamespace import torch from models.common.sampling.generator import ( SamplingGenerator, SamplingParams, SeedManager, _hash_request_seed_to_device_seed, ) from models.common.warmup.warmup_utils import WarmupForwardMixin from models.tt_transformers.tt.common import Mode def _fake_sampling_generator(): calls = [] fake = SimpleNamespace( tt_sampling=SimpleNamespace(max_batch_size=32), seed_manager=SimpleNamespace( max_batch_size=32, apply_slot_remap=lambda remap: calls.append(("remap", list(remap))), ), _slot_state_requires_authoritative_reload=False, reset_sampling_params=lambda params: calls.append(("params", params)), reset_prompt_tokens=lambda tokens, slots=None: calls.append(("prompt", tokens, slots)), reset_output_state=lambda tokens, slots=None: calls.append(("output", tokens, slots)), ) fake.validate_decode_state_commands = lambda **kwargs: SamplingGenerator.validate_decode_state_commands( fake, **kwargs ) fake.commit_decode_state_commands = lambda **kwargs: SamplingGenerator.commit_decode_state_commands(fake, **kwargs) return fake, calls def test_sampling_state_can_reset_without_reloading_params(): fake, calls = _fake_sampling_generator() SamplingGenerator.apply_decode_state( fake, [object()], reload_sampling_params=False, reset_sampling_state=True, prompt_tokens="prompt", output_tokens="output", ) assert calls == [("prompt", "prompt", None), ("output", "output", None)] def test_sampling_state_reset_can_preserve_unlisted_slots(): fake, calls = _fake_sampling_generator() SamplingGenerator.apply_decode_state( fake, [object()], reload_sampling_params=False, reset_sampling_state=True, prompt_tokens="prompt", output_tokens=None, sampling_state_slots=[1, 3], ) assert calls == [("prompt", "prompt", [1, 3]), ("output", None, [1, 3])] def test_non_identity_slot_remap_requires_authoritative_device_state_reload(expect_error): fake, calls = _fake_sampling_generator() remap = [1, 0, *range(2, 32)] SamplingGenerator.apply_slot_remap(fake, remap) assert calls == [("remap", remap)] assert fake._slot_state_requires_authoritative_reload for reload_sampling_params, reset_sampling_state in ((False, False), (True, False), (False, True)): with expect_error(ValueError, "requires reload_sampling_params=True and reset_sampling_state=True"): SamplingGenerator.apply_decode_state( fake, [object()], reload_sampling_params=reload_sampling_params, reset_sampling_state=reset_sampling_state, ) params = SamplingParams(temperature=1.0, top_k=1, top_p=1.0) SamplingGenerator.apply_decode_state( fake, [params], reload_sampling_params=True, reset_sampling_state=True, prompt_tokens="prompt", output_tokens="output", ) assert calls[1][0] == "params" assert calls[2:] == [("prompt", "prompt", None), ("output", "output", None)] assert not fake._slot_state_requires_authoritative_reload def test_partial_sampling_state_rebuild_does_not_clear_slot_remap_invalidation(expect_error): fake, _ = _fake_sampling_generator() remap = [1, 0, *range(2, 32)] params = SamplingParams(temperature=1.0, top_k=1, top_p=1.0) SamplingGenerator.apply_slot_remap(fake, remap) SamplingGenerator.apply_decode_state( fake, [params], reload_sampling_params=True, reset_sampling_state=True, prompt_tokens="prompt", output_tokens="output", sampling_state_slots=[1, 3], ) assert fake._slot_state_requires_authoritative_reload with expect_error(ValueError, "requires reload_sampling_params=True and reset_sampling_state=True"): SamplingGenerator.apply_decode_state( fake, [params], reload_sampling_params=False, reset_sampling_state=False, ) SamplingGenerator.apply_decode_state( fake, [params], reload_sampling_params=True, reset_sampling_state=True, prompt_tokens="prompt", output_tokens="output", ) assert not fake._slot_state_requires_authoritative_reload def test_identity_slot_remap_keeps_device_sampling_state_valid(): fake, calls = _fake_sampling_generator() remap = list(range(32)) SamplingGenerator.apply_slot_remap(fake, remap) assert calls == [("remap", remap)] assert not fake._slot_state_requires_authoritative_reload def test_output_penalty_reset_masks_only_selected_slots(monkeypatch): from models.common.sampling import tt_penalties class FakeTensor: def __init__(self, name): self.name = name self.deallocated = False def deallocate(self): self.deallocated = True allocated = [] multiplies = [] penalty_state = SimpleNamespace( _total_batch=4, _shard_dims_gathered=(0, None), _op_kwargs={}, output_mask=FakeTensor("mask"), output_counts=FakeTensor("counts"), output_counts_gathered=FakeTensor("gathered"), ) def allocate(*, host, shard_dims): result = FakeTensor("keep") allocated.append((host.clone(), shard_dims, result)) return result penalty_state._alloc_int_buffer = allocate monkeypatch.setattr( tt_penalties.ttnn, "mul", lambda value, keep, *, output_tensor, **kwargs: multiplies.append((value, keep, output_tensor)) or output_tensor, ) tt_penalties.TTPenalties.reset_output_tokens(penalty_state, slots=[3, 1]) assert allocated[0][0].reshape(-1).tolist() == [1, 0, 1, 0] assert allocated[0][1] == (0, None) keep = allocated[0][2] assert [(value.name, mask is keep, output.name) for value, mask, output in multiplies] == [ ("mask", True, "mask"), ("counts", True, "counts"), ("gathered", True, "gathered"), ] assert keep.deallocated def test_no_sampling_updates_is_a_true_noop(): fake, calls = _fake_sampling_generator() SamplingGenerator.apply_decode_state( fake, [object()], reload_sampling_params=False, reset_sampling_state=False, ) assert calls == [] def test_sampling_update_commands_are_required_and_have_no_legacy_alias(): params = inspect.signature(SamplingGenerator.apply_decode_state).parameters assert params["reload_sampling_params"].default is inspect.Parameter.empty assert params["reset_sampling_state"].default is inspect.Parameter.empty assert "reset_batch" not in params def test_deepseek_rejects_partial_forward_reload_before_decode(expect_error): from models.demos.deepseek_v3.tt.generator_vllm import DeepseekV3ForCausalLM generator = SimpleNamespace(model_run_config_decode=object()) for reload_inputs, reload_page_table in ((False, False), (False, True), (True, True)): with expect_error(ValueError, "requires a full host-input reload"): DeepseekV3ForCausalLM.decode_forward( generator, reload_inputs=reload_inputs, reload_page_table=reload_page_table, reload_sampling_params=False, reset_sampling_state=False, ) def test_decode_warmup_does_not_reset_absent_request_history(): calls = [] fake = SimpleNamespace( _create_sampling_params=lambda *args, **kwargs: [object()], _create_decode_warmup_inputs=lambda *args: ( torch.zeros((1, 1)), torch.zeros((1,)), torch.zeros((1, 1)), ), decode_forward=lambda **kwargs: calls.append(kwargs), ) WarmupForwardMixin.warmup_model_decode( fake, kv_cache=object(), enable_trace=True, max_batch_size=1, num_blocks=1, can_sample_on_device=True, ) assert len(calls) == 1 assert calls[0]["reload_sampling_params"] is True assert calls[0]["reset_sampling_state"] is False def test_qwen_vl_slot_remap_moves_persistent_rope_deltas(): from models.demos.qwen3_vl.tt.generator import Generator as Qwen3Generator from models.demos.qwen25_vl.tt.generator import Generator as Qwen25Generator for generator_cls in (Qwen25Generator, Qwen3Generator): generator = SimpleNamespace( model=SimpleNamespace( rope_setup=SimpleNamespace( batch_size=4, rope_deltas=torch.tensor([10, 20, 30, 40]), ) ) ) generator_cls.remap_rope_deltas(generator, [3, 1, 2, 3]) assert generator.model.rope_setup.rope_deltas.tolist() == [40, 20, 30, 40] def test_qwen_vl_generator_forwards_slot_remap_to_shared_sampling_owner(): from models.demos.qwen3_vl.tt.generator import Generator as Qwen3Generator from models.demos.qwen25_vl.tt.generator import Generator as Qwen25Generator for generator_cls in (Qwen25Generator, Qwen3Generator): calls = [] generator = SimpleNamespace( _ttt_generator=SimpleNamespace(decode_forward=lambda **kwargs: calls.append(kwargs) or "output") ) result = generator_cls.decode_forward( generator, tokens="tokens", start_pos="positions", slot_remap=[3, 1, 2, 3], reload_inputs=True, reload_page_table=False, reload_sampling_params=False, reset_sampling_state=False, ) assert result == "output" assert calls[0]["slot_remap"] == [3, 1, 2, 3] def test_shared_generator_routes_slot_remap_to_exactly_one_sampling_owner(): from models.tt_transformers.tt.generator import Generator def run(*, sampling_params=None, defer_device_sampling=False, fail_readback=False): events = [] fake = SimpleNamespace( mode=Mode.DECODE, model=[SimpleNamespace(switch_mode=lambda mode: None)], data_parallel=1, _decode_forward_trace_text=lambda **kwargs: events.append("decode") or "logits", sample_decode_on_device=lambda output, **kwargs: events.append(("sample", kwargs["slot_remap"])) or "tokens", read_decode_output=lambda output: events.append("read") or output, process_decode_output_host=lambda output, **kwargs: (_ for _ in ()).throw(RuntimeError("readback")) if fail_readback else events.append("process") or output, _apply_sampling_slot_remap=lambda remap: events.append(("host-remap", remap)), ) try: result = Generator.decode_forward( fake, torch.zeros((1, 1), dtype=torch.int64), torch.zeros((1,), dtype=torch.int64), sampling_params=sampling_params, slot_remap=[0], defer_device_sampling=defer_device_sampling, reload_inputs=True, reload_page_table=False, reload_sampling_params=False, reset_sampling_state=False, ) except RuntimeError: result = None return events, result host_events, _ = run() device_events, _ = run(sampling_params=object()) deferred_events, _ = run(defer_device_sampling=True) failed_host_events, _ = run(fail_readback=True) assert host_events == ["decode", "read", "process", ("host-remap", [0])] assert device_events == ["decode", ("sample", [0]), "read", "process"] assert deferred_events == ["decode"] assert failed_host_events == ["decode", "read"] def test_shared_generator_rebases_and_pads_lane_sampling_remaps(): from models.tt_transformers.tt.generator import Generator calls = [[], []] models = [] for lane in range(2): sampling = SimpleNamespace( seed_manager=SimpleNamespace(max_batch_size=4), apply_slot_remap=lambda remap, lane=lane: calls[lane].append(torch.as_tensor(remap).tolist()), ) models.append(SimpleNamespace(sampling=sampling)) fake = SimpleNamespace(data_parallel=2, model=models) Generator._apply_sampling_slot_remap(fake, torch.tensor([1, 0, 3, 2])) assert calls == [[[1, 0, 2, 3]], [[1, 0, 2, 3]]] def test_shared_generator_rejects_cross_lane_sampling_remap(expect_error): from models.tt_transformers.tt.generator import Generator sampling = SimpleNamespace(seed_manager=SimpleNamespace(max_batch_size=2), apply_slot_remap=lambda remap: None) fake = SimpleNamespace( data_parallel=2, model=[ SimpleNamespace(sampling=sampling), SimpleNamespace(sampling=sampling), ], ) with expect_error(ValueError, "outside its global range"): Generator._apply_sampling_slot_remap(fake, torch.tensor([0, 2, 2, 3])) def test_sglang_bridge_explicitly_requests_host_authoritative_decode(monkeypatch): from models.tt_transformers.tt.generator import Generator from models.tt_transformers.tt.generator_sglang import LlamaForCausalLM calls = [] monkeypatch.setattr( Generator, "decode_forward", lambda self, *args, **kwargs: calls.append(kwargs) or "output", ) result = LlamaForCausalLM.decode_forward(object(), tokens="tokens", start_pos="positions") assert result == "output" assert calls == [ { "tokens": "tokens", "start_pos": "positions", "reload_inputs": True, "reload_page_table": False, "reload_sampling_params": False, "reset_sampling_state": False, } ] def test_lfm_demo_supplies_every_decode_update_command(): source_path = Path("models/demos/multimodal/lfm25_vl/demo/vision_demo.py") tree = ast.parse(source_path.read_text()) required = { "reload_inputs", "reload_page_table", "reload_sampling_params", "reset_sampling_state", } calls = [ node for node in ast.walk(tree) if isinstance(node, ast.Call) and isinstance(node.func, ast.Attribute) and node.func.attr == "decode_forward" and isinstance(node.func.value, ast.Name) and node.func.value.id == "generator" ] assert len(calls) == 2 for call in calls: assert required <= {keyword.arg for keyword in call.keywords} def test_all_known_shared_generator_callers_supply_every_decode_update_command(): required = { "reload_inputs", "reload_page_table", "reload_sampling_params", "reset_sampling_state", } expected_calls = { Path("tt-train/sources/examples/grpo_remote_rollout/utils/ttt_generation_worker.py"): 1, Path("models/experimental/ops/quasar/gpt_oss/demo/text_demo.py"): 2, Path("models/experimental/ops/quasar/gpt_oss/tests/accuracy/test_model.py"): 1, Path("models/experimental/ops/quasar/qwen3_vl/demo/demo.py"): 1, } for source_path, expected_count in expected_calls.items(): tree = ast.parse(source_path.read_text()) calls = [ node for node in ast.walk(tree) if isinstance(node, ast.Call) and isinstance(node.func, ast.Attribute) and node.func.attr == "decode_forward" ] assert len(calls) == expected_count, source_path for call in calls: assert required <= {keyword.arg for keyword in call.keywords}, (source_path, call.lineno) def test_shared_generator_preserves_explicit_commands_after_mainline_seed_fix(): source_path = Path("models/tt_transformers/tt/generator.py") source_text = source_path.read_text() tree = ast.parse(source_text) generator = next(node for node in tree.body if isinstance(node, ast.ClassDef) and node.name == "Generator") decode = next( node for node in generator.body if isinstance(node, ast.FunctionDef) and node.name == "decode_forward" ) trace_decode = next( node for node in generator.body if isinstance(node, ast.FunctionDef) and node.name == "_decode_forward_trace_text" ) required = {"reload_inputs", "reload_page_table", "reload_sampling_params", "reset_sampling_state"} assert required <= {arg.arg for arg in decode.args.kwonlyargs} assert {"reload_inputs", "reload_page_table"} <= {arg.arg for arg in trace_decode.args.kwonlyargs} decode_source = ast.get_source_segment(source_text, decode) trace_source = ast.get_source_segment(source_text, trace_decode) assert decode_source is not None assert trace_source is not None assert "reset_batch" not in decode_source assert "_prev_on_device_sampling" not in decode_source assert "_tt_vllm_always_refresh_decode_trace_inputs" not in trace_source assert "torch.equal" not in trace_source def test_gemma4_override_uses_only_explicit_decode_update_commands(): source_path = Path("models/demos/gemma4/tt/generator.py") tree = ast.parse(source_path.read_text()) mixin = next( node for node in tree.body if isinstance(node, ast.ClassDef) and node.name == "ChunkedPrefillPageTableGuardMixin" ) decode = next(node for node in mixin.body if isinstance(node, ast.FunctionDef) and node.name == "decode_forward") trace_decode = next( node for node in mixin.body if isinstance(node, ast.FunctionDef) and node.name == "_decode_forward_trace_text" ) required = { "reload_inputs", "reload_page_table", "reload_sampling_params", "reset_sampling_state", } decode_params = {arg.arg for arg in decode.args.kwonlyargs} trace_params = {arg.arg for arg in trace_decode.args.kwonlyargs} assert required <= decode_params assert {"reload_inputs", "reload_page_table"} <= trace_params assert "reset_batch" not in {arg.arg for arg in decode.args.args} assert "reset_batch" not in {arg.arg for arg in trace_decode.args.args} source = ast.get_source_segment(source_path.read_text(), trace_decode) assert source is not None assert "_prev_on_device_sampling" not in source assert "_tt_vllm_always_refresh_decode_trace_inputs" not in source assert "torch.equal" not in source def test_gemma4_pli_explicitly_disables_decode_token_feedback(): source = Path("models/demos/gemma4/tt/model.py").read_text() assert "_tt_supports_decode_token_feedback = False" in source assert "self._tt_supports_decode_token_feedback = not self._tt_vllm_always_refresh_decode_trace_inputs" in source def test_explicit_trace_reloads_select_mode_and_gemma4_bucket_buffers(monkeypatch): from models.demos.gemma4.tt import generator as gemma4_generator from models.tt_transformers.tt import common from models.tt_transformers.tt import generator as shared_generator copies = [] def copy_inputs(*, host_tensors, device_tensors): copies.append("full") for host, device in zip(host_tensors, device_tensors): device.copy_(host) def copy_page(host, device): copies.append("page") device.copy_(host) monkeypatch.setattr(common, "copy_host_to_device", copy_inputs) monkeypatch.setattr(shared_generator, "copy_host_to_device", copy_inputs) monkeypatch.setattr(shared_generator.ttnn, "copy_host_to_device_tensor", copy_page) monkeypatch.setattr(shared_generator.ttnn, "execute_trace", lambda *args, **kwargs: None) for cls, buckets in ( (shared_generator.Generator, (1,)), (gemma4_generator.ChunkedPrefillPageTableGuardMixin, (1, 32)), ): def key(mode, batch): return (mode, batch) if len(buckets) > 1 else mode keys = [key(mode, batch) for mode in (False, True) for batch in buckets] buffers = {k: [[torch.full((1,), -1) for _ in range(4)]] for k in keys} fake = SimpleNamespace( data_parallel=1, model=[ SimpleNamespace( prepare_decode_inputs_host=lambda tokens, positions, page: [ tokens[0], positions[:1], positions[:1], page[0], ] ) ], model_args=[SimpleNamespace(mesh_device=object())], trace_ids_decode={k: {0: object()} for k in keys}, trace_inputs_decode=buffers, trace_output_decode={k: object() for k in keys}, ) fake._decode_trace_key = cls._decode_trace_key.__get__(fake) # The plugin commands a full reload on mode/layout transitions. Switch # back to an existing bucket too, where stale resident inputs matter. for mode, batch in [(True, 1), (False, 1), (True, buckets[-1]), (True, 1)]: selected = buffers[key(mode, batch)][0] kwargs = dict( tokens=[torch.full((batch, 1), 11)], current_pos=[torch.full((batch,), 12)], page_table=[torch.full((batch, 1), 13)], on_device_sampling=mode, ) copies.clear() cls._decode_forward_trace_text(fake, **kwargs, reload_inputs=True, reload_page_table=False) assert [int(t.item()) for t in selected] == [11, 12, 12, 13] assert copies == ["full"] selected[0].fill_(21) selected[1].fill_(22) copies.clear() cls._decode_forward_trace_text(fake, **kwargs, reload_inputs=False, reload_page_table=True) assert [int(t.item()) for t in selected] == [21, 22, 12, 13] assert copies == ["page"] copies.clear() cls._decode_forward_trace_text(fake, **kwargs, reload_inputs=False, reload_page_table=False) assert [int(t.item()) for t in selected] == [21, 22, 12, 13] assert copies == [] def test_gemma4_decode_restores_full_layer_page_tables_after_sequential_prefill(): from models.demos.gemma4.tt.generator import ChunkedPrefillPageTableGuardMixin full_tables = [torch.tensor([[1], [2]])] model = SimpleNamespace(switch_mode=lambda mode: None) def decode(stage, **kwargs): assert model._active_page_tables_per_layer is full_tables assert not hasattr(model, "_sequential_batch_page_tables") return stage fake = SimpleNamespace( mode=Mode.PREFILL, model=[model], data_parallel=1, _decode_forward_trace_text=lambda **kwargs: decode("replay", **kwargs), _prepare_decode_trace_variant=lambda **kwargs: decode("prepare", **kwargs), _apply_sampling_slot_remap=lambda remap: None, ) fake._clear_sequential_batch_page_tables = ( lambda: ChunkedPrefillPageTableGuardMixin._clear_sequential_batch_page_tables(fake) ) for enable_trace, prepare_trace, expected in ((True, False, "replay"), (False, True, "prepare")): model._active_page_tables_per_layer = [full_tables[0][:1]] model._sequential_batch_page_tables = full_tables assert ( ChunkedPrefillPageTableGuardMixin.decode_forward( fake, tokens=torch.zeros((2, 1), dtype=torch.int32), start_pos=torch.ones(2, dtype=torch.int32), enable_trace=enable_trace, prepare_trace=prepare_trace, read_from_device=False, reload_inputs=True, reload_page_table=False, reload_sampling_params=False, reset_sampling_state=False, ) == expected ) def test_galaxy_reset_only_formats_seed_slots(monkeypatch): from models.demos.llama3_70b_galaxy.tt import generator as galaxy_generator formatted = SimpleNamespace(seed=[17, 17, 17, 17]) monkeypatch.setattr( galaxy_generator, "format_sampling_params", lambda params, max_batch_size: formatted, ) monkeypatch.setattr( galaxy_generator, "_fill_inactive_params_from_active", lambda params, active_slots, max_batch_size: params, ) reset_calls = [] seed_manager = SimpleNamespace( max_batch_size=4, deactivate_slots_except=lambda slots: None, reset_seed_from_slots=lambda seeds, slots: reset_calls.append((seeds, slots)), align_seed_counters_to_positions=lambda *args: None, reset_seed_from_slots_if_needed=lambda *args: None, get_new_values=lambda slots: None, ) sampling = SimpleNamespace( seed_manager=seed_manager, _slot_state_requires_authoritative_reload=False, reset_sampling_params=lambda params: (_ for _ in ()).throw( AssertionError("reset-only must not upload sampling parameters") ), reset_prompt_tokens=lambda tokens: None, reset_output_state=lambda tokens: None, sample=lambda **kwargs: "tokens", ) sampling.validate_decode_state_commands = lambda **kwargs: SamplingGenerator.validate_decode_state_commands( sampling, **kwargs ) sampling.commit_decode_state_commands = lambda **kwargs: SamplingGenerator.commit_decode_state_commands( sampling, **kwargs ) fake = SimpleNamespace( trace_inputs_decode={True: None}, model=SimpleNamespace(sampling=sampling), model_args=SimpleNamespace(max_batch_size=4), _apply_sampling_slot_remap=lambda remap: None, ) result = galaxy_generator.Generator.sample_decode_on_device( fake, tt_logits="logits", sampling_params=SimpleNamespace(seed=[17]), start_pos=torch.tensor([0, -1, -1, -1]), reload_inputs=True, reload_sampling_params=False, reset_sampling_state=True, ) assert result == "tokens" assert reset_calls == [([17, 17, 17, 17], [0])] def test_galaxy_generator_routes_slot_remap_to_exactly_one_sampling_owner(): from models.demos.llama3_70b_galaxy.tt.generator import Generator def run(*, sampling_params=None, defer_device_sampling=False): events = [] fake = SimpleNamespace( model=SimpleNamespace(is_decode_setup=True), _decode_easy_trace_text=lambda **kwargs: events.append("decode") or ("logits", None), sample_decode_on_device=lambda output, **kwargs: events.append(("sample", kwargs["slot_remap"])) or ("tokens", "logprobs"), read_decode_output=lambda output, **kwargs: events.append("read") or output, process_decode_output_host=lambda output, **kwargs: events.append("process") or output, _apply_sampling_slot_remap=lambda remap: events.append(("host-remap", remap)), ) result = Generator.decode_forward( fake, torch.zeros((1, 1), dtype=torch.int64), torch.zeros((1,), dtype=torch.int64), kv_cache=[object()], sampling_params=sampling_params, slot_remap=[0], defer_device_sampling=defer_device_sampling, reload_inputs=True, reload_page_table=False, reload_sampling_params=False, reset_sampling_state=False, ) assert fake._decode_reload_inputs is True return events, result host_events, _ = run() device_events, _ = run(sampling_params=object()) deferred_events, _ = run(defer_device_sampling=True) assert host_events == ["decode", "read", "process", ("host-remap", [0])] assert device_events == ["decode", ("sample", [0]), "read", "process"] assert deferred_events == ["decode"] def test_galaxy_slot_remap_moves_parameter_shadow_with_seed_state(): from models.demos.llama3_70b_galaxy.tt.generator import Generator seed_remaps = [] fake = SimpleNamespace( model=SimpleNamespace( sampling=SimpleNamespace( seed_manager=SimpleNamespace(max_batch_size=4), apply_slot_remap=lambda remap: seed_remaps.append(remap), ) ), model_args=SimpleNamespace(max_batch_size=4), _slot_sampling_params={ "temperature": [0.1, 0.2, 0.3, 0.4], "top_k": [1, 2, 3, 4], }, ) Generator._apply_sampling_slot_remap(fake, [2, 0, 1, 3]) assert seed_remaps == [[2, 0, 1, 3]] assert fake._slot_sampling_params == { "temperature": [0.3, 0.1, 0.2, 0.4], "top_k": [3, 1, 2, 4], } def test_galaxy_seed_stream_realigns_only_on_authoritative_input_reload(): from models.demos.llama3_70b_galaxy.tt.generator import Generator class RecordingSeedManager(SeedManager): def write_device_seed_values(self, values): pushed.append(values[0]) pushed = [] events = [] manager = RecordingSeedManager(SimpleNamespace(_sampling_dp=1), max_batch_size=32) sampling = SimpleNamespace( seed_manager=manager, validate_decode_state_commands=lambda **kwargs: None, commit_decode_state_commands=lambda **kwargs: None, reset_sampling_params=lambda params: events.append("params"), reset_prompt_tokens=lambda tokens: events.append("prompt"), reset_output_state=lambda tokens: events.append("output"), sample=lambda **kwargs: "tokens", ) generator = SimpleNamespace( trace_inputs_decode={True: None}, model=SimpleNamespace(sampling=sampling), model_args=SimpleNamespace(max_batch_size=32), _apply_sampling_slot_remap=lambda remap: None, _remember_slot_params=lambda params: None, ) params = SamplingParams(temperature=[1.0] * 32, top_k=[32] * 32, top_p=[1.0] * 32, seed=[7] + [None] * 31) def step(position, *, reload_inputs, reset=False): Generator.sample_decode_on_device( generator, "logits", sampling_params=params, start_pos=torch.tensor([position] + [-1] * 31), reload_inputs=reload_inputs, reload_sampling_params=reset, reset_sampling_state=reset, ) step(100, reload_inputs=True, reset=True) for stale_position in (100, 101, 101): step(stale_position, reload_inputs=False) # An authoritative full reload can move the position without changing # sampling parameters or resetting penalty history. step(200, reload_inputs=True) assert pushed == [_hash_request_seed_to_device_seed(7, pos) for pos in (101, 102, 103, 104, 201)] assert events == ["params", "prompt", "output"] def test_unseeded_decode_reset_loads_fresh_device_seed(monkeypatch): seed_buffer = object() manager = SeedManager(SimpleNamespace(_sampling_dp=1, seeds_tt_tensor=seed_buffer), max_batch_size=1) manager.seed_counters = [4] before = manager.rngs[0].getstate() host_seed_tensor = object() uploads = [] monkeypatch.setattr(manager, "_next_unseeded_device_seed", lambda: 123) monkeypatch.setattr( "models.common.sampling.generator.ttnn.from_torch", lambda *args, **kwargs: host_seed_tensor, ) monkeypatch.setattr( "models.common.sampling.generator.ttnn.copy_host_to_device_tensor", lambda host, device: uploads.append((host, device)), ) # The conditional path would see None == None and do nothing. A decode # state reset must be unconditional so the following get_new_values() # enters its init state and uploads a fresh device seed. manager.reset_seed_from_slots([None], [0]) manager.get_new_values([0]) assert manager.seed_counters == [0] assert manager.rngs[0].getstate() != before assert uploads == [(host_seed_tensor, seed_buffer)] assert manager._needs_skip assert not manager._reseted