189 lines
8.3 KiB
Python
189 lines
8.3 KiB
Python
from __future__ import annotations
|
|
|
|
from collections.abc import Sequence
|
|
|
|
import torch
|
|
from torch import Tensor
|
|
|
|
from .types import AlignmentOutput, MODALITIES, SequenceBatch
|
|
|
|
|
|
def uniform_time_intervals(durations: Tensor, grid_size: int) -> Tensor:
|
|
"""Return ``[B, K, 2]`` equal-duration windows in seconds."""
|
|
if durations.ndim != 1 or grid_size < 1:
|
|
raise ValueError("durations must be [B] and grid_size must be positive")
|
|
if bool((durations <= 0).any()) or not bool(torch.isfinite(durations).all()):
|
|
raise ValueError("durations must be finite and positive")
|
|
edges = torch.linspace(
|
|
0.0, 1.0, grid_size + 1, device=durations.device, dtype=durations.dtype
|
|
)[None, :] * durations[:, None]
|
|
return torch.stack((edges[:, :-1], edges[:, 1:]), dim=-1)
|
|
|
|
|
|
def word_intervals_to_grid(word_intervals: Tensor, grid_size: int) -> Tensor:
|
|
"""Resample ordered word spans into K consecutive text-order intervals.
|
|
|
|
Each grid slot covers an equal share of transcript word order. Its time
|
|
boundaries are interpolated from forced-alignment word boundaries, so long
|
|
and short words retain their actual duration on the audio/video timeline.
|
|
"""
|
|
if word_intervals.ndim != 2 or word_intervals.shape[1] != 2:
|
|
raise ValueError("word_intervals must have shape [word_count, 2]")
|
|
if word_intervals.shape[0] == 0 or grid_size < 1:
|
|
raise ValueError("at least one word interval and a positive grid size are required")
|
|
intervals = word_intervals.to(dtype=torch.float32)
|
|
if not bool(torch.isfinite(intervals).all()):
|
|
raise ValueError("word intervals must be finite")
|
|
if bool((intervals[:, 1] < intervals[:, 0]).any()):
|
|
raise ValueError("word interval end must not precede its start")
|
|
if bool((intervals[1:, 0] < intervals[:-1, 0]).any()):
|
|
raise ValueError("word intervals must be ordered by start time")
|
|
|
|
word_count = intervals.shape[0]
|
|
if word_count == 1:
|
|
boundaries = torch.cat((intervals[:1, 0], intervals[:1, 1]))
|
|
else:
|
|
between = (intervals[:-1, 1] + intervals[1:, 0]) / 2
|
|
boundaries = torch.cat((intervals[:1, 0], between, intervals[-1:, 1]))
|
|
boundaries = torch.cummax(boundaries, dim=0).values
|
|
|
|
positions = torch.linspace(
|
|
0, word_count, grid_size + 1, device=intervals.device, dtype=intervals.dtype
|
|
)
|
|
left = positions.floor().long().clamp(max=word_count)
|
|
right = (left + 1).clamp(max=word_count)
|
|
fraction = (positions - left.to(positions.dtype)).unsqueeze(-1)
|
|
time_edges = boundaries[left] + fraction.squeeze(-1) * (boundaries[right] - boundaries[left])
|
|
return torch.stack((time_edges[:-1], time_edges[1:]), dim=-1)
|
|
|
|
|
|
def index_alignment(valid: Tensor, grid_size: int) -> tuple[Tensor, int]:
|
|
"""Map equal ranges of valid sequence order to K grid slots."""
|
|
if valid.ndim != 2 or valid.dtype != torch.bool:
|
|
raise ValueError("valid must be a boolean [B, L] tensor")
|
|
if grid_size < 1:
|
|
raise ValueError("grid_size must be positive")
|
|
batch_size, length = valid.shape
|
|
result = torch.zeros(batch_size, grid_size, length, device=valid.device, dtype=torch.float32)
|
|
fallback_count = 0
|
|
for batch_index in range(batch_size):
|
|
positions = torch.nonzero(valid[batch_index], as_tuple=False).flatten()
|
|
count = positions.numel()
|
|
if count == 0:
|
|
raise ValueError("each sample must contain a valid position")
|
|
for grid_index in range(grid_size):
|
|
start = (grid_index * count) // grid_size
|
|
end = ((grid_index + 1) * count) // grid_size
|
|
if start == end:
|
|
source_index = min(int((grid_index + 0.5) * count / grid_size), count - 1)
|
|
result[batch_index, grid_index, positions[source_index]] = 1.0
|
|
fallback_count += 1
|
|
else:
|
|
chosen = positions[start:end]
|
|
result[batch_index, grid_index, chosen] = 1.0 / chosen.numel()
|
|
return result, fallback_count
|
|
|
|
|
|
def interval_alignment(times: Tensor, valid: Tensor, intervals: Tensor) -> tuple[Tensor, int]:
|
|
"""Create row-normalized interval membership weights with nearest-time fallback."""
|
|
if times.ndim != 2 or valid.shape != times.shape or valid.dtype != torch.bool:
|
|
raise ValueError("times and valid must have matching [B, L] shapes")
|
|
if intervals.ndim != 3 or intervals.shape[0] != times.shape[0] or intervals.shape[2] != 2:
|
|
raise ValueError("intervals must have shape [B, K, 2]")
|
|
batch_size, length = times.shape
|
|
grid_size = intervals.shape[1]
|
|
result = torch.zeros(batch_size, grid_size, length, device=times.device, dtype=torch.float32)
|
|
fallback_count = 0
|
|
for batch_index in range(batch_size):
|
|
valid_positions = torch.nonzero(valid[batch_index], as_tuple=False).flatten()
|
|
valid_times = times[batch_index, valid_positions]
|
|
for grid_index in range(grid_size):
|
|
start, end = intervals[batch_index, grid_index]
|
|
is_last = grid_index == grid_size - 1
|
|
in_window = (valid_times >= start) & (
|
|
(valid_times <= end) if is_last else (valid_times < end)
|
|
)
|
|
chosen = valid_positions[in_window]
|
|
if chosen.numel() > 0:
|
|
result[batch_index, grid_index, chosen] = 1.0 / chosen.numel()
|
|
else:
|
|
center = (start + end) / 2
|
|
nearest = valid_positions[torch.argmin((valid_times - center).abs())]
|
|
result[batch_index, grid_index, nearest] = 1.0
|
|
fallback_count += 1
|
|
return result, fallback_count
|
|
|
|
|
|
def _output_from_weights(
|
|
sequences: dict[str, SequenceBatch],
|
|
weights: dict[str, Tensor],
|
|
fallbacks: dict[str, int],
|
|
) -> AlignmentOutput:
|
|
aligned = {name: torch.bmm(weights[name].to(sequences[name].features.dtype), sequences[name].features)
|
|
for name in MODALITIES}
|
|
output = AlignmentOutput(weights=weights, aligned=aligned, fallback_rows=fallbacks)
|
|
output.validate({name: sequences[name].valid for name in MODALITIES})
|
|
return output
|
|
|
|
|
|
def align_fixed_windows(
|
|
sequences: dict[str, SequenceBatch], durations: Tensor, grid_size: int = 50
|
|
) -> AlignmentOutput:
|
|
"""M2: average each modality inside the same K equal-duration windows."""
|
|
intervals = uniform_time_intervals(durations, grid_size)
|
|
weights: dict[str, Tensor] = {}
|
|
fallbacks: dict[str, int] = {}
|
|
for name in MODALITIES:
|
|
weights[name], fallbacks[name] = interval_alignment(
|
|
sequences[name].times, sequences[name].valid, intervals
|
|
)
|
|
return _output_from_weights(sequences, weights, fallbacks)
|
|
|
|
|
|
def align_forced_timestamps(
|
|
sequences: dict[str, SequenceBatch],
|
|
word_intervals: Sequence[Tensor],
|
|
grid_size: int = 50,
|
|
) -> AlignmentOutput:
|
|
"""M1: use forced word timestamps to define text-ordered grid intervals."""
|
|
batch_size = sequences["text"].features.shape[0]
|
|
if len(word_intervals) != batch_size:
|
|
raise ValueError("provide one ordered word-interval array per sample")
|
|
device = sequences["text"].features.device
|
|
intervals = torch.stack(
|
|
[word_intervals_to_grid(spans.to(device), grid_size) for spans in word_intervals], dim=0
|
|
)
|
|
weights: dict[str, Tensor] = {}
|
|
fallbacks: dict[str, int] = {}
|
|
for name in MODALITIES:
|
|
weights[name], fallbacks[name] = interval_alignment(
|
|
sequences[name].times, sequences[name].valid, intervals
|
|
)
|
|
return _output_from_weights(sequences, weights, fallbacks)
|
|
|
|
|
|
def make_block_mask(
|
|
batch_size: int,
|
|
grid_size: int,
|
|
ratio: float,
|
|
device: torch.device | str,
|
|
generator: torch.Generator | None = None,
|
|
) -> Tensor:
|
|
"""Sample one continuous masked interval per sequence on the common grid."""
|
|
if not 0 < ratio < 1:
|
|
raise ValueError("ratio must be between zero and one")
|
|
if batch_size < 1 or grid_size < 2:
|
|
raise ValueError("batch_size must be positive and grid_size at least two")
|
|
block_length = min(max(1, round(grid_size * ratio)), grid_size - 1)
|
|
mask = torch.zeros(batch_size, grid_size, dtype=torch.bool, device=device)
|
|
starts = torch.randint(
|
|
0,
|
|
grid_size - block_length + 1,
|
|
(batch_size,),
|
|
device=device,
|
|
generator=generator,
|
|
)
|
|
offsets = torch.arange(block_length, device=device)
|
|
mask[torch.arange(batch_size, device=device)[:, None], starts[:, None] + offsets] = True
|
|
return mask
|