建立分批同步基线(基础文件)
This commit is contained in:
@@ -0,0 +1,247 @@
|
||||
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
|
||||
Reference in New Issue
Block a user