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