from __future__ import annotations from collections.abc import Mapping import torch from torch import Tensor, nn from .alignment import index_alignment from .types import AlignmentOutput, MODALITIES, SequenceBatch def _sinusoidal_position_encoding(length: int, dimension: int) -> Tensor: """Build a fixed, small-amplitude sinusoidal code for ordered slots.""" positions = torch.arange(length, dtype=torch.float32).unsqueeze(1) frequencies = torch.exp( torch.arange(0, dimension, 2, dtype=torch.float32) * (-torch.log(torch.tensor(10000.0)) / dimension) ) encoding = torch.zeros(length, dimension, dtype=torch.float32) encoding[:, 0::2] = torch.sin(positions * frequencies) odd_width = encoding[:, 1::2].shape[1] if odd_width: encoding[:, 1::2] = torch.cos(positions * frequencies[:odd_width]) return encoding * (dimension**-0.5) def _temporal_position_encoding(times: Tensor, dimension: int) -> Tensor: """Encode normalized source times with a fixed Fourier feature bank.""" half_width = (dimension + 1) // 2 frequencies = torch.logspace( 0.0, 1.6989700043360187, steps=half_width, device=times.device, dtype=times.dtype, ) angles = (2.0 * torch.pi) * times.unsqueeze(-1) * frequencies encoding = torch.empty(*times.shape, dimension, device=times.device, dtype=times.dtype) encoding[..., 0::2] = torch.sin(angles) if dimension > 1: encoding[..., 1::2] = torch.cos(angles[..., : encoding[..., 1::2].shape[-1]]) return encoding * (dimension**-0.25) def _validate_inputs(sequences: Mapping[str, SequenceBatch], grid_size: int) -> int: if set(sequences) != set(MODALITIES): raise ValueError(f"sequences must contain exactly {MODALITIES}") batch_sizes = {sequences[name].features.shape[0] for name in MODALITIES} if len(batch_sizes) != 1: raise ValueError("all modalities must have the same batch size") if grid_size < 1: raise ValueError("grid_size must be positive") return batch_sizes.pop() class _CrossAttention(nn.Module): def __init__(self, dimension: int, heads: int, dropout: float) -> None: super().__init__() if dimension % heads != 0: raise ValueError("dimension must be divisible by heads") self.attention = nn.MultiheadAttention( # Keep the returned alignment matrix row-stochastic during training. # PyTorch applies attention dropout to returned weights when it is # nonzero, which breaks the shared AlignmentOutput contract. embed_dim=dimension, num_heads=heads, dropout=0.0, batch_first=True ) self.input_dropout = nn.Dropout(dropout) def forward( self, query: Tensor, source: Tensor, valid: Tensor, *, source_position: Tensor | None = None, ) -> tuple[Tensor, Tensor]: key = source if source_position is None else source + source_position values, weights = self.attention( self.input_dropout(query), self.input_dropout(key), self.input_dropout(source), key_padding_mask=~valid, need_weights=True, average_attn_weights=True, ) return values, weights class TextAnchoredCrossAttention(nn.Module): """M3: transcript-order text slots query Audio and Vision sequences.""" def __init__( self, dimensions: Mapping[str, int], grid_size: int = 50, hidden_size: int = 128, heads: int = 4, dropout: float = 0.1, source_time_encoding: bool = False, ) -> None: super().__init__() if set(dimensions) != set(MODALITIES): raise ValueError(f"dimensions must contain exactly {MODALITIES}") self.grid_size = grid_size self.source_time_encoding = source_time_encoding self.hidden_size = hidden_size self.projections = nn.ModuleDict( {name: nn.Linear(dimensions[name], hidden_size) for name in MODALITIES} ) self.audio_attention = _CrossAttention(hidden_size, heads, dropout) self.vision_attention = _CrossAttention(hidden_size, heads, dropout) def forward( self, sequences: Mapping[str, SequenceBatch], durations: Tensor | None = None ) -> AlignmentOutput: _validate_inputs(sequences, self.grid_size) if self.source_time_encoding: if durations is None or durations.shape != (sequences["text"].features.shape[0],): raise ValueError("durations with shape [B] are required for source time encoding") durations = durations.to(device=sequences["text"].times.device).clamp_min(1e-8) projected = { name: self.projections[name](sequences[name].features) for name in MODALITIES } text_weights, text_fallbacks = index_alignment(sequences["text"].valid, self.grid_size) text_query = torch.bmm(text_weights.to(projected["text"].dtype), projected["text"]) source_positions: dict[str, Tensor] = {} if self.source_time_encoding: assert durations is not None text_centers = torch.bmm( text_weights.to(sequences["text"].times.dtype), sequences["text"].times.unsqueeze(-1), ).squeeze(-1) / durations[:, None] text_query = text_query + _temporal_position_encoding( text_centers, self.hidden_size ).to(text_query.dtype) for name in ("audio", "vision"): normalized_times = sequences[name].times / durations[:, None] source_positions[name] = _temporal_position_encoding( normalized_times, self.hidden_size ).to(projected[name].dtype) audio_values, audio_weights = self.audio_attention( text_query, projected["audio"], sequences["audio"].valid, source_position=source_positions.get("audio"), ) vision_values, vision_weights = self.vision_attention( text_query, projected["vision"], sequences["vision"].valid, source_position=source_positions.get("vision"), ) output = AlignmentOutput( weights={"text": text_weights, "audio": audio_weights, "vision": vision_weights}, aligned={"text": text_query, "audio": audio_values, "vision": vision_values}, fallback_rows={"text": text_fallbacks, "audio": 0, "vision": 0}, ) output.validate({name: sequences[name].valid for name in MODALITIES}) return output class SharedLatentTimeline(nn.Module): """M4: K learned shared slots attend independently to all three modalities.""" def __init__( self, dimensions: Mapping[str, int], grid_size: int = 50, hidden_size: int = 128, heads: int = 4, dropout: float = 0.1, absolute_position_encoding: bool = False, source_time_encoding: bool = False, ) -> None: super().__init__() if set(dimensions) != set(MODALITIES): raise ValueError(f"dimensions must contain exactly {MODALITIES}") if hidden_size % heads != 0: raise ValueError("hidden_size must be divisible by heads") self.grid_size = grid_size self.hidden_size = hidden_size self.absolute_position_encoding = absolute_position_encoding self.source_time_encoding = source_time_encoding self.register_buffer( "sinusoidal_positions", _sinusoidal_position_encoding(grid_size, hidden_size), persistent=False, ) self.projections = nn.ModuleDict( {name: nn.Linear(dimensions[name], hidden_size) for name in MODALITIES} ) self.slots = nn.Parameter(torch.empty(grid_size, hidden_size)) nn.init.normal_(self.slots, mean=0.0, std=hidden_size**-0.5) self.attention = nn.ModuleDict( {name: _CrossAttention(hidden_size, heads, dropout) for name in MODALITIES} ) def forward( self, sequences: Mapping[str, SequenceBatch], durations: Tensor | None = None ) -> AlignmentOutput: batch_size = _validate_inputs(sequences, self.grid_size) if self.source_time_encoding: if durations is None or durations.shape != (batch_size,): raise ValueError("durations with shape [B] are required for source time encoding") durations = durations.to(device=sequences["text"].times.device).clamp_min(1e-8) latent_queries = self.slots.unsqueeze(0).expand(batch_size, -1, -1) if self.absolute_position_encoding: latent_queries = latent_queries + self.sinusoidal_positions.unsqueeze(0) if self.source_time_encoding: centers = ( torch.arange( self.grid_size, dtype=sequences["text"].times.dtype, device=sequences["text"].times.device, ) + 0.5 ) / self.grid_size centers = centers.unsqueeze(0).expand(batch_size, -1) latent_queries = latent_queries + _temporal_position_encoding( centers, self.hidden_size ).to(latent_queries.dtype) weights: dict[str, Tensor] = {} aligned: dict[str, Tensor] = {} for name in MODALITIES: source = self.projections[name](sequences[name].features) source_position = None if self.source_time_encoding: assert durations is not None normalized_times = sequences[name].times / durations[:, None] source_position = _temporal_position_encoding( normalized_times, self.hidden_size ).to(source.dtype) aligned[name], weights[name] = self.attention[name]( latent_queries, source, sequences[name].valid, source_position=source_position, ) output = AlignmentOutput( weights=weights, aligned=aligned, fallback_rows={name: 0 for name in MODALITIES}, ) output.validate({name: sequences[name].valid for name in MODALITIES}) return output