abe123's picture
Squash clean branch history
1b7bd7b
Raw
History Blame Contribute Delete
2.34 kB
"""Attention-mask helpers shared by repo-owned decoder LMs."""
from typing import Optional
import torch
def build_decoder_attention_mask(
input_ids: torch.Tensor,
pad_token_id: int,
eos_token_id: int,
sequence_boundary_policy: str,
attention_mask: Optional[torch.Tensor] = None,
segment_boundary_token_id: Optional[int] = None,
bidirectional: bool = False,
) -> torch.Tensor:
if input_ids.ndim != 2:
raise ValueError(
"input_ids must be rank-2 [batch, seq], "
f"got shape {tuple(input_ids.shape)}"
)
_, seq_len = input_ids.shape
if attention_mask is None:
valid_tokens = input_ids != pad_token_id
else:
valid_tokens = attention_mask.bool()
query_mask = valid_tokens.unsqueeze(2)
key_mask = valid_tokens.unsqueeze(1)
if bidirectional:
mask = query_mask & key_mask
else:
directionality = torch.tril(
torch.ones(seq_len, seq_len, dtype=torch.bool, device=input_ids.device)
).unsqueeze(0)
mask = directionality & query_mask & key_mask
if sequence_boundary_policy == "none":
return mask
if sequence_boundary_policy == "segment_document":
if segment_boundary_token_id is None:
raise ValueError(
"segment_boundary_token_id is required when "
"sequence_boundary_policy='segment_document'"
)
# The boundary token starts the next segment: cumsum increments on the
# boundary position, so the marker attends with the following tokens.
segment_ids = torch.cumsum(
input_ids == segment_boundary_token_id, dim=1
)
same_segment = segment_ids.unsqueeze(1) == segment_ids.unsqueeze(2)
return mask & same_segment
if sequence_boundary_policy != "eos_document":
raise ValueError(f"Unsupported sequence_boundary_policy: {sequence_boundary_policy}")
# Next-token training predicts EOS from the preceding document token.
# Once EOS is present as an input token, it starts the next segment so the
# following document is not predicted with prior-document context.
document_ids = torch.cumsum(input_ids == eos_token_id, dim=1)
same_document = document_ids.unsqueeze(1) == document_ids.unsqueeze(2)
return mask & same_document