Spaces:
Running on Zero
Running on Zero
| #!/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, | |
| ) | |
| 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() | |