# SPDX-FileCopyrightText: © 2025 Tenstorrent USA, Inc. # SPDX-License-Identifier: Apache-2.0 """Tests for Penalties1D module.""" from types import SimpleNamespace import pytest import torch import ttnn from models.common.modules.lazy_buffer import LazyBuffer from models.common.modules.sampling import penalties_1d as penalties_module from models.common.modules.sampling.penalties_1d import ( Penalties1D, Penalties1DConfig, PenaltyAccumulator, PenaltyParams, _materialize, _resolve_penalties1d_config, ) from models.common.utility_functions import comp_pcc # 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 _list_collected_penalty_cases() -> list[pytest.param]: """ Collected from TTTv1 demo runs (Phase B of test_case_collection.md). Each entry is: (mesh_shape, vocab_size, batch_size, seq_len, presence, frequency, repetition, pcc, hf_model_name) Source CSVs: sampling_generator_config_collected.csv, sampling_generator_params_collected.csv, penalties_prompt_tokens_collected.csv Deduplicated by (topology, vocab, batch, seq_len, penalties_active). """ # fmt: off return [ # --- (1,1) Mistral7B v32768 --- pytest.param((1, 1), 32768, 32, 128, 0.0, 0.0, 1.0, 0.999, MISTRAL_7B, id="1x1-Mistral7B-v32768-b32-s128-no-pen"), pytest.param((1, 1), 32768, 32, 128, 1.2, 1.2, 1.5, 0.95, MISTRAL_7B, id="1x1-Mistral7B-v32768-b32-s128-pen", marks=_slow), pytest.param((1, 1), 32768, 32, 1024, 0.0, 0.0, 1.0, 0.999, MISTRAL_7B, id="1x1-Mistral7B-v32768-b32-s1024-no-pen", marks=_slow), pytest.param((1, 1), 32768, 32, 1024, 1.2, 1.2, 1.5, 0.95, MISTRAL_7B, id="1x1-Mistral7B-v32768-b32-s1024-pen", marks=_slow), pytest.param((1, 1), 32768, 32, 2048, 0.0, 0.0, 1.0, 0.999, MISTRAL_7B, id="1x1-Mistral7B-v32768-b32-s2048-no-pen", marks=_slow), pytest.param((1, 1), 32768, 32, 2048, 1.2, 1.2, 1.5, 0.95, MISTRAL_7B, id="1x1-Mistral7B-v32768-b32-s2048-pen", marks=_slow), pytest.param((1, 1), 32768, 32, 4096, 0.0, 0.0, 1.0, 0.999, MISTRAL_7B, id="1x1-Mistral7B-v32768-b32-s4096-no-pen", marks=_slow), pytest.param((1, 1), 32768, 32, 4096, 1.2, 1.2, 1.5, 0.95, MISTRAL_7B, id="1x1-Mistral7B-v32768-b32-s4096-pen", marks=_slow), # --- (1,1) Llama8B v128256 --- pytest.param((1, 1), 128256, 1, 125, 0.0, 0.0, 1.0, 0.999, LLAMA_8B, id="1x1-Llama8B-v128256-b1-s125-no-pen"), # --- (1,1) Llama1B v128256 --- pytest.param((1, 1), 128256, 1, 71, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b1-s71-no-pen", marks=_slow), pytest.param((1, 1), 128256, 1, 80, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b1-s80-no-pen", marks=_slow), pytest.param((1, 1), 128256, 1, 115, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b1-s115-no-pen", marks=_slow), pytest.param((1, 1), 128256, 1, 337, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b1-s337-no-pen", marks=_slow), pytest.param((1, 1), 128256, 1, 512, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b1-s512-no-pen", marks=_slow), pytest.param((1, 1), 128256, 1, 785, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b1-s785-no-pen", marks=_slow), pytest.param((1, 1), 128256, 1, 16229, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b1-s16229-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 52, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s52-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 56, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s56-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 57, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s57-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 59, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s59-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 60, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s60-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 61, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s61-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 63, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s63-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 65, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s65-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 66, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s66-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 67, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s67-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 70, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s70-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 71, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s71-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 72, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s72-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 75, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s75-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 77, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s77-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 80, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s80-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 81, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s81-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 84, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s84-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 86, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s86-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 87, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s87-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 91, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s91-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 93, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s93-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 94, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s94-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 96, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s96-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 98, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s98-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 99, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s99-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 101, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s101-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 102, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s102-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 103, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s103-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 104, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s104-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 105, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s105-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 106, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s106-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 109, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s109-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 113, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s113-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 114, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s114-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 115, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s115-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 116, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s116-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 117, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s117-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 119, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s119-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 124, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s124-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 125, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s125-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 128, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s128-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 128, 1.2, 1.2, 1.5, 0.95, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s128-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 337, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s337-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 512, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s512-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 712, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s712-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 785, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s785-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 1024, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s1024-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 1024, 1.2, 1.2, 1.5, 0.95, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s1024-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 2048, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s2048-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 2048, 1.2, 1.2, 1.5, 0.95, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s2048-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 4096, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s4096-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 4096, 1.2, 1.2, 1.5, 0.95, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s4096-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 8192, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s8192-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 8192, 1.2, 1.2, 1.5, 0.95, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s8192-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 16229, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s16229-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 16384, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s16384-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 16384, 1.2, 1.2, 1.5, 0.95, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s16384-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 32768, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s32768-no-pen", marks=_slow), pytest.param((1, 1), 128256, 32, 32768, 1.2, 1.2, 1.5, 0.95, LLAMA_1B, id="1x1-Llama1B-v128256-b32-s32768-pen", marks=_slow), # --- (1,2) Mistral7B v32768 --- pytest.param((1, 2), 32768, 32, 128, 0.0, 0.0, 1.0, 0.999, MISTRAL_7B, id="1x2-Mistral7B-v32768-b32-s128-no-pen"), pytest.param((1, 2), 32768, 32, 128, 1.2, 1.2, 1.5, 0.95, MISTRAL_7B, id="1x2-Mistral7B-v32768-b32-s128-pen", marks=_slow), pytest.param((1, 2), 32768, 32, 1024, 0.0, 0.0, 1.0, 0.999, MISTRAL_7B, id="1x2-Mistral7B-v32768-b32-s1024-no-pen", marks=_slow), pytest.param((1, 2), 32768, 32, 1024, 1.2, 1.2, 1.5, 0.95, MISTRAL_7B, id="1x2-Mistral7B-v32768-b32-s1024-pen", marks=_slow), pytest.param((1, 2), 32768, 32, 2048, 0.0, 0.0, 1.0, 0.999, MISTRAL_7B, id="1x2-Mistral7B-v32768-b32-s2048-no-pen", marks=_slow), pytest.param((1, 2), 32768, 32, 2048, 1.2, 1.2, 1.5, 0.95, MISTRAL_7B, id="1x2-Mistral7B-v32768-b32-s2048-pen", marks=_slow), pytest.param((1, 2), 32768, 32, 4096, 0.0, 0.0, 1.0, 0.999, MISTRAL_7B, id="1x2-Mistral7B-v32768-b32-s4096-no-pen", marks=_slow), pytest.param((1, 2), 32768, 32, 4096, 1.2, 1.2, 1.5, 0.95, MISTRAL_7B, id="1x2-Mistral7B-v32768-b32-s4096-pen", marks=_slow), # --- (1,2) Llama1B v128256 --- pytest.param((1, 2), 128256, 1, 71, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b1-s71-no-pen"), pytest.param((1, 2), 128256, 1, 80, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b1-s80-no-pen", marks=_slow), pytest.param((1, 2), 128256, 1, 115, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b1-s115-no-pen", marks=_slow), pytest.param((1, 2), 128256, 1, 337, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b1-s337-no-pen", marks=_slow), pytest.param((1, 2), 128256, 1, 512, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b1-s512-no-pen", marks=_slow), pytest.param((1, 2), 128256, 1, 785, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b1-s785-no-pen", marks=_slow), pytest.param((1, 2), 128256, 1, 16229, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b1-s16229-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 52, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s52-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 56, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s56-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 57, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s57-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 59, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s59-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 60, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s60-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 61, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s61-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 63, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s63-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 65, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s65-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 66, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s66-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 67, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s67-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 70, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s70-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 71, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s71-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 72, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s72-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 75, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s75-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 77, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s77-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 80, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s80-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 81, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s81-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 84, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s84-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 86, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s86-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 87, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s87-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 91, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s91-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 93, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s93-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 94, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s94-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 96, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s96-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 98, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s98-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 99, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s99-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 101, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s101-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 102, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s102-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 103, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s103-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 104, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s104-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 105, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s105-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 106, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s106-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 109, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s109-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 113, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s113-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 114, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s114-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 115, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s115-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 116, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s116-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 117, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s117-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 119, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s119-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 124, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s124-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 125, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s125-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 128, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s128-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 128, 1.2, 1.2, 1.5, 0.95, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s128-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 337, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s337-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 512, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s512-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 712, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s712-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 785, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s785-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 1024, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s1024-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 1024, 1.2, 1.2, 1.5, 0.95, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s1024-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 2048, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s2048-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 2048, 1.2, 1.2, 1.5, 0.95, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s2048-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 4096, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s4096-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 4096, 1.2, 1.2, 1.5, 0.95, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s4096-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 8192, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s8192-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 8192, 1.2, 1.2, 1.5, 0.95, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s8192-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 16229, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s16229-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 16384, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s16384-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 16384, 1.2, 1.2, 1.5, 0.95, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s16384-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 32768, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s32768-no-pen", marks=_slow), pytest.param((1, 2), 128256, 32, 32768, 1.2, 1.2, 1.5, 0.95, LLAMA_1B, id="1x2-Llama1B-v128256-b32-s32768-pen", marks=_slow), # --- (1,8) Mixtral8x7B v32000 --- pytest.param((1, 8), 32000, 1, 87, 0.0, 0.0, 1.0, 0.999, MIXTRAL_8X7B, id="1x8-Mixtral8x7B-v32000-b1-s87-no-pen"), pytest.param((1, 8), 32000, 1, 393, 0.0, 0.0, 1.0, 0.999, MIXTRAL_8X7B, id="1x8-Mixtral8x7B-v32000-b1-s393-no-pen", marks=_slow), pytest.param((1, 8), 32000, 32, 57, 0.0, 0.0, 1.0, 0.999, MIXTRAL_8X7B, id="1x8-Mixtral8x7B-v32000-b32-s57-no-pen", marks=_slow), pytest.param((1, 8), 32000, 32, 61, 0.0, 0.0, 1.0, 0.999, MIXTRAL_8X7B, id="1x8-Mixtral8x7B-v32000-b32-s61-no-pen", marks=_slow), pytest.param((1, 8), 32000, 32, 64, 0.0, 0.0, 1.0, 0.999, MIXTRAL_8X7B, id="1x8-Mixtral8x7B-v32000-b32-s64-no-pen", marks=_slow), pytest.param((1, 8), 32000, 32, 66, 0.0, 0.0, 1.0, 0.999, MIXTRAL_8X7B, id="1x8-Mixtral8x7B-v32000-b32-s66-no-pen", marks=_slow), pytest.param((1, 8), 32000, 32, 67, 0.0, 0.0, 1.0, 0.999, MIXTRAL_8X7B, id="1x8-Mixtral8x7B-v32000-b32-s67-no-pen", marks=_slow), pytest.param((1, 8), 32000, 32, 69, 0.0, 0.0, 1.0, 0.999, MIXTRAL_8X7B, id="1x8-Mixtral8x7B-v32000-b32-s69-no-pen", marks=_slow), pytest.param((1, 8), 32000, 32, 70, 0.0, 0.0, 1.0, 0.999, MIXTRAL_8X7B, id="1x8-Mixtral8x7B-v32000-b32-s70-no-pen", marks=_slow), pytest.param((1, 8), 32000, 32, 72, 0.0, 0.0, 1.0, 0.999, MIXTRAL_8X7B, id="1x8-Mixtral8x7B-v32000-b32-s72-no-pen", marks=_slow), pytest.param((1, 8), 32000, 32, 73, 0.0, 0.0, 1.0, 0.999, MIXTRAL_8X7B, id="1x8-Mixtral8x7B-v32000-b32-s73-no-pen", marks=_slow), pytest.param((1, 8), 32000, 32, 74, 0.0, 0.0, 1.0, 0.999, MIXTRAL_8X7B, id="1x8-Mixtral8x7B-v32000-b32-s74-no-pen", marks=_slow), pytest.param((1, 8), 32000, 32, 76, 0.0, 0.0, 1.0, 0.999, MIXTRAL_8X7B, id="1x8-Mixtral8x7B-v32000-b32-s76-no-pen", marks=_slow), pytest.param((1, 8), 32000, 32, 79, 0.0, 0.0, 1.0, 0.999, MIXTRAL_8X7B, id="1x8-Mixtral8x7B-v32000-b32-s79-no-pen", marks=_slow), pytest.param((1, 8), 32000, 32, 82, 0.0, 0.0, 1.0, 0.999, MIXTRAL_8X7B, id="1x8-Mixtral8x7B-v32000-b32-s82-no-pen", marks=_slow), pytest.param((1, 8), 32000, 32, 83, 0.0, 0.0, 1.0, 0.999, MIXTRAL_8X7B, id="1x8-Mixtral8x7B-v32000-b32-s83-no-pen", marks=_slow), pytest.param((1, 8), 32000, 32, 87, 0.0, 0.0, 1.0, 0.999, MIXTRAL_8X7B, id="1x8-Mixtral8x7B-v32000-b32-s87-no-pen", marks=_slow), pytest.param((1, 8), 32000, 32, 89, 0.0, 0.0, 1.0, 0.999, MIXTRAL_8X7B, id="1x8-Mixtral8x7B-v32000-b32-s89-no-pen", marks=_slow), pytest.param((1, 8), 32000, 32, 101, 0.0, 0.0, 1.0, 0.999, MIXTRAL_8X7B, id="1x8-Mixtral8x7B-v32000-b32-s101-no-pen", marks=_slow), pytest.param((1, 8), 32000, 32, 128, 0.0, 0.0, 1.0, 0.999, MIXTRAL_8X7B, id="1x8-Mixtral8x7B-v32000-b32-s128-no-pen", marks=_slow), pytest.param((1, 8), 32000, 32, 128, 1.2, 1.2, 1.5, 0.95, MIXTRAL_8X7B, id="1x8-Mixtral8x7B-v32000-b32-s128-pen", marks=_slow), pytest.param((1, 8), 32000, 32, 393, 0.0, 0.0, 1.0, 0.999, MIXTRAL_8X7B, id="1x8-Mixtral8x7B-v32000-b32-s393-no-pen", marks=_slow), pytest.param((1, 8), 32000, 32, 820, 0.0, 0.0, 1.0, 0.999, MIXTRAL_8X7B, id="1x8-Mixtral8x7B-v32000-b32-s820-no-pen", marks=_slow), pytest.param((1, 8), 32000, 32, 1024, 0.0, 0.0, 1.0, 0.999, MIXTRAL_8X7B, id="1x8-Mixtral8x7B-v32000-b32-s1024-no-pen", marks=_slow), pytest.param((1, 8), 32000, 32, 1024, 1.2, 1.2, 1.5, 0.95, MIXTRAL_8X7B, id="1x8-Mixtral8x7B-v32000-b32-s1024-pen", marks=_slow), pytest.param((1, 8), 32000, 32, 2048, 0.0, 0.0, 1.0, 0.999, MIXTRAL_8X7B, id="1x8-Mixtral8x7B-v32000-b32-s2048-no-pen", marks=_slow), pytest.param((1, 8), 32000, 32, 2048, 1.2, 1.2, 1.5, 0.95, MIXTRAL_8X7B, id="1x8-Mixtral8x7B-v32000-b32-s2048-pen", marks=_slow), pytest.param((1, 8), 32000, 32, 4096, 0.0, 0.0, 1.0, 0.999, MIXTRAL_8X7B, id="1x8-Mixtral8x7B-v32000-b32-s4096-no-pen", marks=_slow), pytest.param((1, 8), 32000, 32, 4096, 1.2, 1.2, 1.5, 0.95, MIXTRAL_8X7B, id="1x8-Mixtral8x7B-v32000-b32-s4096-pen", marks=_slow), # --- (1,8) Mistral7B v32768 --- pytest.param((1, 8), 32768, 32, 128, 0.0, 0.0, 1.0, 0.999, MISTRAL_7B, id="1x8-Mistral7B-v32768-b32-s128-no-pen"), pytest.param((1, 8), 32768, 32, 128, 1.2, 1.2, 1.5, 0.95, MISTRAL_7B, id="1x8-Mistral7B-v32768-b32-s128-pen", marks=_slow), pytest.param((1, 8), 32768, 32, 1024, 0.0, 0.0, 1.0, 0.999, MISTRAL_7B, id="1x8-Mistral7B-v32768-b32-s1024-no-pen", marks=_slow), pytest.param((1, 8), 32768, 32, 1024, 1.2, 1.2, 1.5, 0.95, MISTRAL_7B, id="1x8-Mistral7B-v32768-b32-s1024-pen", marks=_slow), pytest.param((1, 8), 32768, 32, 2048, 0.0, 0.0, 1.0, 0.999, MISTRAL_7B, id="1x8-Mistral7B-v32768-b32-s2048-no-pen", marks=_slow), pytest.param((1, 8), 32768, 32, 2048, 1.2, 1.2, 1.5, 0.95, MISTRAL_7B, id="1x8-Mistral7B-v32768-b32-s2048-pen", marks=_slow), pytest.param((1, 8), 32768, 32, 4096, 0.0, 0.0, 1.0, 0.999, MISTRAL_7B, id="1x8-Mistral7B-v32768-b32-s4096-no-pen", marks=_slow), pytest.param((1, 8), 32768, 32, 4096, 1.2, 1.2, 1.5, 0.95, MISTRAL_7B, id="1x8-Mistral7B-v32768-b32-s4096-pen", marks=_slow), # --- (1,8) Llama1B v128256 --- pytest.param((1, 8), 128256, 1, 71, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b1-s71-no-pen"), pytest.param((1, 8), 128256, 1, 80, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b1-s80-no-pen", marks=_slow), pytest.param((1, 8), 128256, 1, 115, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b1-s115-no-pen", marks=_slow), pytest.param((1, 8), 128256, 1, 337, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b1-s337-no-pen", marks=_slow), pytest.param((1, 8), 128256, 1, 512, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b1-s512-no-pen", marks=_slow), pytest.param((1, 8), 128256, 1, 785, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b1-s785-no-pen", marks=_slow), pytest.param((1, 8), 128256, 1, 16229, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b1-s16229-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 52, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s52-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 56, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s56-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 57, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s57-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 59, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s59-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 60, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s60-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 61, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s61-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 63, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s63-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 65, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s65-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 66, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s66-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 67, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s67-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 70, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s70-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 71, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s71-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 72, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s72-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 75, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s75-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 77, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s77-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 80, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s80-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 81, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s81-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 84, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s84-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 86, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s86-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 87, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s87-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 91, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s91-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 93, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s93-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 94, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s94-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 96, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s96-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 98, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s98-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 99, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s99-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 101, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s101-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 102, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s102-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 103, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s103-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 104, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s104-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 105, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s105-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 106, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s106-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 109, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s109-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 113, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s113-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 114, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s114-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 115, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s115-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 116, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s116-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 117, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s117-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 119, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s119-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 124, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s124-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 125, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s125-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 128, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s128-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 128, 1.2, 1.2, 1.5, 0.95, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s128-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 337, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s337-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 512, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s512-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 712, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s712-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 785, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s785-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 1024, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s1024-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 1024, 1.2, 1.2, 1.5, 0.95, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s1024-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 2048, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s2048-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 2048, 1.2, 1.2, 1.5, 0.95, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s2048-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 4096, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s4096-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 4096, 1.2, 1.2, 1.5, 0.95, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s4096-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 8192, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s8192-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 8192, 1.2, 1.2, 1.5, 0.95, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s8192-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 16229, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s16229-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 16384, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s16384-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 16384, 1.2, 1.2, 1.5, 0.95, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s16384-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 32768, 0.0, 0.0, 1.0, 0.999, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s32768-no-pen", marks=_slow), pytest.param((1, 8), 128256, 32, 32768, 1.2, 1.2, 1.5, 0.95, LLAMA_1B, id="1x8-Llama1B-v128256-b32-s32768-pen", marks=_slow), # --- (1,8) Qwen3-32B v151936 --- pytest.param((1, 8), 151936, 32, 128, 1.2, 1.2, 1.5, 0.95, QWEN3_32B, id="1x8-Qwen3-32B-v151936-b32-s128-pen"), # --- (1,8) Qwen2.5-72B v152064 --- pytest.param((1, 8), 152064, 1, 65, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b1-s65-no-pen"), pytest.param((1, 8), 152064, 1, 80, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b1-s80-no-pen", marks=_slow), pytest.param((1, 8), 152064, 1, 109, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b1-s109-no-pen", marks=_slow), pytest.param((1, 8), 152064, 1, 421, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b1-s421-no-pen", marks=_slow), pytest.param((1, 8), 152064, 1, 512, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b1-s512-no-pen", marks=_slow), pytest.param((1, 8), 152064, 1, 780, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b1-s780-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 48, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s48-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 50, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s50-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 51, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s51-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 54, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s54-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 55, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s55-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 56, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s56-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 57, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s57-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 59, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s59-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 60, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s60-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 61, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s61-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 64, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s64-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 65, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s65-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 68, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s68-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 69, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s69-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 71, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s71-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 74, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s74-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 75, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s75-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 76, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s76-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 80, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s80-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 81, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s81-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 85, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s85-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 87, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s87-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 88, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s88-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 90, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s90-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 92, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s92-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 93, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s93-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 95, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s95-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 96, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s96-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 97, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s97-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 98, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s98-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 99, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s99-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 100, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s100-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 101, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s101-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 102, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s102-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 103, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s103-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 107, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s107-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 108, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s108-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 109, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s109-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 110, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s110-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 111, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s111-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 113, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s113-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 118, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s118-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 119, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s119-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 128, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s128-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 128, 1.2, 1.2, 1.5, 0.95, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s128-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 421, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s421-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 512, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s512-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 707, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s707-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 780, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s780-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 1024, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s1024-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 1024, 1.2, 1.2, 1.5, 0.95, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s1024-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 2048, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s2048-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 2048, 1.2, 1.2, 1.5, 0.95, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s2048-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 4096, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s4096-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 4096, 1.2, 1.2, 1.5, 0.95, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s4096-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 8192, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s8192-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 8192, 1.2, 1.2, 1.5, 0.95, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s8192-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 16255, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s16255-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 16384, 0.0, 0.0, 1.0, 0.999, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s16384-no-pen", marks=_slow), pytest.param((1, 8), 152064, 32, 16384, 1.2, 1.2, 1.5, 0.95, QWEN25_72B, id="1x8-Qwen2.5-72B-v152064-b32-s16384-pen", marks=_slow), ] # fmt: on # ============================================================================== # Reference implementation (pure torch) # ============================================================================== def reference_apply_penalties(logits, prompt_mask, output_mask, output_counts, presence, frequency, repetition): """Pure-torch reference for penalty math, following the OpenAI API spec. Algorithm source: vLLM's ``apply_penalties`` in ``vllm/model_executor/layers/utils.py`` (https://github.com/vllm-project/vllm/blob/main/vllm/model_executor/layers/utils.py). - Presence: subtract flat penalty for each token that appeared in output - Frequency: subtract penalty proportional to token occurrence count - Repetition: sign-dependent scaling for tokens in prompt OR output (positive logits divided by penalty, negative logits multiplied by penalty) """ logits = logits.clone().float() output_mask_f = output_mask.float() output_counts_f = output_counts.float() # Presence: logits -= output_mask * presence (vLLM: presence_penalties * output_mask) logits -= output_mask_f * presence # Frequency: logits -= output_counts * frequency (vLLM: frequency_penalties * output_bin_counts) logits -= output_counts_f * frequency # Repetition: sign-dependent scaling (vLLM: apply_repetition_penalties on combined prompt+output mask) combined = ((prompt_mask + output_mask) > 0).float() inv_rep = 1.0 / repetition # If logit > 0: multiply by 1/rep (shrink toward 0). If logit <= 0: multiply by rep (push away from 0). scale = torch.where( logits > 0, torch.where(combined.bool(), inv_rep, torch.ones_like(logits)), torch.where(combined.bool(), repetition, torch.ones_like(logits)), ) logits *= scale return logits # ============================================================================== # Unit tests: Config and dataclasses (no device) # ============================================================================== class TestConfigUnit: def test_config_defaults(self): cfg = Penalties1DConfig(vocab_size=1024) assert cfg.max_batch_size == 32 assert cfg.mesh_device is None assert cfg.sub_core_grids is None assert cfg.prompt_mask is None def test_config_not_resolved_without_mesh_device(self): cfg = Penalties1DConfig(vocab_size=1024) assert not cfg.is_resolved() def test_penalty_params_fields(self): fields = PenaltyParams.__dataclass_fields__ assert set(fields.keys()) == { "prompt_mask", "presence_penalties", "frequency_penalties", "repetition_penalties", "inverse_repetition_penalties", } def test_penalty_accumulator_fields(self): fields = PenaltyAccumulator.__dataclass_fields__ assert set(fields.keys()) == {"output_mask", "output_counts", "output_counts_gathered"} # ============================================================================== # Device tests: Config resolution and Penalties1D # ============================================================================== @pytest.mark.parametrize("ttnn_mesh_device", [(1, 1), (1, 2), (1, 8)], ids=["1x1", "1x2", "1x8"], indirect=True) class TestPenalties1DDevice: @pytest.mark.parametrize("vocab_size", [1024]) def test_resolve_config(self, ttnn_mesh_device, vocab_size): cfg = Penalties1DConfig(vocab_size=vocab_size, mesh_device=ttnn_mesh_device) resolved = _resolve_penalties1d_config(cfg) assert resolved.is_resolved() assert resolved.mesh_device is ttnn_mesh_device @pytest.mark.parametrize("vocab_size", [1024]) def test_load_device_buffers(self, ttnn_mesh_device, vocab_size): pen = Penalties1D(vocab_size=vocab_size, mesh_device=ttnn_mesh_device) pen.load_device_buffers() assert pen._device_buffers_loaded assert isinstance(pen._decode_src, ttnn.Tensor) assert isinstance(pen._zeros, ttnn.Tensor) @pytest.mark.parametrize("vocab_size", [1024]) def test_decode_forward_none_passthrough(self, ttnn_mesh_device, vocab_size): """forward() with None params/accum returns logits unchanged.""" pen = Penalties1D(vocab_size=vocab_size, mesh_device=ttnn_mesh_device) logits_host = torch.randn(32, vocab_size, dtype=torch.bfloat16) logits_tt = ttnn.from_torch(logits_host, device=ttnn_mesh_device, dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT) result = pen.forward(logits_tt, params=None, accum=None) assert result is logits_tt @pytest.mark.parametrize("vocab_size", [1024]) 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 pen = Penalties1D.from_model_args(ttnn_mesh_device, MockArgs()) assert pen.config.vocab_size == vocab_size assert pen.config.mesh_device is ttnn_mesh_device def test_rejects_galaxy(self, ttnn_mesh_device): """from_model_args should reject 2D (Galaxy) topologies.""" class FakeMesh: shape = (2, 4) class MockArgs: padded_vocab_size = 1024 sub_core_grids = None with pytest.raises(ValueError, match="1D mesh topologies"): # allow-pytest.raises: pre-existing Penalties1D.from_model_args(FakeMesh(), MockArgs()) # ============================================================================== # VS Reference tests — full penalty pipeline compared against pure-torch golden # ============================================================================== def _make_penalty_tensors_on_device( ttnn_mesh_device, B, *, prompt_mask_host, output_mask_host, output_counts_host, presence_val, frequency_val, repetition_val, ): """Helper: build PenaltyParams + PenaltyAccumulator on device from host tensors.""" params = PenaltyParams( prompt_mask=ttnn.from_torch( prompt_mask_host, device=ttnn_mesh_device, dtype=ttnn.int32, layout=ttnn.TILE_LAYOUT, memory_config=ttnn.DRAM_MEMORY_CONFIG, ), presence_penalties=ttnn.from_torch( torch.full((B, 1), presence_val), device=ttnn_mesh_device, dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT, ), frequency_penalties=ttnn.from_torch( torch.full((B, 1), frequency_val), device=ttnn_mesh_device, dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT, ), repetition_penalties=ttnn.from_torch( torch.full((B, 1), repetition_val), device=ttnn_mesh_device, dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT, ), inverse_repetition_penalties=ttnn.from_torch( torch.full((B, 1), 1.0 / repetition_val), device=ttnn_mesh_device, dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT, ), ) accum = PenaltyAccumulator( output_mask=ttnn.from_torch( output_mask_host, device=ttnn_mesh_device, dtype=ttnn.int32, layout=ttnn.TILE_LAYOUT, memory_config=ttnn.DRAM_MEMORY_CONFIG, ), output_counts=ttnn.from_torch( output_counts_host, device=ttnn_mesh_device, dtype=ttnn.int32, layout=ttnn.TILE_LAYOUT, memory_config=ttnn.DRAM_MEMORY_CONFIG, ), output_counts_gathered=ttnn.from_torch( output_counts_host, device=ttnn_mesh_device, dtype=ttnn.int32, layout=ttnn.TILE_LAYOUT, memory_config=ttnn.DRAM_MEMORY_CONFIG, ), ) return params, accum def _readback_logits(result_tt, ttnn_mesh_device, B, vocab_size): """Helper: read logits back from device to host torch tensor.""" result_host = ttnn.to_torch( result_tt, mesh_composer=ttnn.ConcatMesh2dToTensor( ttnn_mesh_device, dims=(0, 1), mesh_shape=tuple(ttnn_mesh_device.shape) ), ) return result_host[:B, :vocab_size].float() @pytest.mark.parametrize("ttnn_mesh_device", [(1, 1), (1, 2), (1, 8)], ids=["1x1", "1x2", "1x8"], indirect=True) @pytest.mark.parametrize( "mesh_shape,vocab_size,batch_size,seq_len,presence,frequency,repetition,pcc,hf_model_name", _list_collected_penalty_cases(), ) def test_penalties1d_vs_reference( ttnn_mesh_device, mesh_shape, vocab_size, batch_size, seq_len, presence, frequency, repetition, pcc, hf_model_name, ): """ Test Penalties1D.decode_forward matches the pure-torch reference_apply_penalties. Parametrized across penalty combinations to verify each penalty type independently and in combination. PCC thresholds account for bfloat16 precision. """ torch.manual_seed(42) B = 32 pen = Penalties1D(vocab_size=vocab_size, mesh_device=ttnn_mesh_device) # Build host tensors with realistic patterns logits_host = torch.randn(B, vocab_size, dtype=torch.bfloat16) # prompt_mask: first 50 tokens were in the prompt prompt_mask_host = torch.zeros(B, vocab_size, dtype=torch.int32) prompt_mask_host[:, :50] = 1 # output_mask: tokens 40-60 appeared in output (overlaps with prompt) output_mask_host = torch.zeros(B, vocab_size, dtype=torch.int32) output_mask_host[:, 40:60] = 1 # output_counts: tokens 40-60 appeared 1-3 times output_counts_host = torch.zeros(B, vocab_size, dtype=torch.int32) output_counts_host[:, 40:50] = 2 output_counts_host[:, 50:60] = 1 # --- Reference (pure torch) --- expected = reference_apply_penalties( logits_host, prompt_mask_host, output_mask_host, output_counts_host, presence, frequency, repetition, ) # --- TT device --- logits_tt = ttnn.from_torch( logits_host, device=ttnn_mesh_device, dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT, ) params, accum = _make_penalty_tensors_on_device( ttnn_mesh_device, B, prompt_mask_host=prompt_mask_host, output_mask_host=output_mask_host, output_counts_host=output_counts_host, presence_val=presence, frequency_val=frequency, repetition_val=repetition, ) result_tt = pen.decode_forward(logits_tt, params, accum) result_host = _readback_logits(result_tt, ttnn_mesh_device, B, vocab_size) passing, pcc_msg = comp_pcc(expected, result_host, pcc=pcc) assert passing, f"Penalties1D vs reference failed: {pcc_msg} (threshold={pcc})" @pytest.mark.parametrize("ttnn_mesh_device", [(1, 1), (1, 2), (1, 8)], ids=["1x1", "1x2", "1x8"], indirect=True) def test_penalties1d_changes_argmax(ttnn_mesh_device): """ Heavy repetition penalty should change which token has the highest logit. Setup: token 0 has the highest logit AND appears in prompt+output. With repetition=5.0, the penalty should push token 0's logit down far enough that a different token becomes the argmax. """ torch.manual_seed(123) B = 32 vocab_size = 1024 pen = Penalties1D(vocab_size=vocab_size, mesh_device=ttnn_mesh_device) # Token 0 is the clear winner in raw logits logits_host = torch.randn(B, vocab_size, dtype=torch.bfloat16) logits_host[:, 0] = 10.0 # make token 0 dominant # Token 0 appears in prompt and output prompt_mask_host = torch.zeros(B, vocab_size, dtype=torch.int32) prompt_mask_host[:, 0] = 1 output_mask_host = torch.zeros(B, vocab_size, dtype=torch.int32) output_mask_host[:, 0] = 1 output_counts_host = torch.zeros(B, vocab_size, dtype=torch.int32) output_counts_host[:, 0] = 3 # Original argmax should be token 0 assert logits_host[0].argmax().item() == 0 logits_tt = ttnn.from_torch( logits_host, device=ttnn_mesh_device, dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT, ) params, accum = _make_penalty_tensors_on_device( ttnn_mesh_device, B, prompt_mask_host=prompt_mask_host, output_mask_host=output_mask_host, output_counts_host=output_counts_host, presence_val=2.0, frequency_val=2.0, repetition_val=5.0, ) result_tt = pen.decode_forward(logits_tt, params, accum) result_host = _readback_logits(result_tt, ttnn_mesh_device, B, vocab_size) # After heavy penalties, token 0 should no longer be argmax new_argmax = result_host[0].argmax().item() assert new_argmax != 0, f"Expected penalty to change argmax from 0, but it's still {new_argmax}" # ============================================================================== # Helper: build topology-correct PenaltyParams + PenaltyAccumulator from config # ============================================================================== def _make_proper_params_accum(pen: Penalties1D): """Build PenaltyParams + PenaltyAccumulator from the module's resolved config. Uses the config's LazyBuffer mesh_mappers to guarantee the correct dtype/layout/ sharding for whatever device topology is active. This is required for methods like init_prompt_penalties and update_output_tokens that call _token_bin_counts_and_mask, which expects properly sharded output tensors. """ params = PenaltyParams( prompt_mask=_materialize(pen.config.prompt_mask), presence_penalties=_materialize(pen.config.presence_penalties), frequency_penalties=_materialize(pen.config.frequency_penalties), repetition_penalties=_materialize(pen.config.repetition_penalties), inverse_repetition_penalties=_materialize(pen.config.inverse_repetition_penalties), ) accum = PenaltyAccumulator( output_mask=_materialize(pen.config.output_mask), output_counts=_materialize(pen.config.output_counts), output_counts_gathered=_materialize(pen.config.output_counts_gathered), ) return params, accum # ============================================================================== # Additional unit tests (no device) # ============================================================================== class TestConfigUnitMore: def test_buf_resolved_with_lazy_buffer(self): """_buf_resolved calls buf.is_resolved() for a LazyBuffer (line 97).""" from models.common.modules.lazy_buffer import LazyBuffer lb = LazyBuffer(source=torch.zeros(1), device=None) assert not Penalties1DConfig._buf_resolved(lb) # is_resolved() → False (device=None) def test_buf_resolved_returns_false_for_none(self): """_buf_resolved returns False for None (baseline, complements line 95/97 tests).""" assert not Penalties1DConfig._buf_resolved(None) # ============================================================================== # release() / cleanup contract (no device) # ============================================================================== # Module-owned buffer fields, in the order Penalties1D.release() walks them. _OWNED_BUFFER_NAMES = ( "prompt_mask", "output_mask", "output_counts", "output_counts_gathered", "zeros", "decode_src", "presence_penalties", "frequency_penalties", "repetition_penalties", "inverse_repetition_penalties", ) class _NotingError(RuntimeError): """RuntimeError with an add_note() on every interpreter, so the note branch of the cleanup helpers is exercised on Python 3.10 (CI) as well as 3.11+.""" def __init__(self, message): super().__init__(message) self.notes = () def add_note(self, note): self.notes = self.notes + (note,) def _loaded_penalties_with_fake_buffers(): """A Penalties1D that looks loaded: ten owned LazyBuffers and two slice tensors hold fake device handles. Returns (penalties, buffers, handles-in-release-order).""" buffers = {} values = [] for name in _OWNED_BUFFER_NAMES: buffer = LazyBuffer(source=torch.zeros(1)) value = object() buffer._value = value buffers[name] = buffer values.append(value) penalties = object.__new__(Penalties1D) penalties.config = SimpleNamespace(**buffers) penalties._slice_start = object() penalties._slice_end = object() values.extend((penalties._slice_start, penalties._slice_end)) penalties._decode_src = buffers["decode_src"]._value penalties._zeros = buffers["zeros"]._value penalties._device_buffers_loaded = True return penalties, buffers, values def test_penalties_release_deallocates_owned_lazy_buffers_and_slice_tensors(monkeypatch): released = [] monkeypatch.setattr(penalties_module.ttnn, "deallocate", released.append) penalties, buffers, values = _loaded_penalties_with_fake_buffers() penalties.release() penalties.release() # idempotent: nothing left to deallocate assert released == values assert all(buffer._value is None for buffer in buffers.values()) assert penalties._slice_start is None and penalties._slice_end is None assert penalties._decode_src is None and penalties._zeros is None assert not penalties._device_buffers_loaded def test_penalties_release_is_best_effort_and_retries_only_failed_buffers(monkeypatch, expect_error): penalties, buffers, values = _loaded_penalties_with_fake_buffers() counts_handle = buffers["output_counts"]._value slice_handle = penalties._slice_start counts_error = _NotingError("output_counts deallocate failed once") slice_error = RuntimeError("slice_start deallocate failed once") failures = {counts_handle: counts_error, slice_handle: slice_error} attempts = [] def deallocate(value): attempts.append(value) if value in failures and attempts.count(value) == 1: raise failures[value] monkeypatch.setattr(penalties_module.ttnn, "deallocate", deallocate) with expect_error(RuntimeError, "output_counts deallocate failed once") as caught: penalties.release() # First failure is raised; later failures ride along on it. assert caught.value is counts_error assert caught.value.cleanup_failures == (slice_error,) assert counts_error.notes == ("cleanup also encountered 1 additional failure(s)",) # Every buffer was attempted once; only the failed ones keep their handles. assert attempts == values assert buffers["output_counts"]._value is not None assert all(buffer._value is None for name, buffer in buffers.items() if name != "output_counts") assert penalties._slice_start is not None assert penalties._slice_end is None assert penalties._decode_src is None and penalties._zeros is None assert not penalties._device_buffers_loaded penalties.release() # Retry touches only the two buffers that failed, and now clears them. assert attempts == values + [counts_handle, slice_handle] assert all(buffer._value is None for buffer in buffers.values()) assert penalties._slice_start is None def test_load_device_buffers_failure_releases_partial_state_and_attaches_cleanup_failures(monkeypatch, expect_error): decode_src = LazyBuffer(source=torch.zeros(1)) decode_src._value = object() zeros = LazyBuffer(source=torch.zeros(1)) zeros._value = object() penalties = object.__new__(Penalties1D) penalties.config = SimpleNamespace( decode_src=decode_src, zeros=zeros, mesh_device=SimpleNamespace(shape=(1, 8)), sub_core_grids=None, is_resolved=lambda: True, ) penalties._device_buffers_loaded = False allocation_error = _NotingError("slice tensors failed") cleanup_error = RuntimeError("zeros deallocate failed once") deallocated = [] def deallocate(value): deallocated.append(value) if value is zeros._value and deallocated.count(value) == 1: raise cleanup_error def build_slice_tensors_and_fail(): raise allocation_error monkeypatch.setattr(penalties_module.ttnn, "ShardTensor2dMesh", lambda *args, **kwargs: object()) monkeypatch.setattr(penalties_module.ttnn, "deallocate", deallocate) monkeypatch.setattr(penalties, "_build_slice_tensors", build_slice_tensors_and_fail) with expect_error(RuntimeError, "slice tensors failed") as caught: penalties.load_device_buffers() # The allocation failure is what propagates; the cleanup failure is attached, not raised. assert caught.value is allocation_error assert caught.value.cleanup_failures == (cleanup_error,) assert allocation_error.notes == ("cleanup also encountered 1 failure(s)",) assert not penalties._device_buffers_loaded assert decode_src._value is None # released during cleanup assert zeros._value is not None # deallocate failed once, handle retained for retry penalties.release() assert zeros._value is None assert deallocated.count(decode_src._value) == 0 and len(deallocated) == 3 # A later load succeeds and repopulates the module-owned state. slices = (object(), object()) monkeypatch.setattr(penalties, "_build_slice_tensors", lambda: slices) monkeypatch.setattr(penalties_module, "_materialize", lambda buf: object()) penalties.load_device_buffers() assert penalties._device_buffers_loaded assert (penalties._slice_start, penalties._slice_end) == slices assert penalties._cluster_shape == (1, 8) and penalties._num_devices == 8 def test_cleanup_failure_helpers_tolerate_exceptions_without_add_note(expect_error): """Plain exceptions (no add_note on Python 3.10) still carry cleanup_failures.""" primary = RuntimeError("primary") penalties_module._attach_cleanup_failures(primary, ()) assert not hasattr(primary, "cleanup_failures") penalties_module._attach_cleanup_failures(primary, (ValueError("first"),)) penalties_module._attach_cleanup_failures(primary, (ValueError("second"),)) assert [str(error) for error in primary.cleanup_failures] == ["first", "second"] lone = RuntimeError("lone failure") with expect_error(RuntimeError, "lone failure") as caught: penalties_module._raise_cleanup_failures([lone]) assert caught.value is lone assert not hasattr(lone, "cleanup_failures") head, tail = RuntimeError("head failure"), ValueError("tail failure") with expect_error(RuntimeError, "head failure") as caught: penalties_module._raise_cleanup_failures([head, tail]) assert caught.value is head assert head.cleanup_failures == (tail,) def test_build_slice_tensors_releases_first_slice_when_second_allocation_fails(monkeypatch, expect_error): penalties = object.__new__(Penalties1D) penalties.config = SimpleNamespace(vocab_size=1024, max_batch_size=32, mesh_device=object()) penalties._cluster_shape = (1, 8) penalties._num_devices = 8 first_slice = object() allocation_error = _NotingError("slice_end allocation failed") cleanup_error = RuntimeError("slice_start deallocate failed") host_tensors = [] deallocated = [] def from_torch(tensor, **_kwargs): host_tensors.append(tensor) if len(host_tensors) == 1: return first_slice raise allocation_error def deallocate(value): deallocated.append(value) raise cleanup_error monkeypatch.setattr(penalties_module.ttnn, "ShardTensor2dMesh", lambda *args, **kwargs: object()) monkeypatch.setattr(penalties_module.ttnn, "from_torch", from_torch) monkeypatch.setattr(penalties_module.ttnn, "deallocate", deallocate) with expect_error(RuntimeError, "slice_end allocation failed") as caught: penalties._build_slice_tensors() assert caught.value is allocation_error assert deallocated == [first_slice] # the half-built pair is torn down before re-raising assert caught.value.cleanup_failures == (cleanup_error,) assert allocation_error.notes == ("cleanup also encountered 1 failure(s)",) # Per-device vocab slice bounds for 1024 / 8 devices, padded batch 32. assert host_tensors[0].tolist() == [v for d in range(8) for v in (0, 128 * d)] assert host_tensors[1].tolist() == [v for d in range(8) for v in (32, 128 * (d + 1))] # ============================================================================== # Additional device tests: coverage for previously untested methods # ============================================================================== @pytest.mark.parametrize( ("batch_height", "expected_operation"), [(1, "to_layout"), (32, "tilize")], ) def test_histogram_tilize_preserves_batch32_path_and_pads_smaller_batches( monkeypatch, batch_height, expected_operation, ): """Only non-tile batch heights use the padding-aware layout conversion.""" calls = [] counts = SimpleNamespace(padded_shape=(batch_height, 1024)) result = object() pen = object.__new__(Penalties1D) pen._op_kwargs = {"sub_core_grids": "sub-grid"} pen._use_low_perf_tilize = True monkeypatch.setattr( ttnn, "tilize", lambda tensor, **kwargs: calls.append(("tilize", tensor, kwargs)) or result, ) monkeypatch.setattr( ttnn, "to_layout", lambda tensor, layout, **kwargs: calls.append(("to_layout", tensor, layout, kwargs)) or result, ) assert pen._tilize_counts(counts) is result if expected_operation == "tilize": assert calls == [ ( "tilize", counts, {"sub_core_grids": "sub-grid", "use_low_perf": True}, ) ] else: assert calls == [ ( "to_layout", counts, ttnn.TILE_LAYOUT, {"sub_core_grids": "sub-grid"}, ) ] def test_non_tile_histogram_slices_padded_views_and_preserves_logical_output(monkeypatch): """Batch-1 vocab slicing uses padded aliases of caller-owned state.""" events = [] counts_new_rm = SimpleNamespace(name="counts-new-rm") counts_new_tiled = SimpleNamespace(name="counts-new-tiled") counts = SimpleNamespace(name="counts", shape=(1, 1024), padded_shape=(32, 1024)) counts_padded = SimpleNamespace(name="counts-padded") counts_sliced = SimpleNamespace(name="counts-sliced", shape=(1, 128), padded_shape=(32, 128)) counts_sliced_padded = SimpleNamespace(name="counts-sliced-padded") mask = object() new_tokens = SimpleNamespace(deallocate=lambda: events.append("deallocate-tokens")) pen = object.__new__(Penalties1D) pen.config = SimpleNamespace(max_batch_size=1) pen._zeros = object() pen._op_kwargs = {} pen._slice_start = object() pen._slice_end = object() pen._num_devices = 8 pen._tilize_counts = lambda tensor: events.append(("tilize", tensor)) or counts_new_tiled monkeypatch.setattr( ttnn, "scatter_add", lambda *args, **kwargs: events.append(("scatter", args, kwargs)) or counts_new_rm, ) monkeypatch.setattr( ttnn, "add", lambda lhs, rhs, *, output_tensor, **kwargs: events.append(("add", lhs, rhs, output_tensor, kwargs)) or output_tensor, ) def reshape(tensor, logical_shape, padded_shape, *, skip_padding_fill): events.append(("reshape", tensor, logical_shape, padded_shape, skip_padding_fill)) if tensor is counts: return counts_padded assert tensor is counts_sliced return counts_sliced_padded monkeypatch.setattr(ttnn, "reshape", reshape) def slice_tensor(tensor, start, end, *, output_tensor, slice_dim, num_devices, **kwargs): events.append(("slice", tensor, start, end, output_tensor, slice_dim, num_devices, kwargs)) assert tensor is counts_padded assert end is pen._slice_end assert output_tensor is counts_sliced_padded return output_tensor monkeypatch.setattr(ttnn, "slice", slice_tensor) monkeypatch.setattr( ttnn, "gt", lambda tensor, threshold, *, output_tensor, **kwargs: events.append( ("gt", tensor, threshold, output_tensor, kwargs) ) or output_tensor, ) returned_counts, returned_mask = pen._token_bin_counts_and_mask( new_tokens, object(), counts=counts, mask=mask, counts_sliced=counts_sliced, ) assert returned_counts is counts assert returned_mask is mask assert events[-1] == ("gt", counts_sliced, 0, mask, {}) slices = [event for event in events if isinstance(event, tuple) and event[0] == "slice"] assert slices == [ ( "slice", counts_padded, pen._slice_start, pen._slice_end, counts_sliced_padded, 1, 8, {}, ) ] @pytest.mark.parametrize("ttnn_mesh_device", [(1, 1), (1, 2), (1, 8)], ids=["1x1", "1x2", "1x8"], indirect=True) class TestPenalties1DDeviceExtra: """Coverage for methods not exercised by the reference tests.""" # ------------------------------------------------------------------ # _buf_resolved ttnn.Tensor path (line 95) # ------------------------------------------------------------------ def test_buf_resolved_with_tt_tensor(self, ttnn_mesh_device): """_buf_resolved returns True for a real ttnn.Tensor (line 95).""" tt = ttnn.from_torch( torch.zeros(1, 1, dtype=torch.int32), device=ttnn_mesh_device, dtype=ttnn.int32, layout=ttnn.TILE_LAYOUT, ) assert Penalties1DConfig._buf_resolved(tt) # ------------------------------------------------------------------ # from_config (lines 141-145) # ------------------------------------------------------------------ @pytest.mark.parametrize("vocab_size", [1024]) def test_from_config(self, ttnn_mesh_device, vocab_size): """from_config power-path classmethod (lines 141-145).""" cfg = Penalties1DConfig(vocab_size=vocab_size, mesh_device=ttnn_mesh_device) pen = Penalties1D.from_config(cfg) assert pen.config.vocab_size == vocab_size assert pen.config.mesh_device is ttnn_mesh_device assert not pen._device_buffers_loaded # ------------------------------------------------------------------ # load_device_buffers idempotent guard (line 166) # ------------------------------------------------------------------ @pytest.mark.parametrize("vocab_size", [1024]) def test_load_device_buffers_idempotent(self, ttnn_mesh_device, vocab_size): """Second call to load_device_buffers returns early without re-allocating (line 166).""" pen = Penalties1D(vocab_size=vocab_size, mesh_device=ttnn_mesh_device) pen.load_device_buffers() decode_src_first = pen._decode_src pen.load_device_buffers() # hits early-return at line 166 assert pen._decode_src is decode_src_first # ------------------------------------------------------------------ # init_prompt_penalties + _token_bin_counts_and_mask counts=None path # ------------------------------------------------------------------ @pytest.mark.parametrize("max_batch_size", [1, 32]) @pytest.mark.parametrize("vocab_size", [1024]) def test_init_prompt_penalties(self, ttnn_mesh_device, vocab_size, max_batch_size): """init_prompt_penalties scatters prompt tokens into prompt_mask.""" pen = Penalties1D( vocab_size=vocab_size, mesh_device=ttnn_mesh_device, max_batch_size=max_batch_size, ) pen.load_device_buffers() params, accum = _make_proper_params_accum(pen) prompt_tokens = torch.randint(0, vocab_size, (max_batch_size, 10)) pen.init_prompt_penalties(params, accum, prompt_tokens) # ------------------------------------------------------------------ # forward() prompt init dispatch # ------------------------------------------------------------------ @pytest.mark.parametrize("vocab_size", [1024]) def test_forward_dispatches_to_init_prompt(self, ttnn_mesh_device, vocab_size): """forward() with prompt_tokens kwarg routes to init_prompt_penalties.""" B = 32 pen = Penalties1D(vocab_size=vocab_size, mesh_device=ttnn_mesh_device) pen.load_device_buffers() params, accum = _make_proper_params_accum(pen) logits_tt = ttnn.from_torch( torch.randn(B, vocab_size, dtype=torch.bfloat16), device=ttnn_mesh_device, dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT, ) prompt_tokens = torch.randint(0, vocab_size, (B, 5)) result = pen.forward(logits_tt, params=params, accum=accum, prompt_tokens=prompt_tokens) assert result is logits_tt @pytest.mark.parametrize("vocab_size", [1024]) def test_forward_dispatches_to_decode(self, ttnn_mesh_device, vocab_size): """forward() without prompt_tokens routes to decode_forward (line 280). Uses unsharded (replicated) tensors — same topology as test_penalties1d_vs_reference — so the broadcast between logits [B, V] and penalty masks [B, V] is valid on all mesh shapes. The goal here is line 280 coverage, not correctness (covered elsewhere). """ B = 32 pen = Penalties1D(vocab_size=vocab_size, mesh_device=ttnn_mesh_device) zeros_BV = torch.zeros(B, vocab_size, dtype=torch.int32) params, accum = _make_penalty_tensors_on_device( ttnn_mesh_device, B, prompt_mask_host=zeros_BV, output_mask_host=zeros_BV, output_counts_host=zeros_BV, presence_val=0.0, frequency_val=0.0, repetition_val=1.0, ) logits_tt = ttnn.from_torch( torch.randn(B, vocab_size, dtype=torch.bfloat16), device=ttnn_mesh_device, dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT, ) result = pen.forward(logits_tt, params=params, accum=accum) assert result is not None # ------------------------------------------------------------------ # update_output_tokens: standard decode path (lines 290-294) # and _token_bin_counts_and_mask counts-not-None path (line 419) # ------------------------------------------------------------------ @pytest.mark.parametrize("max_batch_size", [1, 32]) @pytest.mark.parametrize("vocab_size", [1024]) def test_update_output_tokens_standard(self, ttnn_mesh_device, vocab_size, max_batch_size): """update_output_tokens with standard decode-shape [1,1,1,B] (lines 290-294).""" pen = Penalties1D( vocab_size=vocab_size, mesh_device=ttnn_mesh_device, max_batch_size=max_batch_size, ) pen.load_device_buffers() _, accum = _make_proper_params_accum(pen) # Standard sampling output: shape[-1]=B, shape[-2]=1 → if-branch. tokens_tt = ttnn.from_torch( torch.randint(0, vocab_size, (1, 1, 1, max_batch_size), dtype=torch.int32), device=ttnn_mesh_device, dtype=ttnn.int32, layout=ttnn.ROW_MAJOR_LAYOUT, ) pen.update_output_tokens(accum, tokens_tt) # ------------------------------------------------------------------ # update_output_tokens: multi-token else branch (lines 296-303) # ------------------------------------------------------------------ @pytest.mark.parametrize("vocab_size", [1024]) def test_update_output_tokens_multi_token(self, ttnn_mesh_device, vocab_size): """update_output_tokens with multi-token [B,S] shape triggers else-branch (lines 296-303).""" B = 32 pen = Penalties1D(vocab_size=vocab_size, mesh_device=ttnn_mesh_device) pen.load_device_buffers() _, accum = _make_proper_params_accum(pen) # shape[-1]=4 != B=32 → else-branch; src = ones(B, 4) created inline tokens_tt = ttnn.from_torch( torch.randint(0, vocab_size, (B, 4), dtype=torch.int32), device=ttnn_mesh_device, dtype=ttnn.int32, layout=ttnn.ROW_MAJOR_LAYOUT, ) pen.update_output_tokens(accum, tokens_tt) # ------------------------------------------------------------------ # reset_output_tokens: tokens=None (lines 317-324) and with tokens (lines 326-346) # ------------------------------------------------------------------ @pytest.mark.parametrize("vocab_size", [1024]) def test_reset_output_tokens_no_tokens(self, ttnn_mesh_device, vocab_size): """reset_output_tokens(tokens=None) zeros the accum buffers (lines 317-324).""" pen = Penalties1D(vocab_size=vocab_size, mesh_device=ttnn_mesh_device) pen.load_device_buffers() _, accum = _make_proper_params_accum(pen) pen.reset_output_tokens(accum, tokens=None) @pytest.mark.parametrize("vocab_size", [1024]) def test_reset_output_tokens_with_tokens(self, ttnn_mesh_device, vocab_size): """reset_output_tokens(tokens=...) zeros then re-initializes from tokens (lines 326-346).""" B = 32 pen = Penalties1D(vocab_size=vocab_size, mesh_device=ttnn_mesh_device) pen.load_device_buffers() _, accum = _make_proper_params_accum(pen) tokens = torch.randint(0, vocab_size, (B, 5)) pen.reset_output_tokens(accum, tokens=tokens) # ------------------------------------------------------------------ # _pad_batch_to_max: pad (lines 399-401), truncate (402-403), ValueError (396-397) # ------------------------------------------------------------------ def test_pad_batch_to_max_pads_small_batch(self, ttnn_mesh_device): """_pad_batch_to_max pads when B < max_batch_size (lines 399-401).""" pen = Penalties1D(vocab_size=1024, mesh_device=ttnn_mesh_device) pen.load_device_buffers() small = torch.randint(0, 100, (4, 10)) padded = pen._pad_batch_to_max(small, pad_value=-1) assert padded.shape[0] == pen.config.max_batch_size assert (padded[4:] == -1).all() def test_pad_batch_to_max_truncates_large_batch(self, ttnn_mesh_device): """_pad_batch_to_max truncates when B > max_batch_size (lines 402-403).""" pen = Penalties1D(vocab_size=1024, mesh_device=ttnn_mesh_device) pen.load_device_buffers() large = torch.randint(0, 100, (64, 10)) truncated = pen._pad_batch_to_max(large, pad_value=-1) assert truncated.shape[0] == pen.config.max_batch_size def test_pad_batch_to_max_raises_on_non_2d(self, ttnn_mesh_device): """_pad_batch_to_max raises ValueError for non-2D input (lines 396-397).""" pen = Penalties1D(vocab_size=1024, mesh_device=ttnn_mesh_device) pen.load_device_buffers() with pytest.raises(ValueError, match="Expected 2D"): # allow-pytest.raises: pre-existing pen._pad_batch_to_max(torch.zeros(10), pad_value=-1) # ------------------------------------------------------------------ # _resolve_buf: ttnn.Tensor passthrough (lines 474-475) # and LazyBuffer resolve (line 476) # ------------------------------------------------------------------ @pytest.mark.parametrize("vocab_size", [1024]) def test_resolve_buf_tensor_passthrough(self, ttnn_mesh_device, vocab_size): """Pre-existing ttnn.Tensor passes through _resolve_buf unchanged (lines 474-475).""" B = 32 pre_tensor = ttnn.from_torch( torch.zeros(B, vocab_size, dtype=torch.int32), device=ttnn_mesh_device, dtype=ttnn.int32, layout=ttnn.TILE_LAYOUT, ) cfg = Penalties1DConfig(vocab_size=vocab_size, mesh_device=ttnn_mesh_device, prompt_mask=pre_tensor) resolved = _resolve_penalties1d_config(cfg) assert resolved.prompt_mask is pre_tensor @pytest.mark.parametrize("vocab_size", [1024]) def test_resolve_buf_lazy_buffer_passthrough(self, ttnn_mesh_device, vocab_size): """Pre-existing LazyBuffer with device=None gets device filled in (line 476).""" from models.common.modules.lazy_buffer import LazyBuffer B = 32 partial_lb = LazyBuffer( source=torch.zeros(B, vocab_size, dtype=torch.int32), dtype=ttnn.int32, layout=ttnn.TILE_LAYOUT, device=None, memory_config=ttnn.DRAM_MEMORY_CONFIG, ) cfg = Penalties1DConfig(vocab_size=vocab_size, mesh_device=ttnn_mesh_device, prompt_mask=partial_lb) resolved = _resolve_penalties1d_config(cfg) assert isinstance(resolved.prompt_mask, LazyBuffer) assert resolved.prompt_mask.device is ttnn_mesh_device # ------------------------------------------------------------------ # _materialize: ttnn.Tensor passthrough (line 552) # ------------------------------------------------------------------ @pytest.mark.parametrize("vocab_size", [1024]) def test_materialize_tensor_passthrough(self, ttnn_mesh_device, vocab_size): """Pre-existing ttnn.Tensor as decode_src passes through _materialize (line 552).""" B = 32 pre_src = ttnn.from_torch( torch.ones(B, 1, dtype=torch.int32), device=ttnn_mesh_device, dtype=ttnn.int32, layout=ttnn.ROW_MAJOR_LAYOUT, ) cfg = Penalties1DConfig(vocab_size=vocab_size, mesh_device=ttnn_mesh_device, decode_src=pre_src) pen = Penalties1D.from_config(cfg) pen.load_device_buffers() assert pen._decode_src is pre_src