Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
18 commits
Select commit Hold shift + click to select a range
98310f4
[perf] run GenRec with packed FlashAttention
WhiteSwan1 Sep 2, 2026
d85731f
[chore] merge upstream master into GenRec FlashAttention branch
WhiteSwan1 Sep 8, 2026
da3636e
[chore] merge upstream master into GenRec FlashAttention branch
WhiteSwan1 Sep 16, 2026
163e736
[fix] make GenRec a GPU-only, flash-attention-2 path
WhiteSwan1 Sep 16, 2026
4f5e298
[fix] run the genrec model tests on the GPU they now require
WhiteSwan1 Sep 16, 2026
4ec6768
[fix] skip the GenRec tests on the flash_attn wheel, not just on a GPU
WhiteSwan1 Sep 16, 2026
342660d
[feat] choose the GenRec attention kernel from the config
WhiteSwan1 Sep 17, 2026
01ef25a
[fix] return the GenRec tests to the CPU lane, and pin transformers
WhiteSwan1 Sep 17, 2026
6df33ca
[fix] read the attention kernel from our own config, not HF's private…
WhiteSwan1 Sep 17, 2026
3d82728
[fix] give each attention kernel the layout it can actually serve
WhiteSwan1 Sep 18, 2026
233a152
[chore] drop the cuDNN exclusion and a test it outlived
WhiteSwan1 Sep 21, 2026
45b3585
[fix] give the padded path its own positions, and cover it end to end
WhiteSwan1 Sep 21, 2026
d96493e
[chore] merge upstream master into GenRec FlashAttention branch
WhiteSwan1 Sep 21, 2026
27c2534
[chore] ship flash_attn for cp310 and cp312, not just cp311
WhiteSwan1 Sep 22, 2026
66dc5ec
[chore] bump version to 1.4.14
WhiteSwan1 Sep 22, 2026
53ab60b
[bugfix] filter prompt-less rows instead of failing the batch
WhiteSwan1 Sep 22, 2026
5043e15
[refactor] name the pad helper for what it returns
WhiteSwan1 Sep 22, 2026
c0c5835
[refactor] take attn_kernel as the transformers name
WhiteSwan1 Sep 22, 2026
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions requirements/cu126.txt
Original file line number Diff line number Diff line change
Expand Up @@ -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
3 changes: 3 additions & 0 deletions requirements/cu129.txt
Original file line number Diff line number Diff line change
Expand Up @@ -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
3 changes: 3 additions & 0 deletions requirements/cu130.txt
Original file line number Diff line number Diff line change
Expand Up @@ -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
2 changes: 1 addition & 1 deletion requirements/runtime.txt
Original file line number Diff line number Diff line change
Expand Up @@ -23,4 +23,4 @@ tensorboard
torch==2.13.0
torchmetrics==1.0.3
torchrec==1.8.0
transformers
transformers>=4.56
124 changes: 94 additions & 30 deletions tzrec/models/genrec_causal_lm_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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)
Comment thread
WhiteSwan1 marked this conversation as resolved.
# (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.
Expand All @@ -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]
Comment thread
WhiteSwan1 marked this conversation as resolved.
Expand All @@ -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.
Expand Down
Loading
Loading