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