InstructAV2AV / scripts /train.py
suimu's picture
init
e0177dc
Raw
History Blame Contribute Delete
12.7 kB
#!/usr/bin/env python3
"""Unified trainer for instruction-based audio-video editing."""
from __future__ import annotations
import argparse
import logging
import sys
from pathlib import Path
import torch
from accelerate import Accelerator
from accelerate.utils import DistributedDataParallelKwargs, set_seed
from omegaconf import OmegaConf
from tqdm import tqdm
REPO_ROOT = Path(__file__).resolve().parents[2]
if str(REPO_ROOT) not in sys.path:
sys.path.insert(0, str(REPO_ROOT))
from ovi.ovi_fusion_engine import OviFusionEngine
from ovi.utils.av_edit_dataset import AVEditDataset
DEFAULT_CONFIG = REPO_ROOT / "ovi/configs/train/train_av_edit.yaml"
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--config-file", default=str(DEFAULT_CONFIG))
parser.add_argument("--data-manifest", help="Override dataset.metadata_path.")
parser.add_argument("--finetune-path", help="Override the initialization checkpoint.")
parser.add_argument("--output-dir", help="Override output_dir.")
return parser.parse_args()
def configure_logging(is_main_process: bool) -> None:
logging.basicConfig(
level=logging.INFO if is_main_process else logging.ERROR,
format="[%(asctime)s] %(levelname)s: %(message)s",
handlers=[logging.StreamHandler(stream=sys.stdout)],
)
def apply_overrides(config, args: argparse.Namespace) -> None:
if args.data_manifest:
config.dataset.metadata_path = args.data_manifest
if args.finetune_path:
config.finetune_path = args.finetune_path
if args.output_dir:
config.output_dir = args.output_dir
config.av2av_edit = True
config.mode = "t2v"
def resolve_training_mode(config) -> tuple[str, bool, bool]:
has_video = bool(config.get("has_video", True))
has_audio = bool(config.get("has_audio", True))
if has_video and has_audio:
return "av2av", has_video, has_audio
if has_video:
return "video", has_video, has_audio
if has_audio:
return "audio", has_video, has_audio
raise ValueError("At least one of has_video or has_audio must be true.")
def build_dataset(config) -> AVEditDataset:
values = OmegaConf.to_container(config.dataset, resolve=True)
_, has_video, has_audio = resolve_training_mode(config)
values["has_video"] = has_video
values["has_audio"] = has_audio
return AVEditDataset(**values)
def configure_trainable_parameters(model: torch.nn.Module, config) -> None:
include = list(config.training.get("trainable_name_contains", []))
exclude = list(config.training.get("frozen_name_contains", []))
trainable_count = 0
total_count = 0
for name, parameter in model.named_parameters():
selected = not include or any(token in name for token in include)
selected = selected and not any(token in name for token in exclude)
parameter.requires_grad_(selected)
total_count += parameter.numel()
if selected:
trainable_count += parameter.numel()
if trainable_count == 0:
raise ValueError("No trainable parameters remain after applying name filters.")
logging.info(
"Trainable fusion parameters: %.3fB / %.3fB",
trainable_count / 1e9,
total_count / 1e9,
)
@torch.no_grad()
def encode_batch(
engine: OviFusionEngine,
data: dict,
inverse_pair: bool,
has_video: bool,
has_audio: bool,
) -> dict:
text_embeddings = engine.text_model(
[data["instruction"]], engine.text_model.device
)
text_embedding = text_embeddings[0].to(device=engine.device, dtype=engine.target_dtype)
inputs = {"text_embeddings": text_embedding}
if has_audio:
if inverse_pair:
source_audio = data["audio_result_np"]
target_audio = data["audio_ori_np"]
else:
source_audio = data["audio_ori_np"]
target_audio = data["audio_result_np"]
source_audio_tensor = (
torch.from_numpy(source_audio).float().unsqueeze(0).to(engine.device)
)
target_audio_tensor = (
torch.from_numpy(target_audio).float().unsqueeze(0).to(engine.device)
)
inputs["audio_ori_latents"] = (
engine.vae_model_audio.wrapped_encode(source_audio_tensor)
.squeeze(0)
.transpose(0, 1)
)
inputs["audio_result_latents"] = (
engine.vae_model_audio.wrapped_encode(target_audio_tensor)
.squeeze(0)
.transpose(0, 1)
)
if has_video:
if inverse_pair:
source_video = data["video_result_np"]
target_video = data["video_ori_np"]
else:
source_video = data["video_ori_np"]
target_video = data["video_result_np"]
source_video_tensor = (
torch.from_numpy(source_video)
.float()
.unsqueeze(0)
.to(device=engine.device, dtype=engine.target_dtype)
/ 127.5
- 1.0
)
target_video_tensor = (
torch.from_numpy(target_video)
.float()
.unsqueeze(0)
.to(device=engine.device, dtype=engine.target_dtype)
/ 127.5
- 1.0
)
inputs["video_ori_latents"] = (
engine.vae_model_video.wrapped_encode(source_video_tensor)
.to(engine.target_dtype)
.squeeze(0)
)
inputs["video_result_latents"] = (
engine.vae_model_video.wrapped_encode(target_video_tensor)
.to(engine.target_dtype)
.squeeze(0)
)
return inputs
def save_checkpoint(
accelerator: Accelerator,
prepared_model: torch.nn.Module,
output_dir: Path,
name: str,
) -> None:
accelerator.wait_for_everyone()
if not accelerator.is_main_process:
return
unwrapped_model = accelerator.unwrap_model(prepared_model)
state_dict = accelerator.get_state_dict(prepared_model)
trainable_names = {
parameter_name
for parameter_name, parameter in unwrapped_model.named_parameters()
if parameter.requires_grad
}
state_dict = {
key: value.detach().cpu()
for key, value in state_dict.items()
if key in trainable_names
}
output_dir.mkdir(parents=True, exist_ok=True)
checkpoint_path = output_dir / name
accelerator.save(state_dict, checkpoint_path, safe_serialization=True)
logging.info("Saved checkpoint: %s", checkpoint_path)
def main() -> None:
args = parse_args()
config = OmegaConf.load(args.config_file)
apply_overrides(config, args)
training_mode, has_video, has_audio = resolve_training_mode(config)
training = config.training
report_to = training.get("report_to", None)
if isinstance(report_to, str) and report_to.lower() in {"", "none", "null"}:
report_to = None
accelerator = Accelerator(
gradient_accumulation_steps=int(training.get("gradient_accumulation_steps", 1)),
mixed_precision=str(training.get("mixed_precision", "bf16")),
log_with=report_to,
kwargs_handlers=[DistributedDataParallelKwargs(find_unused_parameters=False)],
)
configure_logging(accelerator.is_main_process)
set_seed(int(config.get("seed", 103)), device_specific=True)
if not torch.cuda.is_available():
raise RuntimeError("Training requires a CUDA device.")
device = accelerator.local_process_index
torch.cuda.set_device(device)
output_dir = Path(config.get("output_dir", "./outputs/train_av_edit")).expanduser().resolve()
if accelerator.is_main_process:
output_dir.mkdir(parents=True, exist_ok=True)
OmegaConf.save(config, output_dir / "config_resolved.yaml", resolve=True)
dataset = build_dataset(config)
batch_size = int(training.get("batch_size", 1))
if batch_size != 1:
raise ValueError("AVEditDataset currently requires training.batch_size=1 for variable AV lengths.")
dataloader = torch.utils.data.DataLoader(
dataset,
batch_size=batch_size,
shuffle=True,
num_workers=int(training.get("num_workers", 4)),
pin_memory=True,
collate_fn=lambda batch: batch[0],
)
precision = str(training.get("mixed_precision", "bf16"))
target_dtype = {
"bf16": torch.bfloat16,
"fp16": torch.float16,
"no": torch.float32,
}.get(precision)
if target_dtype is None:
raise ValueError("training.mixed_precision must be one of: bf16, fp16, no.")
engine = OviFusionEngine(config=config, device=device, target_dtype=target_dtype)
configure_trainable_parameters(engine.model, config)
loss_function = {
"av2av": engine.training_loss_av2av,
"video": engine.training_loss_video,
"audio": engine.training_loss_audio,
}[training_mode]
logging.info(
"Training mode: %s (has_video=%s, has_audio=%s)",
training_mode,
has_video,
has_audio,
)
optimizer = torch.optim.AdamW(
(parameter for parameter in engine.model.parameters() if parameter.requires_grad),
lr=float(training.get("learning_rate", 1e-5)),
weight_decay=float(training.get("weight_decay", 0.01)),
eps=float(training.get("adam_epsilon", 1e-8)),
)
scheduler = torch.optim.lr_scheduler.ConstantLR(optimizer, factor=1.0)
engine.model, optimizer, dataloader, scheduler = accelerator.prepare(
engine.model, optimizer, dataloader, scheduler
)
engine.model.train()
if report_to:
accelerator.init_trackers(
project_name=str(config.get("project_name", "instructav2av")),
config=OmegaConf.to_container(config, resolve=True),
init_kwargs={"wandb": {"name": str(training.get("run_name", "av-edit-sft"))}},
)
num_epochs = int(training.get("num_epochs", 1))
max_train_steps = training.get("max_train_steps", None)
max_train_steps = None if max_train_steps is None else int(max_train_steps)
save_steps = int(training.get("save_steps", 500))
max_grad_norm = float(training.get("max_grad_norm", 1.0))
inverse_pair = bool(training.get("inverse_pair", False))
if inverse_pair:
logging.warning("training.inverse_pair=true: source and target AV are intentionally swapped.")
global_step = 0
stop_training = False
optimizer.zero_grad(set_to_none=True)
for epoch in range(num_epochs):
progress = tqdm(
dataloader,
disable=not accelerator.is_local_main_process,
desc=f"Epoch {epoch + 1}/{num_epochs}",
)
for data in progress:
with accelerator.accumulate(engine.model):
inputs = encode_batch(
engine,
data,
inverse_pair=inverse_pair,
has_video=has_video,
has_audio=has_audio,
)
loss = loss_function(**inputs)
accelerator.backward(loss)
if accelerator.sync_gradients and max_grad_norm > 0:
accelerator.clip_grad_norm_(engine.model.parameters(), max_grad_norm)
optimizer.step()
scheduler.step()
optimizer.zero_grad(set_to_none=True)
if accelerator.sync_gradients:
global_step += 1
current_lr = optimizer.param_groups[0]["lr"]
if report_to:
accelerator.log(
{"train/loss": loss.detach().item(), "train/lr": current_lr},
step=global_step,
)
progress.set_postfix(loss=f"{loss.detach().item():.4f}", step=global_step)
if save_steps > 0 and global_step % save_steps == 0:
save_checkpoint(
accelerator,
engine.model,
output_dir,
f"step-{global_step}.safetensors",
)
if max_train_steps is not None and global_step >= max_train_steps:
stop_training = True
break
if stop_training:
break
if global_step == 0:
raise RuntimeError("Training completed without an optimizer step.")
if save_steps <= 0 or global_step % save_steps != 0:
save_checkpoint(
accelerator,
engine.model,
output_dir,
f"step-{global_step}.safetensors",
)
if report_to:
accelerator.end_training()
if __name__ == "__main__":
main()