InstructAV2AV / ovi /utils /av_edit_dataset.py
suimu's picture
init
e0177dc
Raw
History Blame Contribute Delete
10.1 kB
"""Canonical paired dataset for instruction-based audio-video editing."""
from __future__ import annotations
import random
from pathlib import Path
from typing import Any, Optional
import torch
from ovi.utils.av_edit_data import (
get_instruction,
get_video_info,
load_audio,
load_audio_array,
load_manifest,
load_video_array,
pad_or_trim_audio,
resolve_media_path,
snap_num_frames,
)
class AVEditDataset(torch.utils.data.Dataset):
"""Load source/target AV pairs from one canonical manifest.
Required media columns depend on ``has_video`` and ``has_audio``. In AV
mode, audio columns are optional and fall back to the corresponding video
tracks. In audio-only mode, source and target audio columns are required.
"""
def __init__(
self,
metadata_path: str,
base_path: Optional[str] = None,
source_video_column: str = "source_video",
source_audio_column: str = "source_audio",
target_video_column: str = "target_video",
target_audio_column: str = "target_audio",
instruction_column: str = "instruction",
has_video: bool = True,
has_audio: bool = True,
height: Optional[int] = 704,
width: Optional[int] = 1280,
max_pixels: int = 1920 * 1080,
max_num_frames: Optional[int] = 121,
fix_num_frames: Optional[int] = None,
audio_sample_rate: int = 16000,
max_audio_seconds: Optional[float] = None,
fix_audio_seconds: Optional[float] = None,
repeat: int = 1,
max_retries: int = 20,
) -> None:
self.metadata_path = str(Path(metadata_path).expanduser().resolve())
self.base_path = base_path
self.rows = load_manifest(self.metadata_path)
if not self.rows:
raise ValueError(f"Manifest is empty: {self.metadata_path}")
self.source_video_column = source_video_column
self.source_audio_column = source_audio_column
self.target_video_column = target_video_column
self.target_audio_column = target_audio_column
self.instruction_column = instruction_column
self.has_video = bool(has_video)
self.has_audio = bool(has_audio)
self.height = height
self.width = width
self.max_pixels = int(max_pixels)
self.max_num_frames = None if max_num_frames is None else int(max_num_frames)
self.fix_num_frames = None if fix_num_frames is None else int(fix_num_frames)
self.audio_sample_rate = int(audio_sample_rate)
self.max_audio_seconds = (
None if max_audio_seconds is None else float(max_audio_seconds)
)
self.fix_audio_seconds = (
None if fix_audio_seconds is None else float(fix_audio_seconds)
)
self.repeat = int(repeat)
self.max_retries = int(max_retries)
if not self.has_video and not self.has_audio:
raise ValueError("At least one of has_video or has_audio must be true.")
if self.repeat < 1:
raise ValueError("repeat must be at least 1.")
if (
self.has_video
and self.fix_num_frames is not None
and self.fix_num_frames % 4 != 1
):
raise ValueError("fix_num_frames must satisfy num_frames % 4 == 1.")
if self.max_audio_seconds is not None and self.max_audio_seconds <= 0:
raise ValueError("max_audio_seconds must be positive.")
if self.fix_audio_seconds is not None and self.fix_audio_seconds <= 0:
raise ValueError("fix_audio_seconds must be positive.")
if self.max_audio_seconds is not None and self.fix_audio_seconds is not None:
raise ValueError(
"Set only one of max_audio_seconds or fix_audio_seconds."
)
def __len__(self) -> int:
return len(self.rows) * self.repeat
def _path(self, row: dict[str, Any], column: str, required: bool) -> Optional[Path]:
if column not in row:
if required:
raise KeyError(f"Manifest is missing required column '{column}'.")
return None
return resolve_media_path(
row[column],
manifest_path=self.metadata_path,
base_path=self.base_path,
required=required,
)
def _load_row(self, index: int) -> dict[str, Any]:
row = self.rows[index % len(self.rows)].copy()
source_video_path = self._path(
row, self.source_video_column, required=self.has_video
)
target_video_path = self._path(
row, self.target_video_column, required=self.has_video
)
source_audio_path = self._path(
row,
self.source_audio_column,
required=self.has_audio and not self.has_video,
)
target_audio_path = self._path(
row,
self.target_audio_column,
required=self.has_audio and not self.has_video,
)
row["instruction"] = get_instruction(row, self.instruction_column)
duration_seconds = None
if self.has_video:
source_fps, source_frames = get_video_info(source_video_path)
_, target_frames = get_video_info(target_video_path)
available_frames = min(source_frames, target_frames)
if self.fix_num_frames is not None:
if available_frames < self.fix_num_frames:
raise ValueError(
f"Pair has {available_frames} frames but "
f"fix_num_frames={self.fix_num_frames}."
)
num_frames = self.fix_num_frames
else:
if self.max_num_frames is not None:
available_frames = min(available_frames, self.max_num_frames)
num_frames = snap_num_frames(available_frames)
source_video, target_size = load_video_array(
source_video_path,
num_frames=num_frames,
height=self.height,
width=self.width,
max_pixels=self.max_pixels,
)
target_video, _ = load_video_array(
target_video_path,
num_frames=num_frames,
height=self.height,
width=self.width,
max_pixels=self.max_pixels,
target_size=target_size,
)
duration_seconds = num_frames / source_fps
row.update(
{
"video_ori_np": source_video,
"video_result_np": target_video,
"source_video_path": str(source_video_path),
"target_video_path": str(target_video_path),
"num_frames": num_frames,
"fps": source_fps,
}
)
if self.has_audio:
resolved_source_audio = source_audio_path or source_video_path
resolved_target_audio = target_audio_path or target_video_path
if resolved_source_audio is None or resolved_target_audio is None:
raise ValueError("Audio-only training requires source and target audio files.")
if duration_seconds is not None:
num_audio_samples = max(
1, round(duration_seconds * self.audio_sample_rate)
)
source_audio = load_audio_array(
resolved_source_audio,
sample_rate=self.audio_sample_rate,
num_samples=num_audio_samples,
)
target_audio = load_audio_array(
resolved_target_audio,
sample_rate=self.audio_sample_rate,
num_samples=num_audio_samples,
)
else:
source_audio = load_audio(
resolved_source_audio, sample_rate=self.audio_sample_rate
)
target_audio = load_audio(
resolved_target_audio, sample_rate=self.audio_sample_rate
)
if self.fix_audio_seconds is not None:
num_audio_samples = max(
1, round(self.fix_audio_seconds * self.audio_sample_rate)
)
else:
num_audio_samples = max(len(source_audio), len(target_audio))
if self.max_audio_seconds is not None:
num_audio_samples = min(
num_audio_samples,
max(
1,
round(
self.max_audio_seconds * self.audio_sample_rate
),
),
)
source_audio = pad_or_trim_audio(source_audio, num_audio_samples)
target_audio = pad_or_trim_audio(target_audio, num_audio_samples)
row.update(
{
"audio_ori_np": source_audio,
"audio_result_np": target_audio,
"source_audio_path": str(resolved_source_audio),
"target_audio_path": str(resolved_target_audio),
"num_audio_samples": len(source_audio),
"audio_sample_rate": self.audio_sample_rate,
}
)
return row
def __getitem__(self, index: int) -> dict[str, Any]:
errors = []
candidate = index
for _ in range(self.max_retries):
try:
return self._load_row(candidate)
except Exception as error:
errors.append(f"row {candidate % len(self.rows)}: {error}")
candidate = random.randrange(len(self.rows))
details = "\n".join(errors[-3:])
raise RuntimeError(f"Failed to load an AV editing sample after retries:\n{details}")
UnifiedAVEditDataset = AVEditDataset