From 98310f4bf619423f899c1fa3bd2b78472ec39343 Mon Sep 17 00:00:00 2001 From: root <597191244@qq.com> Date: Wed, 2 Sep 2026 11:22:19 +0000 Subject: [PATCH 01/15] [perf] run GenRec with packed FlashAttention --- tzrec/models/genrec_causal_lm_model.py | 74 ++++---- tzrec/models/genrec_causal_lm_model_test.py | 176 +++++++++++++----- tzrec/models/genrec_model.py | 7 +- tzrec/models/genrec_model_test.py | 84 ++++++++- tzrec/modules/dynamic_beam_test.py | 58 +++++- tzrec/prompt/assembler.py | 5 + tzrec/prompt/assembler_test.py | 12 ++ .../genrec_causal_lm_model_mock.config | 1 + tzrec/tests/prompt_integration_test.py | 92 +++++++-- 9 files changed, 418 insertions(+), 91 deletions(-) diff --git a/tzrec/models/genrec_causal_lm_model.py b/tzrec/models/genrec_causal_lm_model.py index c1b6bf2f..b55cbb15 100644 --- a/tzrec/models/genrec_causal_lm_model.py +++ b/tzrec/models/genrec_causal_lm_model.py @@ -20,7 +20,7 @@ not supported; Qwen2.5/Qwen3 are what CI covers. """ -from typing import Any, Dict, List, Optional, Tuple +from typing import Any, Dict, List, Optional, Tuple, cast import torch @@ -143,17 +143,44 @@ def _forward( Logits and labels over the same window, so the shift ``loss`` applies lands on the pairs the window was sized for. """ - padded, mask, labels = self._left_pad_packed_inputs( - embeds, batch, build_labels=True + infos = batch.additional_infos + cu_seqlens = infos[PROMPT_CU_SEQLENS] + starts = cu_seqlens[:-1] + lengths = torch.diff(cu_seqlens) + row_starts = torch.repeat_interleave( + starts, lengths, output_size=embeds.shape[0] ) - suffix = self._prompt.prompt_plan.logits_suffix_len + position_ids = ( + torch.arange(embeds.shape[0], device=embeds.device) - row_starts + ).unsqueeze(0) + + suffix = cast(int, self._prompt.prompt_plan.logits_suffix_len) + suffix_offsets = torch.arange(-suffix, 0, device=embeds.device) + keep_indices = (cu_seqlens[1:, None] + suffix_offsets).reshape(-1) + flash_cu_seqlens = cu_seqlens.to(dtype=torch.int32).contiguous() + max_seqlen = int(infos[PROMPT_MAX_SEQLEN]) outputs = self.lm( - inputs_embeds=padded, - attention_mask=mask, + inputs_embeds=embeds.unsqueeze(0), + attention_mask=None, + position_ids=position_ids, use_cache=False, - logits_to_keep=suffix, + logits_to_keep=keep_indices, + cu_seq_lens_q=flash_cu_seqlens, + cu_seq_lens_k=flash_cu_seqlens, + max_length_q=max_seqlen, + max_length_k=max_seqlen, + ) + logits = outputs.logits.reshape(lengths.numel(), suffix, -1) + + window_ids = infos[PROMPT_INPUT_IDS][keep_indices].reshape( + lengths.numel(), suffix ) - return outputs.logits, labels[:, -suffix:] + columns = torch.arange(suffix, device=embeds.device) + labels = window_ids.masked_fill( + columns[None, :] < suffix - infos[PROMPT_RESPONSE_LENGTHS][:, None], + self._ignore_index, + ) + return logits, labels def _generate(self, embeds: torch.Tensor, batch: Batch) -> torch.Tensor: """Beam-search the SID answer. @@ -165,7 +192,7 @@ def _generate(self, embeds: torch.Tensor, batch: Batch) -> torch.Tensor: Returns: ``(B, num_return, num_levels)`` local codes, best first. """ - padded, mask, _ = self._left_pad_packed_inputs(embeds, batch) + padded, mask = self._left_pad_packed_inputs(embeds, batch) tokens = dynamic_beam_search( self.lm, padded, mask, self._capped_widths, self._bands ) @@ -176,17 +203,15 @@ def _left_pad_packed_inputs( self, embeds: torch.Tensor, batch: Batch, - build_labels: bool = False, - ) -> Tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]]: - """Left-pad packed prompt embeddings for the causal LM. + ) -> Tuple[torch.Tensor, torch.Tensor]: + """Left-pad packed prompt embeddings for beam decode. Args: embeds: packed embeddings, ``(total_tokens, hidden)``. batch: carries the packed prompt metadata. - build_labels: whether to build response-only training labels. Returns: - Padded embeddings, attention mask and optional labels. + Padded embeddings and attention mask. """ infos = batch.additional_infos cu_seqlens = infos[PROMPT_CU_SEQLENS] @@ -202,32 +227,17 @@ def _left_pad_packed_inputs( padded = embeds.new_zeros((batch_size, max_seqlen, hidden)) # mask selects row-major, which is how embeds and input_ids are packed padded[mask] = embeds - if not build_labels: - return padded, mask.long(), None - - input_ids = infos[PROMPT_INPUT_IDS] - response_lengths = infos[PROMPT_RESPONSE_LENGTHS] - labels = torch.full( - (batch_size, max_seqlen), - self._ignore_index, - dtype=input_ids.dtype, - device=embeds.device, - ) - labels[mask] = input_ids - response_mask = columns[None, :] >= (max_seqlen - response_lengths)[:, None] - labels[~response_mask] = self._ignore_index - return padded, mask.long(), labels + return padded, mask.long() @torch.fx.wrap def _fx_wrapped_forward( model: "GenrecCausalLMModel", embeds: torch.Tensor, batch: Batch ) -> Tuple[torch.Tensor, torch.Tensor]: - """Hide the padded forward from FX. + """Hide the packed forward from FX. ``TrainPipelineSparseDist`` symbolically traces the model whenever a - sharded module exists, and ``_left_pad_packed_inputs`` reads - ``max_seqlen`` as a host int. + sharded module exists, and the LM call reads ``max_seqlen`` as a host int. Args: model: the model whose response window to compute. diff --git a/tzrec/models/genrec_causal_lm_model_test.py b/tzrec/models/genrec_causal_lm_model_test.py index e269fc19..422a49a5 100644 --- a/tzrec/models/genrec_causal_lm_model_test.py +++ b/tzrec/models/genrec_causal_lm_model_test.py @@ -10,10 +10,12 @@ # limitations under the License. import unittest -from unittest import mock +from types import SimpleNamespace import torch from parameterized import parameterized +from torch import nn +from transformers.loss.loss_utils import ForCausalLMLoss from tzrec.datasets.utils import Batch from tzrec.models.genrec_causal_lm_model import GenrecCausalLMModel @@ -23,11 +25,8 @@ PROMPT_MAX_SEQLEN, PROMPT_RESPONSE_LENGTHS, ) -from tzrec.tests.prompt_test_util import ( - _CODEBOOK, - GenrecModelTestBase, - offset_sid_codes, -) +from tzrec.protos.models.genrec_model_pb2 import GenrecModelConfig +from tzrec.tests.prompt_test_util import GenrecModelTestBase from tzrec.utils.test_util import parameterized_name_func @@ -37,24 +36,16 @@ class LeftPadPackedInputsTest(unittest.TestCase): def test_packs_rows_of_different_lengths(self) -> None: embeds = torch.arange(18, dtype=torch.float32).reshape(9, 2) cu = torch.tensor([0, 4, 9]) - input_ids = torch.tensor([1, 2, 7, 8, 3, 4, 5, 7, 8]) - response_lengths = torch.tensor([1, 2]) - ignore = -7 batch = Batch( additional_infos={ PROMPT_CU_SEQLENS: cu, - PROMPT_INPUT_IDS: input_ids, PROMPT_MAX_SEQLEN: torch.tensor(7), - PROMPT_RESPONSE_LENGTHS: response_lengths, } ) model = GenrecCausalLMModel.__new__(GenrecCausalLMModel) torch.nn.Module.__init__(model) - model._ignore_index = ignore - padded, mask, out = model._left_pad_packed_inputs( - embeds, batch, build_labels=True - ) + padded, mask = model._left_pad_packed_inputs(embeds, batch) self.assertEqual(padded.shape, (2, 7, 2)) self.assertEqual( @@ -66,15 +57,136 @@ def test_packs_rows_of_different_lengths(self) -> None: torch.testing.assert_close(padded[0, :3], torch.zeros(3, 2)) torch.testing.assert_close(padded[1, :2], torch.zeros(2, 2)) torch.testing.assert_close(padded[:, -1], torch.stack([embeds[3], embeds[8]])) + + +class _CapturingLM(nn.Module): + def __init__(self) -> None: + super().__init__() + self.kwargs = {} + + def forward(self, **kwargs): + self.kwargs = kwargs + count = kwargs["logits_to_keep"].numel() + logits = torch.arange( + count * 5, + dtype=kwargs["inputs_embeds"].dtype, + device=kwargs["inputs_embeds"].device, + ).reshape(1, count, 5) + return SimpleNamespace(logits=logits) + + +class _DifferentiableLM(nn.Module): + def __init__(self, hidden_size: int, vocab_size: int) -> None: + super().__init__() + self.config = SimpleNamespace(vocab_size=vocab_size) + self.lm_head = nn.Linear(hidden_size, vocab_size, bias=False) + self.loss_function = ForCausalLMLoss + + def forward(self, **kwargs): + hidden = kwargs["inputs_embeds"].index_select(1, kwargs["logits_to_keep"]) + return SimpleNamespace(logits=self.lm_head(hidden)) + + +class PackedForwardTest(unittest.TestCase): + def test_passes_varlen_metadata_and_builds_per_row_labels(self) -> None: + model = GenrecCausalLMModel.__new__(GenrecCausalLMModel) + nn.Module.__init__(model) + model.lm = _CapturingLM() + model._ignore_index = -7 + model._prompt = SimpleNamespace( + prompt_plan=SimpleNamespace(logits_suffix_len=4) + ) + embeds = torch.arange(72, dtype=torch.float32).reshape(12, 6) + input_ids = torch.arange(100, 112) + batch = Batch( + additional_infos={ + PROMPT_CU_SEQLENS: torch.tensor([0, 5, 12]), + PROMPT_INPUT_IDS: input_ids, + PROMPT_MAX_SEQLEN: torch.tensor(7), + PROMPT_RESPONSE_LENGTHS: torch.tensor([3, 2]), + } + ) + + logits, labels = model._forward(embeds, batch) + + kwargs = model.lm.kwargs + self.assertEqual(kwargs["inputs_embeds"].shape, (1, 12, 6)) + self.assertIsNone(kwargs["attention_mask"]) self.assertEqual( - out.tolist(), - [[ignore] * 6 + [8], [ignore] * 5 + [7, 8]], + kwargs["position_ids"].tolist(), + [[0, 1, 2, 3, 4, 0, 1, 2, 3, 4, 5, 6]], ) + self.assertEqual(kwargs["logits_to_keep"].tolist(), [1, 2, 3, 4, 8, 9, 10, 11]) + self.assertEqual(kwargs["cu_seq_lens_q"].dtype, torch.int32) + torch.testing.assert_close(kwargs["cu_seq_lens_q"], kwargs["cu_seq_lens_k"]) + self.assertEqual(kwargs["max_length_q"], 7) + self.assertEqual(kwargs["max_length_k"], 7) + self.assertIs(kwargs["use_cache"], False) + self.assertEqual(logits.shape, (2, 4, 5)) + self.assertEqual( + labels.tolist(), + [[-7, 102, 103, 104], [-7, -7, 110, 111]], + ) + + def test_loss_and_gradients_cover_only_valid_response_pairs(self) -> None: + model = GenrecCausalLMModel.__new__(GenrecCausalLMModel) + nn.Module.__init__(model) + model.lm = _DifferentiableLM(hidden_size=6, vocab_size=32) + model._ignore_index = -7 + model._prompt = SimpleNamespace( + prompt_plan=SimpleNamespace(logits_suffix_len=4) + ) + embeds = torch.randn( + 12, 6, generator=torch.Generator().manual_seed(1), requires_grad=True + ) + input_ids = torch.arange(4, 16) + batch = Batch( + additional_infos={ + PROMPT_CU_SEQLENS: torch.tensor([0, 5, 12]), + PROMPT_INPUT_IDS: input_ids, + PROMPT_MAX_SEQLEN: torch.tensor(7), + PROMPT_RESPONSE_LENGTHS: torch.tensor([3, 2]), + } + ) + + logits, labels = model._forward(embeds, batch) + loss = model.loss({"logits": logits, "labels": labels}, batch)["ce_loss"] + expected = nn.functional.cross_entropy( + torch.cat((logits[0, :3], logits[1, 1:3])), + torch.tensor([6, 7, 8, 14, 15]), + ) + torch.testing.assert_close(loss, expected) + + loss.backward() + self.assertIsNotNone(embeds.grad) + grad_norms = embeds.grad.abs().sum(dim=1) + supervised = torch.tensor([1, 2, 3, 9, 10]) + self.assertTrue(bool(torch.all(grad_norms[supervised] > 0))) + unsupervised = torch.ones(12, dtype=torch.bool) + unsupervised[supervised] = False + torch.testing.assert_close( + grad_norms[unsupervised], torch.zeros(7), atol=0, rtol=0 + ) + weight_grad = model.lm.lm_head.weight.grad + self.assertIsNotNone(weight_grad) + self.assertTrue(bool(torch.isfinite(weight_grad).all())) + self.assertGreater(float(weight_grad.abs().sum()), 0) class GenrecCausalLMModelTest(GenrecModelTestBase): """The decode schedule and the training forward, both subclass-owned.""" + def _beam_model( + self, beam_widths=(2, 2, 2), num_return_sequences=2 + ) -> GenrecCausalLMModel: + model = GenrecCausalLMModel.__new__(GenrecCausalLMModel) + nn.Module.__init__(model) + model._prompt = self.compiled_prompt + common = GenrecModelConfig(num_return_sequences=num_return_sequences) + common.beam_widths.extend(beam_widths) + model._read_beam_config(common) + return model + @parameterized.expand( [ [[2, 3, 4], [2, 3, 4]], @@ -84,48 +196,26 @@ class GenrecCausalLMModelTest(GenrecModelTestBase): name_func=parameterized_name_func, ) def test_beam_widths_are_capped_once_at_init(self, beam_widths, expected) -> None: - model = self._model(beam_widths=beam_widths, num_return_sequences=1) + model = self._beam_model(beam_widths=beam_widths, num_return_sequences=1) self.assertEqual(model._capped_widths, expected) space = self.compiled_prompt.sid_space self.assertEqual(model._bands, list(zip(space.band_lo, space.band_hi))) def test_rejects_a_schedule_that_does_not_match_the_codebook(self) -> None: with self.assertRaisesRegex(ValueError, "entries but the codebook has"): - self._model(beam_widths=(2, 2)) + self._beam_model(beam_widths=(2, 2)) def test_rejects_a_non_positive_beam_width(self) -> None: with self.assertRaisesRegex(ValueError, "must be >= 1"): - self._model(beam_widths=(2, 0, 2)) + self._beam_model(beam_widths=(2, 0, 2)) def test_beam_config_uses_final_capped_capacity(self) -> None: with self.assertRaisesRegex(ValueError, "final capped beam width \\(4\\)"): - self._model( + self._beam_model( beam_widths=(1, 1, 100), num_return_sequences=5, ) - def test_training_forward_builds_no_cache(self) -> None: - model = self._model() - inner = model.lm.model.forward - - with mock.patch.object(model.lm.model, "forward", side_effect=inner) as spy: - model.predict( - self._batch( - { - "hist.values": torch.tensor( - offset_sid_codes([0, 1, 2], _CODEBOOK) - ), - "hist.lengths": torch.tensor([3]), - "answer.values": torch.tensor( - offset_sid_codes([1, 2, 3], _CODEBOOK) - ), - "answer.lengths": torch.tensor([3]), - } - ) - ) - - self.assertIs(spy.call_args.kwargs["use_cache"], False) - if __name__ == "__main__": unittest.main() diff --git a/tzrec/models/genrec_model.py b/tzrec/models/genrec_model.py index 7e309cda..8b03efa4 100644 --- a/tzrec/models/genrec_model.py +++ b/tzrec/models/genrec_model.py @@ -127,8 +127,11 @@ def init_backbone( lm_parameter_dtype: dtype of the LM parameters. """ config = AutoConfig.from_pretrained(hf_model_name_or_path) - model = AutoModelForCausalLM.from_config(config) - self.lm = model.to(_PARAM_DTYPE[lm_parameter_dtype]) + self.lm = AutoModelForCausalLM.from_config( + config, + attn_implementation="flash_attention_2", + torch_dtype=_PARAM_DTYPE[lm_parameter_dtype], + ) self._check_backbone_interfaces(hf_model_name_or_path) def _check_backbone_interfaces(self, hf_model_name_or_path: str) -> None: diff --git a/tzrec/models/genrec_model_test.py b/tzrec/models/genrec_model_test.py index 39278856..c22c88bb 100644 --- a/tzrec/models/genrec_model_test.py +++ b/tzrec/models/genrec_model_test.py @@ -11,6 +11,7 @@ import dataclasses import unittest +from unittest import mock import torch from parameterized import parameterized @@ -19,6 +20,7 @@ from tzrec.datasets.utils import Batch from tzrec.models.genrec_model import _PARAM_DTYPE +from tzrec.models.model import TrainWrapper from tzrec.prompt.assembler import ( PROMPT_HOLE_POSITIONS, PROMPT_INPUT_IDS, @@ -36,10 +38,14 @@ ) from tzrec.utils.state_dict_util import init_parameters from tzrec.utils.test_util import ( + mark_ci_scope, + nv_gpu_unavailable, parameterized_name_func, ) +@mark_ci_scope("gpu") +@unittest.skipIf(*nv_gpu_unavailable) class BaseGenrecModelTest(GenrecModelTestBase): """Shared causal-LM behavior, reached through its concrete subclass.""" @@ -72,6 +78,31 @@ def test_rejects_a_prompt_that_declares_no_sid_space(self) -> None: with self.assertRaisesRegex(ValueError, "declares no sid_space"): self._model(compiled_prompt=compiled_prompt) + def test_builds_backbone_with_flash_attention_2_and_target_dtype(self) -> None: + stand_in = AutoModelForCausalLM.from_pretrained(self.backbone) + + with ( + mock.patch.object(stand_in, "to", wraps=stand_in.to) as to_mock, + mock.patch.object( + AutoModelForCausalLM, "from_config", return_value=stand_in + ) as from_config, + ): + model = self._model(lm_parameter_dtype=GenrecModelConfig.BF16) + + from_config.assert_called_once() + args, kwargs = from_config.call_args + self.assertEqual(len(args), 1) + self.assertEqual(args[0].model_type, "qwen2") + self.assertEqual( + kwargs, + { + "attn_implementation": "flash_attention_2", + "torch_dtype": torch.bfloat16, + }, + ) + to_mock.assert_not_called() + self.assertIs(model.lm, stand_in) + def test_shared_projection_name_requires_matching_widths(self) -> None: features = [ create_prompt_feature(_HIST), @@ -155,7 +186,9 @@ def test_projected_slot_follows_a_narrow_lm_dtype(self, lm_parameter_dtype) -> N compiled_prompt=compiled_prompt, lm_parameter_dtype=lm_parameter_dtype, ) - init_parameters(model, device=torch.device("cpu")) + device = torch.device("cuda") + init_parameters(model, device=device) + model.to(device) batch = self._batch( { "hist.values": torch.tensor( @@ -173,7 +206,7 @@ def test_projected_slot_follows_a_narrow_lm_dtype(self, lm_parameter_dtype) -> N values=torch.tensor([5, 9]), lengths=torch.tensor([2]), ), - ) + ).to(device) embeds = model.build_input(batch) self.assertIs(embeds.dtype, _PARAM_DTYPE[lm_parameter_dtype]) @@ -187,6 +220,53 @@ def test_projected_slot_follows_a_narrow_lm_dtype(self, lm_parameter_dtype) -> N self.assertIs(proj.head.weight.dtype, torch.float32) self.assertGreater(float(proj.head.weight.grad.abs().sum()), 0.0) + def test_projected_slot_trains_with_fp32_masters_and_bf16_autocast(self) -> None: + features = [ + create_prompt_feature(_HIST), + create_prompt_feature(projected_feature("prof", 8)), + ] + compiled_prompt = self._compile( + features, + template="History : {{hist}} . Predict {{prof}} :", + response="{{answer}}", + ) + model = self._model(features=features, compiled_prompt=compiled_prompt) + device = torch.device("cuda") + init_parameters(model, device=device) + model.to(device) + batch = self._batch( + { + "hist.values": torch.tensor( + offset_sid_codes([0, 1, 2], _CODEBOOK) + ).reshape(-1, 1), + "hist.lengths": torch.tensor([3]), + "answer.values": torch.tensor(offset_sid_codes([1, 2, 3], _CODEBOOK)), + "answer.lengths": torch.tensor([3]), + "prof.values": torch.tensor([5, 9]), + "prof.lengths": torch.tensor([2]), + }, + compiled_prompt=compiled_prompt, + sparse=KeyedJaggedTensor.from_lengths_sync( + keys=["prof"], + values=torch.tensor([5, 9]), + lengths=torch.tensor([2]), + ), + ).to(device) + + wrapper = TrainWrapper(model, device=device, mixed_precision="BF16") + loss, _ = wrapper(batch) + self.assertTrue(bool(torch.isfinite(loss))) + loss.backward() + + lm_weight = model.lm.model.layers[0].self_attn.q_proj.weight + self.assertIs(lm_weight.dtype, torch.float32) + self.assertIsNotNone(lm_weight.grad) + self.assertTrue(bool(torch.isfinite(lm_weight.grad).all())) + proj_weight = next(iter(model.projections.values())).head.weight + self.assertIs(proj_weight.dtype, torch.float32) + self.assertIsNotNone(proj_weight.grad) + self.assertTrue(bool(torch.isfinite(proj_weight.grad).all())) + def test_metric_averages_the_loss_across_batches(self) -> None: model = self._model() model.init_metric() diff --git a/tzrec/modules/dynamic_beam_test.py b/tzrec/modules/dynamic_beam_test.py index bbb81eac..e6a580fd 100644 --- a/tzrec/modules/dynamic_beam_test.py +++ b/tzrec/modules/dynamic_beam_test.py @@ -15,9 +15,15 @@ import torch from parameterized import parameterized +from transformers import AutoConfig, AutoModelForCausalLM from tzrec.modules.dynamic_beam import capped_beam_widths, dynamic_beam_search -from tzrec.utils.test_util import create_tiny_causal_lm, parameterized_name_func +from tzrec.utils.test_util import ( + create_tiny_causal_lm, + mark_ci_scope, + nv_gpu_unavailable, + parameterized_name_func, +) def _decode(lm, ids, pairs, width=8, beam_widths=None, attention_mask=None): @@ -132,5 +138,55 @@ def test_exhaustive_matches_bruteforce_topk(self) -> None: self.assertEqual(got[0], max(ref, key=lambda combo: ref[combo])) +@mark_ci_scope("gpu") +@unittest.skipIf(*nv_gpu_unavailable) +class DynamicBeamSearchFlashAttentionTest(unittest.TestCase): + @parameterized.expand( + [["qwen2"], ["qwen3"]], + name_func=parameterized_name_func, + ) + def test_ragged_batch_matches_solo_dynamic_cache(self, model_type: str) -> None: + device = torch.device("cuda") + config = AutoConfig.for_model( + model_type, + vocab_size=30, + hidden_size=32, + intermediate_size=64, + num_hidden_layers=2, + num_attention_heads=4, + num_key_value_heads=2, + head_dim=8, + max_position_embeddings=64, + tie_word_embeddings=False, + ) + with torch.random.fork_rng(devices=[]): + torch.manual_seed(0) + lm = AutoModelForCausalLM.from_config( + config, + attn_implementation="flash_attention_2", + torch_dtype=torch.bfloat16, + ).to(device) + lm.eval() + pairs = [(20, 21), (22, 24), (25, 28)] + ids = torch.tensor( + [[5, 6, 7, 8], [0, 9, 10, 11]], + device=device, + ) + attention_mask = torch.tensor( + [[1, 1, 1, 1], [0, 1, 1, 1]], + device=device, + ) + + out = _decode(lm, ids, pairs, attention_mask=attention_mask) + width = out.shape[0] // 2 + solos = [ + torch.tensor([[5, 6, 7, 8]], device=device), + torch.tensor([[9, 10, 11]], device=device), + ] + for index, solo_ids in enumerate(solos): + solo = _decode(lm, solo_ids, pairs) + torch.testing.assert_close(out[index * width : (index + 1) * width], solo) + + if __name__ == "__main__": unittest.main() diff --git a/tzrec/prompt/assembler.py b/tzrec/prompt/assembler.py index 04eb64b3..75e6c00e 100644 --- a/tzrec/prompt/assembler.py +++ b/tzrec/prompt/assembler.py @@ -179,6 +179,11 @@ def _build_packed_prompt( seg_lengths[index] = projected_lengths[seg.name] row_lengths = seg_lengths.sum(axis=0) + body_lengths = seg_lengths[:body_count].sum(axis=0) + empty = np.flatnonzero(body_lengths == 0) + if empty.size: + sample = int(empty[0]) + raise ValueError(f"assembled sample {sample} has an empty prompt body.") cu_seqlens = np.concatenate(([0], np.cumsum(row_lengths))) max_length = self._prompt_plan.max_length if max_length: diff --git a/tzrec/prompt/assembler_test.py b/tzrec/prompt/assembler_test.py index c6475b24..4b8cd9ff 100644 --- a/tzrec/prompt/assembler_test.py +++ b/tzrec/prompt/assembler_test.py @@ -238,6 +238,18 @@ def test_over_long_row_is_an_error_not_a_truncation(self) -> None: with self.assertRaisesRegex(ValueError, "never truncated"): asm.forward(_parsed({"hist": [np.array([1, 6, 11])]})) + def test_empty_prompt_body_is_rejected(self) -> None: + asm = _asm( + (_slot("hist", FillMode.INLINE),), + response=(_slot("answer", FillMode.INLINE),), + ) + parsed = _parsed( + {"hist": [np.array([], dtype=np.int64)], "answer": [np.array([1, 6, 11])]} + ) + + with self.assertRaisesRegex(ValueError, "sample 0 has an empty prompt body"): + asm.forward(parsed) + def test_inline_without_a_sid_space_is_rejected_at_construction(self) -> None: plan = _plan((_slot("hist", FillMode.INLINE),)) with self.assertRaisesRegex(ValueError, "no sid_space was compiled"): diff --git a/tzrec/tests/configs/genrec_causal_lm_model_mock.config b/tzrec/tests/configs/genrec_causal_lm_model_mock.config index f6cf4cf4..a807eda3 100644 --- a/tzrec/tests/configs/genrec_causal_lm_model_mock.config +++ b/tzrec/tests/configs/genrec_causal_lm_model_mock.config @@ -18,6 +18,7 @@ train_config { } num_epochs: 1 save_checkpoints_epochs: 1 + mixed_precision: "BF16" } eval_config { } diff --git a/tzrec/tests/prompt_integration_test.py b/tzrec/tests/prompt_integration_test.py index ed86e8a0..576adab7 100644 --- a/tzrec/tests/prompt_integration_test.py +++ b/tzrec/tests/prompt_integration_test.py @@ -15,28 +15,47 @@ import numpy as np import torch import torch.fx +from parameterized import parameterized +from transformers import AutoConfig from tzrec.datasets.utils import Batch from tzrec.models.model import TrainWrapper +from tzrec.protos.models.genrec_model_pb2 import GenrecModelConfig from tzrec.tests.prompt_test_util import ( GenrecModelTestBase, assemble_into, offset_sid_codes, ) +from tzrec.utils.test_util import ( + mark_ci_scope, + nv_gpu_unavailable, + parameterized_name_func, +) _CODEBOOK = [4, 4, 4] _WORDS = ["History", "Predict", ":", ".", "", "<|im_end|>"] +@mark_ci_scope("gpu") +@unittest.skipIf(*nv_gpu_unavailable) class PromptStackIntegrationTest(GenrecModelTestBase): """compile -> assemble -> model, on the real code path.""" def _batch_from_codes(self, hist, answer): + return self._batch_from_rows([hist], [answer]) + + def _batch_from_rows(self, hist_rows, answer_rows): parsed = { - "hist.values": torch.tensor(offset_sid_codes(hist, _CODEBOOK)), - "hist.lengths": torch.tensor([len(hist)]), - "answer.values": torch.tensor(offset_sid_codes(answer, _CODEBOOK)), - "answer.lengths": torch.tensor([len(answer)]), + "hist.values": torch.from_numpy( + np.concatenate([offset_sid_codes(row, _CODEBOOK) for row in hist_rows]) + ), + "hist.lengths": torch.tensor([len(row) for row in hist_rows]), + "answer.values": torch.from_numpy( + np.concatenate( + [offset_sid_codes(row, _CODEBOOK) for row in answer_rows] + ) + ), + "answer.lengths": torch.tensor([len(row) for row in answer_rows]), } streams = assemble_into(self.compiled_prompt, parsed) batch = Batch() @@ -62,17 +81,68 @@ def test_model_resizes_to_target_vocab_size(self) -> None: self.assertEqual(rows, self.compiled_prompt.sid_space.target_vocab_size) self.assertGreater(rows, self.compiled_prompt.sid_space.band_hi[-1]) - def test_loss_is_finite_and_backpropagates_into_the_backbone(self) -> None: - model = self._model() - batch = self._batch_from_codes([0, 1, 2, 3, 0, 1], [1, 2, 3]) - predictions = model.predict(batch) - loss = model.loss(predictions, batch)["ce_loss"] + @parameterized.expand( + [["qwen2"], ["qwen3"]], + name_func=parameterized_name_func, + ) + def test_packed_rows_match_solo_runs_and_backpropagate( + self, model_type: str + ) -> None: + device = torch.device("cuda") + backbone = os.path.join(self.test_dir, model_type) + AutoConfig.for_model( + model_type, + vocab_size=64, + hidden_size=32, + intermediate_size=64, + num_hidden_layers=2, + num_attention_heads=4, + num_key_value_heads=2, + head_dim=8, + max_position_embeddings=64, + tie_word_embeddings=False, + ).save_pretrained(backbone) + with torch.random.fork_rng(devices=[]): + torch.manual_seed(0) + model = self._model( + lm_parameter_dtype=GenrecModelConfig.BF16, + hf_model_name_or_path=backbone, + ).to(device) + model.eval() + hist_rows = [[0, 1, 2], [3, 0, 1, 2, 3, 0]] + answer_rows = [[1, 2, 3], [2, 3, 0]] + packed_batch = self._batch_from_rows(hist_rows, answer_rows).to(device) + + packed = model.predict(packed_batch) + with torch.no_grad(): + solos = [ + model.predict(self._batch_from_codes(hist, answer).to(device)) + for hist, answer in zip(hist_rows, answer_rows) + ] + changed = model.predict( + self._batch_from_rows([[3, 3, 3], hist_rows[1]], answer_rows).to(device) + ) + + torch.testing.assert_close( + packed["logits"], + torch.cat([result["logits"] for result in solos]), + atol=1e-2, + rtol=1e-2, + ) + torch.testing.assert_close( + packed["labels"], torch.cat([result["labels"] for result in solos]) + ) + torch.testing.assert_close( + packed["logits"][1], changed["logits"][1], atol=0, rtol=0 + ) + + loss = model.loss(packed, packed_batch)["ce_loss"] self.assertTrue(bool(torch.isfinite(loss))) loss.backward() - grad = model.lm.get_input_embeddings().weight.grad self.assertIsNotNone(grad) - self.assertTrue(bool((grad.abs().sum() > 0))) + self.assertTrue(bool(torch.isfinite(grad).all())) + self.assertGreater(float(grad.abs().sum()), 0) def test_training_forward_survives_fx_tracing(self) -> None: model = self._model() From 163e736ccc0ee64a97ad32d6206828e4bba1a85c Mon Sep 17 00:00:00 2001 From: root <597191244@qq.com> Date: Wed, 16 Sep 2026 11:33:19 +0000 Subject: [PATCH 02/15] [fix] make GenRec a GPU-only, flash-attention-2 path GenRec builds its backbone with attn_implementation="flash_attention_2", which needs an NVIDIA GPU and the flash_attn wheel, so the tests that build a model cannot run on the CPU lane (run.py with no --scope runs everything not skipped, so a scope marker alone does not exclude them). - declare flash_attn in requirements/cu126|cu129|cu130.txt, cp311 only, beside fbgemm_gpu_hstu - mark BaseGenRecModelTest, GenRecFrontEndTest, test_training_forward_builds_no_cache and test_genrec_train_eval_export gpu-scoped and skip them without an NVIDIA GPU - raise a clear error when the backbone runs in fp32 with no autocast, instead of letting the flash kernel fail on the dtype Co-Authored-By: Claude Opus 5 (1M context) --- requirements/cu126.txt | 1 + requirements/cu129.txt | 1 + requirements/cu130.txt | 1 + tzrec/models/genrec_causal_lm_model.py | 20 ++++++++++++++++---- tzrec/models/genrec_causal_lm_model_test.py | 4 ++++ tzrec/models/genrec_model_test.py | 8 +++++--- tzrec/tests/genrec_integration_test.py | 3 +++ 7 files changed, 31 insertions(+), 7 deletions(-) diff --git a/requirements/cu126.txt b/requirements/cu126.txt index a27c483a..9c119216 100644 --- a/requirements/cu126.txt +++ b/requirements/cu126.txt @@ -7,4 +7,5 @@ faiss @ https://tzrec.oss-accelerate.aliyuncs.com/third_party/faiss/cu12/faiss-1 fbgemm_gpu_hstu @ https://tzrec.oss-accelerate.aliyuncs.com/third_party/hstu/cu126/fbgemm_gpu_hstu-0.1.0%2B20260914.fece651b.fn5.cu126-cp310-cp310-linux_x86_64.whl ; python_version=="3.10" fbgemm_gpu_hstu @ https://tzrec.oss-accelerate.aliyuncs.com/third_party/hstu/cu126/fbgemm_gpu_hstu-0.1.0%2B20260914.fece651b.fn5.cu126-cp311-cp311-linux_x86_64.whl ; python_version=="3.11" fbgemm_gpu_hstu @ https://tzrec.oss-accelerate.aliyuncs.com/third_party/hstu/cu126/fbgemm_gpu_hstu-0.1.0%2B20260914.fece651b.fn5.cu126-cp312-cp312-linux_x86_64.whl ; python_version=="3.12" +flash_attn @ https://tzrec.oss-accelerate.aliyuncs.com/third_party/flash_attn/cu126/flash_attn-2.8.3.post1%2Bcu126.torch213.sm80.90-cp311-cp311-linux_x86_64.whl ; python_version=="3.11" triton==3.7.1 diff --git a/requirements/cu129.txt b/requirements/cu129.txt index 7abb0df7..7e49b3fa 100644 --- a/requirements/cu129.txt +++ b/requirements/cu129.txt @@ -7,4 +7,5 @@ faiss @ https://tzrec.oss-accelerate.aliyuncs.com/third_party/faiss/cu12/faiss-1 fbgemm_gpu_hstu @ https://tzrec.oss-accelerate.aliyuncs.com/third_party/hstu/cu129/fbgemm_gpu_hstu-0.1.0%2B20260914.fece651b.fn5.cu129-cp310-cp310-linux_x86_64.whl ; python_version=="3.10" fbgemm_gpu_hstu @ https://tzrec.oss-accelerate.aliyuncs.com/third_party/hstu/cu129/fbgemm_gpu_hstu-0.1.0%2B20260914.fece651b.fn5.cu129-cp311-cp311-linux_x86_64.whl ; python_version=="3.11" fbgemm_gpu_hstu @ https://tzrec.oss-accelerate.aliyuncs.com/third_party/hstu/cu129/fbgemm_gpu_hstu-0.1.0%2B20260914.fece651b.fn5.cu129-cp312-cp312-linux_x86_64.whl ; python_version=="3.12" +flash_attn @ https://tzrec.oss-accelerate.aliyuncs.com/third_party/flash_attn/cu129/flash_attn-2.8.3.post1%2Bcu129.torch213.sm80.90.120-cp311-cp311-linux_x86_64.whl ; python_version=="3.11" triton==3.7.1 diff --git a/requirements/cu130.txt b/requirements/cu130.txt index 54fa9c62..62b1646e 100644 --- a/requirements/cu130.txt +++ b/requirements/cu130.txt @@ -7,5 +7,6 @@ faiss @ https://tzrec.oss-accelerate.aliyuncs.com/third_party/faiss/cu13/faiss-1 fbgemm_gpu_hstu @ https://tzrec.oss-accelerate.aliyuncs.com/third_party/hstu/cu130/fbgemm_gpu_hstu-0.1.0%2B20260914.fece651b.fn5.cu130-cp310-cp310-linux_x86_64.whl ; python_version=="3.10" fbgemm_gpu_hstu @ https://tzrec.oss-accelerate.aliyuncs.com/third_party/hstu/cu130/fbgemm_gpu_hstu-0.1.0%2B20260914.fece651b.fn5.cu130-cp311-cp311-linux_x86_64.whl ; python_version=="3.11" fbgemm_gpu_hstu @ https://tzrec.oss-accelerate.aliyuncs.com/third_party/hstu/cu130/fbgemm_gpu_hstu-0.1.0%2B20260914.fece651b.fn5.cu130-cp312-cp312-linux_x86_64.whl ; python_version=="3.12" +flash_attn @ https://tzrec.oss-accelerate.aliyuncs.com/third_party/flash_attn/cu130/flash_attn-2.8.3.post1%2Bcu130.torch213.sm80.90.120-cp311-cp311-linux_x86_64.whl ; python_version=="3.11" torch-tensorrt==2.13.0 triton==3.7.1 diff --git a/tzrec/models/genrec_causal_lm_model.py b/tzrec/models/genrec_causal_lm_model.py index ddea1fac..e9849982 100644 --- a/tzrec/models/genrec_causal_lm_model.py +++ b/tzrec/models/genrec_causal_lm_model.py @@ -143,6 +143,21 @@ def _forward( Logits and labels over the same window, so the shift ``loss`` applies lands on the pairs the window was sized for. """ + # fp32 masters are fine under autocast, so only the forward can tell + if ( + embeds.is_cuda + and embeds.dtype == torch.float32 + and not torch.is_autocast_enabled("cuda") + ): + raise ValueError( + f"{type(self).__name__}: flash_attention_2 needs bf16 or fp16 " + f"activations, but the backbone ran in fp32 with no autocast " + f"active. Set train_config.mixed_precision to BF16 or FP16 to " + f"keep fp32 master weights, or set " + f"the model's common.lm_parameter_dtype to BF16 or FP16 to " + f"narrow the parameters themselves." + ) + infos = batch.additional_infos cu_seqlens = infos[CU_SEQLENS] starts = cu_seqlens[:-1] @@ -155,10 +170,7 @@ def _forward( ).unsqueeze(0) suffix = cast(int, self._prompt.prompt_plan.logits_suffix_len) - # the window walks back from each row's end, so a row shorter than - # it reads the row before it, and for row 0 wraps to the tail of the - # pack. The assembler is scripted for serving and cannot raise, so - # the packed forward owns the check. + # a row shorter than the window reads the row before it if bool((lengths < suffix).any()): raise ValueError( f"{type(self).__name__}: every assembled sample must be at " diff --git a/tzrec/models/genrec_causal_lm_model_test.py b/tzrec/models/genrec_causal_lm_model_test.py index 61f32feb..b89c6e96 100644 --- a/tzrec/models/genrec_causal_lm_model_test.py +++ b/tzrec/models/genrec_causal_lm_model_test.py @@ -38,6 +38,8 @@ create_genrec_test_model, create_genrec_test_tokenizer, make_test_dir, + mark_ci_scope, + nv_gpu_unavailable, parameterized_name_func, ) @@ -285,6 +287,8 @@ def test_beam_config_uses_final_capped_capacity(self) -> None: num_return_sequences=5, ) + @unittest.skipIf(*nv_gpu_unavailable) + @mark_ci_scope("gpu") def test_training_forward_builds_no_cache(self) -> None: model, compiled_prompt = create_genrec_test_model(self.test_dir) batch = Batch() diff --git a/tzrec/models/genrec_model_test.py b/tzrec/models/genrec_model_test.py index 06ec86dc..c47162b6 100644 --- a/tzrec/models/genrec_model_test.py +++ b/tzrec/models/genrec_model_test.py @@ -105,6 +105,8 @@ def _projected_batch(compiled_prompt) -> Batch: ) +@mark_ci_scope("gpu") +@unittest.skipIf(*nv_gpu_unavailable) class BaseGenRecModelTest(unittest.TestCase): """Shared causal-LM behavior, reached through its concrete subclass.""" @@ -308,6 +310,8 @@ def test_training_forward_survives_fx_tracing(self) -> None: torch.fx.symbolic_trace(TrainWrapper(self.model)) +@mark_ci_scope("gpu") +@unittest.skipIf(*nv_gpu_unavailable) class GenRecFrontEndTest(unittest.TestCase): """The served half of the model, under the same wrapper every export uses.""" @@ -395,9 +399,7 @@ def test_front_end_traces_and_scripts(self) -> None: self.assertTrue(torch.equal(out[key], value), key) -# offset SID codes for the (4, 4, 4) codebook, beside the module's _HIST_CODES / -# _LONG_HIST_CODES / _ANSWER_CODES: a second answer, and a rewrite of -# _HIST_CODES that keeps its width so the row after it does not move +# a second answer, and a rewrite of _HIST_CODES that keeps its width _OTHER_ANSWER_CODES = [2, 7, 8] _REWRITTEN_HIST_CODES = [3, 7, 11] diff --git a/tzrec/tests/genrec_integration_test.py b/tzrec/tests/genrec_integration_test.py index 6f63204a..c6dded1b 100644 --- a/tzrec/tests/genrec_integration_test.py +++ b/tzrec/tests/genrec_integration_test.py @@ -39,6 +39,7 @@ gpu_unavailable, make_test_dir, mark_ci_scope, + nv_gpu_unavailable, ) _MOCK_CONFIG = "tzrec/tests/configs/genrec_causal_lm_model_mock.config" @@ -120,6 +121,8 @@ def _request(self, columns, rows: int = 4): out[column + ".lengths"] = torch.tensor([len(row) for row in lists]) return out + @unittest.skipIf(*nv_gpu_unavailable) + @mark_ci_scope("gpu") def test_genrec_train_eval_export(self): trained = self._train_eval_export() export_dir = os.path.join(self.test_dir, "export") From 4f5e2981dc775be5b1b8f8ad5e375e41ca6a98c9 Mon Sep 17 00:00:00 2001 From: root <597191244@qq.com> Date: Wed, 16 Sep 2026 11:56:23 +0000 Subject: [PATCH 03/15] [fix] run the genrec model tests on the GPU they now require flash_attention_2 has no CPU kernel, so a test that builds the model and runs a forward has to put both the model and the batch on cuda; five did not, and failed with "Could not run 'flash_attn::_flash_attn_varlen_forward' with arguments from the 'CPU' backend" even on a GPU host. Two of them also ran the backbone in fp32 with no autocast, which the new dtype check rejects: the loss test now runs under bf16 autocast, matching the fp32-masters setup lm_parameter_dtype documents, and the no-cache test narrows its backbone to BF16. Verified on 2x H20 (torch 2.13.0+cu126, transformers 5.17.0, flash_attn 2.8.3.post1): 41 tests OK, including PackedFlashAttentionTest, which asserts atol=0 row isolation and had never executed before. Co-Authored-By: Claude Opus 5 (1M context) --- tzrec/models/genrec_causal_lm_model_test.py | 7 +++++- tzrec/models/genrec_model_test.py | 25 ++++++++++++--------- 2 files changed, 21 insertions(+), 11 deletions(-) diff --git a/tzrec/models/genrec_causal_lm_model_test.py b/tzrec/models/genrec_causal_lm_model_test.py index b89c6e96..11eace2c 100644 --- a/tzrec/models/genrec_causal_lm_model_test.py +++ b/tzrec/models/genrec_causal_lm_model_test.py @@ -290,7 +290,11 @@ def test_beam_config_uses_final_capped_capacity(self) -> None: @unittest.skipIf(*nv_gpu_unavailable) @mark_ci_scope("gpu") def test_training_forward_builds_no_cache(self) -> None: - model, compiled_prompt = create_genrec_test_model(self.test_dir) + model, compiled_prompt = create_genrec_test_model( + self.test_dir, lm_parameter_dtype=GenRecModelConfig.BF16 + ) + device = torch.device("cuda") + model.to(device) batch = Batch() batch.additional_infos.update( PromptAssembler(compiled_prompt.prompt_plan, compiled_prompt.sid_space)( @@ -303,6 +307,7 @@ def test_training_forward_builds_no_cache(self) -> None: } ) ) + batch = batch.to(device) inner = model.lm.model.forward with mock.patch.object(model.lm.model, "forward", side_effect=inner) as spy: diff --git a/tzrec/models/genrec_model_test.py b/tzrec/models/genrec_model_test.py index c47162b6..6c0a9a59 100644 --- a/tzrec/models/genrec_model_test.py +++ b/tzrec/models/genrec_model_test.py @@ -218,8 +218,10 @@ def test_projected_slot_follows_a_narrow_lm_dtype(self, lm_parameter_dtype) -> N prompt="History : {{hist}} . Predict {{prof}} :", lm_parameter_dtype=lm_parameter_dtype, ) - init_parameters(model, device=torch.device("cpu")) - batch = _projected_batch(compiled_prompt) + device = torch.device("cuda") + init_parameters(model, device=device) + model.to(device) + batch = _projected_batch(compiled_prompt).to(device) embeds = model.build_input(batch) self.assertIs(embeds.dtype, _PARAM_DTYPE[lm_parameter_dtype]) @@ -239,12 +241,12 @@ def test_projected_slot_trains_with_fp32_masters_and_bf16_autocast(self) -> None feature_configs=[_hist(), _projected("prof", 8)], prompt="History : {{hist}} . Predict {{prof}} :", ) - init_parameters(model, device=torch.device("cpu")) - batch = _projected_batch(compiled_prompt) + device = torch.device("cuda") + init_parameters(model, device=device) + model.to(device) + batch = _projected_batch(compiled_prompt).to(device) - wrapper = TrainWrapper( - model, device=torch.device("cpu"), mixed_precision="BF16" - ) + wrapper = TrainWrapper(model, device=device, mixed_precision="BF16") loss, _ = wrapper(batch) self.assertTrue(bool(torch.isfinite(loss))) loss.backward() @@ -288,6 +290,8 @@ def test_model_resizes_to_target_vocab_size(self) -> None: self.assertGreater(rows, self.compiled_prompt.sid_space.band_hi[-1]) def test_loss_is_finite_and_backpropagates_into_the_backbone(self) -> None: + device = torch.device("cuda") + self.model.to(device) batch = _batch( self.compiled_prompt, { @@ -296,9 +300,10 @@ def test_loss_is_finite_and_backpropagates_into_the_backbone(self) -> None: "answer.values": torch.tensor(_ANSWER_CODES), "answer.lengths": torch.tensor([3]), }, - ) - predictions = self.model.predict(batch) - loss = self.model.loss(predictions, batch)["ce_loss"] + ).to(device) + with torch.autocast("cuda", dtype=torch.bfloat16): + predictions = self.model.predict(batch) + loss = self.model.loss(predictions, batch)["ce_loss"] self.assertTrue(bool(torch.isfinite(loss))) loss.backward() From 4ec67686958300f7409926b68700b88fad96099f Mon Sep 17 00:00:00 2001 From: root <597191244@qq.com> Date: Wed, 16 Sep 2026 13:38:32 +0000 Subject: [PATCH 04/15] [fix] skip the GenRec tests on the flash_attn wheel, not just on a GPU nv_gpu_unavailable only reports whether CUDA is present, so a GPU host running a python the flash_attn wheel is not built for -- it ships cp311 only, while the project supports 3.10/3.11/3.12 -- runs these tests and errors in init_backbone instead of skipping. Add flash_attn_unavailable beside the other optional-wheel probes, in the same find_spec idiom as cutlass_hstu_unavailable and faiss_unavailable, and guard the four sites that build a flash_attention_2 backbone. Teach the ci-scope lint the new token so a flash-only skip still has to declare a gpu lane. Co-Authored-By: Claude Opus 5 (1M context) --- tzrec/models/genrec_causal_lm_model_test.py | 2 ++ tzrec/models/genrec_model_test.py | 4 ++++ tzrec/modules/dynamic_beam_test.py | 2 ++ tzrec/tests/ci_scope_coverage_test.py | 1 + tzrec/tests/genrec_integration_test.py | 2 ++ tzrec/utils/test_util.py | 4 ++++ 6 files changed, 15 insertions(+) diff --git a/tzrec/models/genrec_causal_lm_model_test.py b/tzrec/models/genrec_causal_lm_model_test.py index 11eace2c..d6aa88c0 100644 --- a/tzrec/models/genrec_causal_lm_model_test.py +++ b/tzrec/models/genrec_causal_lm_model_test.py @@ -37,6 +37,7 @@ from tzrec.utils.test_util import ( create_genrec_test_model, create_genrec_test_tokenizer, + flash_attn_unavailable, make_test_dir, mark_ci_scope, nv_gpu_unavailable, @@ -288,6 +289,7 @@ def test_beam_config_uses_final_capped_capacity(self) -> None: ) @unittest.skipIf(*nv_gpu_unavailable) + @unittest.skipIf(*flash_attn_unavailable) @mark_ci_scope("gpu") def test_training_forward_builds_no_cache(self) -> None: model, compiled_prompt = create_genrec_test_model( diff --git a/tzrec/models/genrec_model_test.py b/tzrec/models/genrec_model_test.py index 6c0a9a59..f8fb722a 100644 --- a/tzrec/models/genrec_model_test.py +++ b/tzrec/models/genrec_model_test.py @@ -48,6 +48,7 @@ from tzrec.utils.test_util import ( create_genrec_test_model, create_genrec_test_tokenizer, + flash_attn_unavailable, make_test_dir, mark_ci_scope, nv_gpu_unavailable, @@ -107,6 +108,7 @@ def _projected_batch(compiled_prompt) -> Batch: @mark_ci_scope("gpu") @unittest.skipIf(*nv_gpu_unavailable) +@unittest.skipIf(*flash_attn_unavailable) class BaseGenRecModelTest(unittest.TestCase): """Shared causal-LM behavior, reached through its concrete subclass.""" @@ -317,6 +319,7 @@ def test_training_forward_survives_fx_tracing(self) -> None: @mark_ci_scope("gpu") @unittest.skipIf(*nv_gpu_unavailable) +@unittest.skipIf(*flash_attn_unavailable) class GenRecFrontEndTest(unittest.TestCase): """The served half of the model, under the same wrapper every export uses.""" @@ -483,6 +486,7 @@ def _packed_flash_model( @mark_ci_scope("gpu") @unittest.skipIf(*nv_gpu_unavailable) +@unittest.skipIf(*flash_attn_unavailable) class PackedFlashAttentionTest(unittest.TestCase): """The packed varlen forward, against the same rows run one at a time.""" diff --git a/tzrec/modules/dynamic_beam_test.py b/tzrec/modules/dynamic_beam_test.py index e6a580fd..390bff15 100644 --- a/tzrec/modules/dynamic_beam_test.py +++ b/tzrec/modules/dynamic_beam_test.py @@ -20,6 +20,7 @@ from tzrec.modules.dynamic_beam import capped_beam_widths, dynamic_beam_search from tzrec.utils.test_util import ( create_tiny_causal_lm, + flash_attn_unavailable, mark_ci_scope, nv_gpu_unavailable, parameterized_name_func, @@ -140,6 +141,7 @@ def test_exhaustive_matches_bruteforce_topk(self) -> None: @mark_ci_scope("gpu") @unittest.skipIf(*nv_gpu_unavailable) +@unittest.skipIf(*flash_attn_unavailable) class DynamicBeamSearchFlashAttentionTest(unittest.TestCase): @parameterized.expand( [["qwen2"], ["qwen3"]], diff --git a/tzrec/tests/ci_scope_coverage_test.py b/tzrec/tests/ci_scope_coverage_test.py index fb7b7a2e..17a2ba97 100644 --- a/tzrec/tests/ci_scope_coverage_test.py +++ b/tzrec/tests/ci_scope_coverage_test.py @@ -26,6 +26,7 @@ "has_dynamicemb", "has_tensorrt", "torch_fx_tool_unavailable", + "flash_attn_unavailable", "device_count", ) _TZREC_ROOT = os.path.dirname(os.path.dirname(os.path.dirname(__file__))) diff --git a/tzrec/tests/genrec_integration_test.py b/tzrec/tests/genrec_integration_test.py index c6dded1b..5c4c7494 100644 --- a/tzrec/tests/genrec_integration_test.py +++ b/tzrec/tests/genrec_integration_test.py @@ -36,6 +36,7 @@ from tzrec.utils.test_util import ( create_genrec_test_tokenizer, create_tiny_causal_lm, + flash_attn_unavailable, gpu_unavailable, make_test_dir, mark_ci_scope, @@ -122,6 +123,7 @@ def _request(self, columns, rows: int = 4): return out @unittest.skipIf(*nv_gpu_unavailable) + @unittest.skipIf(*flash_attn_unavailable) @mark_ci_scope("gpu") def test_genrec_train_eval_export(self): trained = self._train_eval_export() diff --git a/tzrec/utils/test_util.py b/tzrec/utils/test_util.py index 7d221e5b..0bb9650d 100644 --- a/tzrec/utils/test_util.py +++ b/tzrec/utils/test_util.py @@ -61,6 +61,10 @@ importlib.util.find_spec("faiss") is None, "faiss is not installed (required for SID residual K-Means)", ) +flash_attn_unavailable: Tuple[bool, str] = ( + importlib.util.find_spec("flash_attn") is None, + "flash_attn wheel is not installed (required for GenRec packed attention)", +) def get_compare_tolerance( From 342660d8ec6ce3624640b2679eaeacc74f2fcdeb Mon Sep 17 00:00:00 2001 From: root <597191244@qq.com> Date: Thu, 17 Sep 2026 03:37:27 +0000 Subject: [PATCH 05/15] [feat] choose the GenRec attention kernel from the config flash_attention_2 was hard-coded, so building a GenRec model needed an NVIDIA GPU and the flash_attn wheel -- which ships cp311 only -- even to construct one for export or on a python the wheel is not built for. Add common.attn_implementation, SDPA or FLASH_ATTENTION_2, defaulting to SDPA, in the shape of the existing Kernel enum: the backend whose dependency is always present is the default, and the optional-wheel one is explicit opt-in. There is no AUTO and no silent fallback -- configuring FLASH_ATTENTION_2 without the wheel raises and names the install path, since a quiet downgrade would hide the regression this path exists to fix. The packed forward is unchanged: sdpa carries the same row boundaries through position_ids, and the varlen kwargs pass through it inert. Selecting SDPA on CUDA also disables the cuDNN sdpa backend, which is eligible for the packed mask and returns NaN losses on it. The fp32 check now fires only on the flash path, which is the only one that rejects fp32, and the packed-vs-solo row-isolation test runs on both kernels. Co-Authored-By: Claude Opus 5 (1M context) --- tzrec/models/genrec_causal_lm_model.py | 5 +-- tzrec/models/genrec_causal_lm_model_test.py | 32 +++++++++++++++-- tzrec/models/genrec_model.py | 39 +++++++++++++++++++-- tzrec/models/genrec_model_test.py | 34 ++++++++++++++---- tzrec/protos/models/genrec_model.proto | 10 ++++++ tzrec/utils/test_util.py | 4 +++ 6 files changed, 110 insertions(+), 14 deletions(-) diff --git a/tzrec/models/genrec_causal_lm_model.py b/tzrec/models/genrec_causal_lm_model.py index e9849982..a5e0d8d7 100644 --- a/tzrec/models/genrec_causal_lm_model.py +++ b/tzrec/models/genrec_causal_lm_model.py @@ -143,9 +143,10 @@ def _forward( Logits and labels over the same window, so the shift ``loss`` applies lands on the pairs the window was sized for. """ - # fp32 masters are fine under autocast, so only the forward can tell + # only the flash kernel rejects fp32, and fp32 masters are fine under + # autocast, so neither the config nor the dtype alone can tell if ( - embeds.is_cuda + self.lm.config._attn_implementation == "flash_attention_2" and embeds.dtype == torch.float32 and not torch.is_autocast_enabled("cuda") ): diff --git a/tzrec/models/genrec_causal_lm_model_test.py b/tzrec/models/genrec_causal_lm_model_test.py index d6aa88c0..0231ace0 100644 --- a/tzrec/models/genrec_causal_lm_model_test.py +++ b/tzrec/models/genrec_causal_lm_model_test.py @@ -107,9 +107,10 @@ def test_packs_rows_of_different_lengths(self) -> None: class _CapturingLM(nn.Module): - def __init__(self) -> None: + def __init__(self, attn_implementation: str = "sdpa") -> None: super().__init__() self.kwargs = {} + self.config = SimpleNamespace(_attn_implementation=attn_implementation) def forward(self, **kwargs): self.kwargs = kwargs @@ -125,7 +126,9 @@ def forward(self, **kwargs): class _DifferentiableLM(nn.Module): def __init__(self, hidden_size: int, vocab_size: int) -> None: super().__init__() - self.config = SimpleNamespace(vocab_size=vocab_size) + self.config = SimpleNamespace( + vocab_size=vocab_size, _attn_implementation="sdpa" + ) self.lm_head = nn.Linear(hidden_size, vocab_size, bias=False) self.loss_function = ForCausalLMLoss @@ -196,6 +199,31 @@ def test_a_row_shorter_than_the_window_is_rejected(self) -> None: with self.assertRaisesRegex(ValueError, "at least logits_suffix_len"): model._forward(torch.zeros(12, 6), batch) + def test_fp32_without_autocast_is_rejected_on_the_flash_path(self) -> None: + """The flash kernel takes bf16/fp16 only; sdpa is happy in fp32.""" + model = GenRecCausalLMModel.__new__(GenRecCausalLMModel) + nn.Module.__init__(model) + model.lm = _CapturingLM("flash_attention_2") + model._ignore_index = -7 + model._prompt = SimpleNamespace( + prompt_plan=SimpleNamespace(logits_suffix_len=4) + ) + batch = Batch( + additional_infos={ + CU_SEQLENS: torch.tensor([0, 5, 12], dtype=torch.int32), + INPUT_IDS: torch.arange(100, 112), + MAX_SEQLEN: torch.tensor(7), + RESPONSE_LENGTHS: torch.tensor([3, 2]), + } + ) + + with self.assertRaisesRegex(ValueError, "needs bf16 or fp16"): + model._forward(torch.zeros(12, 6, dtype=torch.float32), batch) + + model.lm = _CapturingLM("sdpa") + logits, _ = model._forward(torch.zeros(12, 6, dtype=torch.float32), batch) + self.assertEqual(logits.shape, (2, 4, 5)) + def test_loss_and_gradients_cover_only_valid_response_pairs(self) -> None: model = GenRecCausalLMModel.__new__(GenRecCausalLMModel) nn.Module.__init__(model) diff --git a/tzrec/models/genrec_model.py b/tzrec/models/genrec_model.py index bf8abe97..c0c49838 100644 --- a/tzrec/models/genrec_model.py +++ b/tzrec/models/genrec_model.py @@ -18,6 +18,7 @@ HuggingFace weights. """ +import importlib.util import inspect from typing import Any, Dict, List, Optional, Sequence, Tuple @@ -53,6 +54,11 @@ GenRecModelConfig.FP16: torch.float16, } +_ATTN_IMPL: Dict[int, str] = { + GenRecModelConfig.SDPA: "sdpa", + GenRecModelConfig.FLASH_ATTENTION_2: "flash_attention_2", +} + _REQUIRED_LM_ATTRS: Tuple[str, ...] = ( "loss_function", "get_input_embeddings", @@ -98,7 +104,11 @@ def __init__( self._ignore_index = int(cfg.common.ignore_index) self.lm: nn.Module - self.init_backbone(cfg.hf_model_name_or_path, cfg.common.lm_parameter_dtype) + self.init_backbone( + cfg.hf_model_name_or_path, + cfg.common.lm_parameter_dtype, + cfg.common.attn_implementation, + ) # Every run replaces this initialization from pretrained or DCP weights. self.lm.resize_token_embeddings( compiled_prompt.sid_space.target_vocab_size, mean_resizing=False @@ -120,7 +130,10 @@ def init_input(self) -> None: self.init_projections() def init_backbone( - self, hf_model_name_or_path: str, lm_parameter_dtype: int + self, + hf_model_name_or_path: str, + lm_parameter_dtype: int, + attn_implementation: int = GenRecModelConfig.SDPA, ) -> None: """Assign ``self.lm`` from config, so HF weights load only on cold start. @@ -128,11 +141,31 @@ def init_backbone( hf_model_name_or_path: hub id or local directory naming the architecture and cold-start weights. lm_parameter_dtype: dtype of the LM parameters. + attn_implementation: attention kernel to build the backbone with. + + Raises: + ImportError: FLASH_ATTENTION_2 is configured but the wheel is + absent, so the packed forward has no kernel to run on. """ + impl = _ATTN_IMPL[attn_implementation] + if impl == "flash_attention_2": + if importlib.util.find_spec("flash_attn") is None: + raise ImportError( + f"{type(self).__name__}: attn_implementation is " + f"FLASH_ATTENTION_2 but the flash_attn wheel is not " + f"installed. Install it from " + f"https://tzrec.oss-accelerate.aliyuncs.com/third_party/" + f"flash_attn/${{DEVICE}}/ (cu126/cu129/cu130), or set " + f"attn_implementation to SDPA." + ) + elif torch.cuda.is_available(): + # the packed mask makes CUDNN_ATTENTION eligible, and it returns + # NaN losses on this shape; the other sdpa backends are fine + torch.backends.cuda.enable_cudnn_sdp(False) config = AutoConfig.from_pretrained(hf_model_name_or_path) self.lm = AutoModelForCausalLM.from_config( config, - attn_implementation="flash_attention_2", + attn_implementation=impl, torch_dtype=_PARAM_DTYPE[lm_parameter_dtype], ) self._check_backbone_interfaces(hf_model_name_or_path) diff --git a/tzrec/models/genrec_model_test.py b/tzrec/models/genrec_model_test.py index f8fb722a..ee7ee470 100644 --- a/tzrec/models/genrec_model_test.py +++ b/tzrec/models/genrec_model_test.py @@ -140,7 +140,16 @@ def test_rejects_a_model_built_without_a_prompt(self) -> None: with self.assertRaisesRegex(ValueError, "needs a compiled prompt"): _create_model(model_config, [], ["answer"], compiled_prompt=None) - def test_builds_backbone_with_flash_attention_2_and_target_dtype(self) -> None: + @parameterized.expand( + [ + [GenRecModelConfig.SDPA, "sdpa"], + [GenRecModelConfig.FLASH_ATTENTION_2, "flash_attention_2"], + ], + name_func=parameterized_name_func, + ) + def test_builds_backbone_with_the_configured_kernel_and_dtype( + self, attn_implementation, expected_impl + ) -> None: # from_config is mocked, so this pins the kwargs init_backbone sends # without building a second backbone; setUp still builds a real one. stand_in = AutoModelForCausalLM.from_pretrained( @@ -154,7 +163,9 @@ def test_builds_backbone_with_flash_attention_2_and_target_dtype(self) -> None: ) as from_config, ): model, _ = create_genrec_test_model( - self.test_dir, lm_parameter_dtype=GenRecModelConfig.BF16 + self.test_dir, + lm_parameter_dtype=GenRecModelConfig.BF16, + attn_implementation=attn_implementation, ) from_config.assert_called_once() @@ -164,7 +175,7 @@ def test_builds_backbone_with_flash_attention_2_and_target_dtype(self) -> None: self.assertEqual( kwargs, { - "attn_implementation": "flash_attention_2", + "attn_implementation": expected_impl, "torch_dtype": torch.bfloat16, }, ) @@ -428,7 +439,7 @@ def _packed_batch(compiled_prompt, hist_rows, answer_rows) -> Batch: def _packed_flash_model( - test_dir: str, model_type: str + test_dir: str, model_type: str, attn_implementation: int ) -> Tuple[BaseModel, CompiledPrompt]: """Build a bf16 genrec model over a tiny ``model_type`` backbone. @@ -441,6 +452,7 @@ def _packed_flash_model( Args: test_dir (str): scratch directory the backbone is written under. model_type (str): the hugging-face ``model_type`` of the backbone. + attn_implementation (int): ``GenRecModelConfig.AttnImpl`` to build with. Returns: Tuple[BaseModel, CompiledPrompt]: the model and the prompt it was @@ -476,6 +488,7 @@ def _packed_flash_model( lm_config.common.num_return_sequences = 2 # flash attention runs on fp16/bf16 only, and this arm carries no autocast lm_config.common.lm_parameter_dtype = GenRecModelConfig.BF16 + lm_config.common.attn_implementation = attn_implementation with torch.random.fork_rng(devices=[]): torch.manual_seed(0) model = _create_model( @@ -494,11 +507,16 @@ def setUp(self) -> None: self.test_dir = make_test_dir() @parameterized.expand( - [["qwen2"], ["qwen3"]], + [ + ["qwen2", GenRecModelConfig.SDPA], + ["qwen3", GenRecModelConfig.SDPA], + ["qwen2", GenRecModelConfig.FLASH_ATTENTION_2], + ["qwen3", GenRecModelConfig.FLASH_ATTENTION_2], + ], name_func=parameterized_name_func, ) def test_packed_rows_match_solo_runs_and_backpropagate( - self, model_type: str + self, model_type: str, attn_implementation: int ) -> None: """A packed batch reads as the rows do alone, and nothing crosses rows. @@ -509,7 +527,9 @@ def test_packed_rows_match_solo_runs_and_backpropagate( identical. """ device = torch.device("cuda") - model, compiled_prompt = _packed_flash_model(self.test_dir, model_type) + model, compiled_prompt = _packed_flash_model( + self.test_dir, model_type, attn_implementation + ) model.to(device) model.eval() hist_rows = [_HIST_CODES, _LONG_HIST_CODES] diff --git a/tzrec/protos/models/genrec_model.proto b/tzrec/protos/models/genrec_model.proto index d8d86edd..61e3a858 100644 --- a/tzrec/protos/models/genrec_model.proto +++ b/tzrec/protos/models/genrec_model.proto @@ -24,6 +24,16 @@ message GenRecModelConfig { // LM parameters. FP32 avoids bf16-ULP underflow of Adam's small updates; // bf16 compute comes from mixed_precision, not from this. optional ParamDtype lm_parameter_dtype = 5 [default = FP32]; + + // Attention kernel the backbone is built with. FLASH_ATTENTION_2 packs the + // batch through the varlen kernel and is what training should use, but it + // needs the flash_attn wheel and a GPU; SDPA carries the same row + // boundaries through position_ids and runs anywhere. + enum AttnImpl { + SDPA = 0; + FLASH_ATTENTION_2 = 1; + } + optional AttnImpl attn_implementation = 6 [default = SDPA]; } message GenRecCausalLMModel { diff --git a/tzrec/utils/test_util.py b/tzrec/utils/test_util.py index 0bb9650d..55ea7230 100644 --- a/tzrec/utils/test_util.py +++ b/tzrec/utils/test_util.py @@ -395,6 +395,7 @@ def create_genrec_test_model( beam_widths: Sequence[int] = (2, 2, 2), num_return_sequences: int = 2, lm_parameter_dtype: Optional["GenRecModelConfig.ParamDtype"] = None, + attn_implementation: Optional["GenRecModelConfig.AttnImpl"] = None, ) -> Tuple[BaseModel, CompiledPrompt]: """Build a GenRecCausalLMModel over a tiny backbone and a compiled prompt. @@ -412,6 +413,7 @@ def create_genrec_test_model( beam_widths (Sequence[int]): per-level beam widths. num_return_sequences (int): sequences returned per sample. lm_parameter_dtype (optional): ``GenRecModelConfig.ParamDtype`` value. + attn_implementation (optional): ``GenRecModelConfig.AttnImpl`` value. Returns: Tuple[BaseModel, CompiledPrompt]: the model and the prompt it was built on. @@ -449,6 +451,8 @@ def create_genrec_test_model( lm_config.common.num_return_sequences = num_return_sequences if lm_parameter_dtype is not None: lm_config.common.lm_parameter_dtype = lm_parameter_dtype + if attn_implementation is not None: + lm_config.common.attn_implementation = attn_implementation model = _create_model( model_config, features, ["answer"], compiled_prompt=compiled_prompt ) From 01ef25a8b3146f2a74d900cbdeba5e9aee8cffb4 Mon Sep 17 00:00:00 2001 From: root <597191244@qq.com> Date: Thu, 17 Sep 2026 03:47:11 +0000 Subject: [PATCH 06/15] [fix] return the GenRec tests to the CPU lane, and pin transformers With SDPA the default, building a GenRec model needs neither a GPU nor the flash_attn wheel, so the tests that only construct and run a model belong back on the lane that runs every PR. This reverts the cuda moves and the gpu markers on BaseGenRecModelTest, GenRecFrontEndTest, test_training_forward_builds_no_cache and test_genrec_train_eval_export. The packed-vs-solo row-isolation case splits in two: PackedSdpaAttentionTest runs it in fp32 on CPU, and PackedFlashAttentionTest runs the same body in bf16 on cuda behind the GPU and wheel probes. The leakage probe now guards the feature on every PR rather than only where the wheel exists. That makes the transformers floor load-bearing rather than advisory: masking_utils.find_packed_sequence_indices, which carries the row boundaries on every non-flash kernel, lands in 4.56. On 4.51.2 the sdpa case fails -- rows attend across each other -- so requirements now say transformers>=4.56. Verified with no GPU and no wheel: 54 tests, OK, 4 skipped (the two flash cases and the two flash beam cases). Co-Authored-By: Claude Opus 5 (1M context) --- requirements/runtime.txt | 2 +- tzrec/models/genrec_causal_lm_model_test.py | 13 +-- tzrec/models/genrec_model.py | 4 +- tzrec/models/genrec_model_test.py | 116 ++++++++++---------- tzrec/tests/genrec_integration_test.py | 5 - 5 files changed, 65 insertions(+), 75 deletions(-) diff --git a/requirements/runtime.txt b/requirements/runtime.txt index e850ea10..3caae27d 100644 --- a/requirements/runtime.txt +++ b/requirements/runtime.txt @@ -23,4 +23,4 @@ tensorboard torch==2.13.0 torchmetrics==1.0.3 torchrec==1.8.0 -transformers +transformers>=4.56 diff --git a/tzrec/models/genrec_causal_lm_model_test.py b/tzrec/models/genrec_causal_lm_model_test.py index 0231ace0..2c8eb709 100644 --- a/tzrec/models/genrec_causal_lm_model_test.py +++ b/tzrec/models/genrec_causal_lm_model_test.py @@ -37,10 +37,7 @@ from tzrec.utils.test_util import ( create_genrec_test_model, create_genrec_test_tokenizer, - flash_attn_unavailable, make_test_dir, - mark_ci_scope, - nv_gpu_unavailable, parameterized_name_func, ) @@ -316,15 +313,8 @@ def test_beam_config_uses_final_capped_capacity(self) -> None: num_return_sequences=5, ) - @unittest.skipIf(*nv_gpu_unavailable) - @unittest.skipIf(*flash_attn_unavailable) - @mark_ci_scope("gpu") def test_training_forward_builds_no_cache(self) -> None: - model, compiled_prompt = create_genrec_test_model( - self.test_dir, lm_parameter_dtype=GenRecModelConfig.BF16 - ) - device = torch.device("cuda") - model.to(device) + model, compiled_prompt = create_genrec_test_model(self.test_dir) batch = Batch() batch.additional_infos.update( PromptAssembler(compiled_prompt.prompt_plan, compiled_prompt.sid_space)( @@ -337,7 +327,6 @@ def test_training_forward_builds_no_cache(self) -> None: } ) ) - batch = batch.to(device) inner = model.lm.model.forward with mock.patch.object(model.lm.model, "forward", side_effect=inner) as spy: diff --git a/tzrec/models/genrec_model.py b/tzrec/models/genrec_model.py index c0c49838..082e5d63 100644 --- a/tzrec/models/genrec_model.py +++ b/tzrec/models/genrec_model.py @@ -18,8 +18,8 @@ HuggingFace weights. """ -import importlib.util import inspect +from importlib.util import find_spec from typing import Any, Dict, List, Optional, Sequence, Tuple import torch @@ -149,7 +149,7 @@ def init_backbone( """ impl = _ATTN_IMPL[attn_implementation] if impl == "flash_attention_2": - if importlib.util.find_spec("flash_attn") is None: + if find_spec("flash_attn") is None: raise ImportError( f"{type(self).__name__}: attn_implementation is " f"FLASH_ATTENTION_2 but the flash_attn wheel is not " diff --git a/tzrec/models/genrec_model_test.py b/tzrec/models/genrec_model_test.py index ee7ee470..a0bb1bdd 100644 --- a/tzrec/models/genrec_model_test.py +++ b/tzrec/models/genrec_model_test.py @@ -106,9 +106,6 @@ def _projected_batch(compiled_prompt) -> Batch: ) -@mark_ci_scope("gpu") -@unittest.skipIf(*nv_gpu_unavailable) -@unittest.skipIf(*flash_attn_unavailable) class BaseGenRecModelTest(unittest.TestCase): """Shared causal-LM behavior, reached through its concrete subclass.""" @@ -161,6 +158,8 @@ def test_builds_backbone_with_the_configured_kernel_and_dtype( mock.patch.object( AutoModelForCausalLM, "from_config", return_value=stand_in ) as from_config, + # the backbone is mocked, so the wheel probe has nothing to check + mock.patch("tzrec.models.genrec_model.find_spec", return_value=object()), ): model, _ = create_genrec_test_model( self.test_dir, @@ -231,10 +230,8 @@ def test_projected_slot_follows_a_narrow_lm_dtype(self, lm_parameter_dtype) -> N prompt="History : {{hist}} . Predict {{prof}} :", lm_parameter_dtype=lm_parameter_dtype, ) - device = torch.device("cuda") - init_parameters(model, device=device) - model.to(device) - batch = _projected_batch(compiled_prompt).to(device) + init_parameters(model, device=torch.device("cpu")) + batch = _projected_batch(compiled_prompt) embeds = model.build_input(batch) self.assertIs(embeds.dtype, _PARAM_DTYPE[lm_parameter_dtype]) @@ -254,12 +251,12 @@ def test_projected_slot_trains_with_fp32_masters_and_bf16_autocast(self) -> None feature_configs=[_hist(), _projected("prof", 8)], prompt="History : {{hist}} . Predict {{prof}} :", ) - device = torch.device("cuda") - init_parameters(model, device=device) - model.to(device) - batch = _projected_batch(compiled_prompt).to(device) + init_parameters(model, device=torch.device("cpu")) + batch = _projected_batch(compiled_prompt) - wrapper = TrainWrapper(model, device=device, mixed_precision="BF16") + wrapper = TrainWrapper( + model, device=torch.device("cpu"), mixed_precision="BF16" + ) loss, _ = wrapper(batch) self.assertTrue(bool(torch.isfinite(loss))) loss.backward() @@ -303,8 +300,6 @@ def test_model_resizes_to_target_vocab_size(self) -> None: self.assertGreater(rows, self.compiled_prompt.sid_space.band_hi[-1]) def test_loss_is_finite_and_backpropagates_into_the_backbone(self) -> None: - device = torch.device("cuda") - self.model.to(device) batch = _batch( self.compiled_prompt, { @@ -313,10 +308,9 @@ def test_loss_is_finite_and_backpropagates_into_the_backbone(self) -> None: "answer.values": torch.tensor(_ANSWER_CODES), "answer.lengths": torch.tensor([3]), }, - ).to(device) - with torch.autocast("cuda", dtype=torch.bfloat16): - predictions = self.model.predict(batch) - loss = self.model.loss(predictions, batch)["ce_loss"] + ) + predictions = self.model.predict(batch) + loss = self.model.loss(predictions, batch)["ce_loss"] self.assertTrue(bool(torch.isfinite(loss))) loss.backward() @@ -328,9 +322,6 @@ def test_training_forward_survives_fx_tracing(self) -> None: torch.fx.symbolic_trace(TrainWrapper(self.model)) -@mark_ci_scope("gpu") -@unittest.skipIf(*nv_gpu_unavailable) -@unittest.skipIf(*flash_attn_unavailable) class GenRecFrontEndTest(unittest.TestCase): """The served half of the model, under the same wrapper every export uses.""" @@ -438,8 +429,8 @@ def _packed_batch(compiled_prompt, hist_rows, answer_rows) -> Batch: ) -def _packed_flash_model( - test_dir: str, model_type: str, attn_implementation: int +def _packed_model( + test_dir: str, model_type: str, attn_implementation: int, lm_parameter_dtype: int ) -> Tuple[BaseModel, CompiledPrompt]: """Build a bf16 genrec model over a tiny ``model_type`` backbone. @@ -453,6 +444,7 @@ def _packed_flash_model( test_dir (str): scratch directory the backbone is written under. model_type (str): the hugging-face ``model_type`` of the backbone. attn_implementation (int): ``GenRecModelConfig.AttnImpl`` to build with. + lm_parameter_dtype (int): ``GenRecModelConfig.ParamDtype`` to build in. Returns: Tuple[BaseModel, CompiledPrompt]: the model and the prompt it was @@ -486,8 +478,7 @@ def _packed_flash_model( lm_config.hf_model_name_or_path = backbone lm_config.common.beam_widths.extend([2, 2, 2]) lm_config.common.num_return_sequences = 2 - # flash attention runs on fp16/bf16 only, and this arm carries no autocast - lm_config.common.lm_parameter_dtype = GenRecModelConfig.BF16 + lm_config.common.lm_parameter_dtype = lm_parameter_dtype lm_config.common.attn_implementation = attn_implementation with torch.random.fork_rng(devices=[]): torch.manual_seed(0) @@ -497,38 +488,29 @@ def _packed_flash_model( return model, compiled_prompt -@mark_ci_scope("gpu") -@unittest.skipIf(*nv_gpu_unavailable) -@unittest.skipIf(*flash_attn_unavailable) -class PackedFlashAttentionTest(unittest.TestCase): - """The packed varlen forward, against the same rows run one at a time.""" +class _PackedRowsCase: + """A packed batch must read exactly as its rows do one at a time. + + The batch is one concatenated stream with ``cu_seq_lens`` marking the + boundaries, so the failure this guards against is a row attending into the + row before it. Rewriting row 0 to different codes of the same width leaves + row 1 at the same packed offsets: its logits have to stay bit identical. + """ + + device = torch.device("cpu") + attn_implementation = GenRecModelConfig.SDPA + lm_parameter_dtype = GenRecModelConfig.FP32 def setUp(self) -> None: self.test_dir = make_test_dir() - @parameterized.expand( - [ - ["qwen2", GenRecModelConfig.SDPA], - ["qwen3", GenRecModelConfig.SDPA], - ["qwen2", GenRecModelConfig.FLASH_ATTENTION_2], - ["qwen3", GenRecModelConfig.FLASH_ATTENTION_2], - ], - name_func=parameterized_name_func, - ) - def test_packed_rows_match_solo_runs_and_backpropagate( - self, model_type: str, attn_implementation: int - ) -> None: - """A packed batch reads as the rows do alone, and nothing crosses rows. - - The batch is one concatenated stream with ``cu_seq_lens`` marking the - boundaries, so the failure this guards against is a row attending into - the row before it. Rewriting row 0 to different codes of the same width - leaves row 1 at the same packed offsets: its logits have to stay bit - identical. - """ - device = torch.device("cuda") - model, compiled_prompt = _packed_flash_model( - self.test_dir, model_type, attn_implementation + def _assert_packed_matches_solo(self, model_type: str) -> None: + device = self.device + model, compiled_prompt = _packed_model( + self.test_dir, + model_type, + self.attn_implementation, + self.lm_parameter_dtype, ) model.to(device) model.eval() @@ -574,5 +556,29 @@ def test_packed_rows_match_solo_runs_and_backpropagate( self.assertGreater(float(grad.abs().sum()), 0) -if __name__ == "__main__": - unittest.main() +class PackedSdpaAttentionTest(_PackedRowsCase, unittest.TestCase): + """The packed forward under sdpa, which needs no GPU and no wheel.""" + + @parameterized.expand([["qwen2"], ["qwen3"]], name_func=parameterized_name_func) + def test_packed_rows_match_solo_runs_and_backpropagate( + self, model_type: str + ) -> None: + self._assert_packed_matches_solo(model_type) + + +@mark_ci_scope("gpu") +@unittest.skipIf(*nv_gpu_unavailable) +@unittest.skipIf(*flash_attn_unavailable) +class PackedFlashAttentionTest(_PackedRowsCase, unittest.TestCase): + """The same rows through the varlen flash kernel.""" + + device = torch.device("cuda") + attn_implementation = GenRecModelConfig.FLASH_ATTENTION_2 + # the flash kernel takes fp16/bf16 only, and this arm carries no autocast + lm_parameter_dtype = GenRecModelConfig.BF16 + + @parameterized.expand([["qwen2"], ["qwen3"]], name_func=parameterized_name_func) + def test_packed_rows_match_solo_runs_and_backpropagate( + self, model_type: str + ) -> None: + self._assert_packed_matches_solo(model_type) diff --git a/tzrec/tests/genrec_integration_test.py b/tzrec/tests/genrec_integration_test.py index 5c4c7494..6f63204a 100644 --- a/tzrec/tests/genrec_integration_test.py +++ b/tzrec/tests/genrec_integration_test.py @@ -36,11 +36,9 @@ from tzrec.utils.test_util import ( create_genrec_test_tokenizer, create_tiny_causal_lm, - flash_attn_unavailable, gpu_unavailable, make_test_dir, mark_ci_scope, - nv_gpu_unavailable, ) _MOCK_CONFIG = "tzrec/tests/configs/genrec_causal_lm_model_mock.config" @@ -122,9 +120,6 @@ def _request(self, columns, rows: int = 4): out[column + ".lengths"] = torch.tensor([len(row) for row in lists]) return out - @unittest.skipIf(*nv_gpu_unavailable) - @unittest.skipIf(*flash_attn_unavailable) - @mark_ci_scope("gpu") def test_genrec_train_eval_export(self): trained = self._train_eval_export() export_dir = os.path.join(self.test_dir, "export") From 6df33ca06eaa661b1476034397a83d9fa76277d3 Mon Sep 17 00:00:00 2001 From: root <597191244@qq.com> Date: Thu, 17 Sep 2026 09:55:37 +0000 Subject: [PATCH 07/15] [fix] read the attention kernel from our own config, not HF's private one The fp32 check asked self.lm.config._attn_implementation, which is a private property over _attn_implementation_internal and the only thing transformers exposes -- on a dependency with no upper bound. A rename there would make the comparison silently False and the check would stop firing. init_backbone already knows which kernel it was asked for, and since FLASH_ATTENTION_2 raises rather than falling back when the wheel is absent, what was asked for is what was built. Record it and read that instead. Co-Authored-By: Claude Opus 5 (1M context) --- tzrec/models/genrec_causal_lm_model.py | 4 ++-- tzrec/models/genrec_causal_lm_model_test.py | 15 ++++++++------- tzrec/models/genrec_model.py | 4 ++++ 3 files changed, 14 insertions(+), 9 deletions(-) diff --git a/tzrec/models/genrec_causal_lm_model.py b/tzrec/models/genrec_causal_lm_model.py index a5e0d8d7..6a1c2e98 100644 --- a/tzrec/models/genrec_causal_lm_model.py +++ b/tzrec/models/genrec_causal_lm_model.py @@ -144,9 +144,9 @@ def _forward( applies lands on the pairs the window was sized for. """ # only the flash kernel rejects fp32, and fp32 masters are fine under - # autocast, so neither the config nor the dtype alone can tell + # autocast, so neither the kernel nor the dtype alone can tell if ( - self.lm.config._attn_implementation == "flash_attention_2" + self._attn_impl == "flash_attention_2" and embeds.dtype == torch.float32 and not torch.is_autocast_enabled("cuda") ): diff --git a/tzrec/models/genrec_causal_lm_model_test.py b/tzrec/models/genrec_causal_lm_model_test.py index 2c8eb709..8441a681 100644 --- a/tzrec/models/genrec_causal_lm_model_test.py +++ b/tzrec/models/genrec_causal_lm_model_test.py @@ -104,10 +104,9 @@ def test_packs_rows_of_different_lengths(self) -> None: class _CapturingLM(nn.Module): - def __init__(self, attn_implementation: str = "sdpa") -> None: + def __init__(self) -> None: super().__init__() self.kwargs = {} - self.config = SimpleNamespace(_attn_implementation=attn_implementation) def forward(self, **kwargs): self.kwargs = kwargs @@ -123,9 +122,7 @@ def forward(self, **kwargs): class _DifferentiableLM(nn.Module): def __init__(self, hidden_size: int, vocab_size: int) -> None: super().__init__() - self.config = SimpleNamespace( - vocab_size=vocab_size, _attn_implementation="sdpa" - ) + self.config = SimpleNamespace(vocab_size=vocab_size) self.lm_head = nn.Linear(hidden_size, vocab_size, bias=False) self.loss_function = ForCausalLMLoss @@ -139,6 +136,7 @@ def test_passes_varlen_metadata_and_builds_per_row_labels(self) -> None: model = GenRecCausalLMModel.__new__(GenRecCausalLMModel) nn.Module.__init__(model) model.lm = _CapturingLM() + model._attn_impl = "sdpa" model._ignore_index = -7 model._prompt = SimpleNamespace( prompt_plan=SimpleNamespace(logits_suffix_len=4) @@ -180,6 +178,7 @@ def test_a_row_shorter_than_the_window_is_rejected(self) -> None: model = GenRecCausalLMModel.__new__(GenRecCausalLMModel) nn.Module.__init__(model) model.lm = _CapturingLM() + model._attn_impl = "sdpa" model._ignore_index = -7 model._prompt = SimpleNamespace( prompt_plan=SimpleNamespace(logits_suffix_len=4) @@ -200,7 +199,8 @@ def test_fp32_without_autocast_is_rejected_on_the_flash_path(self) -> None: """The flash kernel takes bf16/fp16 only; sdpa is happy in fp32.""" model = GenRecCausalLMModel.__new__(GenRecCausalLMModel) nn.Module.__init__(model) - model.lm = _CapturingLM("flash_attention_2") + model.lm = _CapturingLM() + model._attn_impl = "flash_attention_2" model._ignore_index = -7 model._prompt = SimpleNamespace( prompt_plan=SimpleNamespace(logits_suffix_len=4) @@ -217,7 +217,7 @@ def test_fp32_without_autocast_is_rejected_on_the_flash_path(self) -> None: with self.assertRaisesRegex(ValueError, "needs bf16 or fp16"): model._forward(torch.zeros(12, 6, dtype=torch.float32), batch) - model.lm = _CapturingLM("sdpa") + model._attn_impl = "sdpa" logits, _ = model._forward(torch.zeros(12, 6, dtype=torch.float32), batch) self.assertEqual(logits.shape, (2, 4, 5)) @@ -225,6 +225,7 @@ def test_loss_and_gradients_cover_only_valid_response_pairs(self) -> None: model = GenRecCausalLMModel.__new__(GenRecCausalLMModel) nn.Module.__init__(model) model.lm = _DifferentiableLM(hidden_size=6, vocab_size=32) + model._attn_impl = "sdpa" model._ignore_index = -7 model._prompt = SimpleNamespace( prompt_plan=SimpleNamespace(logits_suffix_len=4) diff --git a/tzrec/models/genrec_model.py b/tzrec/models/genrec_model.py index 082e5d63..737ead25 100644 --- a/tzrec/models/genrec_model.py +++ b/tzrec/models/genrec_model.py @@ -104,6 +104,7 @@ def __init__( self._ignore_index = int(cfg.common.ignore_index) self.lm: nn.Module + self._attn_impl: str self.init_backbone( cfg.hf_model_name_or_path, cfg.common.lm_parameter_dtype, @@ -148,6 +149,9 @@ def init_backbone( absent, so the packed forward has no kernel to run on. """ impl = _ATTN_IMPL[attn_implementation] + # the forward needs it too, and transformers exposes only the private + # config._attn_implementation + self._attn_impl = impl if impl == "flash_attention_2": if find_spec("flash_attn") is None: raise ImportError( From 3d82728a4606dd6c32fbc0a61ae22745b6900e22 Mon Sep 17 00:00:00 2001 From: root <597191244@qq.com> Date: Fri, 18 Sep 2026 03:33:12 +0000 Subject: [PATCH 08/15] [fix] give each attention kernel the layout it can actually serve _forward packed every row into one sequence regardless of the configured kernel, but only the flash varlen kernel reads cu_seq_lens. On sdpa the boundaries come back as a dense (1, 1, T, T) mask, so the batch costs (sum L)^2 where the left-padded path it replaced cost sum(B * L^2) -- B times more attention work than before this branch, on the default kernel. Measured at 8 rows of 128: packed 1,048,576 score-pairs and a 1 MiB mask, padded 131,072 and no mask at all. So branch: FLASH_ATTENTION_2 keeps the packed forward, every other kernel gets _left_pad_packed_inputs, which _generate already used. The supervised window is the same absolute positions either way, so the labels are built once and only the logits split. Verified the two branches agree to 5.96e-08 with identical labels. Drop both flash guards. The fp32 check and the wheel probe each replaced an error transformers already raises -- "FlashAttention only support fp16 and bf16 data type" and "the package for FlashAttention2 doesn't seem to be installed" -- so they bought a nicer message and cost a per-step check and an import probe. lm_parameter_dtype is left alone rather than silently narrowed to BF16 when flash is configured. Name the enum after the Kernel it mirrors: AttnImpl was the only *Impl of the thirteen in tzrec/protos, where eight are bare nouns. AttnKernel, field attn_kernel, and _attn_kernel now holds that enum instead of an HF string. Alongside, in the forward: keep the window 2-D rather than flattening and reshaping back, fold `columns` into `suffix_offsets`, drop a .contiguous() that cannot copy, and replace a process-wide enable_cudnn_sdp(False) -- which disabled the backend for every other module in the process -- with a scoped sdpa_kernel exclusion at the call that needs it. And in the tests: move the row-isolation suite beside the forward it exercises; delete three assertions that only pinned a mock, the file's own stub, or a use_cache kwarg asserted next door; give create_tiny_causal_lm a model_type, split create_genrec_test_prompt out of create_genrec_test_model, and collapse _packed_model, _compiled_prompt and six copied arrange blocks onto them. Co-Authored-By: Claude Opus 5 (1M context) --- tzrec/models/genrec_causal_lm_model.py | 115 ++++--- tzrec/models/genrec_causal_lm_model_test.py | 315 +++++++++++--------- tzrec/models/genrec_model.py | 38 +-- tzrec/models/genrec_model_test.py | 216 ++------------ tzrec/modules/dynamic_beam_test.py | 27 +- tzrec/protos/models/genrec_model.proto | 10 +- tzrec/utils/test_util.py | 121 ++++++-- 7 files changed, 373 insertions(+), 469 deletions(-) diff --git a/tzrec/models/genrec_causal_lm_model.py b/tzrec/models/genrec_causal_lm_model.py index 6a1c2e98..d24edb9c 100644 --- a/tzrec/models/genrec_causal_lm_model.py +++ b/tzrec/models/genrec_causal_lm_model.py @@ -23,6 +23,7 @@ from typing import Any, Dict, List, Optional, Tuple, cast import torch +from torch.nn.attention import SDPBackend, sdpa_kernel from tzrec.datasets.utils import Batch from tzrec.features.feature import BaseFeature @@ -143,69 +144,103 @@ def _forward( Logits and labels over the same window, so the shift ``loss`` applies lands on the pairs the window was sized for. """ - # only the flash kernel rejects fp32, and fp32 masters are fine under - # autocast, so neither the kernel nor the dtype alone can tell - if ( - self._attn_impl == "flash_attention_2" - and embeds.dtype == torch.float32 - and not torch.is_autocast_enabled("cuda") - ): - raise ValueError( - f"{type(self).__name__}: flash_attention_2 needs bf16 or fp16 " - f"activations, but the backbone ran in fp32 with no autocast " - f"active. Set train_config.mixed_precision to BF16 or FP16 to " - f"keep fp32 master weights, or set " - f"the model's common.lm_parameter_dtype to BF16 or FP16 to " - f"narrow the parameters themselves." - ) - infos = batch.additional_infos cu_seqlens = infos[CU_SEQLENS] - starts = cu_seqlens[:-1] lengths = torch.diff(cu_seqlens) - row_starts = torch.repeat_interleave( - starts, lengths, output_size=embeds.shape[0] - ) - position_ids = ( - torch.arange(embeds.shape[0], device=embeds.device) - row_starts - ).unsqueeze(0) - suffix = cast(int, self._prompt.prompt_plan.logits_suffix_len) # a row shorter than the window reads the row before it if bool((lengths < suffix).any()): raise ValueError( f"{type(self).__name__}: every assembled sample must be at " - f"least logits_suffix_len ({suffix}) tokens long; a sample " - f"whose prompt body is empty is not. Drop the rows whose " - f"prompt features are all empty, or give the template " - f"static text." + f"least logits_suffix_len ({suffix}) tokens long. Drop the " + f"rows whose prompt features are all empty, or give the " + f"template static text." ) suffix_offsets = torch.arange(-suffix, 0, device=embeds.device) - keep_indices = (cu_seqlens[1:, None] + suffix_offsets).reshape(-1) - # varlen flash-attention wants int32 boundaries; the assembler - # already emits them, so this is normally a no-op - flash_cu_seqlens = cu_seqlens.to(dtype=torch.int32).contiguous() + # (rows, suffix): the window each row is supervised over, the same + # absolute positions whichever layout the kernel is given + keep = cu_seqlens[1:, None] + suffix_offsets + labels = infos[INPUT_IDS][keep].masked_fill( + suffix_offsets[None, :] < -infos[RESPONSE_LENGTHS][:, None], + self._ignore_index, + ) + # CUDNN_ATTENTION is eligible for the packed mask and returns NaN + # losses on it; exclude it here rather than disabling it process-wide + with sdpa_kernel( + [ + SDPBackend.FLASH_ATTENTION, + SDPBackend.EFFICIENT_ATTENTION, + SDPBackend.MATH, + ] + ): + if self._attn_kernel == GenRecModelConfig.FLASH_ATTENTION_2: + logits = self._packed_logits(embeds, batch, keep) + else: + logits = self._padded_logits(embeds, batch, suffix) + return logits, labels + + def _packed_logits( + self, embeds: torch.Tensor, batch: Batch, keep: torch.Tensor + ) -> torch.Tensor: + """Score every row in one varlen call, with no padding at all. + + Only the flash varlen kernel reads ``cu_seq_lens``; every other kernel + would rebuild the row boundaries as a dense ``(1, 1, T, T)`` mask and + pay ``(sum L)^2`` instead of ``sum L^2``. + + Args: + embeds: the assembled prompt embeddings, packed. + batch: carries the row boundaries and the collator's width. + keep: ``(rows, suffix)`` absolute positions of the response window. + + Returns: + ``(rows, suffix, vocab)`` logits over the response window. + """ + infos = batch.additional_infos + cu_seqlens = infos[CU_SEQLENS] + row_starts = torch.repeat_interleave( + cu_seqlens[:-1], torch.diff(cu_seqlens), output_size=embeds.shape[0] + ) + position_ids = ( + torch.arange(embeds.shape[0], device=embeds.device) - row_starts + ).unsqueeze(0) + # varlen flash-attention wants int32 boundaries + flash_cu_seqlens = cu_seqlens.to(torch.int32) max_seqlen = int(infos[MAX_SEQLEN]) outputs = self.lm( inputs_embeds=embeds.unsqueeze(0), attention_mask=None, position_ids=position_ids, use_cache=False, - logits_to_keep=keep_indices, + logits_to_keep=keep.reshape(-1), cu_seq_lens_q=flash_cu_seqlens, cu_seq_lens_k=flash_cu_seqlens, max_length_q=max_seqlen, max_length_k=max_seqlen, ) - logits = outputs.logits.reshape(lengths.numel(), suffix, -1) + return outputs.logits.reshape(keep.shape[0], keep.shape[1], -1) - window_ids = infos[INPUT_IDS][keep_indices].reshape(lengths.numel(), suffix) - columns = torch.arange(suffix, device=embeds.device) - labels = window_ids.masked_fill( - columns[None, :] < suffix - infos[RESPONSE_LENGTHS][:, None], - self._ignore_index, + def _padded_logits( + self, embeds: torch.Tensor, batch: Batch, suffix: int + ) -> torch.Tensor: + """Score left-padded rows, which every kernel can serve. + + Args: + embeds: the assembled prompt embeddings, packed. + batch: carries the row boundaries and the collator's width. + suffix: width of the supervised window. + + Returns: + ``(rows, suffix, vocab)`` logits over the response window. + """ + padded, mask = self._left_pad_packed_inputs(embeds, batch) + outputs = self.lm( + inputs_embeds=padded, + attention_mask=mask, + use_cache=False, + logits_to_keep=suffix, ) - return logits, labels + return outputs.logits def _generate(self, embeds: torch.Tensor, batch: Batch) -> torch.Tensor: """Beam-search the SID answer. diff --git a/tzrec/models/genrec_causal_lm_model_test.py b/tzrec/models/genrec_causal_lm_model_test.py index 8441a681..135b77ae 100644 --- a/tzrec/models/genrec_causal_lm_model_test.py +++ b/tzrec/models/genrec_causal_lm_model_test.py @@ -9,10 +9,9 @@ # See the License for the specific language governing permissions and # limitations under the License. -import os import unittest from types import SimpleNamespace -from unittest import mock +from typing import Optional, Sequence import torch from parameterized import parameterized @@ -20,7 +19,6 @@ from transformers.loss.loss_utils import ForCausalLMLoss from tzrec.datasets.utils import Batch -from tzrec.features.feature import FgMode, create_features from tzrec.models.genrec_causal_lm_model import GenRecCausalLMModel from tzrec.prompt.assembler import ( CU_SEQLENS, @@ -29,51 +27,18 @@ RESPONSE_LENGTHS, PromptAssembler, ) -from tzrec.prompt.compile import compile_prompt -from tzrec.prompt.types import CompiledPrompt -from tzrec.protos import feature_pb2 from tzrec.protos.models.genrec_model_pb2 import GenRecModelConfig -from tzrec.protos.prompt_pb2 import PromptConfig from tzrec.utils.test_util import ( create_genrec_test_model, - create_genrec_test_tokenizer, + create_genrec_test_prompt, + flash_attn_unavailable, make_test_dir, + mark_ci_scope, + nv_gpu_unavailable, parameterized_name_func, ) -def _compiled_prompt(test_dir: str) -> CompiledPrompt: - """Compile the genrec test prompt without building a backbone. - - ``create_genrec_test_model`` also constructs the LM. The decode-schedule - tests only read the SID space, so they compile the prompt on its own and - keep ``_read_beam_config`` reachable without an HF backbone. - - Args: - test_dir (str): scratch directory for the tokenizer. - - Returns: - CompiledPrompt: the prompt over the ``(4, 4, 4)`` codebook. - """ - features = create_features( - [ - feature_pb2.FeatureConfig( - sequence_raw_feature=feature_pb2.RawFeature( - feature_name="hist", expression="user:hist" - ) - ) - ], - fg_mode=FgMode.FG_NONE, - ) - prompt_config = PromptConfig( - tokenizer_path=create_genrec_test_tokenizer(os.path.join(test_dir, "tok.json")), - prompt="History : {{hist}} . Predict :", - response="{{answer}}", - ) - prompt_config.sid_space.codebook.extend([4, 4, 4]) - return compile_prompt(prompt_config, features, ["answer"]) - - class LeftPadPackedInputsTest(unittest.TestCase): """The one adapter where padding lives.""" @@ -110,13 +75,11 @@ def __init__(self) -> None: def forward(self, **kwargs): self.kwargs = kwargs + embeds = kwargs["inputs_embeds"] count = kwargs["logits_to_keep"].numel() - logits = torch.arange( - count * 5, - dtype=kwargs["inputs_embeds"].dtype, - device=kwargs["inputs_embeds"].device, - ).reshape(1, count, 5) - return SimpleNamespace(logits=logits) + return SimpleNamespace( + logits=torch.zeros(1, count, 5, dtype=embeds.dtype, device=embeds.device) + ) class _DifferentiableLM(nn.Module): @@ -131,26 +94,51 @@ def forward(self, **kwargs): return SimpleNamespace(logits=self.lm_head(hidden)) +def _stub_model( + lm: Optional[nn.Module] = None, + attn_kernel: int = GenRecModelConfig.FLASH_ATTENTION_2, +) -> GenRecCausalLMModel: + """A model with just the attributes ``_forward`` reads, over a stub LM. + + Defaults to the varlen kernel: these cases are about the packed layout, + which is the branch only that kernel takes. + """ + model = GenRecCausalLMModel.__new__(GenRecCausalLMModel) + nn.Module.__init__(model) + model.lm = _CapturingLM() if lm is None else lm + model._attn_kernel = attn_kernel + model._ignore_index = -7 + model._prompt = SimpleNamespace(prompt_plan=SimpleNamespace(logits_suffix_len=4)) + return model + + +def _varlen_batch( + cu_seqlens: Sequence[int] = (0, 5, 12), + max_seqlen: int = 7, + response_lengths: Sequence[int] = (3, 2), + first_input_id: int = 100, + cu_dtype: torch.dtype = torch.int64, +) -> Batch: + """The four varlen keys ``_forward`` reads, over ``cu_seqlens[-1]`` tokens. + + ``PromptAssembler`` emits ``cu_seqlens`` as int32; the int64 default is + what makes ``_forward``'s cast to int32 observable. + """ + return Batch( + additional_infos={ + CU_SEQLENS: torch.tensor(cu_seqlens, dtype=cu_dtype), + INPUT_IDS: torch.arange(first_input_id, first_input_id + cu_seqlens[-1]), + MAX_SEQLEN: torch.tensor(max_seqlen), + RESPONSE_LENGTHS: torch.tensor(response_lengths), + } + ) + + class PackedForwardTest(unittest.TestCase): def test_passes_varlen_metadata_and_builds_per_row_labels(self) -> None: - model = GenRecCausalLMModel.__new__(GenRecCausalLMModel) - nn.Module.__init__(model) - model.lm = _CapturingLM() - model._attn_impl = "sdpa" - model._ignore_index = -7 - model._prompt = SimpleNamespace( - prompt_plan=SimpleNamespace(logits_suffix_len=4) - ) + model = _stub_model() embeds = torch.arange(72, dtype=torch.float32).reshape(12, 6) - input_ids = torch.arange(100, 112) - batch = Batch( - additional_infos={ - CU_SEQLENS: torch.tensor([0, 5, 12]), - INPUT_IDS: input_ids, - MAX_SEQLEN: torch.tensor(7), - RESPONSE_LENGTHS: torch.tensor([3, 2]), - } - ) + batch = _varlen_batch() logits, labels = model._forward(embeds, batch) @@ -175,73 +163,23 @@ def test_passes_varlen_metadata_and_builds_per_row_labels(self) -> None: def test_a_row_shorter_than_the_window_is_rejected(self) -> None: """A row under ``logits_suffix_len`` would index into its neighbour.""" - model = GenRecCausalLMModel.__new__(GenRecCausalLMModel) - nn.Module.__init__(model) - model.lm = _CapturingLM() - model._attn_impl = "sdpa" - model._ignore_index = -7 - model._prompt = SimpleNamespace( - prompt_plan=SimpleNamespace(logits_suffix_len=4) - ) - batch = Batch( - additional_infos={ - CU_SEQLENS: torch.tensor([0, 3, 12], dtype=torch.int32), - INPUT_IDS: torch.arange(100, 112), - MAX_SEQLEN: torch.tensor(9), - RESPONSE_LENGTHS: torch.tensor([3, 3]), - } + model = _stub_model() + batch = _varlen_batch( + cu_seqlens=(0, 3, 12), + max_seqlen=9, + response_lengths=(3, 3), + cu_dtype=torch.int32, ) with self.assertRaisesRegex(ValueError, "at least logits_suffix_len"): model._forward(torch.zeros(12, 6), batch) - def test_fp32_without_autocast_is_rejected_on_the_flash_path(self) -> None: - """The flash kernel takes bf16/fp16 only; sdpa is happy in fp32.""" - model = GenRecCausalLMModel.__new__(GenRecCausalLMModel) - nn.Module.__init__(model) - model.lm = _CapturingLM() - model._attn_impl = "flash_attention_2" - model._ignore_index = -7 - model._prompt = SimpleNamespace( - prompt_plan=SimpleNamespace(logits_suffix_len=4) - ) - batch = Batch( - additional_infos={ - CU_SEQLENS: torch.tensor([0, 5, 12], dtype=torch.int32), - INPUT_IDS: torch.arange(100, 112), - MAX_SEQLEN: torch.tensor(7), - RESPONSE_LENGTHS: torch.tensor([3, 2]), - } - ) - - with self.assertRaisesRegex(ValueError, "needs bf16 or fp16"): - model._forward(torch.zeros(12, 6, dtype=torch.float32), batch) - - model._attn_impl = "sdpa" - logits, _ = model._forward(torch.zeros(12, 6, dtype=torch.float32), batch) - self.assertEqual(logits.shape, (2, 4, 5)) - def test_loss_and_gradients_cover_only_valid_response_pairs(self) -> None: - model = GenRecCausalLMModel.__new__(GenRecCausalLMModel) - nn.Module.__init__(model) - model.lm = _DifferentiableLM(hidden_size=6, vocab_size=32) - model._attn_impl = "sdpa" - model._ignore_index = -7 - model._prompt = SimpleNamespace( - prompt_plan=SimpleNamespace(logits_suffix_len=4) - ) + model = _stub_model(_DifferentiableLM(hidden_size=6, vocab_size=32)) embeds = torch.randn( 12, 6, generator=torch.Generator().manual_seed(1), requires_grad=True ) - input_ids = torch.arange(4, 16) - batch = Batch( - additional_infos={ - CU_SEQLENS: torch.tensor([0, 5, 12]), - INPUT_IDS: input_ids, - MAX_SEQLEN: torch.tensor(7), - RESPONSE_LENGTHS: torch.tensor([3, 2]), - } - ) + batch = _varlen_batch(first_input_id=4) logits, labels = model._forward(embeds, batch) loss = model.loss({"logits": logits, "labels": labels}, batch)["ce_loss"] @@ -261,10 +199,6 @@ def test_loss_and_gradients_cover_only_valid_response_pairs(self) -> None: torch.testing.assert_close( grad_norms[unsupervised], torch.zeros(7), atol=0, rtol=0 ) - weight_grad = model.lm.lm_head.weight.grad - self.assertIsNotNone(weight_grad) - self.assertTrue(bool(torch.isfinite(weight_grad).all())) - self.assertGreater(float(weight_grad.abs().sum()), 0) class GenRecCausalLMModelTest(unittest.TestCase): @@ -272,7 +206,7 @@ class GenRecCausalLMModelTest(unittest.TestCase): def setUp(self) -> None: self.test_dir = make_test_dir() - self.compiled_prompt = _compiled_prompt(self.test_dir) + self.compiled_prompt, _ = create_genrec_test_prompt(self.test_dir) def _beam_model( self, beam_widths=(2, 2, 2), num_return_sequences=2 @@ -314,26 +248,123 @@ def test_beam_config_uses_final_capped_capacity(self) -> None: num_return_sequences=5, ) - def test_training_forward_builds_no_cache(self) -> None: - model, compiled_prompt = create_genrec_test_model(self.test_dir) - batch = Batch() - batch.additional_infos.update( - PromptAssembler(compiled_prompt.prompt_plan, compiled_prompt.sid_space)( - { - # offset SID codes for the (4, 4, 4) codebook - "hist.values": torch.tensor([0, 5, 10]), - "hist.lengths": torch.tensor([3]), - "answer.values": torch.tensor([1, 6, 11]), - "answer.lengths": torch.tensor([3]), - } + +# offset SID codes for the (4, 4, 4) codebook: level_offsets[l] + code +_HIST_CODES = [0, 5, 10] +_LONG_HIST_CODES = [0, 5, 10, 3, 4, 9] +_ANSWER_CODES = [1, 6, 11] + + +# a second answer, and a rewrite of _HIST_CODES that keeps its width +_OTHER_ANSWER_CODES = [2, 7, 8] +_REWRITTEN_HIST_CODES = [3, 7, 11] + + +def _packed_batch(compiled_prompt, hist_rows, answer_rows) -> Batch: + """Several rows in one batch, packed the way the collator packs them.""" + batch = Batch() + batch.additional_infos.update( + PromptAssembler(compiled_prompt.prompt_plan, compiled_prompt.sid_space)( + { + "hist.values": torch.tensor( + [code for row in hist_rows for code in row] + ), + "hist.lengths": torch.tensor([len(row) for row in hist_rows]), + "answer.values": torch.tensor( + [code for row in answer_rows for code in row] + ), + "answer.lengths": torch.tensor([len(row) for row in answer_rows]), + } + ) + ) + return batch + + +class _RowIsolationCase: + """Every row must read exactly as it does alone, whatever the layout. + + The failure this guards against is a row attending into the row before it, + which the two kernels avoid by different means: the varlen kernel is handed + ``cu_seq_lens``, every other kernel is handed left-padded rows. Rewriting + row 0 to different codes of the same width leaves row 1 at the same + offsets, so its logits have to stay bit identical either way. + """ + + device = torch.device("cpu") + attn_kernel = GenRecModelConfig.SDPA + lm_parameter_dtype = GenRecModelConfig.FP32 + + def setUp(self) -> None: + self.test_dir = make_test_dir() + + @parameterized.expand([["qwen2"], ["qwen3"]], name_func=parameterized_name_func) + def test_rows_match_solo_runs_and_backpropagate(self, model_type: str) -> None: + device = self.device + model, compiled_prompt = create_genrec_test_model( + self.test_dir, + model_type=model_type, + attn_kernel=self.attn_kernel, + lm_parameter_dtype=self.lm_parameter_dtype, + init_seed=0, + ) + model.to(device) + model.eval() + hist_rows = [_HIST_CODES, _LONG_HIST_CODES] + answer_rows = [_ANSWER_CODES, _OTHER_ANSWER_CODES] + packed_batch = _packed_batch(compiled_prompt, hist_rows, answer_rows).to(device) + + packed = model.predict(packed_batch) + with torch.no_grad(): + solos = [ + model.predict( + _packed_batch(compiled_prompt, [hist], [answer]).to(device) + ) + for hist, answer in zip(hist_rows, answer_rows) + ] + changed = model.predict( + _packed_batch( + compiled_prompt, + [_REWRITTEN_HIST_CODES, hist_rows[1]], + answer_rows, + ).to(device) ) + + torch.testing.assert_close( + packed["logits"], + torch.cat([result["logits"] for result in solos]), + atol=1e-2, + rtol=1e-2, + ) + torch.testing.assert_close( + packed["labels"], torch.cat([result["labels"] for result in solos]) + ) + torch.testing.assert_close( + packed["logits"][1], changed["logits"][1], atol=0, rtol=0 ) - inner = model.lm.model.forward - with mock.patch.object(model.lm.model, "forward", side_effect=inner) as spy: - model.predict(batch) + loss = model.loss(packed, packed_batch)["ce_loss"] + self.assertTrue(bool(torch.isfinite(loss))) + loss.backward() + grad = model.lm.get_input_embeddings().weight.grad + self.assertIsNotNone(grad) + self.assertTrue(bool(torch.isfinite(grad).all())) + self.assertGreater(float(grad.abs().sum()), 0) + + +class SdpaRowIsolationTest(_RowIsolationCase, unittest.TestCase): + """The left-padded layout sdpa is given, which needs no GPU and no wheel.""" + + +@mark_ci_scope("gpu") +@unittest.skipIf(*nv_gpu_unavailable) +@unittest.skipIf(*flash_attn_unavailable) +class FlashAttentionRowIsolationTest(_RowIsolationCase, unittest.TestCase): + """The same rows packed into one stream, through the varlen flash kernel.""" - self.assertIs(spy.call_args.kwargs["use_cache"], False) + device = torch.device("cuda") + attn_kernel = GenRecModelConfig.FLASH_ATTENTION_2 + # the flash kernel takes fp16/bf16 only, and this arm carries no autocast + lm_parameter_dtype = GenRecModelConfig.BF16 if __name__ == "__main__": diff --git a/tzrec/models/genrec_model.py b/tzrec/models/genrec_model.py index 737ead25..4f6e5c9c 100644 --- a/tzrec/models/genrec_model.py +++ b/tzrec/models/genrec_model.py @@ -19,7 +19,6 @@ """ import inspect -from importlib.util import find_spec from typing import Any, Dict, List, Optional, Sequence, Tuple import torch @@ -54,7 +53,7 @@ GenRecModelConfig.FP16: torch.float16, } -_ATTN_IMPL: Dict[int, str] = { +_ATTN_KERNEL: Dict[int, str] = { GenRecModelConfig.SDPA: "sdpa", GenRecModelConfig.FLASH_ATTENTION_2: "flash_attention_2", } @@ -104,11 +103,11 @@ def __init__( self._ignore_index = int(cfg.common.ignore_index) self.lm: nn.Module - self._attn_impl: str + self._attn_kernel: int self.init_backbone( cfg.hf_model_name_or_path, cfg.common.lm_parameter_dtype, - cfg.common.attn_implementation, + cfg.common.attn_kernel, ) # Every run replaces this initialization from pretrained or DCP weights. self.lm.resize_token_embeddings( @@ -134,7 +133,7 @@ def init_backbone( self, hf_model_name_or_path: str, lm_parameter_dtype: int, - attn_implementation: int = GenRecModelConfig.SDPA, + attn_kernel: int, ) -> None: """Assign ``self.lm`` from config, so HF weights load only on cold start. @@ -142,34 +141,15 @@ def init_backbone( hf_model_name_or_path: hub id or local directory naming the architecture and cold-start weights. lm_parameter_dtype: dtype of the LM parameters. - attn_implementation: attention kernel to build the backbone with. - - Raises: - ImportError: FLASH_ATTENTION_2 is configured but the wheel is - absent, so the packed forward has no kernel to run on. + attn_kernel: attention kernel to build the backbone with. """ - impl = _ATTN_IMPL[attn_implementation] - # the forward needs it too, and transformers exposes only the private - # config._attn_implementation - self._attn_impl = impl - if impl == "flash_attention_2": - if find_spec("flash_attn") is None: - raise ImportError( - f"{type(self).__name__}: attn_implementation is " - f"FLASH_ATTENTION_2 but the flash_attn wheel is not " - f"installed. Install it from " - f"https://tzrec.oss-accelerate.aliyuncs.com/third_party/" - f"flash_attn/${{DEVICE}}/ (cu126/cu129/cu130), or set " - f"attn_implementation to SDPA." - ) - elif torch.cuda.is_available(): - # the packed mask makes CUDNN_ATTENTION eligible, and it returns - # NaN losses on this shape; the other sdpa backends are fine - torch.backends.cuda.enable_cudnn_sdp(False) + # the forward picks its layout from this: only the varlen kernel reads + # cu_seq_lens, so every other kernel has to be handed padded rows + self._attn_kernel = attn_kernel config = AutoConfig.from_pretrained(hf_model_name_or_path) self.lm = AutoModelForCausalLM.from_config( config, - attn_implementation=impl, + attn_implementation=_ATTN_KERNEL[attn_kernel], torch_dtype=_PARAM_DTYPE[lm_parameter_dtype], ) self._check_backbone_interfaces(hf_model_name_or_path) diff --git a/tzrec/models/genrec_model_test.py b/tzrec/models/genrec_model_test.py index a0bb1bdd..b4e688c8 100644 --- a/tzrec/models/genrec_model_test.py +++ b/tzrec/models/genrec_model_test.py @@ -18,10 +18,9 @@ import torch.fx from parameterized import parameterized from torchrec import KeyedJaggedTensor -from transformers import AutoConfig, AutoModelForCausalLM +from transformers import AutoModelForCausalLM from tzrec.datasets.utils import BASE_DATA_GROUP, Batch -from tzrec.features.feature import FgMode, create_features from tzrec.main import _create_model from tzrec.models.genrec_model import ( _PARAM_DTYPE, @@ -36,22 +35,16 @@ INPUT_IDS, PromptAssembler, ) -from tzrec.prompt.compile import compile_prompt from tzrec.prompt.hole_keys import HOLE_KEYS, HoleKeyBuilder -from tzrec.prompt.types import CompiledPrompt from tzrec.protos import feature_pb2 from tzrec.protos.model_pb2 import ModelConfig from tzrec.protos.models.genrec_model_pb2 import GenRecModelConfig -from tzrec.protos.prompt_pb2 import PromptConfig, PromptSlot +from tzrec.protos.prompt_pb2 import PromptSlot from tzrec.utils.fx_util import symbolic_trace from tzrec.utils.state_dict_util import init_parameters from tzrec.utils.test_util import ( create_genrec_test_model, - create_genrec_test_tokenizer, - flash_attn_unavailable, make_test_dir, - mark_ci_scope, - nv_gpu_unavailable, parameterized_name_func, ) @@ -145,7 +138,7 @@ def test_rejects_a_model_built_without_a_prompt(self) -> None: name_func=parameterized_name_func, ) def test_builds_backbone_with_the_configured_kernel_and_dtype( - self, attn_implementation, expected_impl + self, attn_kernel, expected_impl ) -> None: # from_config is mocked, so this pins the kwargs init_backbone sends # without building a second backbone; setUp still builds a real one. @@ -158,16 +151,15 @@ def test_builds_backbone_with_the_configured_kernel_and_dtype( mock.patch.object( AutoModelForCausalLM, "from_config", return_value=stand_in ) as from_config, - # the backbone is mocked, so the wheel probe has nothing to check - mock.patch("tzrec.models.genrec_model.find_spec", return_value=object()), ): - model, _ = create_genrec_test_model( + create_genrec_test_model( self.test_dir, lm_parameter_dtype=GenRecModelConfig.BF16, - attn_implementation=attn_implementation, + attn_kernel=attn_kernel, ) - from_config.assert_called_once() + # the fixture backbone is built through from_config too, so this + # reads the last call, which is init_backbone's args, kwargs = from_config.call_args self.assertEqual(len(args), 1) self.assertEqual(args[0].model_type, "qwen2") @@ -179,7 +171,6 @@ def test_builds_backbone_with_the_configured_kernel_and_dtype( }, ) to_mock.assert_not_called() - self.assertIs(model.lm, stand_in) def test_shared_projection_name_requires_matching_widths(self) -> None: with self.assertRaisesRegex(ValueError, "cannot share a module"): @@ -197,15 +188,20 @@ def test_shared_projection_name_requires_matching_widths(self) -> None: ], ) - def test_projected_slot_overwrites_sentinels_and_backpropagates(self) -> None: + def _projected_model(self, **kwargs) -> Tuple[BaseModel, Batch]: + """A model with one projected ``prof`` slot, materialized, and its batch.""" model, compiled_prompt = create_genrec_test_model( self.test_dir, feature_configs=[_hist(), _projected("prof", 8)], prompt="History : {{hist}} . Predict {{prof}} :", + **kwargs, ) # the embedding table is built on meta until something materializes it init_parameters(model, device=torch.device("cpu")) - batch = _projected_batch(compiled_prompt) + return model, _projected_batch(compiled_prompt) + + def test_projected_slot_overwrites_sentinels_and_backpropagates(self) -> None: + model, batch = self._projected_model() embeds = model.build_input(batch) raw = model.lm.get_input_embeddings()(batch.additional_infos[INPUT_IDS]) @@ -224,14 +220,7 @@ def test_projected_slot_overwrites_sentinels_and_backpropagates(self) -> None: name_func=parameterized_name_func, ) def test_projected_slot_follows_a_narrow_lm_dtype(self, lm_parameter_dtype) -> None: - model, compiled_prompt = create_genrec_test_model( - self.test_dir, - feature_configs=[_hist(), _projected("prof", 8)], - prompt="History : {{hist}} . Predict {{prof}} :", - lm_parameter_dtype=lm_parameter_dtype, - ) - init_parameters(model, device=torch.device("cpu")) - batch = _projected_batch(compiled_prompt) + model, batch = self._projected_model(lm_parameter_dtype=lm_parameter_dtype) embeds = model.build_input(batch) self.assertIs(embeds.dtype, _PARAM_DTYPE[lm_parameter_dtype]) @@ -409,176 +398,5 @@ def test_front_end_traces_and_scripts(self) -> None: self.assertTrue(torch.equal(out[key], value), key) -# a second answer, and a rewrite of _HIST_CODES that keeps its width -_OTHER_ANSWER_CODES = [2, 7, 8] -_REWRITTEN_HIST_CODES = [3, 7, 11] - - -def _packed_batch(compiled_prompt, hist_rows, answer_rows) -> Batch: - """Several rows in one batch, packed the way the collator packs them.""" - return _batch( - compiled_prompt, - { - "hist.values": torch.tensor([code for row in hist_rows for code in row]), - "hist.lengths": torch.tensor([len(row) for row in hist_rows]), - "answer.values": torch.tensor( - [code for row in answer_rows for code in row] - ), - "answer.lengths": torch.tensor([len(row) for row in answer_rows]), - }, - ) - - -def _packed_model( - test_dir: str, model_type: str, attn_implementation: int, lm_parameter_dtype: int -) -> Tuple[BaseModel, CompiledPrompt]: - """Build a bf16 genrec model over a tiny ``model_type`` backbone. - - ``create_genrec_test_model`` always writes a Qwen2 backbone, and the packed - forward only holds where the backbone threads the varlen keyword arguments - through to its attention, so the architecture is chosen here instead. Only - the config is written: ``init_backbone`` builds from it and this test never - restores pretrained weights. - - Args: - test_dir (str): scratch directory the backbone is written under. - model_type (str): the hugging-face ``model_type`` of the backbone. - attn_implementation (int): ``GenRecModelConfig.AttnImpl`` to build with. - lm_parameter_dtype (int): ``GenRecModelConfig.ParamDtype`` to build in. - - Returns: - Tuple[BaseModel, CompiledPrompt]: the model and the prompt it was - built on. - """ - backbone = os.path.join(test_dir, model_type) - AutoConfig.for_model( - model_type, - vocab_size=64, - hidden_size=32, - intermediate_size=64, - num_hidden_layers=2, - num_attention_heads=4, - num_key_value_heads=2, - head_dim=8, - max_position_embeddings=64, - tie_word_embeddings=False, - ).save_pretrained(backbone) - - features = create_features([_hist()], fg_mode=FgMode.FG_NONE) - prompt_config = PromptConfig( - tokenizer_path=create_genrec_test_tokenizer(os.path.join(test_dir, "tok.json")), - prompt="History : {{hist}} . Predict :", - response="{{answer}}", - ) - prompt_config.sid_space.codebook.extend([4, 4, 4]) - compiled_prompt = compile_prompt(prompt_config, features, ["answer"]) - - model_config = ModelConfig() - lm_config = model_config.genrec_causal_lm_model - lm_config.hf_model_name_or_path = backbone - lm_config.common.beam_widths.extend([2, 2, 2]) - lm_config.common.num_return_sequences = 2 - lm_config.common.lm_parameter_dtype = lm_parameter_dtype - lm_config.common.attn_implementation = attn_implementation - with torch.random.fork_rng(devices=[]): - torch.manual_seed(0) - model = _create_model( - model_config, features, ["answer"], compiled_prompt=compiled_prompt - ) - return model, compiled_prompt - - -class _PackedRowsCase: - """A packed batch must read exactly as its rows do one at a time. - - The batch is one concatenated stream with ``cu_seq_lens`` marking the - boundaries, so the failure this guards against is a row attending into the - row before it. Rewriting row 0 to different codes of the same width leaves - row 1 at the same packed offsets: its logits have to stay bit identical. - """ - - device = torch.device("cpu") - attn_implementation = GenRecModelConfig.SDPA - lm_parameter_dtype = GenRecModelConfig.FP32 - - def setUp(self) -> None: - self.test_dir = make_test_dir() - - def _assert_packed_matches_solo(self, model_type: str) -> None: - device = self.device - model, compiled_prompt = _packed_model( - self.test_dir, - model_type, - self.attn_implementation, - self.lm_parameter_dtype, - ) - model.to(device) - model.eval() - hist_rows = [_HIST_CODES, _LONG_HIST_CODES] - answer_rows = [_ANSWER_CODES, _OTHER_ANSWER_CODES] - packed_batch = _packed_batch(compiled_prompt, hist_rows, answer_rows).to(device) - - packed = model.predict(packed_batch) - with torch.no_grad(): - solos = [ - model.predict( - _packed_batch(compiled_prompt, [hist], [answer]).to(device) - ) - for hist, answer in zip(hist_rows, answer_rows) - ] - changed = model.predict( - _packed_batch( - compiled_prompt, - [_REWRITTEN_HIST_CODES, hist_rows[1]], - answer_rows, - ).to(device) - ) - - torch.testing.assert_close( - packed["logits"], - torch.cat([result["logits"] for result in solos]), - atol=1e-2, - rtol=1e-2, - ) - torch.testing.assert_close( - packed["labels"], torch.cat([result["labels"] for result in solos]) - ) - torch.testing.assert_close( - packed["logits"][1], changed["logits"][1], atol=0, rtol=0 - ) - - loss = model.loss(packed, packed_batch)["ce_loss"] - self.assertTrue(bool(torch.isfinite(loss))) - loss.backward() - grad = model.lm.get_input_embeddings().weight.grad - self.assertIsNotNone(grad) - self.assertTrue(bool(torch.isfinite(grad).all())) - self.assertGreater(float(grad.abs().sum()), 0) - - -class PackedSdpaAttentionTest(_PackedRowsCase, unittest.TestCase): - """The packed forward under sdpa, which needs no GPU and no wheel.""" - - @parameterized.expand([["qwen2"], ["qwen3"]], name_func=parameterized_name_func) - def test_packed_rows_match_solo_runs_and_backpropagate( - self, model_type: str - ) -> None: - self._assert_packed_matches_solo(model_type) - - -@mark_ci_scope("gpu") -@unittest.skipIf(*nv_gpu_unavailable) -@unittest.skipIf(*flash_attn_unavailable) -class PackedFlashAttentionTest(_PackedRowsCase, unittest.TestCase): - """The same rows through the varlen flash kernel.""" - - device = torch.device("cuda") - attn_implementation = GenRecModelConfig.FLASH_ATTENTION_2 - # the flash kernel takes fp16/bf16 only, and this arm carries no autocast - lm_parameter_dtype = GenRecModelConfig.BF16 - - @parameterized.expand([["qwen2"], ["qwen3"]], name_func=parameterized_name_func) - def test_packed_rows_match_solo_runs_and_backpropagate( - self, model_type: str - ) -> None: - self._assert_packed_matches_solo(model_type) +if __name__ == "__main__": + unittest.main() diff --git a/tzrec/modules/dynamic_beam_test.py b/tzrec/modules/dynamic_beam_test.py index 390bff15..537f0898 100644 --- a/tzrec/modules/dynamic_beam_test.py +++ b/tzrec/modules/dynamic_beam_test.py @@ -15,7 +15,6 @@ import torch from parameterized import parameterized -from transformers import AutoConfig, AutoModelForCausalLM from tzrec.modules.dynamic_beam import capped_beam_widths, dynamic_beam_search from tzrec.utils.test_util import ( @@ -149,26 +148,12 @@ class DynamicBeamSearchFlashAttentionTest(unittest.TestCase): ) def test_ragged_batch_matches_solo_dynamic_cache(self, model_type: str) -> None: device = torch.device("cuda") - config = AutoConfig.for_model( - model_type, - vocab_size=30, - hidden_size=32, - intermediate_size=64, - num_hidden_layers=2, - num_attention_heads=4, - num_key_value_heads=2, - head_dim=8, - max_position_embeddings=64, - tie_word_embeddings=False, - ) - with torch.random.fork_rng(devices=[]): - torch.manual_seed(0) - lm = AutoModelForCausalLM.from_config( - config, - attn_implementation="flash_attention_2", - torch_dtype=torch.bfloat16, - ).to(device) - lm.eval() + lm = create_tiny_causal_lm( + 30, + model_type=model_type, + attn_implementation="flash_attention_2", + torch_dtype=torch.bfloat16, + ).to(device) pairs = [(20, 21), (22, 24), (25, 28)] ids = torch.tensor( [[5, 6, 7, 8], [0, 9, 10, 11]], diff --git a/tzrec/protos/models/genrec_model.proto b/tzrec/protos/models/genrec_model.proto index 61e3a858..e0a32a64 100644 --- a/tzrec/protos/models/genrec_model.proto +++ b/tzrec/protos/models/genrec_model.proto @@ -25,15 +25,13 @@ message GenRecModelConfig { // bf16 compute comes from mixed_precision, not from this. optional ParamDtype lm_parameter_dtype = 5 [default = FP32]; - // Attention kernel the backbone is built with. FLASH_ATTENTION_2 packs the - // batch through the varlen kernel and is what training should use, but it - // needs the flash_attn wheel and a GPU; SDPA carries the same row - // boundaries through position_ids and runs anywhere. - enum AttnImpl { + // FLASH_ATTENTION_2 is what training should use; it needs the flash_attn + // wheel and a GPU. SDPA runs anywhere. + enum AttnKernel { SDPA = 0; FLASH_ATTENTION_2 = 1; } - optional AttnImpl attn_implementation = 6 [default = SDPA]; + optional AttnKernel attn_kernel = 6 [default = SDPA]; } message GenRecCausalLMModel { diff --git a/tzrec/utils/test_util.py b/tzrec/utils/test_util.py index 55ea7230..6c270f93 100644 --- a/tzrec/utils/test_util.py +++ b/tzrec/utils/test_util.py @@ -14,7 +14,7 @@ import os import tempfile from enum import Enum -from typing import Any, Dict, List, Optional, Sequence, Tuple, Union +from typing import TYPE_CHECKING, Any, Dict, List, Optional, Sequence, Tuple, Union import numpy as np import pandas as pd @@ -31,6 +31,9 @@ from tzrec.protos.models.genrec_model_pb2 import GenRecModelConfig from tzrec.protos.prompt_pb2 import PromptSlot from tzrec.utils.export_util import split_model + +if TYPE_CHECKING: + from tzrec.features.feature import BaseFeature from tzrec.utils.fx_util import symbolic_trace nv_gpu_unavailable: Tuple[bool, str] = ( @@ -132,8 +135,11 @@ def create_tiny_causal_lm( vocab_size: int, seed: int = 0, tie_word_embeddings: bool = False, + model_type: str = "qwen2", + attn_implementation: str = "sdpa", + torch_dtype: torch.dtype = torch.float32, ) -> nn.Module: - """A 2-layer Qwen2 causal LM cheap enough to build inside a unit test. + """A 2-layer causal LM cheap enough to build inside a unit test. Seeded so two builds agree, and in ``eval()`` so dropout cannot make a decode non-deterministic. @@ -142,25 +148,33 @@ def create_tiny_causal_lm( vocab_size (int): rows in the embedding table. seed (int): torch seed the random init draws from. tie_word_embeddings (bool): tie ``lm_head`` to the input embedding. + model_type (str): hugging-face ``model_type`` of the architecture. + attn_implementation (str): attention kernel to build it with. + torch_dtype (torch.dtype): dtype of the parameters. Returns: - an eval-mode ``Qwen2ForCausalLM``. + an eval-mode causal LM of ``model_type``. """ - from transformers import Qwen2Config, Qwen2ForCausalLM + from transformers import AutoConfig, AutoModelForCausalLM with torch.random.fork_rng(devices=[]): torch.manual_seed(seed) - config = Qwen2Config( + config = AutoConfig.for_model( + model_type, vocab_size=vocab_size, hidden_size=32, intermediate_size=64, num_hidden_layers=2, num_attention_heads=4, num_key_value_heads=2, + # qwen3 defaults head_dim to 128 rather than hidden_size // heads + head_dim=8, max_position_embeddings=64, tie_word_embeddings=tie_word_embeddings, ) - return Qwen2ForCausalLM(config).eval() + return AutoModelForCausalLM.from_config( + config, attn_implementation=attn_implementation, torch_dtype=torch_dtype + ).eval() # pyre-ignore [2] @@ -386,46 +400,35 @@ def create_genrec_test_tokenizer( return path -def create_genrec_test_model( +def create_genrec_test_prompt( test_dir: str, feature_configs: Optional[List[feature_pb2.FeatureConfig]] = None, prompt: str = "History : {{hist}} . Predict :", response: str = "{{answer}}", slots: Sequence[PromptSlot] = (), - beam_widths: Sequence[int] = (2, 2, 2), - num_return_sequences: int = 2, - lm_parameter_dtype: Optional["GenRecModelConfig.ParamDtype"] = None, - attn_implementation: Optional["GenRecModelConfig.AttnImpl"] = None, -) -> Tuple[BaseModel, CompiledPrompt]: - """Build a GenRecCausalLMModel over a tiny backbone and a compiled prompt. +) -> Tuple[CompiledPrompt, List["BaseFeature"]]: + """Compile a prompt over a tiny tokenizer, without building a backbone. - The backbone and the tokenizer are written under ``test_dir``. The default - features are one ``hist`` raw sequence and the codebook is ``(4, 4, 4)``, so - an offset SID code is ``level_offsets[l] + code`` with offsets ``(0, 4, 8)``. + The default features are one ``hist`` raw sequence and the codebook is + ``(4, 4, 4)``, so an offset SID code is ``level_offsets[l] + code`` with + offsets ``(0, 4, 8)``. Args: - test_dir (str): scratch directory. + test_dir (str): scratch directory the tokenizer is written under. feature_configs (list, optional): feature configs; the ``hist`` raw sequence when None. prompt (str): the prompt template. response (str): the response template. slots (Sequence[PromptSlot]): explicit slot declarations. - beam_widths (Sequence[int]): per-level beam widths. - num_return_sequences (int): sequences returned per sample. - lm_parameter_dtype (optional): ``GenRecModelConfig.ParamDtype`` value. - attn_implementation (optional): ``GenRecModelConfig.AttnImpl`` value. Returns: - Tuple[BaseModel, CompiledPrompt]: the model and the prompt it was built on. + Tuple[CompiledPrompt, List["BaseFeature"]]: the compiled prompt and the + features it was compiled against. """ from tzrec.features.feature import FgMode, create_features - from tzrec.main import _create_model from tzrec.prompt.compile import compile_prompt - from tzrec.protos.model_pb2 import ModelConfig from tzrec.protos.prompt_pb2 import PromptConfig - backbone = os.path.join(test_dir, "backbone") - create_tiny_causal_lm(64).save_pretrained(backbone) if feature_configs is None: feature_configs = [ feature_pb2.FeatureConfig( @@ -442,7 +445,54 @@ def create_genrec_test_model( ) prompt_config.sid_space.codebook.extend([4, 4, 4]) prompt_config.slots.extend(slots) - compiled_prompt = compile_prompt(prompt_config, features, ["answer"]) + return compile_prompt(prompt_config, features, ["answer"]), features + + +def create_genrec_test_model( + test_dir: str, + feature_configs: Optional[List[feature_pb2.FeatureConfig]] = None, + prompt: str = "History : {{hist}} . Predict :", + response: str = "{{answer}}", + slots: Sequence[PromptSlot] = (), + beam_widths: Sequence[int] = (2, 2, 2), + num_return_sequences: int = 2, + lm_parameter_dtype: Optional["GenRecModelConfig.ParamDtype"] = None, + attn_kernel: Optional["GenRecModelConfig.AttnKernel"] = None, + model_type: str = "qwen2", + init_seed: Optional[int] = None, +) -> Tuple[BaseModel, CompiledPrompt]: + """Build a GenRecCausalLMModel over a tiny backbone and a compiled prompt. + + The backbone and the tokenizer are written under ``test_dir``. The default + features are one ``hist`` raw sequence and the codebook is ``(4, 4, 4)``, so + an offset SID code is ``level_offsets[l] + code`` with offsets ``(0, 4, 8)``. + + Args: + test_dir (str): scratch directory. + feature_configs (list, optional): feature configs; the ``hist`` raw + sequence when None. + prompt (str): the prompt template. + response (str): the response template. + slots (Sequence[PromptSlot]): explicit slot declarations. + beam_widths (Sequence[int]): per-level beam widths. + num_return_sequences (int): sequences returned per sample. + lm_parameter_dtype (optional): ``GenRecModelConfig.ParamDtype`` value. + attn_kernel (optional): ``GenRecModelConfig.AttnKernel`` value. + model_type (str): hugging-face ``model_type`` of the backbone. + init_seed (int, optional): seed the random init draws from, when the + caller needs two builds to agree. + + Returns: + Tuple[BaseModel, CompiledPrompt]: the model and the prompt it was built on. + """ + from tzrec.main import _create_model + from tzrec.protos.model_pb2 import ModelConfig + + backbone = os.path.join(test_dir, "backbone") + create_tiny_causal_lm(64, model_type=model_type).save_pretrained(backbone) + compiled_prompt, features = create_genrec_test_prompt( + test_dir, feature_configs, prompt, response, slots + ) model_config = ModelConfig() lm_config = model_config.genrec_causal_lm_model @@ -451,9 +501,16 @@ def create_genrec_test_model( lm_config.common.num_return_sequences = num_return_sequences if lm_parameter_dtype is not None: lm_config.common.lm_parameter_dtype = lm_parameter_dtype - if attn_implementation is not None: - lm_config.common.attn_implementation = attn_implementation - model = _create_model( - model_config, features, ["answer"], compiled_prompt=compiled_prompt - ) + if attn_kernel is not None: + lm_config.common.attn_kernel = attn_kernel + if init_seed is None: + model = _create_model( + model_config, features, ["answer"], compiled_prompt=compiled_prompt + ) + else: + with torch.random.fork_rng(devices=[]): + torch.manual_seed(init_seed) + model = _create_model( + model_config, features, ["answer"], compiled_prompt=compiled_prompt + ) return model, compiled_prompt From 233a1521a5c6890c34be6c4aff91d89022af2363 Mon Sep 17 00:00:00 2001 From: root <597191244@qq.com> Date: Mon, 21 Sep 2026 08:13:41 +0000 Subject: [PATCH 09/15] [chore] drop the cuDNN exclusion and a test it outlived The sdpa_kernel block excluded CUDNN_ATTENTION around the LM call, to keep the NaN losses cuDNN returns on a packed mask. It defended a configuration this branch no longer produces: sdpa takes the left-padded path now, and the NaN was only ever reported against packed-sdpa. It was also dead on the flash branch, which never enters sdpa dispatch at all. An environment that still needs it can set enable_cudnn_sdp(False) itself. Delete test_loss_is_finite_and_backpropagates_into_the_backbone: its four assertions are a subset of SdpaRowIsolationTest's, on the same production path. Checked against 34 production mutations -- 11 kill the loss test and all 11 kill the sibling too, including two shaped to exploit the only axes where they differ (a batch-of-one gradient kill and a train-mode zeroing, both caught by the projected-slot tests). Three mutations kill the sibling while the loss test passes, so it was the weaker of the two. Also drop scaffolding nothing exercises: _stub_model's attn_kernel keyword, which no call site passes, and a cu_dtype override on a test that raises before the cast it was meant to observe. Two docstrings pointed at code that has since moved. Co-Authored-By: Claude Opus 5 (1M context) --- tzrec/models/genrec_causal_lm_model.py | 18 ++++-------------- tzrec/models/genrec_causal_lm_model_test.py | 10 +++------- tzrec/models/genrec_model_test.py | 19 ------------------- tzrec/prompt/assembler_test.py | 4 ++-- 4 files changed, 9 insertions(+), 42 deletions(-) diff --git a/tzrec/models/genrec_causal_lm_model.py b/tzrec/models/genrec_causal_lm_model.py index d24edb9c..4289bc76 100644 --- a/tzrec/models/genrec_causal_lm_model.py +++ b/tzrec/models/genrec_causal_lm_model.py @@ -23,7 +23,6 @@ from typing import Any, Dict, List, Optional, Tuple, cast import torch -from torch.nn.attention import SDPBackend, sdpa_kernel from tzrec.datasets.utils import Batch from tzrec.features.feature import BaseFeature @@ -164,19 +163,10 @@ def _forward( suffix_offsets[None, :] < -infos[RESPONSE_LENGTHS][:, None], self._ignore_index, ) - # CUDNN_ATTENTION is eligible for the packed mask and returns NaN - # losses on it; exclude it here rather than disabling it process-wide - with sdpa_kernel( - [ - SDPBackend.FLASH_ATTENTION, - SDPBackend.EFFICIENT_ATTENTION, - SDPBackend.MATH, - ] - ): - if self._attn_kernel == GenRecModelConfig.FLASH_ATTENTION_2: - logits = self._packed_logits(embeds, batch, keep) - else: - logits = self._padded_logits(embeds, batch, suffix) + if self._attn_kernel == GenRecModelConfig.FLASH_ATTENTION_2: + logits = self._packed_logits(embeds, batch, keep) + else: + logits = self._padded_logits(embeds, batch, suffix) return logits, labels def _packed_logits( diff --git a/tzrec/models/genrec_causal_lm_model_test.py b/tzrec/models/genrec_causal_lm_model_test.py index 135b77ae..bda3613e 100644 --- a/tzrec/models/genrec_causal_lm_model_test.py +++ b/tzrec/models/genrec_causal_lm_model_test.py @@ -94,10 +94,7 @@ def forward(self, **kwargs): return SimpleNamespace(logits=self.lm_head(hidden)) -def _stub_model( - lm: Optional[nn.Module] = None, - attn_kernel: int = GenRecModelConfig.FLASH_ATTENTION_2, -) -> GenRecCausalLMModel: +def _stub_model(lm: Optional[nn.Module] = None) -> GenRecCausalLMModel: """A model with just the attributes ``_forward`` reads, over a stub LM. Defaults to the varlen kernel: these cases are about the packed layout, @@ -106,7 +103,7 @@ def _stub_model( model = GenRecCausalLMModel.__new__(GenRecCausalLMModel) nn.Module.__init__(model) model.lm = _CapturingLM() if lm is None else lm - model._attn_kernel = attn_kernel + model._attn_kernel = GenRecModelConfig.FLASH_ATTENTION_2 model._ignore_index = -7 model._prompt = SimpleNamespace(prompt_plan=SimpleNamespace(logits_suffix_len=4)) return model @@ -122,7 +119,7 @@ def _varlen_batch( """The four varlen keys ``_forward`` reads, over ``cu_seqlens[-1]`` tokens. ``PromptAssembler`` emits ``cu_seqlens`` as int32; the int64 default is - what makes ``_forward``'s cast to int32 observable. + what makes ``_packed_logits``' cast to int32 observable. """ return Batch( additional_infos={ @@ -168,7 +165,6 @@ def test_a_row_shorter_than_the_window_is_rejected(self) -> None: cu_seqlens=(0, 3, 12), max_seqlen=9, response_lengths=(3, 3), - cu_dtype=torch.int32, ) with self.assertRaisesRegex(ValueError, "at least logits_suffix_len"): diff --git a/tzrec/models/genrec_model_test.py b/tzrec/models/genrec_model_test.py index b4e688c8..e130a693 100644 --- a/tzrec/models/genrec_model_test.py +++ b/tzrec/models/genrec_model_test.py @@ -288,25 +288,6 @@ def test_model_resizes_to_target_vocab_size(self) -> None: self.assertEqual(rows, self.compiled_prompt.sid_space.target_vocab_size) self.assertGreater(rows, self.compiled_prompt.sid_space.band_hi[-1]) - def test_loss_is_finite_and_backpropagates_into_the_backbone(self) -> None: - batch = _batch( - self.compiled_prompt, - { - "hist.values": torch.tensor(_LONG_HIST_CODES), - "hist.lengths": torch.tensor([6]), - "answer.values": torch.tensor(_ANSWER_CODES), - "answer.lengths": torch.tensor([3]), - }, - ) - predictions = self.model.predict(batch) - loss = self.model.loss(predictions, batch)["ce_loss"] - self.assertTrue(bool(torch.isfinite(loss))) - loss.backward() - - grad = self.model.lm.get_input_embeddings().weight.grad - self.assertIsNotNone(grad) - self.assertTrue(bool((grad.abs().sum() > 0))) - def test_training_forward_survives_fx_tracing(self) -> None: torch.fx.symbolic_trace(TrainWrapper(self.model)) diff --git a/tzrec/prompt/assembler_test.py b/tzrec/prompt/assembler_test.py index f3fa2f4b..797a8b48 100644 --- a/tzrec/prompt/assembler_test.py +++ b/tzrec/prompt/assembler_test.py @@ -226,8 +226,8 @@ def test_an_empty_prompt_body_leaves_only_the_response(self) -> None: """The walk no longer measures rows, so its consumer must. A row whose body assembles to nothing is shorter than the supervised - suffix window, and the packed forward's window would then reach back - into the sample before it. The assembler emits the row regardless -- + suffix window, whose logits would then reach back into the sample + before it, on either attention layout. The assembler emits the row regardless -- it trusts the parsed batch -- so the guard belongs to the model; this pins the shape that guard has to reject. """ From 45b35851267c8866e277ae4a242fa5272036fdaf Mon Sep 17 00:00:00 2001 From: root <597191244@qq.com> Date: Mon, 21 Sep 2026 09:26:47 +0000 Subject: [PATCH 10/15] [fix] give the padded path its own positions, and cover it end to end Five follow-ups from review. _padded_logits sent no position_ids, so HF assigned arange(max_seqlen) over the padding too and a length-5 row in a width-12 batch started at position 7. Both sibling paths spell the positions out -- _packed_logits subtracts the row starts, dynamic_beam uses the mask cumsum -- and training was the only one that did not. RoPE is relative so the outputs were identical, but a backbone with absolute positions would have trained shifted and decoded from zero. Use the same mask cumsum the decode path uses. _attn_kernel moves to __init__, beside the config read it comes from. init_backbone is a documented override point whose purpose is letting a subclass supply its own self.lm; such a subclass never reached the assignment, and the forward branches on it. Add an integration case that runs the pipeline on FLASH_ATTENTION_2. The kernel now defaults to SDPA, so train_eval -> eval -> export only ever exercised the padded layout -- nothing drove the packed one through the train pipeline, FX tracing or export. The config override threads through _prepare_config, so the sdpa case keeps the CPU lane to itself. Drop _varlen_batch's cu_dtype, which no caller passes, and hoist the SID code fixtures into test_util beside the prompt helper that fixes their offsets -- they were duplicated verbatim in two modules, each with its own copy of the (4, 4, 4) codebook comment. Co-Authored-By: Claude Opus 5 (1M context) --- tzrec/models/genrec_causal_lm_model.py | 4 ++ tzrec/models/genrec_causal_lm_model_test.py | 28 ++++++-------- tzrec/models/genrec_model.py | 6 +-- tzrec/models/genrec_model_test.py | 18 ++++----- tzrec/tests/genrec_integration_test.py | 43 ++++++++++++++++++--- tzrec/utils/test_util.py | 7 ++++ 6 files changed, 72 insertions(+), 34 deletions(-) diff --git a/tzrec/models/genrec_causal_lm_model.py b/tzrec/models/genrec_causal_lm_model.py index 4289bc76..930c7873 100644 --- a/tzrec/models/genrec_causal_lm_model.py +++ b/tzrec/models/genrec_causal_lm_model.py @@ -224,9 +224,13 @@ def _padded_logits( ``(rows, suffix, vocab)`` logits over the response window. """ padded, mask = self._left_pad_packed_inputs(embeds, batch) + # left padding offsets every row, so spell the positions out rather + # than letting HF assign arange(max_seqlen) over the padding too + position_ids = (mask.cumsum(-1) - 1).clamp(min=0) outputs = self.lm( inputs_embeds=padded, attention_mask=mask, + position_ids=position_ids, use_cache=False, logits_to_keep=suffix, ) diff --git a/tzrec/models/genrec_causal_lm_model_test.py b/tzrec/models/genrec_causal_lm_model_test.py index bda3613e..163adfff 100644 --- a/tzrec/models/genrec_causal_lm_model_test.py +++ b/tzrec/models/genrec_causal_lm_model_test.py @@ -29,6 +29,9 @@ ) from tzrec.protos.models.genrec_model_pb2 import GenRecModelConfig from tzrec.utils.test_util import ( + GENREC_ANSWER_CODES, + GENREC_HIST_CODES, + GENREC_LONG_HIST_CODES, create_genrec_test_model, create_genrec_test_prompt, flash_attn_unavailable, @@ -114,16 +117,15 @@ def _varlen_batch( max_seqlen: int = 7, response_lengths: Sequence[int] = (3, 2), first_input_id: int = 100, - cu_dtype: torch.dtype = torch.int64, ) -> Batch: """The four varlen keys ``_forward`` reads, over ``cu_seqlens[-1]`` tokens. - ``PromptAssembler`` emits ``cu_seqlens`` as int32; the int64 default is - what makes ``_packed_logits``' cast to int32 observable. + ``PromptAssembler`` emits ``cu_seqlens`` as int32; int64 here is what + makes ``_packed_logits``' cast to int32 do visible work. """ return Batch( additional_infos={ - CU_SEQLENS: torch.tensor(cu_seqlens, dtype=cu_dtype), + CU_SEQLENS: torch.tensor(cu_seqlens), INPUT_IDS: torch.arange(first_input_id, first_input_id + cu_seqlens[-1]), MAX_SEQLEN: torch.tensor(max_seqlen), RESPONSE_LENGTHS: torch.tensor(response_lengths), @@ -245,15 +247,9 @@ def test_beam_config_uses_final_capped_capacity(self) -> None: ) -# offset SID codes for the (4, 4, 4) codebook: level_offsets[l] + code -_HIST_CODES = [0, 5, 10] -_LONG_HIST_CODES = [0, 5, 10, 3, 4, 9] -_ANSWER_CODES = [1, 6, 11] - - -# a second answer, and a rewrite of _HIST_CODES that keeps its width -_OTHER_ANSWER_CODES = [2, 7, 8] -_REWRITTEN_HIST_CODES = [3, 7, 11] +# a second answer, and a rewrite of GENREC_HIST_CODES that keeps its width +_OTHERGENREC_ANSWER_CODES = [2, 7, 8] +_REWRITTENGENREC_HIST_CODES = [3, 7, 11] def _packed_batch(compiled_prompt, hist_rows, answer_rows) -> Batch: @@ -305,8 +301,8 @@ def test_rows_match_solo_runs_and_backpropagate(self, model_type: str) -> None: ) model.to(device) model.eval() - hist_rows = [_HIST_CODES, _LONG_HIST_CODES] - answer_rows = [_ANSWER_CODES, _OTHER_ANSWER_CODES] + hist_rows = [GENREC_HIST_CODES, GENREC_LONG_HIST_CODES] + answer_rows = [GENREC_ANSWER_CODES, _OTHERGENREC_ANSWER_CODES] packed_batch = _packed_batch(compiled_prompt, hist_rows, answer_rows).to(device) packed = model.predict(packed_batch) @@ -320,7 +316,7 @@ def test_rows_match_solo_runs_and_backpropagate(self, model_type: str) -> None: changed = model.predict( _packed_batch( compiled_prompt, - [_REWRITTEN_HIST_CODES, hist_rows[1]], + [_REWRITTENGENREC_HIST_CODES, hist_rows[1]], answer_rows, ).to(device) ) diff --git a/tzrec/models/genrec_model.py b/tzrec/models/genrec_model.py index 4f6e5c9c..d589f064 100644 --- a/tzrec/models/genrec_model.py +++ b/tzrec/models/genrec_model.py @@ -104,6 +104,9 @@ def __init__( self._ignore_index = int(cfg.common.ignore_index) self.lm: nn.Module self._attn_kernel: int + # the forward picks its layout from this: only the varlen kernel reads + # cu_seq_lens, so every other kernel has to be handed padded rows + self._attn_kernel = cfg.common.attn_kernel self.init_backbone( cfg.hf_model_name_or_path, cfg.common.lm_parameter_dtype, @@ -143,9 +146,6 @@ def init_backbone( lm_parameter_dtype: dtype of the LM parameters. attn_kernel: attention kernel to build the backbone with. """ - # the forward picks its layout from this: only the varlen kernel reads - # cu_seq_lens, so every other kernel has to be handed padded rows - self._attn_kernel = attn_kernel config = AutoConfig.from_pretrained(hf_model_name_or_path) self.lm = AutoModelForCausalLM.from_config( config, diff --git a/tzrec/models/genrec_model_test.py b/tzrec/models/genrec_model_test.py index e130a693..9072448b 100644 --- a/tzrec/models/genrec_model_test.py +++ b/tzrec/models/genrec_model_test.py @@ -43,16 +43,14 @@ from tzrec.utils.fx_util import symbolic_trace from tzrec.utils.state_dict_util import init_parameters from tzrec.utils.test_util import ( + GENREC_ANSWER_CODES, + GENREC_HIST_CODES, + GENREC_LONG_HIST_CODES, create_genrec_test_model, make_test_dir, parameterized_name_func, ) -# offset SID codes for the (4, 4, 4) codebook: level_offsets[l] + code -_HIST_CODES = [0, 5, 10] -_LONG_HIST_CODES = [0, 5, 10, 3, 4, 9] -_ANSWER_CODES = [1, 6, 11] - def _hist() -> feature_pb2.FeatureConfig: return feature_pb2.FeatureConfig( @@ -86,9 +84,9 @@ def _projected_batch(compiled_prompt) -> Batch: return _batch( compiled_prompt, { - "hist.values": torch.tensor(_HIST_CODES).reshape(-1, 1), + "hist.values": torch.tensor(GENREC_HIST_CODES).reshape(-1, 1), "hist.lengths": torch.tensor([3]), - "answer.values": torch.tensor(_ANSWER_CODES), + "answer.values": torch.tensor(GENREC_ANSWER_CODES), "answer.lengths": torch.tensor([3]), "prof.values": torch.tensor([5, 9]), "prof.lengths": torch.tensor([2]), @@ -306,9 +304,9 @@ def setUp(self) -> None: # the parsed dict as the data parser emits it: a dense sequence of codes # and a sparse behaviour sequence self.data = { - "hist.values": torch.tensor(_LONG_HIST_CODES, dtype=torch.float32).reshape( - -1, 1 - ), + "hist.values": torch.tensor( + GENREC_LONG_HIST_CODES, dtype=torch.float32 + ).reshape(-1, 1), "hist.lengths": torch.tensor([6]), "beh.values": torch.tensor([3, 9]), "beh.lengths": torch.tensor([2]), diff --git a/tzrec/tests/genrec_integration_test.py b/tzrec/tests/genrec_integration_test.py index 6f63204a..a1fe38ec 100644 --- a/tzrec/tests/genrec_integration_test.py +++ b/tzrec/tests/genrec_integration_test.py @@ -14,6 +14,7 @@ import os import shutil import unittest +from typing import Optional from unittest import mock import numpy as np @@ -31,14 +32,17 @@ ) from tzrec.prompt.compile import compile_prompt from tzrec.prompt.hole_keys import HOLE_KEYS, HoleKeyBuilder +from tzrec.protos.models.genrec_model_pb2 import GenRecModelConfig from tzrec.tests import utils from tzrec.utils import config_util from tzrec.utils.test_util import ( create_genrec_test_tokenizer, create_tiny_causal_lm, + flash_attn_unavailable, gpu_unavailable, make_test_dir, mark_ci_scope, + nv_gpu_unavailable, ) _MOCK_CONFIG = "tzrec/tests/configs/genrec_causal_lm_model_mock.config" @@ -65,8 +69,16 @@ def tearDown(self): if self.success and os.path.exists(self.test_dir): shutil.rmtree(self.test_dir) - def _prepare_config(self) -> str: - """Write the tiny backbone, tokenizer, manifest and data; return the config.""" + def _prepare_config(self, attn_kernel: Optional[int] = None) -> str: + """Write the tiny backbone, tokenizer, manifest and data; return the config. + + Args: + attn_kernel (int, optional): ``GenRecModelConfig.AttnKernel`` to run + the pipeline on; the config's own default when None. + + Returns: + str: path to the written pipeline config. + """ backbone = os.path.join(self.test_dir, "backbone") create_tiny_causal_lm(64).save_pretrained(backbone) tokenizer = create_genrec_test_tokenizer( @@ -87,13 +99,23 @@ def _prepare_config(self) -> str: config.prompt_config.sid_space.manifest_path = manifest text_format.Merge(_BEH, config.feature_configs.add()) config.prompt_config.prompt = "History : {{hist}} . {{beh}} Predict :" + if attn_kernel is not None: + config.model_config.genrec_causal_lm_model.common.attn_kernel = attn_kernel config_path = os.path.join(self.test_dir, "genrec.config") config_util.save_message(config, config_path) return config_path - def _train_eval_export(self) -> str: - """Run the pipeline; return the trained ``pipeline.config`` path.""" - config_path = self._prepare_config() + def _train_eval_export(self, attn_kernel: Optional[int] = None) -> str: + """Run the pipeline; return the trained ``pipeline.config`` path. + + Args: + attn_kernel (int, optional): ``GenRecModelConfig.AttnKernel`` to run + the pipeline on; the config's own default when None. + + Returns: + str: path to the trained ``pipeline.config``. + """ + config_path = self._prepare_config(attn_kernel) self.success = utils.test_train_eval(config_path, self.test_dir) trained = os.path.join(self.test_dir, "pipeline.config") if self.success: @@ -194,6 +216,17 @@ def test_genrec_train_eval_export(self): ) self.assertEqual(tokenizer.eos_token_id, compiled.sid_space.eos_token_id) + @unittest.skipIf(*nv_gpu_unavailable) + @unittest.skipIf(*flash_attn_unavailable) + @mark_ci_scope("gpu") + def test_genrec_train_eval_export_on_the_varlen_kernel(self): + """The packed forward, through the train pipeline and out to an export. + + The default kernel is SDPA, so the sibling case above covers the padded + layout; this is the only end-to-end run of the packed one. + """ + self._train_eval_export(attn_kernel=GenRecModelConfig.FLASH_ATTENTION_2) + @unittest.skipIf(*gpu_unavailable) @mark_ci_scope("gpu") def test_genrec_export_distributed_embedding(self): diff --git a/tzrec/utils/test_util.py b/tzrec/utils/test_util.py index 6c270f93..0fb513e5 100644 --- a/tzrec/utils/test_util.py +++ b/tzrec/utils/test_util.py @@ -400,6 +400,13 @@ def create_genrec_test_tokenizer( return path +# offset SID codes for the (4, 4, 4) codebook create_genrec_test_prompt fixes: +# level_offsets[l] + code, with offsets (0, 4, 8) +GENREC_HIST_CODES = [0, 5, 10] +GENREC_LONG_HIST_CODES = [0, 5, 10, 3, 4, 9] +GENREC_ANSWER_CODES = [1, 6, 11] + + def create_genrec_test_prompt( test_dir: str, feature_configs: Optional[List[feature_pb2.FeatureConfig]] = None, From 27c25344d9ba83c05fb501c77b4315daa947562c Mon Sep 17 00:00:00 2001 From: root <597191244@qq.com> Date: Tue, 22 Sep 2026 02:35:43 +0000 Subject: [PATCH 11/15] [chore] ship flash_attn for cp310 and cp312, not just cp311 The 20260921 build covers all three supported Pythons on every CUDA minor, so the pins no longer leave 3.10 and 3.12 without a flash path. Also moves the variant into the version's local segment the way the sibling wheels write it: build date, CUDA minor, torch minor, SM list, upstream commit. cu129 and cu130 gain sm100/sm120; cu126 still builds sm80/sm90 only. Co-Authored-By: Claude Opus 5 (1M context) --- requirements/cu126.txt | 4 +++- requirements/cu129.txt | 4 +++- requirements/cu130.txt | 4 +++- 3 files changed, 9 insertions(+), 3 deletions(-) diff --git a/requirements/cu126.txt b/requirements/cu126.txt index a70d9af4..7bc6065f 100644 --- a/requirements/cu126.txt +++ b/requirements/cu126.txt @@ -7,5 +7,7 @@ faiss @ https://tzrec.oss-accelerate.aliyuncs.com/third_party/faiss/cu12/faiss-1 fbgemm_gpu_hstu @ https://tzrec.oss-accelerate.aliyuncs.com/third_party/hstu/cu126/fbgemm_gpu_hstu-0.1.0%2B20260914.fece651b.fn5.cu126-cp310-cp310-linux_x86_64.whl ; python_version=="3.10" fbgemm_gpu_hstu @ https://tzrec.oss-accelerate.aliyuncs.com/third_party/hstu/cu126/fbgemm_gpu_hstu-0.1.0%2B20260914.fece651b.fn5.cu126-cp311-cp311-linux_x86_64.whl ; python_version=="3.11" fbgemm_gpu_hstu @ https://tzrec.oss-accelerate.aliyuncs.com/third_party/hstu/cu126/fbgemm_gpu_hstu-0.1.0%2B20260914.fece651b.fn5.cu126-cp312-cp312-linux_x86_64.whl ; python_version=="3.12" -flash_attn @ https://tzrec.oss-accelerate.aliyuncs.com/third_party/flash_attn/cu126/flash_attn-2.8.3.post1%2Bcu126.torch213.sm80.90-cp311-cp311-linux_x86_64.whl ; python_version=="3.11" +flash_attn @ https://tzrec.oss-accelerate.aliyuncs.com/third_party/flash_attn/cu126/flash_attn-2.8.3%2B20260921.cu126.torch213.sm80.90.g060c9188beec-cp310-cp310-linux_x86_64.whl ; python_version=="3.10" +flash_attn @ https://tzrec.oss-accelerate.aliyuncs.com/third_party/flash_attn/cu126/flash_attn-2.8.3%2B20260921.cu126.torch213.sm80.90.g060c9188beec-cp311-cp311-linux_x86_64.whl ; python_version=="3.11" +flash_attn @ https://tzrec.oss-accelerate.aliyuncs.com/third_party/flash_attn/cu126/flash_attn-2.8.3%2B20260921.cu126.torch213.sm80.90.g060c9188beec-cp312-cp312-linux_x86_64.whl ; python_version=="3.12" triton==3.7.1 diff --git a/requirements/cu129.txt b/requirements/cu129.txt index 2f1f0e83..c6ee3f1d 100644 --- a/requirements/cu129.txt +++ b/requirements/cu129.txt @@ -7,5 +7,7 @@ faiss @ https://tzrec.oss-accelerate.aliyuncs.com/third_party/faiss/cu12/faiss-1 fbgemm_gpu_hstu @ https://tzrec.oss-accelerate.aliyuncs.com/third_party/hstu/cu129/fbgemm_gpu_hstu-0.1.0%2B20260914.fece651b.fn5.cu129-cp310-cp310-linux_x86_64.whl ; python_version=="3.10" fbgemm_gpu_hstu @ https://tzrec.oss-accelerate.aliyuncs.com/third_party/hstu/cu129/fbgemm_gpu_hstu-0.1.0%2B20260914.fece651b.fn5.cu129-cp311-cp311-linux_x86_64.whl ; python_version=="3.11" fbgemm_gpu_hstu @ https://tzrec.oss-accelerate.aliyuncs.com/third_party/hstu/cu129/fbgemm_gpu_hstu-0.1.0%2B20260914.fece651b.fn5.cu129-cp312-cp312-linux_x86_64.whl ; python_version=="3.12" -flash_attn @ https://tzrec.oss-accelerate.aliyuncs.com/third_party/flash_attn/cu129/flash_attn-2.8.3.post1%2Bcu129.torch213.sm80.90.120-cp311-cp311-linux_x86_64.whl ; python_version=="3.11" +flash_attn @ https://tzrec.oss-accelerate.aliyuncs.com/third_party/flash_attn/cu129/flash_attn-2.8.3%2B20260921.cu129.torch213.sm80.90.100.120.g060c9188beec-cp310-cp310-linux_x86_64.whl ; python_version=="3.10" +flash_attn @ https://tzrec.oss-accelerate.aliyuncs.com/third_party/flash_attn/cu129/flash_attn-2.8.3%2B20260921.cu129.torch213.sm80.90.100.120.g060c9188beec-cp311-cp311-linux_x86_64.whl ; python_version=="3.11" +flash_attn @ https://tzrec.oss-accelerate.aliyuncs.com/third_party/flash_attn/cu129/flash_attn-2.8.3%2B20260921.cu129.torch213.sm80.90.100.120.g060c9188beec-cp312-cp312-linux_x86_64.whl ; python_version=="3.12" triton==3.7.1 diff --git a/requirements/cu130.txt b/requirements/cu130.txt index 84f790f0..a96aa8c3 100644 --- a/requirements/cu130.txt +++ b/requirements/cu130.txt @@ -7,6 +7,8 @@ faiss @ https://tzrec.oss-accelerate.aliyuncs.com/third_party/faiss/cu13/faiss-1 fbgemm_gpu_hstu @ https://tzrec.oss-accelerate.aliyuncs.com/third_party/hstu/cu130/fbgemm_gpu_hstu-0.1.0%2B20260914.fece651b.fn5.cu130-cp310-cp310-linux_x86_64.whl ; python_version=="3.10" fbgemm_gpu_hstu @ https://tzrec.oss-accelerate.aliyuncs.com/third_party/hstu/cu130/fbgemm_gpu_hstu-0.1.0%2B20260914.fece651b.fn5.cu130-cp311-cp311-linux_x86_64.whl ; python_version=="3.11" fbgemm_gpu_hstu @ https://tzrec.oss-accelerate.aliyuncs.com/third_party/hstu/cu130/fbgemm_gpu_hstu-0.1.0%2B20260914.fece651b.fn5.cu130-cp312-cp312-linux_x86_64.whl ; python_version=="3.12" -flash_attn @ https://tzrec.oss-accelerate.aliyuncs.com/third_party/flash_attn/cu130/flash_attn-2.8.3.post1%2Bcu130.torch213.sm80.90.120-cp311-cp311-linux_x86_64.whl ; python_version=="3.11" +flash_attn @ https://tzrec.oss-accelerate.aliyuncs.com/third_party/flash_attn/cu130/flash_attn-2.8.3%2B20260921.cu130.torch213.sm80.90.100.120.g060c9188beec-cp310-cp310-linux_x86_64.whl ; python_version=="3.10" +flash_attn @ https://tzrec.oss-accelerate.aliyuncs.com/third_party/flash_attn/cu130/flash_attn-2.8.3%2B20260921.cu130.torch213.sm80.90.100.120.g060c9188beec-cp311-cp311-linux_x86_64.whl ; python_version=="3.11" +flash_attn @ https://tzrec.oss-accelerate.aliyuncs.com/third_party/flash_attn/cu130/flash_attn-2.8.3%2B20260921.cu130.torch213.sm80.90.100.120.g060c9188beec-cp312-cp312-linux_x86_64.whl ; python_version=="3.12" torch-tensorrt==2.13.0 triton==3.7.1 From 66dc5ec86109f124d023450250f1830322e3c918 Mon Sep 17 00:00:00 2001 From: root <597191244@qq.com> Date: Tue, 22 Sep 2026 03:11:05 +0000 Subject: [PATCH 12/15] [chore] bump version to 1.4.14 Co-Authored-By: Claude Opus 5 (1M context) --- tzrec/version.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tzrec/version.py b/tzrec/version.py index 4709ed01..4071ecb5 100644 --- a/tzrec/version.py +++ b/tzrec/version.py @@ -9,4 +9,4 @@ # See the License for the specific language governing permissions and # limitations under the License. -__version__ = "1.4.13" +__version__ = "1.4.14" From 53ab60ba2420a225280f6b7ba18742ca7899c542 Mon Sep 17 00:00:00 2001 From: root <597191244@qq.com> Date: Tue, 22 Sep 2026 07:25:31 +0000 Subject: [PATCH 13/15] [bugfix] filter prompt-less rows instead of failing the batch The first response token is predicted by the column before it, so a row whose prompt assembles to nothing has no predictor for it: the packed arm took that logit from the row before, the padded arm from a pad position, and neither label was masked, because masking keys off response_length rather than row length. The loss stayed finite and plausible, so nothing surfaced it. The guard covering this compared row length against logits_suffix_len, which is response_max + 1. That rejected any row below the response segment's maximum -- including a row with a full prompt and an empty answer, which is harmless -- and it read a device tensor back to the host on every step. PromptAssembler now zeroes such a row's response_length on the host, which masks its whole window and leaves the packed stream aligned with the batch's features; a batch where nothing survives still raises, because the loss returns NaN when it averages over no tokens. _padded_logits keeps a narrower check of its own, on the padded width, since logits_to_keep clamps silently and fails later inside the loss. Renames _packed_logits to _varlen_logits, which only the flash kernel reaches, and adds the cross-kernel case pinning that the two layouts agree. Co-Authored-By: Claude Opus 5 (1M context) --- tzrec/models/genrec_causal_lm_model.py | 22 ++- tzrec/models/genrec_causal_lm_model_test.py | 162 +++++++++++++++----- tzrec/prompt/assembler.py | 12 ++ tzrec/prompt/assembler_test.py | 41 +++-- 4 files changed, 169 insertions(+), 68 deletions(-) diff --git a/tzrec/models/genrec_causal_lm_model.py b/tzrec/models/genrec_causal_lm_model.py index 930c7873..f4474e19 100644 --- a/tzrec/models/genrec_causal_lm_model.py +++ b/tzrec/models/genrec_causal_lm_model.py @@ -145,16 +145,7 @@ def _forward( """ infos = batch.additional_infos cu_seqlens = infos[CU_SEQLENS] - lengths = torch.diff(cu_seqlens) suffix = cast(int, self._prompt.prompt_plan.logits_suffix_len) - # a row shorter than the window reads the row before it - if bool((lengths < suffix).any()): - raise ValueError( - f"{type(self).__name__}: every assembled sample must be at " - f"least logits_suffix_len ({suffix}) tokens long. Drop the " - f"rows whose prompt features are all empty, or give the " - f"template static text." - ) suffix_offsets = torch.arange(-suffix, 0, device=embeds.device) # (rows, suffix): the window each row is supervised over, the same # absolute positions whichever layout the kernel is given @@ -164,12 +155,12 @@ def _forward( self._ignore_index, ) if self._attn_kernel == GenRecModelConfig.FLASH_ATTENTION_2: - logits = self._packed_logits(embeds, batch, keep) + logits = self._varlen_logits(embeds, batch, keep) else: logits = self._padded_logits(embeds, batch, suffix) return logits, labels - def _packed_logits( + def _varlen_logits( self, embeds: torch.Tensor, batch: Batch, keep: torch.Tensor ) -> torch.Tensor: """Score every row in one varlen call, with no padding at all. @@ -224,6 +215,13 @@ def _padded_logits( ``(rows, suffix, vocab)`` logits over the response window. """ padded, mask = self._left_pad_packed_inputs(embeds, batch) + if padded.shape[1] < suffix: + raise ValueError( + f"{type(self).__name__}: no sample reaches the supervised " + f"window -- the padded width is {padded.shape[1]}, the window " + f"is {suffix}, and logits_to_keep would silently return the " + f"narrower one." + ) # left padding offsets every row, so spell the positions out rather # than letting HF assign arange(max_seqlen) over the padding too position_ids = (mask.cumsum(-1) - 1).clamp(min=0) @@ -258,7 +256,7 @@ def _left_pad_packed_inputs( embeds: torch.Tensor, batch: Batch, ) -> Tuple[torch.Tensor, torch.Tensor]: - """Left-pad packed prompt embeddings for beam decode. + """Left-pad packed prompt embeddings for beam decode and the padded forward. Args: embeds: packed embeddings, ``(total_tokens, hidden)``. diff --git a/tzrec/models/genrec_causal_lm_model_test.py b/tzrec/models/genrec_causal_lm_model_test.py index 163adfff..43078965 100644 --- a/tzrec/models/genrec_causal_lm_model_test.py +++ b/tzrec/models/genrec_causal_lm_model_test.py @@ -42,8 +42,21 @@ ) -class LeftPadPackedInputsTest(unittest.TestCase): - """The one adapter where padding lives.""" +class PaddedForwardTest(unittest.TestCase): + """The layout every kernel but the varlen one is given.""" + + def test_a_padded_width_under_the_window_is_rejected(self) -> None: + """``logits_to_keep`` would return fewer columns than there are labels.""" + model = _stub_model() + model._attn_kernel = GenRecModelConfig.SDPA + batch = _varlen_batch( + cu_seqlens=(0, 3, 6), + max_seqlen=3, + response_lengths=(2, 2), + ) + + with self.assertRaisesRegex(ValueError, "no sample reaches the supervised"): + model._forward(torch.zeros(6, 6), batch) def test_packs_rows_of_different_lengths(self) -> None: embeds = torch.arange(18, dtype=torch.float32).reshape(9, 2) @@ -121,7 +134,7 @@ def _varlen_batch( """The four varlen keys ``_forward`` reads, over ``cu_seqlens[-1]`` tokens. ``PromptAssembler`` emits ``cu_seqlens`` as int32; int64 here is what - makes ``_packed_logits``' cast to int32 do visible work. + makes ``_varlen_logits``' cast to int32 do visible work. """ return Batch( additional_infos={ @@ -160,17 +173,24 @@ def test_passes_varlen_metadata_and_builds_per_row_labels(self) -> None: [[-7, 102, 103, 104], [-7, -7, 110, 111]], ) - def test_a_row_shorter_than_the_window_is_rejected(self) -> None: - """A row under ``logits_suffix_len`` would index into its neighbour.""" + def test_a_row_the_assembler_zeroed_supervises_nothing(self) -> None: + """A prompt-less row reaches here with ``response_lengths`` at zero. + + ``PromptAssembler`` zeroes it rather than failing the batch, so every + column of that row's window has to mask out -- including the ones whose + ``keep`` indices fall in the row before it. + """ model = _stub_model() batch = _varlen_batch( cu_seqlens=(0, 3, 12), max_seqlen=9, - response_lengths=(3, 3), + response_lengths=(0, 3), ) - with self.assertRaisesRegex(ValueError, "at least logits_suffix_len"): - model._forward(torch.zeros(12, 6), batch) + _, labels = model._forward(torch.zeros(12, 6), batch) + + self.assertEqual(labels[0].tolist(), [-7, -7, -7, -7]) + self.assertEqual(labels[1].tolist(), [-7, 109, 110, 111]) def test_loss_and_gradients_cover_only_valid_response_pairs(self) -> None: model = _stub_model(_DifferentiableLM(hidden_size=6, vocab_size=32)) @@ -272,37 +292,57 @@ def _packed_batch(compiled_prompt, hist_rows, answer_rows) -> Batch: return batch -class _RowIsolationCase: - """Every row must read exactly as it does alone, whatever the layout. +# the two rows every layout case runs: different histories, different answers +_TWO_ROWS = ( + [GENREC_HIST_CODES, GENREC_LONG_HIST_CODES], + [GENREC_ANSWER_CODES, _OTHERGENREC_ANSWER_CODES], +) - The failure this guards against is a row attending into the row before it, - which the two kernels avoid by different means: the varlen kernel is handed - ``cu_seq_lens``, every other kernel is handed left-padded rows. Rewriting - row 0 to different codes of the same width leaves row 1 at the same - offsets, so its logits have to stay bit identical either way. - """ - device = torch.device("cpu") - attn_kernel = GenRecModelConfig.SDPA - lm_parameter_dtype = GenRecModelConfig.FP32 +class LayoutInvarianceTest(unittest.TestCase): + """Neither batching nor the kernel may change what a row scores. + + Two invariants, one fixture. Row isolation: a row must not attend into the + row before it, which the kernels avoid by different means -- the varlen + kernel is handed ``cu_seq_lens``, every other kernel left-padded rows. + Rewriting row 0 to different codes of the same width leaves row 1 at the + same offsets, so its logits have to stay bit identical either way. Layout + equivalence: the two arms must then agree with each other, which comparing + each against its own solo runs never establishes. + """ def setUp(self) -> None: self.test_dir = make_test_dir() - @parameterized.expand([["qwen2"], ["qwen3"]], name_func=parameterized_name_func) - def test_rows_match_solo_runs_and_backpropagate(self, model_type: str) -> None: - device = self.device + def _eval_model( + self, + model_type: str, + attn_kernel: "GenRecModelConfig.AttnKernel", + device: torch.device, + lm_parameter_dtype: "GenRecModelConfig.ParamDtype", + ): + """The same backbone from the same seed, under one kernel.""" model, compiled_prompt = create_genrec_test_model( self.test_dir, model_type=model_type, - attn_kernel=self.attn_kernel, - lm_parameter_dtype=self.lm_parameter_dtype, + attn_kernel=attn_kernel, + lm_parameter_dtype=lm_parameter_dtype, init_seed=0, ) - model.to(device) - model.eval() - hist_rows = [GENREC_HIST_CODES, GENREC_LONG_HIST_CODES] - answer_rows = [GENREC_ANSWER_CODES, _OTHERGENREC_ANSWER_CODES] + return model.to(device).eval(), compiled_prompt + + def _assert_rows_match_solo_runs( + self, + model_type: str, + attn_kernel: "GenRecModelConfig.AttnKernel", + device: torch.device, + lm_parameter_dtype: "GenRecModelConfig.ParamDtype", + ) -> None: + """Run the isolation invariant under one kernel.""" + model, compiled_prompt = self._eval_model( + model_type, attn_kernel, device, lm_parameter_dtype + ) + hist_rows, answer_rows = _TWO_ROWS packed_batch = _packed_batch(compiled_prompt, hist_rows, answer_rows).to(device) packed = model.predict(packed_batch) @@ -342,21 +382,61 @@ def test_rows_match_solo_runs_and_backpropagate(self, model_type: str) -> None: self.assertTrue(bool(torch.isfinite(grad).all())) self.assertGreater(float(grad.abs().sum()), 0) + @parameterized.expand([["qwen2"], ["qwen3"]], name_func=parameterized_name_func) + def test_sdpa_rows_match_solo_runs(self, model_type: str) -> None: + self._assert_rows_match_solo_runs( + model_type, + GenRecModelConfig.SDPA, + torch.device("cpu"), + GenRecModelConfig.FP32, + ) -class SdpaRowIsolationTest(_RowIsolationCase, unittest.TestCase): - """The left-padded layout sdpa is given, which needs no GPU and no wheel.""" - - -@mark_ci_scope("gpu") -@unittest.skipIf(*nv_gpu_unavailable) -@unittest.skipIf(*flash_attn_unavailable) -class FlashAttentionRowIsolationTest(_RowIsolationCase, unittest.TestCase): - """The same rows packed into one stream, through the varlen flash kernel.""" + # mark_ci_scope must sit below expand: expand returns None, so tagging + # above it raises at import + @parameterized.expand([["qwen2"], ["qwen3"]], name_func=parameterized_name_func) + @unittest.skipIf(*nv_gpu_unavailable) + @unittest.skipIf(*flash_attn_unavailable) + @mark_ci_scope("gpu") + def test_flash_rows_match_solo_runs(self, model_type: str) -> None: + self._assert_rows_match_solo_runs( + model_type, + GenRecModelConfig.FLASH_ATTENTION_2, + torch.device("cuda"), + # the flash kernel takes fp16/bf16 only, and this carries no autocast + GenRecModelConfig.BF16, + ) - device = torch.device("cuda") - attn_kernel = GenRecModelConfig.FLASH_ATTENTION_2 - # the flash kernel takes fp16/bf16 only, and this arm carries no autocast - lm_parameter_dtype = GenRecModelConfig.BF16 + @parameterized.expand([["qwen2"], ["qwen3"]], name_func=parameterized_name_func) + @unittest.skipIf(*nv_gpu_unavailable) + @unittest.skipIf(*flash_attn_unavailable) + @mark_ci_scope("gpu") + def test_the_packed_and_padded_arms_agree(self, model_type: str) -> None: + """The kernel is a speed choice, so it must not change the result. + + The two cases above each compare a kernel against solo runs of itself, + which pins neither arm to the other. Both run BF16 here because the + flash kernel takes no fp32. + """ + device = torch.device("cuda") + outputs = [] + for attn_kernel in ( + GenRecModelConfig.SDPA, + GenRecModelConfig.FLASH_ATTENTION_2, + ): + model, compiled_prompt = self._eval_model( + model_type, attn_kernel, device, GenRecModelConfig.BF16 + ) + batch = _packed_batch(compiled_prompt, *_TWO_ROWS).to(device) + with torch.no_grad(): + outputs.append(model.predict(batch)) + + padded, packed = outputs + # the same absolute window either way, so the labels are built once and + # only the logits path differs + torch.testing.assert_close(packed["labels"], padded["labels"]) + torch.testing.assert_close( + packed["logits"], padded["logits"], atol=1e-2, rtol=1e-2 + ) if __name__ == "__main__": diff --git a/tzrec/prompt/assembler.py b/tzrec/prompt/assembler.py index b7c1e33e..474d59af 100644 --- a/tzrec/prompt/assembler.py +++ b/tzrec/prompt/assembler.py @@ -309,6 +309,18 @@ def forward(self, batch: Dict[str, torch.Tensor]) -> Dict[str, torch.Tensor]: response_lengths = response_lengths + seg_lens[i] out.index_copy_(0, torch.cat(dests, dim=0), torch.cat(seg_values, dim=0)) + if self.num_segments > self.num_body: + # the first response token is predicted by the one before it, so a + # row that is nothing but its response supervises nothing + response_lengths = response_lengths * (row_total > response_lengths) + if not bool((response_lengths > 0).any()): + raise ValueError( + "PromptAssembler: nothing in this batch can be supervised " + "-- every sample either carries no response or no token " + "before it. Give the template static text, or drop the " + "samples whose prompt features are all empty." + ) + hole_positions = torch.cat(hole_parts, dim=0) # how many of those holes each projected occurrence owns, so a host can # cut the flat streams back into per-slot spans without the plan diff --git a/tzrec/prompt/assembler_test.py b/tzrec/prompt/assembler_test.py index 797a8b48..cbbec266 100644 --- a/tzrec/prompt/assembler_test.py +++ b/tzrec/prompt/assembler_test.py @@ -222,14 +222,12 @@ def test_response_is_optional_and_its_length_is_recorded(self) -> None: ) self.assertEqual(prompt_only[RESPONSE_LENGTHS].tolist(), [0]) - def test_an_empty_prompt_body_leaves_only_the_response(self) -> None: - """The walk no longer measures rows, so its consumer must. - - A row whose body assembles to nothing is shorter than the supervised - suffix window, whose logits would then reach back into the sample - before it, on either attention layout. The assembler emits the row regardless -- - it trusts the parsed batch -- so the guard belongs to the model; this - pins the shape that guard has to reject. + def test_a_row_with_no_prompt_body_supervises_nothing(self) -> None: + """Its first response token has no token before it to be predicted by. + + The row stays in the packed stream -- dropping it would desync the + batch's features from its ids -- so the walk zeroes its response length + instead, which masks every column of its window downstream. """ asm = _asm( (_slot("hist", FillMode.INLINE),), @@ -238,18 +236,31 @@ def test_an_empty_prompt_body_leaves_only_the_response(self) -> None: out = asm( _parsed( { - "hist": [np.array([], dtype=np.int64)], - "answer": [np.array([1, 6, 11])], + "hist": [np.array([1, 6, 11]), np.array([], dtype=np.int64)], + "answer": [np.array([2, 7, 10]), np.array([1, 6, 11])], } ) ) - row_lengths = torch.diff(out[CU_SEQLENS]) - self.assertEqual(row_lengths.tolist(), out[RESPONSE_LENGTHS].tolist()) - self.assertEqual( - out[INPUT_IDS].tolist(), - [_BASE_VOCAB_SIZE + 1, _BASE_VOCAB_SIZE + 6, _BASE_VOCAB_SIZE + 11], + self.assertEqual(torch.diff(out[CU_SEQLENS]).tolist(), [6, 3]) + self.assertEqual(out[RESPONSE_LENGTHS].tolist(), [3, 0]) + self.assertEqual(out[INPUT_IDS].numel(), 9) + + def test_a_batch_that_supervises_nothing_is_refused(self) -> None: + """Masking every row would leave the loss averaging over no tokens.""" + asm = _asm( + (_slot("hist", FillMode.INLINE),), + response=(_slot("answer", FillMode.INLINE),), ) + with self.assertRaisesRegex(ValueError, "nothing in this batch"): + asm( + _parsed( + { + "hist": [np.array([], dtype=np.int64)], + "answer": [np.array([1, 6, 11])], + } + ) + ) def test_a_multi_value_history_walks_like_the_flat_layout(self) -> None: """Items with key_lengths and one code per position are one stream.""" From 5043e1581e4262a638e21569c946822b82f07f01 Mon Sep 17 00:00:00 2001 From: root <597191244@qq.com> Date: Tue, 22 Sep 2026 07:31:46 +0000 Subject: [PATCH 14/15] [refactor] name the pad helper for what it returns _left_pad_packed_inputs put its loudest noun on the input, so a call site read `padded = ..._packed_inputs(...)`. The summary line now names the rectangle it produces and why the padding is on the left, leaving the packed input to the Args. Co-Authored-By: Claude Opus 5 (1M context) --- tzrec/models/genrec_causal_lm_model.py | 10 ++++++---- tzrec/models/genrec_causal_lm_model_test.py | 2 +- 2 files changed, 7 insertions(+), 5 deletions(-) diff --git a/tzrec/models/genrec_causal_lm_model.py b/tzrec/models/genrec_causal_lm_model.py index f4474e19..40a64962 100644 --- a/tzrec/models/genrec_causal_lm_model.py +++ b/tzrec/models/genrec_causal_lm_model.py @@ -214,7 +214,7 @@ def _padded_logits( Returns: ``(rows, suffix, vocab)`` logits over the response window. """ - padded, mask = self._left_pad_packed_inputs(embeds, batch) + padded, mask = self._left_pad(embeds, batch) if padded.shape[1] < suffix: raise ValueError( f"{type(self).__name__}: no sample reaches the supervised " @@ -244,19 +244,21 @@ def _generate(self, embeds: torch.Tensor, batch: Batch) -> torch.Tensor: Returns: ``(B, num_return, num_levels)`` local codes, best first. """ - padded, mask = self._left_pad_packed_inputs(embeds, batch) + padded, mask = self._left_pad(embeds, batch) tokens = dynamic_beam_search( self.lm, padded, mask, self._capped_widths, self._bands ) codes = self._tokens_to_local_codes(tokens, padded.shape[0]) return codes[:, : self._num_return_sequences, :] - def _left_pad_packed_inputs( + def _left_pad( self, embeds: torch.Tensor, batch: Batch, ) -> Tuple[torch.Tensor, torch.Tensor]: - """Left-pad packed prompt embeddings for beam decode and the padded forward. + """Pad the packed rows into a rectangle, newest token last. + + Used by beam decode and by the padded training forward. Args: embeds: packed embeddings, ``(total_tokens, hidden)``. diff --git a/tzrec/models/genrec_causal_lm_model_test.py b/tzrec/models/genrec_causal_lm_model_test.py index 43078965..0f29ea3f 100644 --- a/tzrec/models/genrec_causal_lm_model_test.py +++ b/tzrec/models/genrec_causal_lm_model_test.py @@ -70,7 +70,7 @@ def test_packs_rows_of_different_lengths(self) -> None: model = GenRecCausalLMModel.__new__(GenRecCausalLMModel) torch.nn.Module.__init__(model) - padded, mask = model._left_pad_packed_inputs(embeds, batch) + padded, mask = model._left_pad(embeds, batch) self.assertEqual(padded.shape, (2, 7, 2)) self.assertEqual( From c0c5835012ff2eebbef0ca55690bd25cfc2cd58e Mon Sep 17 00:00:00 2001 From: root <597191244@qq.com> Date: Tue, 22 Sep 2026 08:35:45 +0000 Subject: [PATCH 15/15] [refactor] take attn_kernel as the transformers name The value is passed through to attn_implementation unchanged, so the enum only bought a second spelling of names transformers already defines, plus a dict to translate between them. A string keeps the proto namespace clear and lets any implementation transformers accepts be configured without another proto change. The cost is that a misspelling now parses and fails when the backbone is built, rather than at config load; transformers names the accepted values there. Co-Authored-By: Claude Opus 5 (1M context) --- tzrec/models/genrec_causal_lm_model.py | 2 +- tzrec/models/genrec_causal_lm_model_test.py | 16 ++++++++-------- tzrec/models/genrec_model.py | 13 ++++--------- tzrec/models/genrec_model_test.py | 8 ++++---- tzrec/protos/models/genrec_model.proto | 8 +------- tzrec/tests/genrec_integration_test.py | 11 +++++------ tzrec/utils/test_util.py | 4 ++-- 7 files changed, 25 insertions(+), 37 deletions(-) diff --git a/tzrec/models/genrec_causal_lm_model.py b/tzrec/models/genrec_causal_lm_model.py index 40a64962..a19d3811 100644 --- a/tzrec/models/genrec_causal_lm_model.py +++ b/tzrec/models/genrec_causal_lm_model.py @@ -154,7 +154,7 @@ def _forward( suffix_offsets[None, :] < -infos[RESPONSE_LENGTHS][:, None], self._ignore_index, ) - if self._attn_kernel == GenRecModelConfig.FLASH_ATTENTION_2: + if self._attn_kernel == "flash_attention_2": logits = self._varlen_logits(embeds, batch, keep) else: logits = self._padded_logits(embeds, batch, suffix) diff --git a/tzrec/models/genrec_causal_lm_model_test.py b/tzrec/models/genrec_causal_lm_model_test.py index 0f29ea3f..7fbf4963 100644 --- a/tzrec/models/genrec_causal_lm_model_test.py +++ b/tzrec/models/genrec_causal_lm_model_test.py @@ -48,7 +48,7 @@ class PaddedForwardTest(unittest.TestCase): def test_a_padded_width_under_the_window_is_rejected(self) -> None: """``logits_to_keep`` would return fewer columns than there are labels.""" model = _stub_model() - model._attn_kernel = GenRecModelConfig.SDPA + model._attn_kernel = "sdpa" batch = _varlen_batch( cu_seqlens=(0, 3, 6), max_seqlen=3, @@ -119,7 +119,7 @@ def _stub_model(lm: Optional[nn.Module] = None) -> GenRecCausalLMModel: model = GenRecCausalLMModel.__new__(GenRecCausalLMModel) nn.Module.__init__(model) model.lm = _CapturingLM() if lm is None else lm - model._attn_kernel = GenRecModelConfig.FLASH_ATTENTION_2 + model._attn_kernel = "flash_attention_2" model._ignore_index = -7 model._prompt = SimpleNamespace(prompt_plan=SimpleNamespace(logits_suffix_len=4)) return model @@ -317,7 +317,7 @@ def setUp(self) -> None: def _eval_model( self, model_type: str, - attn_kernel: "GenRecModelConfig.AttnKernel", + attn_kernel: str, device: torch.device, lm_parameter_dtype: "GenRecModelConfig.ParamDtype", ): @@ -334,7 +334,7 @@ def _eval_model( def _assert_rows_match_solo_runs( self, model_type: str, - attn_kernel: "GenRecModelConfig.AttnKernel", + attn_kernel: str, device: torch.device, lm_parameter_dtype: "GenRecModelConfig.ParamDtype", ) -> None: @@ -386,7 +386,7 @@ def _assert_rows_match_solo_runs( def test_sdpa_rows_match_solo_runs(self, model_type: str) -> None: self._assert_rows_match_solo_runs( model_type, - GenRecModelConfig.SDPA, + "sdpa", torch.device("cpu"), GenRecModelConfig.FP32, ) @@ -400,7 +400,7 @@ def test_sdpa_rows_match_solo_runs(self, model_type: str) -> None: def test_flash_rows_match_solo_runs(self, model_type: str) -> None: self._assert_rows_match_solo_runs( model_type, - GenRecModelConfig.FLASH_ATTENTION_2, + "flash_attention_2", torch.device("cuda"), # the flash kernel takes fp16/bf16 only, and this carries no autocast GenRecModelConfig.BF16, @@ -420,8 +420,8 @@ def test_the_packed_and_padded_arms_agree(self, model_type: str) -> None: device = torch.device("cuda") outputs = [] for attn_kernel in ( - GenRecModelConfig.SDPA, - GenRecModelConfig.FLASH_ATTENTION_2, + "sdpa", + "flash_attention_2", ): model, compiled_prompt = self._eval_model( model_type, attn_kernel, device, GenRecModelConfig.BF16 diff --git a/tzrec/models/genrec_model.py b/tzrec/models/genrec_model.py index d589f064..985d0532 100644 --- a/tzrec/models/genrec_model.py +++ b/tzrec/models/genrec_model.py @@ -53,11 +53,6 @@ GenRecModelConfig.FP16: torch.float16, } -_ATTN_KERNEL: Dict[int, str] = { - GenRecModelConfig.SDPA: "sdpa", - GenRecModelConfig.FLASH_ATTENTION_2: "flash_attention_2", -} - _REQUIRED_LM_ATTRS: Tuple[str, ...] = ( "loss_function", "get_input_embeddings", @@ -103,7 +98,7 @@ def __init__( self._ignore_index = int(cfg.common.ignore_index) self.lm: nn.Module - self._attn_kernel: int + self._attn_kernel: str # the forward picks its layout from this: only the varlen kernel reads # cu_seq_lens, so every other kernel has to be handed padded rows self._attn_kernel = cfg.common.attn_kernel @@ -136,7 +131,7 @@ def init_backbone( self, hf_model_name_or_path: str, lm_parameter_dtype: int, - attn_kernel: int, + attn_kernel: str, ) -> None: """Assign ``self.lm`` from config, so HF weights load only on cold start. @@ -144,12 +139,12 @@ def init_backbone( hf_model_name_or_path: hub id or local directory naming the architecture and cold-start weights. lm_parameter_dtype: dtype of the LM parameters. - attn_kernel: attention kernel to build the backbone with. + attn_kernel: transformers attn_implementation name. """ config = AutoConfig.from_pretrained(hf_model_name_or_path) self.lm = AutoModelForCausalLM.from_config( config, - attn_implementation=_ATTN_KERNEL[attn_kernel], + attn_implementation=attn_kernel, torch_dtype=_PARAM_DTYPE[lm_parameter_dtype], ) self._check_backbone_interfaces(hf_model_name_or_path) diff --git a/tzrec/models/genrec_model_test.py b/tzrec/models/genrec_model_test.py index 9072448b..f62c142a 100644 --- a/tzrec/models/genrec_model_test.py +++ b/tzrec/models/genrec_model_test.py @@ -130,13 +130,13 @@ def test_rejects_a_model_built_without_a_prompt(self) -> None: @parameterized.expand( [ - [GenRecModelConfig.SDPA, "sdpa"], - [GenRecModelConfig.FLASH_ATTENTION_2, "flash_attention_2"], + ["sdpa"], + ["flash_attention_2"], ], name_func=parameterized_name_func, ) def test_builds_backbone_with_the_configured_kernel_and_dtype( - self, attn_kernel, expected_impl + self, attn_kernel: str ) -> None: # from_config is mocked, so this pins the kwargs init_backbone sends # without building a second backbone; setUp still builds a real one. @@ -164,7 +164,7 @@ def test_builds_backbone_with_the_configured_kernel_and_dtype( self.assertEqual( kwargs, { - "attn_implementation": expected_impl, + "attn_implementation": attn_kernel, "torch_dtype": torch.bfloat16, }, ) diff --git a/tzrec/protos/models/genrec_model.proto b/tzrec/protos/models/genrec_model.proto index e0a32a64..583caa3e 100644 --- a/tzrec/protos/models/genrec_model.proto +++ b/tzrec/protos/models/genrec_model.proto @@ -25,13 +25,7 @@ message GenRecModelConfig { // bf16 compute comes from mixed_precision, not from this. optional ParamDtype lm_parameter_dtype = 5 [default = FP32]; - // FLASH_ATTENTION_2 is what training should use; it needs the flash_attn - // wheel and a GPU. SDPA runs anywhere. - enum AttnKernel { - SDPA = 0; - FLASH_ATTENTION_2 = 1; - } - optional AttnKernel attn_kernel = 6 [default = SDPA]; + optional string attn_kernel = 6 [default = "sdpa"]; } message GenRecCausalLMModel { diff --git a/tzrec/tests/genrec_integration_test.py b/tzrec/tests/genrec_integration_test.py index a1fe38ec..fd456b89 100644 --- a/tzrec/tests/genrec_integration_test.py +++ b/tzrec/tests/genrec_integration_test.py @@ -32,7 +32,6 @@ ) from tzrec.prompt.compile import compile_prompt from tzrec.prompt.hole_keys import HOLE_KEYS, HoleKeyBuilder -from tzrec.protos.models.genrec_model_pb2 import GenRecModelConfig from tzrec.tests import utils from tzrec.utils import config_util from tzrec.utils.test_util import ( @@ -69,11 +68,11 @@ def tearDown(self): if self.success and os.path.exists(self.test_dir): shutil.rmtree(self.test_dir) - def _prepare_config(self, attn_kernel: Optional[int] = None) -> str: + def _prepare_config(self, attn_kernel: Optional[str] = None) -> str: """Write the tiny backbone, tokenizer, manifest and data; return the config. Args: - attn_kernel (int, optional): ``GenRecModelConfig.AttnKernel`` to run + attn_kernel (str, optional): transformers attn_implementation to run the pipeline on; the config's own default when None. Returns: @@ -105,11 +104,11 @@ def _prepare_config(self, attn_kernel: Optional[int] = None) -> str: config_util.save_message(config, config_path) return config_path - def _train_eval_export(self, attn_kernel: Optional[int] = None) -> str: + def _train_eval_export(self, attn_kernel: Optional[str] = None) -> str: """Run the pipeline; return the trained ``pipeline.config`` path. Args: - attn_kernel (int, optional): ``GenRecModelConfig.AttnKernel`` to run + attn_kernel (str, optional): transformers attn_implementation to run the pipeline on; the config's own default when None. Returns: @@ -225,7 +224,7 @@ def test_genrec_train_eval_export_on_the_varlen_kernel(self): The default kernel is SDPA, so the sibling case above covers the padded layout; this is the only end-to-end run of the packed one. """ - self._train_eval_export(attn_kernel=GenRecModelConfig.FLASH_ATTENTION_2) + self._train_eval_export(attn_kernel="flash_attention_2") @unittest.skipIf(*gpu_unavailable) @mark_ci_scope("gpu") diff --git a/tzrec/utils/test_util.py b/tzrec/utils/test_util.py index 0fb513e5..1abbace3 100644 --- a/tzrec/utils/test_util.py +++ b/tzrec/utils/test_util.py @@ -464,7 +464,7 @@ def create_genrec_test_model( beam_widths: Sequence[int] = (2, 2, 2), num_return_sequences: int = 2, lm_parameter_dtype: Optional["GenRecModelConfig.ParamDtype"] = None, - attn_kernel: Optional["GenRecModelConfig.AttnKernel"] = None, + attn_kernel: Optional[str] = None, model_type: str = "qwen2", init_seed: Optional[int] = None, ) -> Tuple[BaseModel, CompiledPrompt]: @@ -484,7 +484,7 @@ def create_genrec_test_model( beam_widths (Sequence[int]): per-level beam widths. num_return_sequences (int): sequences returned per sample. lm_parameter_dtype (optional): ``GenRecModelConfig.ParamDtype`` value. - attn_kernel (optional): ``GenRecModelConfig.AttnKernel`` value. + attn_kernel (optional): transformers attn_implementation name. model_type (str): hugging-face ``model_type`` of the backbone. init_seed (int, optional): seed the random init draws from, when the caller needs two builds to agree.