Files
modeling_zhaocui/deep_learning/Q1/q1/alignment.py
T

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