diff --git a/requirements/cu126.txt b/requirements/cu126.txt index 73188a36..7bc6065f 100644 --- a/requirements/cu126.txt +++ b/requirements/cu126.txt @@ -7,4 +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%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 21084d35..c6ee3f1d 100644 --- a/requirements/cu129.txt +++ b/requirements/cu129.txt @@ -7,4 +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%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 4a98f53a..a96aa8c3 100644 --- a/requirements/cu130.txt +++ b/requirements/cu130.txt @@ -7,5 +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%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 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.py b/tzrec/models/genrec_causal_lm_model.py index 037582ee..a19d3811 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,96 @@ 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[CU_SEQLENS] + suffix = cast(int, self._prompt.prompt_plan.logits_suffix_len) + 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 + keep = cu_seqlens[1:, None] + suffix_offsets + labels = infos[INPUT_IDS][keep].masked_fill( + suffix_offsets[None, :] < -infos[RESPONSE_LENGTHS][:, None], + self._ignore_index, + ) + if self._attn_kernel == "flash_attention_2": + logits = self._varlen_logits(embeds, batch, keep) + else: + logits = self._padded_logits(embeds, batch, suffix) + return logits, labels + + 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. + + 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.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, ) - suffix = self._prompt.prompt_plan.logits_suffix_len + return outputs.logits.reshape(keep.shape[0], keep.shape[1], -1) + + 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(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) outputs = self.lm( inputs_embeds=padded, attention_mask=mask, + position_ids=position_ids, use_cache=False, logits_to_keep=suffix, ) - return outputs.logits, labels[:, -suffix:] + return outputs.logits def _generate(self, embeds: torch.Tensor, batch: Batch) -> torch.Tensor: """Beam-search the SID answer. @@ -165,28 +244,28 @@ 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, - 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]: + """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)``. 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[CU_SEQLENS] @@ -202,32 +281,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[INPUT_IDS] - response_lengths = infos[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 8c9da960..7fbf4963 100644 --- a/tzrec/models/genrec_causal_lm_model_test.py +++ b/tzrec/models/genrec_causal_lm_model_test.py @@ -10,10 +10,13 @@ # limitations under the License. import unittest -from unittest import mock +from types import SimpleNamespace +from typing import Optional, Sequence 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 @@ -24,37 +27,50 @@ RESPONSE_LENGTHS, PromptAssembler, ) +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, make_test_dir, + mark_ci_scope, + nv_gpu_unavailable, parameterized_name_func, ) -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 = "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) 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={ CU_SEQLENS: cu, - INPUT_IDS: input_ids, MAX_SEQLEN: torch.tensor(7), - 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(embeds, batch) self.assertEqual(padded.shape, (2, 7, 2)) self.assertEqual( @@ -66,9 +82,140 @@ 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 + embeds = kwargs["inputs_embeds"] + count = kwargs["logits_to_keep"].numel() + return SimpleNamespace( + logits=torch.zeros(1, count, 5, dtype=embeds.dtype, device=embeds.device) + ) + + +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)) + + +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, + 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 = "flash_attention_2" + 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, +) -> Batch: + """The four varlen keys ``_forward`` reads, over ``cu_seqlens[-1]`` tokens. + + ``PromptAssembler`` emits ``cu_seqlens`` as int32; int64 here is what + makes ``_varlen_logits``' cast to int32 do visible work. + """ + return Batch( + additional_infos={ + 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), + } + ) + + +class PackedForwardTest(unittest.TestCase): + def test_passes_varlen_metadata_and_builds_per_row_labels(self) -> None: + model = _stub_model() + embeds = torch.arange(72, dtype=torch.float32).reshape(12, 6) + batch = _varlen_batch() + + 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_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=(0, 3), + ) + + _, 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)) + embeds = torch.randn( + 12, 6, generator=torch.Generator().manual_seed(1), requires_grad=True + ) + batch = _varlen_batch(first_input_id=4) + + 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 ) @@ -77,6 +224,18 @@ class GenRecCausalLMModelTest(unittest.TestCase): def setUp(self) -> None: self.test_dir = make_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 + ) -> 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( [ @@ -87,47 +246,197 @@ def setUp(self) -> None: name_func=parameterized_name_func, ) def test_beam_widths_are_capped_once_at_init(self, beam_widths, expected) -> None: - model, compiled_prompt = create_genrec_test_model( - self.test_dir, 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 = compiled_prompt.sid_space + 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"): - create_genrec_test_model(self.test_dir, 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"): - create_genrec_test_model(self.test_dir, 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\\)"): - create_genrec_test_model( - self.test_dir, beam_widths=(1, 1, 100), num_return_sequences=5 + self._beam_model( + beam_widths=(1, 1, 100), + 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]), - } + +# 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: + """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 + + +# 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], +) + + +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() + + def _eval_model( + self, + model_type: str, + attn_kernel: str, + 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=attn_kernel, + lm_parameter_dtype=lm_parameter_dtype, + init_seed=0, + ) + return model.to(device).eval(), compiled_prompt + + def _assert_rows_match_solo_runs( + self, + model_type: str, + attn_kernel: str, + 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) + 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, + [_REWRITTENGENREC_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) - self.assertIs(spy.call_args.kwargs["use_cache"], False) + @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, + "sdpa", + torch.device("cpu"), + GenRecModelConfig.FP32, + ) + + # 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, + "flash_attention_2", + torch.device("cuda"), + # the flash kernel takes fp16/bf16 only, and this carries no autocast + 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 ( + "sdpa", + "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/models/genrec_model.py b/tzrec/models/genrec_model.py index ebe15408..985d0532 100644 --- a/tzrec/models/genrec_model.py +++ b/tzrec/models/genrec_model.py @@ -98,7 +98,15 @@ 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._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 + self.init_backbone( + cfg.hf_model_name_or_path, + cfg.common.lm_parameter_dtype, + cfg.common.attn_kernel, + ) # 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 +128,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_kernel: str, ) -> None: """Assign ``self.lm`` from config, so HF weights load only on cold start. @@ -128,10 +139,14 @@ 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: transformers attn_implementation name. """ 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=attn_kernel, + 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 fb6dd618..f62c142a 100644 --- a/tzrec/models/genrec_model_test.py +++ b/tzrec/models/genrec_model_test.py @@ -11,6 +11,8 @@ import os import unittest +from typing import Tuple +from unittest import mock import torch import torch.fx @@ -26,7 +28,7 @@ GenRecFrontEnd, project_slots, ) -from tzrec.models.model import ScriptWrapper, TrainWrapper +from tzrec.models.model import BaseModel, ScriptWrapper, TrainWrapper from tzrec.prompt.assembler import ( HOLE_POSITIONS, HOLE_SLOT_COUNTS, @@ -41,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( @@ -84,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]), @@ -128,6 +128,48 @@ 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) + @parameterized.expand( + [ + ["sdpa"], + ["flash_attention_2"], + ], + name_func=parameterized_name_func, + ) + def test_builds_backbone_with_the_configured_kernel_and_dtype( + 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. + stand_in = AutoModelForCausalLM.from_pretrained( + os.path.join(self.test_dir, "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, + ): + create_genrec_test_model( + self.test_dir, + lm_parameter_dtype=GenRecModelConfig.BF16, + attn_kernel=attn_kernel, + ) + + # 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") + self.assertEqual( + kwargs, + { + "attn_implementation": attn_kernel, + "torch_dtype": torch.bfloat16, + }, + ) + to_mock.assert_not_called() + def test_shared_projection_name_requires_matching_widths(self) -> None: with self.assertRaisesRegex(ValueError, "cannot share a module"): create_genrec_test_model( @@ -144,15 +186,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]) @@ -171,14 +218,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]) @@ -192,6 +232,31 @@ 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: + model, compiled_prompt = create_genrec_test_model( + self.test_dir, + feature_configs=[_hist(), _projected("prof", 8)], + prompt="History : {{hist}} . Predict {{prof}} :", + ) + init_parameters(model, device=torch.device("cpu")) + batch = _projected_batch(compiled_prompt) + + wrapper = TrainWrapper( + model, device=torch.device("cpu"), 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: self.model.init_metric() for value in (1.0, 3.0): @@ -221,25 +286,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)) @@ -258,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/modules/dynamic_beam_test.py b/tzrec/modules/dynamic_beam_test.py index bbb81eac..537f0898 100644 --- a/tzrec/modules/dynamic_beam_test.py +++ b/tzrec/modules/dynamic_beam_test.py @@ -17,7 +17,13 @@ from parameterized import parameterized 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, + flash_attn_unavailable, + 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,42 @@ 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) +@unittest.skipIf(*flash_attn_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") + 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]], + 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 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 3e30e333..cbbec266 100644 --- a/tzrec/prompt/assembler_test.py +++ b/tzrec/prompt/assembler_test.py @@ -222,6 +222,46 @@ def test_response_is_optional_and_its_length_is_recorded(self) -> None: ) self.assertEqual(prompt_only[RESPONSE_LENGTHS].tolist(), [0]) + 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),), + response=(_slot("answer", FillMode.INLINE),), + ) + out = asm( + _parsed( + { + "hist": [np.array([1, 6, 11]), np.array([], dtype=np.int64)], + "answer": [np.array([2, 7, 10]), np.array([1, 6, 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.""" asm = _asm((Static((7,)), _slot("hist", FillMode.INLINE))) diff --git a/tzrec/protos/models/genrec_model.proto b/tzrec/protos/models/genrec_model.proto index d8d86edd..583caa3e 100644 --- a/tzrec/protos/models/genrec_model.proto +++ b/tzrec/protos/models/genrec_model.proto @@ -24,6 +24,8 @@ 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]; + + optional string attn_kernel = 6 [default = "sdpa"]; } message GenRecCausalLMModel { 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/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/genrec_integration_test.py b/tzrec/tests/genrec_integration_test.py index 6f63204a..fd456b89 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 @@ -36,9 +37,11 @@ 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 +68,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[str] = None) -> str: + """Write the tiny backbone, tokenizer, manifest and data; return the config. + + Args: + attn_kernel (str, optional): transformers attn_implementation 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 +98,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[str] = None) -> str: + """Run the pipeline; return the trained ``pipeline.config`` path. + + Args: + attn_kernel (str, optional): transformers attn_implementation 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 +215,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="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 7d221e5b..1abbace3 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] = ( @@ -61,6 +64,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( @@ -128,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. @@ -138,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] @@ -382,44 +400,42 @@ def create_genrec_test_tokenizer( return path -def create_genrec_test_model( +# 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, 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, -) -> 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. 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( @@ -436,7 +452,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[str] = 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): 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. + + 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 @@ -445,7 +508,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 - 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 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"