Files

248 lines
10 KiB
Python

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