"""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