248 lines
10 KiB
Python
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
|