Bad-Apple_Unified-Model-Test / bad_apple_nn.py
Ma7ee7's picture
Upload 5 files
93b70dc verified
Raw
History Blame Contribute Delete
107 kB
from __future__ import annotations
import argparse
import contextlib
import json
import math
import random
import subprocess
import wave
from dataclasses import asdict, dataclass
from pathlib import Path
import numpy as np
import torch
from torch import nn
from torch.nn import functional as F
from torch.utils.checkpoint import checkpoint
import cv2
cv2.setNumThreads(1)
@dataclass(frozen=True)
class VideoSpec:
width: int
height: int
fps: float
frame_stride: int
threshold: int | None
frame_count: int
@property
def duration(self) -> float:
return self.frame_count / self.fps
@dataclass(frozen=True)
class AudioSpec:
sample_rate: int
sample_count: int
channels: int = 1
@property
def duration(self) -> float:
return self.sample_count / self.sample_rate
RenderThreshold = int | str | None
def safe_name(
video_path: Path,
width: int,
height: int,
stride: int,
threshold: int | None,
max_frames: int | None,
) -> str:
threshold_name = "gray" if threshold is None else f"thr{threshold}"
frame_name = "full" if max_frames is None else f"n{max_frames}"
return f"{video_path.stem}_{width}x{height}_s{stride}_{threshold_name}_{frame_name}"
def parse_data_threshold(value: str | int | None) -> int | None:
if value is None:
return None
if isinstance(value, int):
return None if value < 0 else value
text = value.strip().lower()
if text in {"none", "gray", "grayscale", "soft", "-1"}:
return None
threshold = int(text)
return None if threshold < 0 else threshold
def parse_render_threshold(value: str | int | None) -> RenderThreshold:
if value is None:
return None
if isinstance(value, int):
return None if value < 0 else value
text = value.strip().lower()
if text in {"none", "gray", "grayscale", "soft", "-1"}:
return None
if text in {"auto", "otsu", "adaptive", "calibrated"}:
return text
threshold = int(text)
return None if threshold < 0 else threshold
def read_source_resolution(video_path: Path) -> tuple[int, int]:
capture = cv2.VideoCapture(str(video_path))
if not capture.isOpened():
raise RuntimeError(f"Could not open video: {video_path}")
width = int(capture.get(cv2.CAP_PROP_FRAME_WIDTH))
height = int(capture.get(cv2.CAP_PROP_FRAME_HEIGHT))
capture.release()
if width <= 0 or height <= 0:
raise RuntimeError(f"Could not read source resolution from {video_path}")
return width, height
def probe_source_audio(video_path: Path) -> tuple[int, int] | None:
command = [
"ffprobe",
"-v",
"error",
"-select_streams",
"a:0",
"-show_entries",
"stream=sample_rate,channels",
"-of",
"json",
str(video_path),
]
try:
result = subprocess.run(command, check=True, capture_output=True, text=True)
except (FileNotFoundError, subprocess.CalledProcessError):
return None
streams = json.loads(result.stdout or "{}").get("streams", [])
if not streams:
return None
stream = streams[0]
sample_rate = int(stream.get("sample_rate") or 0)
channels = int(stream.get("channels") or 0)
if sample_rate <= 0 or channels <= 0:
return None
return sample_rate, channels
def read_video_frames(
video_path: Path,
width: int,
height: int,
frame_stride: int,
threshold: int | None,
max_frames: int | None,
) -> tuple[np.ndarray, float]:
capture = cv2.VideoCapture(str(video_path))
if not capture.isOpened():
raise RuntimeError(f"Could not open video: {video_path}")
source_fps = capture.get(cv2.CAP_PROP_FPS) or 30.0
frames: list[np.ndarray] = []
source_index = 0
while True:
ok, frame = capture.read()
if not ok:
break
if source_index % frame_stride == 0:
gray = cv2.cvtColor(frame, cv2.COLOR_BGR2GRAY)
small = cv2.resize(gray, (width, height), interpolation=cv2.INTER_AREA)
if threshold is not None:
small = np.where(small >= threshold, 255, 0).astype(np.uint8)
frames.append(small)
if max_frames is not None and len(frames) >= max_frames:
break
source_index += 1
capture.release()
if not frames:
raise RuntimeError("No frames were extracted from the video.")
effective_fps = source_fps / frame_stride
return np.stack(frames, axis=0), effective_fps
def read_audio_samples(video_path: Path, sample_rate: int, channels: int, duration: float) -> np.ndarray:
command = [
"ffmpeg",
"-v",
"error",
"-i",
str(video_path),
"-vn",
"-ac",
str(channels),
"-ar",
str(sample_rate),
"-t",
f"{duration:.6f}",
"-f",
"f32le",
"pipe:1",
]
result = subprocess.run(command, check=True, capture_output=True)
samples = np.frombuffer(result.stdout, dtype=np.float32).copy()
if samples.size == 0:
raise RuntimeError(f"No audio was extracted from {video_path}")
samples = samples.reshape(-1, channels)
return np.clip(samples, -1.0, 1.0)
def prepare_dataset(args: argparse.Namespace) -> tuple[Path, Path, Path | None, Path | None]:
video_path = Path(args.video).resolve()
cache_dir = Path(args.cache_dir).resolve()
cache_dir.mkdir(parents=True, exist_ok=True)
if getattr(args, "match_source_resolution", False):
args.width, args.height = read_source_resolution(video_path)
audio_sample_rate = args.audio_sample_rate
audio_channels = args.audio_channels
if getattr(args, "match_source_audio", False):
probed_audio = probe_source_audio(video_path)
if probed_audio is not None:
audio_sample_rate, audio_channels = probed_audio
else:
print("could not probe source audio; using configured audio sample rate/channels")
args.audio_sample_rate = audio_sample_rate
args.audio_channels = audio_channels
name = safe_name(video_path, args.width, args.height, args.frame_stride, args.threshold, args.max_frames)
frames_path = cache_dir / f"{name}.npy"
spec_path = cache_dir / f"{name}.json"
if not (frames_path.exists() and spec_path.exists()) or args.force:
frames, fps = read_video_frames(
video_path=video_path,
width=args.width,
height=args.height,
frame_stride=args.frame_stride,
threshold=args.threshold,
max_frames=args.max_frames,
)
np.save(frames_path, frames)
spec = VideoSpec(
width=args.width,
height=args.height,
fps=fps,
frame_stride=args.frame_stride,
threshold=args.threshold,
frame_count=int(frames.shape[0]),
)
spec_path.write_text(json.dumps(asdict(spec), indent=2), encoding="utf-8")
print(f"cached {frames.shape[0]} frames at {frames.shape[2]}x{frames.shape[1]} -> {frames_path}")
audio_path = None
audio_spec_path = None
if getattr(args, "train_audio", True):
spec = load_video_spec(spec_path)
audio_stem = f"{name}_audio_{audio_sample_rate}hz_{audio_channels}ch"
audio_path = cache_dir / f"{audio_stem}.npy"
audio_spec_path = cache_dir / f"{audio_stem}.json"
if not (audio_path.exists() and audio_spec_path.exists()) or args.force:
audio = read_audio_samples(video_path, audio_sample_rate, audio_channels, spec.duration)
np.save(audio_path, audio)
audio_spec = AudioSpec(
sample_rate=audio_sample_rate,
sample_count=int(audio.shape[0]),
channels=audio_channels,
)
audio_spec_path.write_text(json.dumps(asdict(audio_spec), indent=2), encoding="utf-8")
print(
f"cached {audio.shape[0]} audio samples at "
f"{audio_sample_rate} Hz, {audio_channels} channel(s) -> {audio_path}"
)
return frames_path, spec_path, audio_path, audio_spec_path
class FourierFeatures(nn.Module):
def __init__(self, dims: int, bands: int):
super().__init__()
freq = 2.0 ** torch.arange(bands, dtype=torch.float32) * math.pi
self.dims = dims
self.register_buffer("freq", freq)
@property
def out_dim(self) -> int:
return self.dims + self.dims * 2 * int(self.freq.numel())
def forward(self, coords: torch.Tensor) -> torch.Tensor:
encoded = [coords]
angles = coords[..., None] * self.freq
encoded.append(torch.sin(angles).flatten(-2))
encoded.append(torch.cos(angles).flatten(-2))
return torch.cat(encoded, dim=-1)
class CoordinateMLP(nn.Module):
def __init__(self, dims: int, hidden: int, layers: int, bands: int, out_activation: str | None):
super().__init__()
self.features = FourierFeatures(dims=dims, bands=bands)
blocks: list[nn.Module] = []
in_dim = self.features.out_dim
for _ in range(layers):
blocks.append(nn.Linear(in_dim, hidden))
blocks.append(nn.SiLU())
in_dim = hidden
blocks.append(nn.Linear(in_dim, 1))
self.net = nn.Sequential(*blocks)
self.out_activation = out_activation
def forward(self, coords: torch.Tensor) -> torch.Tensor:
values = self.net(self.features(coords)).squeeze(-1)
if self.out_activation == "tanh":
return torch.tanh(values)
return values
class LegacyBadAppleMultiModalNet(nn.Module):
def __init__(
self,
video_hidden: int,
video_layers: int,
video_bands: int,
audio_hidden: int,
audio_layers: int,
audio_bands: int,
train_audio: bool,
):
super().__init__()
self.video = CoordinateMLP(3, video_hidden, video_layers, video_bands, out_activation=None)
self.audio_enabled = train_audio
self.audio_channels = 1
self.audio = None
if train_audio:
self.audio = CoordinateMLP(1, audio_hidden, audio_layers, audio_bands, out_activation="tanh")
class UnifiedV1BadAppleMultiModalNet(nn.Module):
def __init__(
self,
video_hidden: int,
video_layers: int,
video_bands: int,
audio_hidden: int,
audio_layers: int,
audio_bands: int,
train_audio: bool,
audio_channels: int,
modality_embedding_dim: int,
):
super().__init__()
hidden = max(video_hidden, audio_hidden if train_audio else video_hidden)
layers = max(video_layers, audio_layers if train_audio else video_layers)
self.audio_enabled = train_audio
self.audio_channels = max(1, audio_channels)
self.video_bands = video_bands
self.audio_bands = audio_bands
self.modality_embedding_dim = modality_embedding_dim
self.coord_features = FourierFeatures(dims=3, bands=video_bands)
self.time_features = FourierFeatures(dims=1, bands=max(video_bands, audio_bands))
self.modality_embedding = nn.Embedding(2, modality_embedding_dim)
trunk_in = self.coord_features.out_dim + self.time_features.out_dim + modality_embedding_dim
blocks: list[nn.Module] = []
in_dim = trunk_in
for _ in range(layers):
blocks.append(nn.Linear(in_dim, hidden))
blocks.append(nn.SiLU())
in_dim = hidden
self.trunk = nn.Sequential(*blocks)
self.video_head = nn.Linear(hidden, 1)
self.audio_head = nn.Linear(hidden, self.audio_channels) if train_audio else None
def _coords3(self, coords: torch.Tensor) -> torch.Tensor:
if coords.shape[-1] == 3:
return coords
if coords.shape[-1] != 1:
raise ValueError(f"expected 1D or 3D coordinates, got shape {tuple(coords.shape)}")
zeros = torch.zeros((coords.shape[0], 2), device=coords.device, dtype=coords.dtype)
return torch.cat((coords, zeros), dim=-1)
def modality_features(self, coords: torch.Tensor, modality_id: int) -> torch.Tensor:
coords3 = self._coords3(coords)
time = coords3[:, :1]
modality = torch.full((coords3.shape[0],), modality_id, device=coords3.device, dtype=torch.long)
encoded = torch.cat(
(
self.coord_features(coords3),
self.time_features(time),
self.modality_embedding(modality),
),
dim=-1,
)
return self.trunk(encoded)
def video(self, coords: torch.Tensor) -> torch.Tensor:
return self.video_head(self.modality_features(coords, 0)).squeeze(-1)
def audio(self, coords: torch.Tensor) -> torch.Tensor:
if self.audio_head is None:
raise RuntimeError("This checkpoint does not contain an audio model.")
return torch.tanh(self.audio_head(self.modality_features(coords, 1)))
class TemporalLatentFeatures(nn.Module):
"""A shared, linearly interpolated time-memory used by both modalities."""
def __init__(self, anchors: int, dim: int):
super().__init__()
self.anchors = anchors
self.dim = dim
self.embedding = nn.Embedding(anchors, dim)
nn.init.normal_(self.embedding.weight, mean=0.0, std=0.02)
def forward(self, time: torch.Tensor) -> torch.Tensor:
position = ((time.squeeze(-1) + 1.0) * 0.5 * (self.anchors - 1)).clamp(0, self.anchors - 1)
lower = position.floor().long()
upper = (lower + 1).clamp(max=self.anchors - 1)
fraction = (position - lower.float()).unsqueeze(-1)
return torch.lerp(self.embedding(lower), self.embedding(upper), fraction)
class QATLinear(nn.Linear):
"""Linear layer with optional per-output-channel INT8 fake quantization."""
def __init__(self, in_features: int, out_features: int, bias: bool = True):
super().__init__(in_features, out_features, bias=bias)
self.qat_enabled = False
def forward(self, values: torch.Tensor) -> torch.Tensor:
if not self.qat_enabled:
return F.linear(values, self.weight, self.bias)
max_abs = self.weight.detach().abs().amax(dim=1, keepdim=True).clamp_min(1e-8)
scale = max_abs / 127.0
quantized = torch.round(self.weight / scale).clamp(-127, 127) * scale
fake_quantized = self.weight + (quantized - self.weight).detach()
return F.linear(values, fake_quantized, self.bias)
def set_qat_enabled(model: nn.Module, enabled: bool) -> None:
for module in model.modules():
if isinstance(module, QATLinear):
module.qat_enabled = enabled
class ResidualBlock(nn.Module):
def __init__(self, hidden: int):
super().__init__()
self.norm = nn.LayerNorm(hidden)
self.expand = QATLinear(hidden, hidden * 2)
self.project = QATLinear(hidden * 2, hidden)
def forward(self, values: torch.Tensor) -> torch.Tensor:
residual = values
values = self.norm(values)
values = F.silu(self.expand(values))
return residual + self.project(values)
class BadAppleMultiModalNet(nn.Module):
"""The current unified coordinate field used for new checkpoints."""
def __init__(
self,
video_hidden: int,
video_layers: int,
video_bands: int,
audio_hidden: int,
audio_layers: int,
audio_bands: int,
train_audio: bool,
audio_channels: int,
modality_embedding_dim: int,
temporal_latent_anchors: int,
temporal_latent_dim: int,
pixel_centers: bool = True,
gradient_checkpointing: bool = False,
):
super().__init__()
hidden = max(video_hidden, audio_hidden if train_audio else video_hidden)
layers = max(video_layers, audio_layers if train_audio else video_layers)
self.audio_enabled = train_audio
self.audio_channels = max(1, audio_channels)
self.video_bands = video_bands
self.audio_bands = audio_bands
self.modality_embedding_dim = modality_embedding_dim
self.temporal_latent_anchors = temporal_latent_anchors
self.temporal_latent_dim = temporal_latent_dim
self.pixel_centers = pixel_centers
self.gradient_checkpointing = gradient_checkpointing
self.coord_features = FourierFeatures(dims=3, bands=video_bands)
self.time_features = FourierFeatures(dims=1, bands=max(video_bands, audio_bands))
self.time_memory = TemporalLatentFeatures(temporal_latent_anchors, temporal_latent_dim)
self.modality_embedding = nn.Embedding(2, modality_embedding_dim)
trunk_in = (
self.coord_features.out_dim
+ self.time_features.out_dim
+ temporal_latent_dim
+ modality_embedding_dim
)
self.input_layer = QATLinear(trunk_in, hidden)
self.blocks = nn.ModuleList(ResidualBlock(hidden) for _ in range(max(1, layers - 1)))
self.output_norm = nn.LayerNorm(hidden)
self.video_head = QATLinear(hidden, 1)
self.audio_head = QATLinear(hidden, self.audio_channels) if train_audio else None
def _coords3(self, coords: torch.Tensor) -> torch.Tensor:
if coords.shape[-1] == 3:
return coords
if coords.shape[-1] != 1:
raise ValueError(f"expected 1D or 3D coordinates, got shape {tuple(coords.shape)}")
zeros = torch.zeros((coords.shape[0], 2), device=coords.device, dtype=coords.dtype)
return torch.cat((coords, zeros), dim=-1)
def modality_features(self, coords: torch.Tensor, modality_id: int) -> torch.Tensor:
coords3 = self._coords3(coords)
time = coords3[:, :1]
modality = torch.full((coords3.shape[0],), modality_id, device=coords3.device, dtype=torch.long)
encoded = torch.cat(
(
self.coord_features(coords3),
self.time_features(time),
self.time_memory(time),
self.modality_embedding(modality),
),
dim=-1,
)
values = F.silu(self.input_layer(encoded))
for block in self.blocks:
if self.gradient_checkpointing and self.training and torch.is_grad_enabled():
values = checkpoint(block, values, use_reentrant=False)
else:
values = block(values)
return self.output_norm(values)
def video(self, coords: torch.Tensor) -> torch.Tensor:
return self.video_head(self.modality_features(coords, 0)).squeeze(-1)
def audio(self, coords: torch.Tensor) -> torch.Tensor:
if self.audio_head is None:
raise RuntimeError("This checkpoint does not contain an audio model.")
return torch.tanh(self.audio_head(self.modality_features(coords, 1)))
def has_audio_model(model: nn.Module) -> bool:
if getattr(model, "audio_enabled", False):
return True
audio = getattr(model, "audio", None)
return audio is not None and not callable(audio)
def predict_video(model: nn.Module, coords: torch.Tensor) -> torch.Tensor:
return model.video(coords)
def predict_audio(model: nn.Module, coords: torch.Tensor) -> torch.Tensor:
if not has_audio_model(model):
raise RuntimeError("This checkpoint does not contain an audio model.")
values = model.audio(coords)
if values.dim() == 1:
values = values.unsqueeze(-1)
return values
def load_video_spec(spec_path: Path) -> VideoSpec:
raw = json.loads(spec_path.read_text(encoding="utf-8"))
return VideoSpec(**raw)
def load_audio_spec(spec_path: Path) -> AudioSpec:
raw = json.loads(spec_path.read_text(encoding="utf-8"))
return AudioSpec(**raw)
def choose_device(name: str) -> torch.device:
if name == "auto":
return torch.device("cuda" if torch.cuda.is_available() else "cpu")
return torch.device(name)
def set_seed(seed: int | None) -> None:
if seed is None:
return
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
if torch.cuda.is_available():
torch.cuda.manual_seed_all(seed)
def build_edge_indices(frames: np.ndarray, edge_threshold: int) -> np.ndarray:
frame_count, height, width = frames.shape
pieces: list[np.ndarray] = []
index_dtype = np.uint32 if frames.size <= np.iinfo(np.uint32).max else np.int64
for frame_index in range(frame_count):
frame = np.asarray(frames[frame_index], dtype=np.uint8)
edge_mask = np.zeros((height, width), dtype=bool)
horizontal = np.abs(frame[:, 1:].astype(np.int16) - frame[:, :-1].astype(np.int16)) >= edge_threshold
vertical = np.abs(frame[1:, :].astype(np.int16) - frame[:-1, :].astype(np.int16)) >= edge_threshold
edge_mask[:, 1:] |= horizontal
edge_mask[:, :-1] |= horizontal
edge_mask[1:, :] |= vertical
edge_mask[:-1, :] |= vertical
yx = np.argwhere(edge_mask)
if yx.size:
linear = frame_index * height * width + yx[:, 0] * width + yx[:, 1]
pieces.append(linear.astype(index_dtype, copy=False))
if not pieces:
return np.empty(0, dtype=index_dtype)
return np.concatenate(pieces, axis=0)
def build_motion_frame_weights(frames: np.ndarray, power: float = 0.5) -> np.ndarray:
"""Favor frames where the silhouette changes without excluding quiet scenes."""
if frames.shape[0] <= 1:
return np.ones(1, dtype=np.float32)
changes = np.empty(frames.shape[0], dtype=np.float32)
previous = np.asarray(frames[0], dtype=np.int16)
for frame_index in range(1, frames.shape[0]):
current = np.asarray(frames[frame_index], dtype=np.int16)
changes[frame_index] = np.abs(current - previous).mean(dtype=np.float64)
previous = current
changes[0] = changes[1]
weights = np.power(changes + 1e-3, power)
return (weights / weights.sum()).astype(np.float32)
def video_white_fraction(frames: np.ndarray) -> float:
white_pixels = 0
for frame_index in range(frames.shape[0]):
white_pixels += int(np.count_nonzero(frames[frame_index] >= 128))
return white_pixels / max(int(frames.size), 1)
def video_coords_from_indices(
frame_count: int,
height: int,
width: int,
t: torch.Tensor,
y: torch.Tensor,
x: torch.Tensor,
pixel_centers: bool = True,
) -> torch.Tensor:
if pixel_centers:
x_norm = (x.float() + 0.5) / width
y_norm = (y.float() + 0.5) / height
else:
x_norm = x.float() / max(width - 1, 1)
y_norm = y.float() / max(height - 1, 1)
coords = torch.stack(
(
t.float() / max(frame_count - 1, 1),
x_norm,
y_norm,
),
dim=-1,
)
return coords * 2.0 - 1.0
def sample_video_batch(
frames: np.ndarray,
batch_size: int,
device: torch.device,
edge_indices: np.ndarray | None = None,
edge_fraction: float = 0.0,
frame_weights: torch.Tensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
frame_count, height, width = frames.shape
edge_count = 0
if edge_indices is not None and edge_indices.size > 0 and edge_fraction > 0:
edge_count = min(batch_size, int(round(batch_size * edge_fraction)))
uniform_count = batch_size - edge_count
parts_t: list[torch.Tensor] = []
parts_y: list[torch.Tensor] = []
parts_x: list[torch.Tensor] = []
if uniform_count > 0:
if frame_weights is None:
parts_t.append(torch.randint(0, frame_count, (uniform_count,), device=device))
else:
parts_t.append(torch.multinomial(frame_weights, uniform_count, replacement=True))
parts_y.append(torch.randint(0, height, (uniform_count,), device=device))
parts_x.append(torch.randint(0, width, (uniform_count,), device=device))
if edge_count > 0 and edge_indices is not None:
picked = edge_indices[np.random.randint(0, edge_indices.shape[0], size=edge_count)].astype(
np.int64,
copy=False,
)
frame_pixels = height * width
picked_t = picked // frame_pixels
picked_remainder = picked % frame_pixels
parts_t.append(torch.from_numpy(picked_t).to(device=device, dtype=torch.long))
parts_y.append(torch.from_numpy(picked_remainder // width).to(device=device, dtype=torch.long))
parts_x.append(torch.from_numpy(picked_remainder % width).to(device=device, dtype=torch.long))
t = torch.cat(parts_t)
y = torch.cat(parts_y)
x = torch.cat(parts_x)
targets_np = frames[t.cpu().numpy(), y.cpu().numpy(), x.cpu().numpy()]
targets = torch.from_numpy(targets_np).to(device=device, dtype=torch.float32) / 255.0
return video_coords_from_indices(frame_count, height, width, t, y, x), targets
def sample_video_pair_batch(
frames: np.ndarray,
batch_size: int,
device: torch.device,
mode: str,
frame_weights: torch.Tensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor] | None:
frame_count, height, width = frames.shape
if batch_size <= 0:
return None
if mode == "temporal":
if frame_count < 2:
return None
if frame_weights is None:
t0 = torch.randint(0, frame_count - 1, (batch_size,), device=device)
else:
t0 = torch.multinomial(frame_weights[:-1], batch_size, replacement=True)
t1 = t0 + 1
y0 = torch.randint(0, height, (batch_size,), device=device)
x0 = torch.randint(0, width, (batch_size,), device=device)
y1 = y0
x1 = x0
elif mode == "spatial":
if height < 2 or width < 2:
return None
t0 = torch.randint(0, frame_count, (batch_size,), device=device)
t1 = t0
y0 = torch.randint(0, height, (batch_size,), device=device)
x0 = torch.randint(0, width, (batch_size,), device=device)
use_x = torch.rand((batch_size,), device=device) < 0.5
x1 = torch.where(use_x, torch.clamp(x0 + 1, max=width - 1), x0)
y1 = torch.where(use_x, y0, torch.clamp(y0 + 1, max=height - 1))
else:
raise ValueError(f"unknown pair batch mode: {mode}")
target0_np = frames[t0.cpu().numpy(), y0.cpu().numpy(), x0.cpu().numpy()]
target1_np = frames[t1.cpu().numpy(), y1.cpu().numpy(), x1.cpu().numpy()]
target0 = torch.from_numpy(target0_np).to(device=device, dtype=torch.float32) / 255.0
target1 = torch.from_numpy(target1_np).to(device=device, dtype=torch.float32) / 255.0
coords0 = video_coords_from_indices(frame_count, height, width, t0, y0, x0)
coords1 = video_coords_from_indices(frame_count, height, width, t1, y1, x1)
return coords0, coords1, target0, target1
def sample_audio_batch(
audio: np.ndarray,
batch_size: int,
device: torch.device,
chunk_samples: int,
) -> tuple[torch.Tensor, torch.Tensor]:
sample_count = audio.shape[0]
if chunk_samples <= 1:
indices = torch.randint(0, sample_count, (batch_size,), device=device)
else:
chunk_count = max(1, batch_size // chunk_samples)
max_start = max(1, sample_count - chunk_samples)
starts = torch.randint(0, max_start, (chunk_count,), device=device)
offsets = torch.arange(chunk_samples, device=device)
indices = (starts[:, None] + offsets[None, :]).reshape(-1)
targets_np = audio[indices.cpu().numpy()]
targets = torch.from_numpy(targets_np).to(device=device, dtype=torch.float32)
if targets.dim() == 1:
targets = targets.unsqueeze(-1)
coords = indices.float().unsqueeze(-1) / max(sample_count - 1, 1)
return coords * 2.0 - 1.0, targets
def audio_stft_loss(pred_chunks: torch.Tensor, target_chunks: torch.Tensor, n_fft: int, hop_length: int) -> torch.Tensor:
if pred_chunks.shape[1] < n_fft:
return pred_chunks.new_tensor(0.0)
pred_chunks = pred_chunks.float()
target_chunks = target_chunks.float()
pred_signals = pred_chunks.permute(0, 2, 1).reshape(-1, pred_chunks.shape[1])
target_signals = target_chunks.permute(0, 2, 1).reshape(-1, target_chunks.shape[1])
window = torch.hann_window(n_fft, device=pred_chunks.device, dtype=pred_chunks.dtype)
pred_spec = torch.stft(
pred_signals,
n_fft=n_fft,
hop_length=hop_length,
window=window,
return_complex=True,
)
target_spec = torch.stft(
target_signals,
n_fft=n_fft,
hop_length=hop_length,
window=window,
return_complex=True,
)
return F.l1_loss(torch.log1p(pred_spec.abs()), torch.log1p(target_spec.abs()))
def soft_dice_loss(logits: torch.Tensor, targets: torch.Tensor, epsilon: float = 1e-6) -> torch.Tensor:
probabilities = torch.sigmoid(logits)
intersection = (probabilities * targets).sum()
denominator = probabilities.sum() + targets.sum()
return 1.0 - (2.0 * intersection + epsilon) / (denominator + epsilon)
def make_evaluation_batch(
frames: np.ndarray,
batch_size: int,
device: torch.device,
seed: int,
pixel_centers: bool = True,
) -> tuple[torch.Tensor, torch.Tensor]:
"""A fixed, reproducible sample for checkpoint selection and calibration."""
frame_count, height, width = frames.shape
rng = np.random.default_rng(seed)
t_np = rng.integers(0, frame_count, size=batch_size, dtype=np.int64)
y_np = rng.integers(0, height, size=batch_size, dtype=np.int64)
x_np = rng.integers(0, width, size=batch_size, dtype=np.int64)
targets_np = frames[t_np, y_np, x_np]
t = torch.from_numpy(t_np).to(device)
y = torch.from_numpy(y_np).to(device)
x = torch.from_numpy(x_np).to(device)
targets = torch.from_numpy(targets_np).to(device=device, dtype=torch.float32) / 255.0
return video_coords_from_indices(frame_count, height, width, t, y, x, pixel_centers=pixel_centers), targets
@torch.inference_mode()
def calibrate_render_threshold(
model: nn.Module,
frames: np.ndarray,
batch_size: int,
device: torch.device,
gamma: float,
contrast: float,
seed: int,
) -> tuple[int, float]:
"""Find one global threshold that maximizes silhouette IoU over fixed pixels."""
was_training = model.training
model.eval()
coords, targets = make_evaluation_batch(
frames,
batch_size,
device,
seed,
pixel_centers=getattr(model, "pixel_centers", False),
)
probabilities = torch.sigmoid(predict_video(model, coords)).detach().cpu().numpy()
adjusted = adjust_image(probabilities, gamma=gamma, contrast=contrast)
target_binary = targets.detach().cpu().numpy() >= 0.5
if not target_binary.any():
if was_training:
model.train()
return 256, 1.0
if target_binary.all():
if was_training:
model.train()
return 0, 1.0
thresholds = np.arange(1, 256, dtype=np.uint8)
scores = []
for threshold in thresholds:
prediction = adjusted >= threshold / 255.0
union = np.logical_or(prediction, target_binary).sum()
scores.append(float(np.logical_and(prediction, target_binary).sum() / max(union, 1)))
best_index = int(np.argmax(scores))
if was_training:
model.train()
return int(thresholds[best_index]), scores[best_index]
def resolve_calibrated_threshold(
requested: RenderThreshold,
model: nn.Module,
frames_path: str | Path | None,
device: torch.device,
batch_size: int,
gamma: float,
contrast: float,
seed: int,
) -> RenderThreshold:
if requested != "calibrated":
return requested
if frames_path is None or not Path(frames_path).exists():
print("could not calibrate threshold because the frame cache is unavailable; using 128")
return 128
frames = np.load(Path(frames_path), mmap_mode="r")
threshold, iou = calibrate_render_threshold(
model=model,
frames=frames,
batch_size=batch_size,
device=device,
gamma=gamma,
contrast=contrast,
seed=seed,
)
print(f"calibrated global render threshold={threshold} silhouette_iou={iou:.3f}")
return threshold
def lr_scale_for_step(step: int, args: argparse.Namespace) -> float:
warmup_steps = max(0, args.warmup_steps)
if warmup_steps > 0 and step <= warmup_steps:
return max(args.warmup_start_ratio, step / warmup_steps)
decay_steps = max(1, args.steps - warmup_steps)
progress = min(max((step - warmup_steps) / decay_steps, 0.0), 1.0)
min_ratio = args.min_lr_ratio
if args.lr_decay == "none":
return 1.0
if args.lr_decay == "linear":
return min_ratio + (1.0 - min_ratio) * (1.0 - progress)
if args.lr_decay == "cosine":
cosine = 0.5 * (1.0 + math.cos(math.pi * progress))
return min_ratio + (1.0 - min_ratio) * cosine
raise ValueError(f"unknown lr decay: {args.lr_decay}")
def set_optimizer_lr(optimizer: torch.optim.Optimizer, lr: float) -> None:
for group in optimizer.param_groups:
group["lr"] = lr
def frame_coords(
frame_index: int,
spec: VideoSpec,
scale: int,
device: torch.device,
supersample: int = 1,
pixel_centers: bool = True,
) -> torch.Tensor:
out_w = spec.width * scale * supersample
out_h = spec.height * scale * supersample
yy, xx = torch.meshgrid(
torch.arange(out_h, device=device),
torch.arange(out_w, device=device),
indexing="ij",
)
t = torch.full_like(xx, frame_index, dtype=torch.float32)
if pixel_centers:
x_norm = (xx.float() + 0.5) / out_w
y_norm = (yy.float() + 0.5) / out_h
else:
x_norm = xx.float() / max(out_w - 1, 1)
y_norm = yy.float() / max(out_h - 1, 1)
coords = torch.stack(
(
t / max(spec.frame_count - 1, 1),
x_norm,
y_norm,
),
dim=-1,
)
return coords.reshape(-1, 3) * 2.0 - 1.0
def audio_coords(start_sample: int, sample_count: int, full_count: int, device: torch.device) -> torch.Tensor:
indices = torch.arange(start_sample, start_sample + sample_count, device=device, dtype=torch.float32)
return (indices / max(full_count - 1, 1)).unsqueeze(-1) * 2.0 - 1.0
def adjust_image(image: np.ndarray, gamma: float, contrast: float) -> np.ndarray:
adjusted = np.clip((image - 0.5) * contrast + 0.5, 0.0, 1.0)
if gamma > 0 and abs(gamma - 1.0) > 1e-6:
adjusted = np.power(adjusted, 1.0 / gamma)
return adjusted
def apply_render_threshold(
image_u8: np.ndarray,
threshold: RenderThreshold,
previous_binary: np.ndarray | None,
hysteresis: int,
) -> tuple[np.ndarray, np.ndarray | None]:
if threshold is None:
return image_u8, None
threshold_value: float | None = None
if threshold == "auto" or threshold == "otsu":
threshold_value, binary = cv2.threshold(image_u8, 0, 255, cv2.THRESH_BINARY + cv2.THRESH_OTSU)
elif threshold == "adaptive":
block_size = max(3, (min(image_u8.shape) // 12) | 1)
binary = cv2.adaptiveThreshold(
image_u8,
255,
cv2.ADAPTIVE_THRESH_GAUSSIAN_C,
cv2.THRESH_BINARY,
block_size,
2,
)
else:
threshold_value = float(threshold)
binary = np.where(image_u8 >= int(threshold), 255, 0).astype(np.uint8)
if previous_binary is not None and threshold_value is not None and hysteresis > 0:
distance = np.abs(image_u8.astype(np.int16) - int(round(threshold_value)))
ambiguous = distance <= hysteresis
binary[ambiguous] = previous_binary[ambiguous]
return binary.astype(np.uint8), binary.astype(np.uint8)
def render_video_frame(
model: nn.Module,
spec: VideoSpec,
frame_index: int,
device: torch.device,
scale: int,
supersample: int,
render_batch: int,
threshold: RenderThreshold,
gamma: float,
contrast: float,
previous_binary: np.ndarray | None,
threshold_hysteresis: int,
) -> tuple[np.ndarray, np.ndarray | None]:
sample_scale = scale * supersample
sample_w = spec.width * sample_scale
sample_h = spec.height * sample_scale
out_w = spec.width * scale
out_h = spec.height * scale
coords = frame_coords(
frame_index,
spec,
scale,
device,
supersample=supersample,
pixel_centers=getattr(model, "pixel_centers", False),
)
chunks = []
for chunk in coords.split(render_batch):
chunks.append(torch.sigmoid(predict_video(model, chunk)).detach().cpu())
image = torch.cat(chunks).reshape(sample_h, sample_w).numpy()
if supersample > 1:
image = cv2.resize(image, (out_w, out_h), interpolation=cv2.INTER_AREA)
image = adjust_image(image, gamma=gamma, contrast=contrast)
image_u8 = (image * 255.0).clip(0, 255).astype(np.uint8)
return apply_render_threshold(image_u8, threshold, previous_binary, threshold_hysteresis)
@torch.inference_mode()
def render_contact_sheet(
model: BadAppleMultiModalNet,
spec: VideoSpec,
output_path: Path,
device: torch.device,
tiles: int,
scale: int,
render_batch: int,
threshold: RenderThreshold,
) -> Path:
was_training = model.training
model.eval()
indices = np.linspace(0, spec.frame_count - 1, tiles, dtype=int)
tile_w = spec.width * scale
tile_h = spec.height * scale
canvas = np.zeros((tile_h, tile_w * tiles), dtype=np.uint8)
for column, frame_index in enumerate(indices):
coords = frame_coords(
int(frame_index),
spec,
scale,
device,
pixel_centers=getattr(model, "pixel_centers", False),
)
chunks = []
for chunk in coords.split(render_batch):
chunks.append(torch.sigmoid(predict_video(model, chunk)).detach().cpu())
image = torch.cat(chunks).reshape(tile_h, tile_w).numpy()
image_u8 = (image * 255.0).clip(0, 255).astype(np.uint8)
image_u8, _ = apply_render_threshold(image_u8, threshold, None, 0)
canvas[:, column * tile_w : (column + 1) * tile_w] = image_u8
output_path.parent.mkdir(parents=True, exist_ok=True)
cv2.imwrite(str(output_path), canvas)
if was_training:
model.train()
return output_path
@torch.inference_mode()
def render_audio_wav(
model: BadAppleMultiModalNet,
audio_spec: AudioSpec,
output_path: Path,
device: torch.device,
render_batch: int,
smooth_samples: int,
normalize: bool,
peak: float,
fade_ms: float,
start_seconds: float = 0.0,
duration_seconds: float | None = None,
) -> Path:
if not has_audio_model(model):
raise RuntimeError("This checkpoint does not contain an audio model.")
was_training = model.training
model.eval()
start_sample = int(start_seconds * audio_spec.sample_rate)
if duration_seconds is None:
sample_count = audio_spec.sample_count - start_sample
else:
sample_count = int(duration_seconds * audio_spec.sample_rate)
sample_count = max(0, min(sample_count, audio_spec.sample_count - start_sample))
pieces = []
for start in range(0, sample_count, render_batch):
count = min(render_batch, sample_count - start)
coords = audio_coords(start_sample + start, count, audio_spec.sample_count, device)
pieces.append(predict_audio(model, coords).detach().cpu())
audio = torch.cat(pieces).numpy() if pieces else np.zeros((0, audio_spec.channels), dtype=np.float32)
if audio.ndim == 1:
audio = audio[:, None]
if smooth_samples > 1 and audio.size >= smooth_samples:
kernel = np.ones(smooth_samples, dtype=np.float32) / smooth_samples
for channel in range(audio.shape[1]):
audio[:, channel] = np.convolve(audio[:, channel], kernel, mode="same")
if fade_ms > 0 and audio.shape[0] > 0:
fade_samples = min(audio.shape[0] // 2, int(audio_spec.sample_rate * fade_ms / 1000.0))
if fade_samples > 0:
fade = np.linspace(0.0, 1.0, fade_samples, dtype=np.float32)[:, None]
audio[:fade_samples] *= fade
audio[-fade_samples:] *= fade[::-1]
if normalize and audio.size > 0:
max_abs = float(np.max(np.abs(audio)))
if max_abs > 1e-8:
audio = audio / max_abs * min(max(peak, 0.0), 1.0)
audio_i16 = (np.clip(audio, -1.0, 1.0) * 32767.0).astype(np.int16)
output_path.parent.mkdir(parents=True, exist_ok=True)
with wave.open(str(output_path), "wb") as wav:
wav.setnchannels(audio.shape[1] if audio.ndim == 2 else audio_spec.channels)
wav.setsampwidth(2)
wav.setframerate(audio_spec.sample_rate)
wav.writeframes(audio_i16.tobytes())
if was_training:
model.train()
return output_path
@torch.inference_mode()
def render_video_file(
model: BadAppleMultiModalNet,
spec: VideoSpec,
output_path: Path,
device: torch.device,
scale: int,
render_batch: int,
threshold: RenderThreshold,
render_log_every: int,
encoder: str,
crf: int,
preset: str,
tune: str,
supersample: int,
gamma: float,
contrast: float,
threshold_hysteresis: int,
start_frame: int = 0,
frame_count: int | None = None,
) -> Path:
was_training = model.training
model.eval()
frame_count = spec.frame_count - start_frame if frame_count is None else frame_count
frame_count = max(0, min(frame_count, spec.frame_count - start_frame))
output_path.parent.mkdir(parents=True, exist_ok=True)
out_w = spec.width * scale
out_h = spec.height * scale
writer = None
ffmpeg_process: subprocess.Popen[bytes] | None = None
if encoder == "opencv":
writer = cv2.VideoWriter(
str(output_path),
cv2.VideoWriter_fourcc(*"mp4v"),
spec.fps,
(out_w, out_h),
isColor=False,
)
if not writer.isOpened():
raise RuntimeError(f"Could not open video writer: {output_path}")
elif encoder == "ffmpeg":
command = [
"ffmpeg",
"-y",
"-v",
"error",
"-f",
"rawvideo",
"-pix_fmt",
"gray",
"-s",
f"{out_w}x{out_h}",
"-r",
f"{spec.fps:.6f}",
"-i",
"pipe:0",
"-an",
"-c:v",
"libx264",
"-preset",
preset,
"-crf",
str(crf),
"-tune",
tune,
"-pix_fmt",
"yuv420p",
"-movflags",
"+faststart",
str(output_path),
]
ffmpeg_process = subprocess.Popen(command, stdin=subprocess.PIPE)
else:
raise ValueError(f"unknown video encoder: {encoder}")
print(f"rendering {frame_count} frames at {out_w}x{out_h} -> {output_path}")
previous_binary: np.ndarray | None = None
for offset in range(frame_count):
frame_index = start_frame + offset
image_u8, previous_binary = render_video_frame(
model=model,
spec=spec,
frame_index=frame_index,
device=device,
scale=scale,
supersample=supersample,
render_batch=render_batch,
threshold=threshold,
gamma=gamma,
contrast=contrast,
previous_binary=previous_binary,
threshold_hysteresis=threshold_hysteresis,
)
if writer is not None:
writer.write(image_u8)
elif ffmpeg_process is not None and ffmpeg_process.stdin is not None:
ffmpeg_process.stdin.write(image_u8.tobytes())
if offset == 0 or (offset + 1) % render_log_every == 0 or offset + 1 == frame_count:
print(f"rendered {offset + 1:>5}/{frame_count}")
if writer is not None:
writer.release()
if ffmpeg_process is not None:
if ffmpeg_process.stdin is not None:
ffmpeg_process.stdin.close()
return_code = ffmpeg_process.wait()
if return_code != 0:
raise RuntimeError(f"ffmpeg video encoding failed with exit code {return_code}")
if was_training:
model.train()
return output_path
def mux_wav_audio(
rendered_video: Path,
audio_wav: Path,
output_video: Path,
copy_video: bool,
audio_bitrate: str,
) -> None:
command = [
"ffmpeg",
"-y",
"-v",
"error",
"-i",
str(rendered_video),
"-i",
str(audio_wav),
"-map",
"0:v:0",
"-map",
"1:a:0",
"-c:v",
"copy" if copy_video else "libx264",
"-c:a",
"aac",
"-b:a",
audio_bitrate,
"-movflags",
"+faststart",
str(output_video),
]
subprocess.run(command, check=True)
def mux_source_audio(
rendered_video: Path,
source_video: Path,
output_video: Path,
copy_video: bool,
duration_seconds: float | None,
) -> None:
command = [
"ffmpeg",
"-y",
"-v",
"error",
"-i",
str(rendered_video),
"-i",
str(source_video),
"-map",
"0:v:0",
"-map",
"1:a:0",
"-c:v",
"copy" if copy_video else "libx264",
"-c:a",
"copy",
]
if duration_seconds is not None:
command.extend(["-t", f"{duration_seconds:.6f}"])
command.extend(["-movflags", "+faststart"])
command.append(str(output_video))
subprocess.run(command, check=True)
def build_model_from_args(args: argparse.Namespace) -> BadAppleMultiModalNet:
return BadAppleMultiModalNet(
video_hidden=args.hidden,
video_layers=args.layers,
video_bands=args.bands,
audio_hidden=args.audio_hidden,
audio_layers=args.audio_layers,
audio_bands=args.audio_bands,
train_audio=args.train_audio,
audio_channels=args.audio_channels,
modality_embedding_dim=args.modality_embedding_dim,
temporal_latent_anchors=args.temporal_latent_anchors,
temporal_latent_dim=args.temporal_latent_dim,
pixel_centers=True,
gradient_checkpointing=args.gradient_checkpointing,
)
def build_model_from_checkpoint(ckpt: dict) -> nn.Module:
if ckpt.get("architecture") in {"unified-v2", "unified-v3", "unified-v3-int8"}:
audio_spec = ckpt.get("audio_spec")
audio_channels = ckpt.get("audio_channels")
if audio_channels is None and audio_spec is not None:
audio_channels = audio_spec.get("channels", 1)
return BadAppleMultiModalNet(
video_hidden=ckpt["video_hidden"],
video_layers=ckpt["video_layers"],
video_bands=ckpt["video_bands"],
audio_hidden=ckpt["audio_hidden"],
audio_layers=ckpt["audio_layers"],
audio_bands=ckpt["audio_bands"],
train_audio=audio_spec is not None,
audio_channels=audio_channels or 1,
modality_embedding_dim=ckpt.get("modality_embedding_dim", 8),
temporal_latent_anchors=ckpt.get("temporal_latent_anchors", 512),
temporal_latent_dim=ckpt.get("temporal_latent_dim", 32),
pixel_centers=ckpt.get("pixel_centers", False),
gradient_checkpointing=False,
)
if ckpt.get("architecture") != "unified":
return LegacyBadAppleMultiModalNet(
video_hidden=ckpt["video_hidden"],
video_layers=ckpt["video_layers"],
video_bands=ckpt["video_bands"],
audio_hidden=ckpt["audio_hidden"],
audio_layers=ckpt["audio_layers"],
audio_bands=ckpt["audio_bands"],
train_audio=ckpt["audio_spec"] is not None,
)
audio_spec = ckpt.get("audio_spec")
audio_channels = ckpt.get("audio_channels")
if audio_channels is None and audio_spec is not None:
audio_channels = audio_spec.get("channels", 1)
return UnifiedV1BadAppleMultiModalNet(
video_hidden=ckpt["video_hidden"],
video_layers=ckpt["video_layers"],
video_bands=ckpt["video_bands"],
audio_hidden=ckpt["audio_hidden"],
audio_layers=ckpt["audio_layers"],
audio_bands=ckpt["audio_bands"],
train_audio=ckpt["audio_spec"] is not None,
audio_channels=audio_channels or 1,
modality_embedding_dim=ckpt.get("modality_embedding_dim", 8),
)
def autocast_context(device: torch.device, enabled: bool) -> contextlib.AbstractContextManager:
if enabled and device.type == "cuda":
return torch.autocast(device_type=device.type)
return contextlib.nullcontext()
def make_grad_scaler(use_amp: bool) -> torch.amp.GradScaler:
if hasattr(torch, "amp") and hasattr(torch.amp, "GradScaler"):
return torch.amp.GradScaler("cuda", enabled=use_amp)
return torch.cuda.amp.GradScaler(enabled=use_amp)
def state_dict_to_cpu(state_dict: dict[str, torch.Tensor]) -> dict[str, torch.Tensor]:
return {key: value.detach().cpu() for key, value in state_dict.items()}
def pack_int8_state_dict(model: nn.Module, state_dict: dict[str, torch.Tensor]) -> dict:
linear_weights = {
f"{name}.weight" if name else "weight"
for name, module in model.named_modules()
if isinstance(module, QATLinear)
}
packed: dict[str, object] = {}
for key, value in state_dict.items():
tensor = value.detach().cpu()
if key in linear_weights:
max_abs = tensor.float().abs().amax(dim=1).clamp_min(1e-8)
scale = max_abs / 127.0
quantized = torch.round(tensor.float() / scale[:, None]).clamp(-127, 127).to(torch.int8)
packed[key] = {"int8": quantized, "scale": scale.to(torch.float16)}
elif tensor.is_floating_point():
max_abs = float(tensor.float().abs().max()) if tensor.numel() else 0.0
packed[key] = tensor.float() if max_abs > 65000.0 else tensor.to(torch.float16)
else:
packed[key] = tensor
return packed
def unpack_int8_state_dict(packed: dict) -> dict[str, torch.Tensor]:
unpacked: dict[str, torch.Tensor] = {}
for key, value in packed.items():
if isinstance(value, dict) and "int8" in value:
unpacked[key] = value["int8"].float() * value["scale"].float()[:, None]
elif isinstance(value, torch.Tensor) and value.is_floating_point():
unpacked[key] = value.float()
else:
unpacked[key] = value
return unpacked
def checkpoint_model_state(ckpt: dict, use_ema: bool) -> tuple[dict[str, torch.Tensor], str]:
if ckpt.get("quantized_model") is not None:
return unpack_int8_state_dict(ckpt["quantized_model"]), "quantized_model"
state_key = "ema_model" if use_ema and ckpt.get("ema_model") is not None else "model"
return ckpt[state_key], state_key
def init_ema_state(model: nn.Module) -> dict[str, torch.Tensor]:
return {key: value.detach().clone() for key, value in model.state_dict().items()}
@torch.no_grad()
def update_ema_state(model: nn.Module, ema_state: dict[str, torch.Tensor], decay: float) -> None:
current = model.state_dict()
for key, value in current.items():
if value.is_floating_point():
ema_state[key].mul_(decay).add_(value.detach(), alpha=1.0 - decay)
else:
ema_state[key].copy_(value)
def checkpoint_payload(
model: nn.Module,
args: argparse.Namespace,
video_spec: VideoSpec,
audio_spec: AudioSpec | None,
frames_path: Path,
audio_path: Path | None,
ema_state: dict[str, torch.Tensor] | None,
step: int,
best_val_loss: float | None,
) -> dict:
return {
"architecture": "unified-v3",
"model": state_dict_to_cpu(model.state_dict()),
"ema_model": state_dict_to_cpu(ema_state) if ema_state is not None else None,
"step": step,
"best_val_loss": best_val_loss,
"video_hidden": args.hidden,
"video_layers": args.layers,
"video_bands": args.bands,
"audio_hidden": args.audio_hidden,
"audio_layers": args.audio_layers,
"audio_bands": args.audio_bands,
"audio_channels": audio_spec.channels if audio_spec is not None else args.audio_channels,
"modality_embedding_dim": args.modality_embedding_dim,
"temporal_latent_anchors": args.temporal_latent_anchors,
"temporal_latent_dim": args.temporal_latent_dim,
"pixel_centers": True,
"video_spec": asdict(video_spec),
"audio_spec": asdict(audio_spec) if audio_spec is not None else None,
"frames_cache": str(frames_path),
"audio_cache": str(audio_path) if audio_path is not None else None,
"source_video": str(Path(args.video).resolve()),
}
def save_checkpoint(path: Path, payload: dict) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
tmp_path = path.with_name(f"{path.name}.tmp")
torch.save(payload, tmp_path)
tmp_path.replace(path)
@torch.inference_mode()
def evaluate_video_quality(
model: nn.Module,
coords: torch.Tensor,
targets: torch.Tensor,
pos_weight: torch.Tensor | None,
) -> tuple[float, float, float]:
was_training = model.training
model.eval()
logits = predict_video(model, coords)
loss = F.binary_cross_entropy_with_logits(logits, targets, pos_weight=pos_weight)
pred = (torch.sigmoid(logits) >= 0.5).float()
acc = (pred == (targets >= 0.5)).float().mean()
target_binary = targets >= 0.5
pred_binary = pred >= 0.5
intersection = torch.logical_and(pred_binary, target_binary).sum()
union = torch.logical_or(pred_binary, target_binary).sum()
iou = intersection.float() / union.clamp_min(1).float()
if was_training:
model.train()
return float(loss.item()), float(acc.item()), float(iou.item())
def train(args: argparse.Namespace) -> Path:
set_seed(args.seed)
frames_path, spec_path, audio_path, audio_spec_path = prepare_dataset(args)
frames = np.load(frames_path, mmap_mode="r")
video_spec = load_video_spec(spec_path)
audio = np.load(audio_path, mmap_mode="r") if audio_path is not None else None
audio_spec = load_audio_spec(audio_spec_path) if audio_spec_path is not None else None
device = choose_device(args.device)
model = build_model_from_args(args).to(device)
if args.resume is not None:
resumed = torch.load(Path(args.resume).resolve(), map_location="cpu")
state_key = "ema_model" if args.resume_ema and resumed.get("ema_model") is not None else "model"
model.load_state_dict(resumed[state_key])
print(f"resumed model weights from {args.resume} ({state_key})")
optimizer = torch.optim.AdamW(model.parameters(), lr=args.lr, weight_decay=args.weight_decay)
use_amp = args.amp and device.type == "cuda"
scaler = make_grad_scaler(use_amp)
ema_state = init_ema_state(model) if args.ema_decay > 0 else None
video_pos_weight = None
if args.balanced_video_loss:
white_fraction = video_white_fraction(frames)
if 1e-4 < white_fraction < 1.0 - 1e-4:
video_pos_weight = torch.tensor((1.0 - white_fraction) / white_fraction, device=device)
print(f"balanced video loss: white_fraction={white_fraction:.4f} pos_weight={video_pos_weight.item():.2f}")
else:
print(f"balanced video loss disabled: white_fraction={white_fraction:.4f} has only one class")
edge_indices = None
if args.edge_sample_fraction > 0:
edge_indices = build_edge_indices(frames, args.edge_threshold)
print(
f"edge sampler: {edge_indices.shape[0]:,} candidate edge pixels "
f"({edge_indices.nbytes / (1024 ** 2):.1f} MiB)"
)
motion_weights = None
if args.motion_sample_fraction > 0:
motion_weights = torch.from_numpy(build_motion_frame_weights(frames)).to(device)
uniform_weight = torch.full_like(motion_weights, 1.0 / motion_weights.numel())
motion_weights = torch.lerp(uniform_weight, motion_weights, args.motion_sample_fraction)
print(f"motion sampler: enabled at {args.motion_sample_fraction:.0%} strength")
evaluation_coords, evaluation_targets = make_evaluation_batch(
frames,
args.val_batch_size,
device,
args.seed + 17,
)
print(
f"training on {device}: {video_spec.frame_count} frames, "
f"{video_spec.width}x{video_spec.height}, {sum(p.numel() for p in model.parameters()):,} parameters"
)
print(
f"batches: video={args.batch_size:,} pairs={args.video_pair_batch_size:,} "
f"audio={args.audio_batch_size:,} checkpointing={args.gradient_checkpointing}"
)
if audio_spec is not None:
print(
f"audio target: {audio_spec.sample_count} samples at "
f"{audio_spec.sample_rate} Hz, {audio_spec.channels} channel(s)"
)
print(
f"lr schedule: base={args.lr:g} warmup={args.warmup_steps} "
f"decay={args.lr_decay} min_ratio={args.min_lr_ratio:g}"
)
output_dir = Path(args.output_dir).resolve()
output_dir.mkdir(parents=True, exist_ok=True)
photo_dir = output_dir / "photo_samples"
video_dir = output_dir / "video_samples"
audio_dir = output_dir / "audio_samples"
ckpt_path = output_dir / f"{frames_path.stem}_multimodal.pt"
best_ckpt_path = output_dir / f"{frames_path.stem}_multimodal_best.pt"
best_val_loss: float | None = None
best_iou = -1.0
model.train()
for step in range(1, args.steps + 1):
current_lr = args.lr * lr_scale_for_step(step, args)
set_optimizer_lr(optimizer, current_lr)
video_coords, video_targets = sample_video_batch(
frames,
args.batch_size,
device,
edge_indices=edge_indices,
edge_fraction=args.edge_sample_fraction,
frame_weights=motion_weights,
)
audio_loss = None
video_temporal_loss = None
video_spatial_loss = None
align_loss = None
optimizer.zero_grad(set_to_none=True)
with autocast_context(device, use_amp):
video_logits = predict_video(model, video_coords)
video_loss = F.binary_cross_entropy_with_logits(video_logits, video_targets, pos_weight=video_pos_weight)
loss = video_loss
if args.video_dice_loss_weight > 0:
loss = loss + args.video_dice_loss_weight * soft_dice_loss(video_logits, video_targets)
if args.video_temporal_loss_weight > 0:
temporal_batch = sample_video_pair_batch(
frames,
args.video_pair_batch_size,
device,
mode="temporal",
frame_weights=motion_weights,
)
if temporal_batch is not None:
coords0, coords1, target0, target1 = temporal_batch
pred_delta = torch.sigmoid(predict_video(model, coords1)) - torch.sigmoid(
predict_video(model, coords0)
)
target_delta = target1 - target0
video_temporal_loss = F.mse_loss(pred_delta, target_delta)
loss = loss + args.video_temporal_loss_weight * video_temporal_loss
if args.video_spatial_loss_weight > 0:
spatial_batch = sample_video_pair_batch(
frames,
args.video_pair_batch_size,
device,
mode="spatial",
)
if spatial_batch is not None:
coords0, coords1, target0, target1 = spatial_batch
pred_delta = torch.sigmoid(predict_video(model, coords1)) - torch.sigmoid(
predict_video(model, coords0)
)
target_delta = target1 - target0
video_spatial_loss = F.mse_loss(pred_delta, target_delta)
loss = loss + args.video_spatial_loss_weight * video_spatial_loss
train_audio_this_step = (
audio is not None
and has_audio_model(model)
and step >= args.audio_start_step
and args.audio_loss_weight > 0
)
if train_audio_this_step:
audio_coords_batch, audio_targets = sample_audio_batch(
audio,
args.audio_batch_size,
device,
args.audio_chunk_samples,
)
audio_pred = predict_audio(model, audio_coords_batch)
audio_loss = F.mse_loss(audio_pred, audio_targets)
if args.audio_chunk_samples > 1:
usable = (audio_pred.shape[0] // args.audio_chunk_samples) * args.audio_chunk_samples
pred_chunks = audio_pred[:usable].reshape(-1, args.audio_chunk_samples, audio_pred.shape[-1])
target_chunks = audio_targets[:usable].reshape(-1, args.audio_chunk_samples, audio_targets.shape[-1])
if args.audio_derivative_loss_weight > 0:
pred_delta = pred_chunks[:, 1:] - pred_chunks[:, :-1]
target_delta = target_chunks[:, 1:] - target_chunks[:, :-1]
audio_loss = audio_loss + args.audio_derivative_loss_weight * F.mse_loss(
pred_delta,
target_delta,
)
if args.audio_stft_loss_weight > 0 and usable > 0:
audio_loss = audio_loss + args.audio_stft_loss_weight * audio_stft_loss(
pred_chunks,
target_chunks,
n_fft=args.audio_stft_n_fft,
hop_length=args.audio_stft_hop_length,
)
audio_ramp = min(1.0, (step - args.audio_start_step + 1) / args.audio_loss_ramp_steps)
loss = loss + args.audio_loss_weight * audio_ramp * audio_loss
if (
args.cross_modal_loss_weight > 0
and audio is not None
and has_audio_model(model)
and hasattr(model, "modality_features")
):
times = torch.rand((args.cross_modal_batch_size, 1), device=device) * 2.0 - 1.0
zeros = torch.zeros((args.cross_modal_batch_size, 2), device=device)
video_features = model.modality_features(torch.cat((times, zeros), dim=-1), 0)
audio_features = model.modality_features(times, 1)
align_loss = F.mse_loss(video_features, audio_features)
loss = loss + args.cross_modal_loss_weight * align_loss
scaler.scale(loss).backward()
if args.grad_clip > 0:
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(model.parameters(), args.grad_clip)
scaler.step(optimizer)
scaler.update()
if ema_state is not None:
update_ema_state(model, ema_state, args.ema_decay)
if step == 1 or step % args.log_every == 0 or step == args.steps:
with torch.no_grad():
pred = (torch.sigmoid(video_logits) >= 0.5).float()
acc = (pred == (video_targets >= 0.5)).float().mean().item()
audio_text = "" if audio_loss is None else f" audio_loss={audio_loss.item():.5f}"
temporal_text = "" if video_temporal_loss is None else f" temporal={video_temporal_loss.item():.5f}"
spatial_text = "" if video_spatial_loss is None else f" spatial={video_spatial_loss.item():.5f}"
align_text = "" if align_loss is None else f" align={align_loss.item():.5f}"
print(
f"step {step:>6}/{args.steps} lr={current_lr:.3g} loss={loss.item():.5f} "
f"video_loss={video_loss.item():.5f}{audio_text}{temporal_text}"
f"{spatial_text}{align_text} pixel_acc={acc:.3f}"
)
if args.val_every > 0 and (step == 1 or step % args.val_every == 0 or step == args.steps):
val_loss, val_acc, val_iou = evaluate_video_quality(
model,
evaluation_coords,
evaluation_targets,
video_pos_weight,
)
print(f"evaluation video_loss={val_loss:.5f} pixel_acc={val_acc:.3f} silhouette_iou={val_iou:.3f}")
if val_iou > best_iou or (math.isclose(val_iou, best_iou) and (best_val_loss is None or val_loss < best_val_loss)):
best_val_loss = val_loss
best_iou = val_iou
save_checkpoint(
best_ckpt_path,
checkpoint_payload(
model=model,
args=args,
video_spec=video_spec,
audio_spec=audio_spec,
frames_path=frames_path,
audio_path=audio_path,
ema_state=ema_state,
step=step,
best_val_loss=best_val_loss,
),
)
print(f"saved best checkpoint -> {best_ckpt_path}")
if args.photo_sample_every > 0 and (step % args.photo_sample_every == 0 or step == args.steps):
sample_path = photo_dir / f"{frames_path.stem}_step_{step:06d}.png"
photo_threshold = args.photo_sample_threshold
if photo_threshold == "calibrated":
photo_threshold, photo_iou = calibrate_render_threshold(
model, frames, args.val_batch_size, device, args.render_gamma, args.render_contrast, args.seed + 17
)
print(f"calibrated photo threshold={photo_threshold} silhouette_iou={photo_iou:.3f}")
render_contact_sheet(
model=model,
spec=video_spec,
output_path=sample_path,
device=device,
tiles=args.photo_sample_tiles,
scale=args.photo_sample_scale,
render_batch=args.render_batch,
threshold=photo_threshold,
)
print(f"photo sample -> {sample_path}")
if args.video_sample_every > 0 and (step % args.video_sample_every == 0 or step == args.steps):
sample_name = f"{frames_path.stem}_step_{step:06d}"
raw_video = video_dir / f"{sample_name}_raw.mp4"
final_video = video_dir / f"{sample_name}.mp4"
if args.video_sample_seconds <= 0:
sample_frames = video_spec.frame_count
else:
sample_frames = min(video_spec.frame_count, int(args.video_sample_seconds * video_spec.fps))
video_threshold = args.video_sample_threshold
if video_threshold == "calibrated":
video_threshold, video_iou = calibrate_render_threshold(
model, frames, args.val_batch_size, device, args.render_gamma, args.render_contrast, args.seed + 17
)
print(f"calibrated video threshold={video_threshold} silhouette_iou={video_iou:.3f}")
render_video_file(
model=model,
spec=video_spec,
output_path=raw_video,
device=device,
scale=args.video_sample_scale,
render_batch=args.render_batch,
threshold=video_threshold,
render_log_every=args.render_log_every,
encoder=args.video_encoder,
crf=args.video_crf,
preset=args.video_preset,
tune=args.video_tune,
supersample=args.video_sample_supersample,
gamma=args.render_gamma,
contrast=args.render_contrast,
threshold_hysteresis=args.render_threshold_hysteresis,
start_frame=0,
frame_count=sample_frames,
)
if audio_spec is not None and has_audio_model(model) and step >= args.audio_start_step:
audio_wav = audio_dir / f"{sample_name}.wav"
render_audio_wav(
model=model,
audio_spec=audio_spec,
output_path=audio_wav,
device=device,
render_batch=args.audio_render_batch,
smooth_samples=args.audio_render_smooth_samples,
normalize=args.audio_normalize,
peak=args.audio_peak,
fade_ms=args.audio_fade_ms,
duration_seconds=sample_frames / video_spec.fps,
)
mux_wav_audio(
raw_video,
audio_wav,
final_video,
copy_video=args.mux_copy_video,
audio_bitrate=args.audio_bitrate,
)
print(f"video sample with generated audio -> {final_video}")
else:
print(f"video sample -> {raw_video}")
save_checkpoint(
ckpt_path,
checkpoint_payload(
model=model,
args=args,
video_spec=video_spec,
audio_spec=audio_spec,
frames_path=frames_path,
audio_path=audio_path,
ema_state=ema_state,
step=args.steps,
best_val_loss=best_val_loss,
),
)
print(f"saved checkpoint -> {ckpt_path}")
if args.render:
render_from_checkpoint(ckpt_path, args)
return ckpt_path
def distilled_checkpoint_payload(
model: nn.Module,
args: argparse.Namespace,
teacher_path: Path,
teacher_ckpt: dict,
video_spec: VideoSpec,
audio_spec: AudioSpec | None,
frames_path: Path,
audio_path: Path | None,
step: int,
best_score: float | None,
quantized: bool,
) -> dict:
payload = checkpoint_payload(
model=model,
args=args,
video_spec=video_spec,
audio_spec=audio_spec,
frames_path=frames_path,
audio_path=audio_path,
ema_state=None,
step=step,
best_val_loss=best_score,
)
payload["distilled_from"] = str(teacher_path)
payload["teacher_step"] = teacher_ckpt.get("step")
payload["qat"] = bool(args.qat)
if quantized:
payload["architecture"] = "unified-v3-int8"
payload["quantized_model"] = pack_int8_state_dict(model, model.state_dict())
payload.pop("model", None)
payload.pop("ema_model", None)
return payload
@torch.inference_mode()
def evaluate_distilled_model(
student: nn.Module,
teacher: nn.Module,
video_coords: torch.Tensor,
video_targets: torch.Tensor,
audio_coords_batch: torch.Tensor | None,
audio_weight: float,
) -> tuple[float, float, float, float | None]:
was_training = student.training
student.eval()
teacher_video = torch.sigmoid(predict_video(teacher, video_coords))
student_video = torch.sigmoid(predict_video(student, video_coords))
video_mse = F.mse_loss(student_video, teacher_video)
student_binary = student_video >= 0.5
teacher_binary = teacher_video >= 0.5
agreement = (student_binary == teacher_binary).float().mean()
target_binary = video_targets >= 0.5
intersection = torch.logical_and(student_binary, target_binary).sum()
union = torch.logical_or(student_binary, target_binary).sum()
iou = intersection.float() / union.clamp_min(1).float()
audio_mse = None
score = video_mse
if audio_coords_batch is not None and has_audio_model(student) and has_audio_model(teacher):
teacher_audio = predict_audio(teacher, audio_coords_batch)
student_audio = predict_audio(student, audio_coords_batch)
audio_mse = F.mse_loss(student_audio, teacher_audio)
score = score + audio_weight * audio_mse
if was_training:
student.train()
return (
float(score.item()),
float(agreement.item()),
float(iou.item()),
None if audio_mse is None else float(audio_mse.item()),
)
def distill(args: argparse.Namespace) -> Path:
set_seed(args.seed)
teacher_path = Path(args.teacher).resolve()
teacher_ckpt = torch.load(teacher_path, map_location="cpu")
video_spec = VideoSpec(**teacher_ckpt["video_spec"])
audio_spec = AudioSpec(**teacher_ckpt["audio_spec"]) if teacher_ckpt.get("audio_spec") is not None else None
frames_path = Path(teacher_ckpt["frames_cache"])
audio_path = Path(teacher_ckpt["audio_cache"]) if teacher_ckpt.get("audio_cache") else None
if not frames_path.exists():
raise FileNotFoundError(f"Teacher frame cache is unavailable: {frames_path}")
if audio_spec is not None and (audio_path is None or not audio_path.exists()):
raise FileNotFoundError(f"Teacher audio cache is unavailable: {audio_path}")
frames = np.load(frames_path, mmap_mode="r")
audio = np.load(audio_path, mmap_mode="r") if audio_path is not None else None
args.video = teacher_ckpt.get("source_video", args.video)
args.train_audio = audio_spec is not None
args.audio_channels = audio_spec.channels if audio_spec is not None else 1
device = choose_device(args.device)
teacher = build_model_from_checkpoint(teacher_ckpt).to(device)
teacher_state, teacher_state_name = checkpoint_model_state(teacher_ckpt, use_ema=True)
teacher.load_state_dict(teacher_state)
teacher.eval()
for parameter in teacher.parameters():
parameter.requires_grad_(False)
student = build_model_from_args(args).to(device)
optimizer = torch.optim.AdamW(student.parameters(), lr=args.lr, weight_decay=args.weight_decay)
use_amp = args.amp and device.type == "cuda"
scaler = make_grad_scaler(use_amp)
qat_start_step = args.qat_start_step or max(1, int(args.steps * 0.8))
edge_indices = None
if args.edge_sample_fraction > 0:
edge_indices = build_edge_indices(frames, args.edge_threshold)
motion_weights = None
if args.motion_sample_fraction > 0:
motion_weights = torch.from_numpy(build_motion_frame_weights(frames)).to(device)
uniform_weight = torch.full_like(motion_weights, 1.0 / motion_weights.numel())
motion_weights = torch.lerp(uniform_weight, motion_weights, args.motion_sample_fraction)
evaluation_coords, evaluation_targets = make_evaluation_batch(
frames,
args.val_batch_size,
device,
args.seed + 29,
)
evaluation_audio_coords = None
if audio is not None:
evaluation_audio_coords, _ = sample_audio_batch(
audio,
args.audio_val_batch_size,
device,
chunk_samples=1,
)
output_dir = Path(args.output_dir).resolve()
output_dir.mkdir(parents=True, exist_ok=True)
base_name = teacher_path.stem.removesuffix("_multimodal_best").removesuffix("_multimodal")
latest_path = output_dir / f"{base_name}_distilled_latest.pt"
best_path = output_dir / f"{base_name}_distilled_best.pt"
int8_path = output_dir / f"{base_name}_distilled_int8.pt"
final_int8_path = output_dir / f"{base_name}_distilled_int8_final.pt"
best_score: float | None = None
best_qat_score: float | None = None
student_params = sum(parameter.numel() for parameter in student.parameters())
teacher_params = sum(parameter.numel() for parameter in teacher.parameters())
print(
f"distilling {teacher_state_name} on {device}: teacher={teacher_params:,} params "
f"student={student_params:,} params ({student_params / teacher_params:.1%})"
)
print(f"QAT starts at step {qat_start_step:,}; INT8 target={int8_path}")
student.train()
qat_active = False
for step in range(1, args.steps + 1):
if args.qat and not qat_active and step >= qat_start_step:
qat_active = True
set_qat_enabled(student, True)
print(f"enabled INT8 fake quantization at step {step:,}")
current_lr = args.lr * lr_scale_for_step(step, args)
set_optimizer_lr(optimizer, current_lr)
video_coords, video_targets = sample_video_batch(
frames,
args.batch_size,
device,
edge_indices=edge_indices,
edge_fraction=args.edge_sample_fraction,
frame_weights=motion_weights,
)
optimizer.zero_grad(set_to_none=True)
with torch.no_grad(), autocast_context(device, use_amp):
teacher_video_logits = predict_video(teacher, video_coords)
with autocast_context(device, use_amp):
student_video_logits = predict_video(student, video_coords)
temperature = args.distill_temperature
teacher_soft = torch.sigmoid(teacher_video_logits / temperature)
soft_video_loss = F.binary_cross_entropy_with_logits(
student_video_logits / temperature,
teacher_soft,
) * (temperature ** 2)
hard_video_loss = F.binary_cross_entropy_with_logits(student_video_logits, video_targets)
video_loss = (1.0 - args.hard_target_weight) * soft_video_loss + args.hard_target_weight * hard_video_loss
loss = video_loss
audio_loss = None
if audio is not None and has_audio_model(student):
audio_coords_batch, audio_targets = sample_audio_batch(
audio,
args.audio_batch_size,
device,
args.audio_chunk_samples,
)
with torch.no_grad():
teacher_audio = predict_audio(teacher, audio_coords_batch)
student_audio = predict_audio(student, audio_coords_batch)
soft_audio_loss = F.mse_loss(student_audio, teacher_audio)
hard_audio_loss = F.mse_loss(student_audio, audio_targets)
audio_loss = (
(1.0 - args.hard_target_weight) * soft_audio_loss
+ args.hard_target_weight * hard_audio_loss
)
usable = (student_audio.shape[0] // args.audio_chunk_samples) * args.audio_chunk_samples
if usable > 0 and args.audio_chunk_samples > 1:
student_chunks = student_audio[:usable].reshape(-1, args.audio_chunk_samples, student_audio.shape[-1])
teacher_chunks = teacher_audio[:usable].reshape(-1, args.audio_chunk_samples, teacher_audio.shape[-1])
if args.audio_derivative_loss_weight > 0:
student_delta = student_chunks[:, 1:] - student_chunks[:, :-1]
teacher_delta = teacher_chunks[:, 1:] - teacher_chunks[:, :-1]
audio_loss = audio_loss + args.audio_derivative_loss_weight * F.mse_loss(
student_delta,
teacher_delta,
)
if args.audio_stft_loss_weight > 0:
audio_loss = audio_loss + args.audio_stft_loss_weight * audio_stft_loss(
student_chunks,
teacher_chunks,
args.audio_stft_n_fft,
args.audio_stft_hop_length,
)
loss = loss + args.audio_loss_weight * audio_loss
scaler.scale(loss).backward()
if args.grad_clip > 0:
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(student.parameters(), args.grad_clip)
scaler.step(optimizer)
scaler.update()
if step == 1 or step % args.log_every == 0 or step == args.steps:
audio_text = "" if audio_loss is None else f" audio={audio_loss.item():.5f}"
print(
f"distill {step:>6}/{args.steps} lr={current_lr:.3g} loss={loss.item():.5f} "
f"video={video_loss.item():.5f}{audio_text} qat={qat_active}"
)
if args.val_every > 0 and (step % args.val_every == 0 or step == args.steps):
score, agreement, iou, audio_mse = evaluate_distilled_model(
student,
teacher,
evaluation_coords,
evaluation_targets,
evaluation_audio_coords,
args.audio_loss_weight,
)
audio_text = "" if audio_mse is None else f" audio_teacher_mse={audio_mse:.6f}"
print(
f"distill evaluation score={score:.6f} teacher_agreement={agreement:.4f} "
f"silhouette_iou={iou:.4f}{audio_text}"
)
payload = distilled_checkpoint_payload(
student,
args,
teacher_path,
teacher_ckpt,
video_spec,
audio_spec,
frames_path,
audio_path,
step,
score,
quantized=False,
)
save_checkpoint(latest_path, payload)
if best_score is None or score < best_score:
best_score = score
save_checkpoint(best_path, payload)
print(f"saved best distilled checkpoint -> {best_path}")
if qat_active and (best_qat_score is None or score < best_qat_score):
best_qat_score = score
int8_payload = distilled_checkpoint_payload(
student,
args,
teacher_path,
teacher_ckpt,
video_spec,
audio_spec,
frames_path,
audio_path,
step,
score,
quantized=True,
)
save_checkpoint(int8_path, int8_payload)
print(f"saved INT8 distilled checkpoint -> {int8_path}")
student.train()
final_payload = distilled_checkpoint_payload(
student,
args,
teacher_path,
teacher_ckpt,
video_spec,
audio_spec,
frames_path,
audio_path,
args.steps,
best_score,
quantized=args.qat,
)
final_path = final_int8_path if args.qat else latest_path
save_checkpoint(final_path, final_payload)
print(f"saved distilled model -> {final_path} ({final_path.stat().st_size / (1024 ** 2):.2f} MiB)")
selected_path = int8_path if args.qat and int8_path.exists() else final_path
if args.render:
render_from_checkpoint(selected_path, args)
return selected_path
@torch.inference_mode()
def render_from_checkpoint(ckpt_path: Path, args: argparse.Namespace) -> Path:
ckpt = torch.load(ckpt_path, map_location="cpu")
video_spec = VideoSpec(**ckpt["video_spec"])
audio_spec = AudioSpec(**ckpt["audio_spec"]) if ckpt["audio_spec"] is not None else None
device = choose_device(args.device)
model = build_model_from_checkpoint(ckpt).to(device)
model_state, state_key = checkpoint_model_state(ckpt, args.use_ema)
model.load_state_dict(model_state)
model.eval()
render_threshold = resolve_calibrated_threshold(
requested=args.render_threshold,
model=model,
frames_path=ckpt.get("frames_cache"),
device=device,
batch_size=args.threshold_calibration_batch,
gamma=args.render_gamma,
contrast=args.render_contrast,
seed=1234,
)
output_dir = Path(args.output_dir).resolve()
output_dir.mkdir(parents=True, exist_ok=True)
raw_video = output_dir / f"{ckpt_path.stem}_nn_raw.mp4"
final_video = output_dir / f"{ckpt_path.stem}_nn.mp4"
render_video_file(
model=model,
spec=video_spec,
output_path=raw_video,
device=device,
scale=args.render_scale,
render_batch=args.render_batch,
threshold=render_threshold,
render_log_every=args.render_log_every,
encoder=args.video_encoder,
crf=args.video_crf,
preset=args.video_preset,
tune=args.video_tune,
supersample=args.render_supersample,
gamma=args.render_gamma,
contrast=args.render_contrast,
threshold_hysteresis=args.render_threshold_hysteresis,
)
if args.audio_source == "source":
source_video = Path(ckpt.get("source_video") or args.video).resolve()
mux_source_audio(
raw_video,
source_video,
final_video,
copy_video=args.mux_copy_video,
duration_seconds=video_spec.duration,
)
print(f"saved video with source audio -> {final_video}")
return final_video
if args.audio_source == "generated" and audio_spec is not None and has_audio_model(model):
audio_wav = output_dir / f"{ckpt_path.stem}_nn.wav"
render_audio_wav(
model=model,
audio_spec=audio_spec,
output_path=audio_wav,
device=device,
render_batch=args.audio_render_batch,
smooth_samples=args.audio_render_smooth_samples,
normalize=args.audio_normalize,
peak=args.audio_peak,
fade_ms=args.audio_fade_ms,
duration_seconds=video_spec.duration,
)
mux_wav_audio(
raw_video,
audio_wav,
final_video,
copy_video=args.mux_copy_video,
audio_bitrate=args.audio_bitrate,
)
print(f"saved video with generated audio -> {final_video}")
return final_video
print(f"saved video -> {raw_video}")
return raw_video
def make_preview(args: argparse.Namespace) -> Path:
ckpt_path = Path(args.checkpoint).resolve()
ckpt = torch.load(ckpt_path, map_location="cpu")
video_spec = VideoSpec(**ckpt["video_spec"])
device = choose_device(args.device)
model = build_model_from_checkpoint(ckpt).to(device)
model_state, state_key = checkpoint_model_state(ckpt, args.use_ema)
model.load_state_dict(model_state)
model.eval()
preview_threshold = resolve_calibrated_threshold(
requested=args.render_threshold,
model=model,
frames_path=ckpt.get("frames_cache"),
device=device,
batch_size=args.threshold_calibration_batch,
gamma=1.0,
contrast=1.0,
seed=1234,
)
output_dir = Path(args.output_dir).resolve()
output_dir.mkdir(parents=True, exist_ok=True)
preview_path = output_dir / f"{ckpt_path.stem}_preview.png"
render_contact_sheet(
model=model,
spec=video_spec,
output_path=preview_path,
device=device,
tiles=args.tiles,
scale=args.render_scale,
render_batch=args.render_batch,
threshold=preview_threshold,
)
print(f"saved preview -> {preview_path}")
return preview_path
def add_common_args(parser: argparse.ArgumentParser) -> None:
parser.add_argument("--video", default="Bad_Apple.mp4", help="Source video path.")
parser.add_argument("--output-dir", default="outputs", help="Directory for checkpoints and renders.")
parser.add_argument("--device", default="auto", help="auto, cpu, cuda, or a PyTorch device like cuda:0.")
def add_data_args(parser: argparse.ArgumentParser) -> None:
parser.add_argument("--cache-dir", default="cache")
parser.add_argument("--width", type=int, default=160)
parser.add_argument("--height", type=int, default=120)
parser.add_argument("--match-source-resolution", action="store_true")
parser.add_argument("--frame-stride", type=int, default=1, help="Use every Nth source frame.")
parser.add_argument(
"--threshold",
type=parse_data_threshold,
default=None,
help="Training target threshold. Use none/gray/-1 to keep soft grayscale targets.",
)
parser.add_argument("--max-frames", type=int, default=None)
parser.add_argument("--force", action="store_true")
parser.add_argument("--audio-sample-rate", type=int, default=16000)
parser.add_argument("--audio-channels", type=int, default=2)
parser.add_argument("--match-source-audio", action="store_true")
parser.add_argument("--train-audio", action=argparse.BooleanOptionalAction, default=True)
def add_render_args(parser: argparse.ArgumentParser) -> None:
parser.add_argument("--low-memory", action="store_true", help="Use smaller training/render chunks.")
parser.add_argument("--render-scale", type=int, default=3)
parser.add_argument("--render-batch", type=int, default=65536)
parser.add_argument("--render-log-every", type=int, default=100)
parser.add_argument(
"--render-threshold",
type=parse_render_threshold,
default="calibrated",
help="Output threshold: calibrated, auto, adaptive, integer, or none/gray/-1.",
)
parser.add_argument("--threshold-calibration-batch", type=int, default=65536)
parser.add_argument("--render-supersample", type=int, default=1)
parser.add_argument("--render-gamma", type=float, default=1.0)
parser.add_argument("--render-contrast", type=float, default=1.0)
parser.add_argument("--render-threshold-hysteresis", type=int, default=4)
parser.add_argument("--video-encoder", choices=("ffmpeg", "opencv"), default="ffmpeg")
parser.add_argument("--video-crf", type=int, default=16)
parser.add_argument("--video-preset", default="slow")
parser.add_argument("--video-tune", default="animation")
parser.add_argument("--mux-copy-video", action=argparse.BooleanOptionalAction, default=True)
parser.add_argument("--generated-audio", action=argparse.BooleanOptionalAction, default=True)
parser.add_argument("--audio-source", choices=("generated", "source", "none"), default=None)
parser.add_argument("--audio-bitrate", default="192k")
parser.add_argument("--audio-render-batch", type=int, default=65536)
parser.add_argument("--audio-render-smooth-samples", type=int, default=5, help="Moving-average smoothing for generated audio. Use 1 to disable.")
parser.add_argument("--audio-normalize", action=argparse.BooleanOptionalAction, default=True)
parser.add_argument("--audio-peak", type=float, default=0.95)
parser.add_argument("--audio-fade-ms", type=float, default=5.0)
parser.add_argument("--use-ema", action=argparse.BooleanOptionalAction, default=True)
def build_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(description="Train a tiny multimodal coordinate net to reconstruct Bad Apple.")
subparsers = parser.add_subparsers(dest="command", required=True)
prep = subparsers.add_parser("prepare", help="Extract and cache small video frames and audio.")
add_common_args(prep)
add_data_args(prep)
train_parser = subparsers.add_parser("train", help="Train the multimodal video+audio model.")
add_common_args(train_parser)
add_data_args(train_parser)
train_parser.add_argument("--hidden", type=int, default=256)
train_parser.add_argument("--layers", type=int, default=6)
train_parser.add_argument("--bands", type=int, default=13)
train_parser.add_argument("--audio-hidden", type=int, default=256)
train_parser.add_argument("--audio-layers", type=int, default=6)
train_parser.add_argument("--audio-bands", type=int, default=22)
train_parser.add_argument("--modality-embedding-dim", type=int, default=8)
train_parser.add_argument("--temporal-latent-anchors", type=int, default=512)
train_parser.add_argument("--temporal-latent-dim", type=int, default=32)
train_parser.add_argument("--gradient-checkpointing", action=argparse.BooleanOptionalAction, default=True)
train_parser.add_argument("--steps", type=int, default=100000)
train_parser.add_argument("--batch-size", type=int, default=32768)
train_parser.add_argument("--audio-batch-size", type=int, default=16384)
train_parser.add_argument("--audio-loss-weight", type=float, default=0.15)
train_parser.add_argument("--audio-start-step", type=int, default=5000)
train_parser.add_argument("--audio-loss-ramp-steps", type=int, default=5000)
train_parser.add_argument("--audio-chunk-samples", type=int, default=2048, help="Train audio on contiguous chunks. Use 1 for random isolated samples.")
train_parser.add_argument("--audio-derivative-loss-weight", type=float, default=0.5, help="Extra continuity loss for generated audio chunks.")
train_parser.add_argument("--audio-stft-loss-weight", type=float, default=0.02)
train_parser.add_argument("--audio-stft-n-fft", type=int, default=512)
train_parser.add_argument("--audio-stft-hop-length", type=int, default=128)
train_parser.add_argument("--cross-modal-loss-weight", type=float, default=0.0)
train_parser.add_argument("--cross-modal-batch-size", type=int, default=1024)
train_parser.add_argument("--video-temporal-loss-weight", type=float, default=0.1)
train_parser.add_argument("--video-spatial-loss-weight", type=float, default=0.05)
train_parser.add_argument("--video-dice-loss-weight", type=float, default=0.1)
train_parser.add_argument("--video-pair-batch-size", type=int, default=8192)
train_parser.add_argument("--edge-sample-fraction", type=float, default=0.35)
train_parser.add_argument("--edge-threshold", type=int, default=24)
train_parser.add_argument("--motion-sample-fraction", type=float, default=0.35)
train_parser.add_argument("--balanced-video-loss", action=argparse.BooleanOptionalAction, default=True)
train_parser.add_argument("--lr", type=float, default=2e-3)
train_parser.add_argument("--warmup-steps", type=int, default=1000)
train_parser.add_argument("--warmup-start-ratio", type=float, default=0.05)
train_parser.add_argument("--lr-decay", choices=("cosine", "linear", "none"), default="cosine")
train_parser.add_argument("--min-lr-ratio", type=float, default=0.05)
train_parser.add_argument("--weight-decay", type=float, default=1e-4)
train_parser.add_argument("--grad-clip", type=float, default=1.0)
train_parser.add_argument("--ema-decay", type=float, default=0.995)
train_parser.add_argument("--val-every", type=int, default=1000)
train_parser.add_argument("--val-batch-size", type=int, default=16384)
train_parser.add_argument("--amp", action=argparse.BooleanOptionalAction, default=True)
train_parser.add_argument("--seed", type=int, default=1234)
train_parser.add_argument("--resume", default=None)
train_parser.add_argument("--resume-ema", action=argparse.BooleanOptionalAction, default=True)
train_parser.add_argument("--log-every", type=int, default=250)
train_parser.add_argument("--render", action="store_true")
add_render_args(train_parser)
train_parser.add_argument("--photo-sample-every", type=int, default=500, help="Write PNG samples every N steps. Use 0 to disable.")
train_parser.add_argument("--photo-sample-tiles", type=int, default=8)
train_parser.add_argument("--photo-sample-scale", type=int, default=2)
train_parser.add_argument(
"--photo-sample-threshold",
type=parse_render_threshold,
default="auto",
help="Set to none/gray/-1 to keep grayscale photo samples.",
)
train_parser.add_argument("--video-sample-every", type=int, default=1000, help="Write MP4 samples every N steps. Use 0 to disable.")
train_parser.add_argument("--video-sample-seconds", type=float, default=4.0, help="Length of each training MP4 sample. Use 0 for full video samples.")
train_parser.add_argument("--video-sample-scale", type=int, default=2)
train_parser.add_argument("--video-sample-supersample", type=int, default=1)
train_parser.add_argument(
"--video-sample-threshold",
type=parse_render_threshold,
default="auto",
help="Set to none/gray/-1 to keep grayscale video samples.",
)
train_parser.add_argument("--sample-every", type=int, default=None, help=argparse.SUPPRESS)
distill_parser = subparsers.add_parser("distill", help="Distill a checkpoint into a QAT INT8 student.")
add_common_args(distill_parser)
distill_parser.set_defaults(output_dir="outputs/distilled")
distill_parser.add_argument("teacher", help="Teacher checkpoint path.")
distill_parser.add_argument("--hidden", type=int, default=192)
distill_parser.add_argument("--layers", type=int, default=4)
distill_parser.add_argument("--bands", type=int, default=11)
distill_parser.add_argument("--audio-hidden", type=int, default=192)
distill_parser.add_argument("--audio-layers", type=int, default=4)
distill_parser.add_argument("--audio-bands", type=int, default=18)
distill_parser.add_argument("--modality-embedding-dim", type=int, default=8)
distill_parser.add_argument("--temporal-latent-anchors", type=int, default=384)
distill_parser.add_argument("--temporal-latent-dim", type=int, default=24)
distill_parser.add_argument("--gradient-checkpointing", action=argparse.BooleanOptionalAction, default=True)
distill_parser.add_argument("--steps", type=int, default=30000)
distill_parser.add_argument("--batch-size", type=int, default=8192)
distill_parser.add_argument("--audio-batch-size", type=int, default=4096)
distill_parser.add_argument("--audio-val-batch-size", type=int, default=4096)
distill_parser.add_argument("--audio-chunk-samples", type=int, default=2048)
distill_parser.add_argument("--audio-loss-weight", type=float, default=1.0)
distill_parser.add_argument("--audio-derivative-loss-weight", type=float, default=0.5)
distill_parser.add_argument("--audio-stft-loss-weight", type=float, default=0.02)
distill_parser.add_argument("--audio-stft-n-fft", type=int, default=512)
distill_parser.add_argument("--audio-stft-hop-length", type=int, default=128)
distill_parser.add_argument("--hard-target-weight", type=float, default=0.05)
distill_parser.add_argument("--distill-temperature", type=float, default=2.0)
distill_parser.add_argument("--qat", action=argparse.BooleanOptionalAction, default=True)
distill_parser.add_argument("--qat-start-step", type=int, default=None)
distill_parser.add_argument("--edge-sample-fraction", type=float, default=0.35)
distill_parser.add_argument("--edge-threshold", type=int, default=24)
distill_parser.add_argument("--motion-sample-fraction", type=float, default=0.35)
distill_parser.add_argument("--lr", type=float, default=1e-3)
distill_parser.add_argument("--warmup-steps", type=int, default=500)
distill_parser.add_argument("--warmup-start-ratio", type=float, default=0.05)
distill_parser.add_argument("--lr-decay", choices=("cosine", "linear", "none"), default="cosine")
distill_parser.add_argument("--min-lr-ratio", type=float, default=0.05)
distill_parser.add_argument("--weight-decay", type=float, default=1e-4)
distill_parser.add_argument("--grad-clip", type=float, default=1.0)
distill_parser.add_argument("--val-every", type=int, default=500)
distill_parser.add_argument("--val-batch-size", type=int, default=4096)
distill_parser.add_argument("--log-every", type=int, default=100)
distill_parser.add_argument("--amp", action=argparse.BooleanOptionalAction, default=True)
distill_parser.add_argument("--seed", type=int, default=1234)
distill_parser.add_argument("--render", action="store_true")
add_render_args(distill_parser)
render = subparsers.add_parser("render", help="Render a trained checkpoint back to MP4.")
add_common_args(render)
render.add_argument("checkpoint")
add_render_args(render)
preview = subparsers.add_parser("preview", help="Make a still contact sheet from a trained checkpoint.")
add_common_args(preview)
preview.add_argument("checkpoint")
preview.add_argument("--tiles", type=int, default=8)
preview.add_argument("--render-scale", type=int, default=2)
preview.add_argument("--render-batch", type=int, default=65536)
preview.add_argument("--low-memory", action="store_true", help="Use smaller render chunks.")
preview.add_argument(
"--render-threshold",
type=parse_render_threshold,
default="calibrated",
help="Output threshold: calibrated, auto, adaptive, integer, or none/gray/-1.",
)
preview.add_argument("--threshold-calibration-batch", type=int, default=65536)
preview.add_argument("--use-ema", action=argparse.BooleanOptionalAction, default=True)
return parser
def normalize_args(args: argparse.Namespace) -> argparse.Namespace:
if hasattr(args, "threshold"):
args.threshold = parse_data_threshold(args.threshold)
for name in ("render_threshold", "photo_sample_threshold", "video_sample_threshold"):
if hasattr(args, name):
setattr(args, name, parse_render_threshold(getattr(args, name)))
if getattr(args, "sample_every", None) is not None:
args.photo_sample_every = args.sample_every
if getattr(args, "low_memory", False):
for name, limit in (
("batch_size", 8192),
("audio_batch_size", 4096),
("audio_val_batch_size", 4096),
("video_pair_batch_size", 2048),
("cross_modal_batch_size", 256),
("val_batch_size", 4096),
("render_batch", 16384),
("threshold_calibration_batch", 16384),
("audio_render_batch", 16384),
):
if hasattr(args, name):
setattr(args, name, min(getattr(args, name), limit))
if hasattr(args, "audio_source") and args.audio_source is None:
args.audio_source = "generated" if getattr(args, "generated_audio", True) else "none"
if hasattr(args, "audio_channels"):
args.audio_channels = max(1, args.audio_channels)
if hasattr(args, "audio_chunk_samples"):
args.audio_chunk_samples = max(1, args.audio_chunk_samples)
if hasattr(args, "audio_val_batch_size"):
args.audio_val_batch_size = max(1, args.audio_val_batch_size)
if hasattr(args, "audio_render_smooth_samples"):
args.audio_render_smooth_samples = max(1, args.audio_render_smooth_samples)
if hasattr(args, "audio_start_step"):
args.audio_start_step = max(1, args.audio_start_step)
if hasattr(args, "audio_loss_ramp_steps"):
args.audio_loss_ramp_steps = max(1, args.audio_loss_ramp_steps)
if hasattr(args, "audio_stft_n_fft"):
args.audio_stft_n_fft = max(16, args.audio_stft_n_fft)
if hasattr(args, "audio_stft_hop_length"):
args.audio_stft_hop_length = max(1, args.audio_stft_hop_length)
if hasattr(args, "edge_sample_fraction"):
args.edge_sample_fraction = min(max(args.edge_sample_fraction, 0.0), 1.0)
if hasattr(args, "motion_sample_fraction"):
args.motion_sample_fraction = min(max(args.motion_sample_fraction, 0.0), 1.0)
if hasattr(args, "edge_threshold"):
args.edge_threshold = max(1, args.edge_threshold)
if hasattr(args, "cross_modal_batch_size"):
args.cross_modal_batch_size = max(1, args.cross_modal_batch_size)
if hasattr(args, "video_pair_batch_size"):
args.video_pair_batch_size = max(1, args.video_pair_batch_size)
if hasattr(args, "val_batch_size"):
args.val_batch_size = max(1, args.val_batch_size)
if hasattr(args, "threshold_calibration_batch"):
args.threshold_calibration_batch = max(1, args.threshold_calibration_batch)
if hasattr(args, "temporal_latent_anchors"):
args.temporal_latent_anchors = max(2, args.temporal_latent_anchors)
if hasattr(args, "temporal_latent_dim"):
args.temporal_latent_dim = max(1, args.temporal_latent_dim)
if hasattr(args, "video_dice_loss_weight"):
args.video_dice_loss_weight = max(0.0, args.video_dice_loss_weight)
if hasattr(args, "ema_decay"):
args.ema_decay = min(max(args.ema_decay, 0.0), 0.999999)
if hasattr(args, "grad_clip"):
args.grad_clip = max(0.0, args.grad_clip)
if hasattr(args, "render_supersample"):
args.render_supersample = max(1, args.render_supersample)
if hasattr(args, "video_sample_supersample"):
args.video_sample_supersample = max(1, args.video_sample_supersample)
if hasattr(args, "render_gamma"):
args.render_gamma = max(1e-3, args.render_gamma)
if hasattr(args, "render_contrast"):
args.render_contrast = max(0.0, args.render_contrast)
if hasattr(args, "render_threshold_hysteresis"):
args.render_threshold_hysteresis = max(0, args.render_threshold_hysteresis)
if hasattr(args, "video_crf"):
args.video_crf = min(max(args.video_crf, 0), 51)
if hasattr(args, "audio_peak"):
args.audio_peak = min(max(args.audio_peak, 0.0), 1.0)
if hasattr(args, "audio_fade_ms"):
args.audio_fade_ms = max(0.0, args.audio_fade_ms)
if hasattr(args, "min_lr_ratio"):
args.min_lr_ratio = min(max(args.min_lr_ratio, 0.0), 1.0)
if hasattr(args, "warmup_start_ratio"):
args.warmup_start_ratio = min(max(args.warmup_start_ratio, 0.0), 1.0)
if hasattr(args, "hard_target_weight"):
args.hard_target_weight = min(max(args.hard_target_weight, 0.0), 1.0)
if hasattr(args, "distill_temperature"):
args.distill_temperature = max(args.distill_temperature, 1e-3)
if hasattr(args, "qat_start_step") and args.qat_start_step is not None:
args.qat_start_step = max(1, args.qat_start_step)
return args
def main() -> None:
parser = build_parser()
args = normalize_args(parser.parse_args())
if args.command == "prepare":
prepare_dataset(args)
elif args.command == "train":
train(args)
elif args.command == "distill":
distill(args)
elif args.command == "render":
render_from_checkpoint(Path(args.checkpoint).resolve(), args)
elif args.command == "preview":
make_preview(args)
else:
parser.error(f"unknown command: {args.command}")
if __name__ == "__main__":
main()