echo-infinity / inference /inference.py
multimodalart's picture
multimodalart HF Staff
Upload folder using huggingface_hub
3e936b2 verified
Raw
History Blame Contribute Delete
9.81 kB
import sys
from pathlib import Path
_PROJECT_ROOT = str(Path(__file__).resolve().parents[1])
if _PROJECT_ROOT in sys.path:
sys.path.remove(_PROJECT_ROOT)
sys.path.insert(0, _PROJECT_ROOT)
import argparse
import torch
import os
from omegaconf import OmegaConf
from tqdm import tqdm
from torchvision import transforms
from torchvision.io import write_video
from einops import rearrange
import torch.distributed as dist
from torch.utils.data import DataLoader, SequentialSampler
from torch.utils.data.distributed import DistributedSampler
from pipeline import CausalInferencePipeline
from utils.dataset import TextDataset
from utils.misc import set_seed
from utils.memory import get_cuda_free_memory_gb, DynamicSwapInstaller
import json
import glob
parser = argparse.ArgumentParser()
parser.add_argument('--config_path', type=str, help='Path to the config file')
parser.add_argument('--checkpoint_path', type=str, default=None, help='Override config generator_ckpt')
parser.add_argument('--lora_ckpt', type=str, default=None, help='Override config lora_ckpt')
parser.add_argument('--data_path', type=str, default=None, help='Override config data_path')
parser.add_argument('--output_folder', type=str, default=None, help='Override config output_folder')
parser.add_argument('--use_ema', action='store_true', help='Override config use_ema')
parser.add_argument('--seed', type=int, default=None, help='Override config seed')
parser.add_argument('--num_samples', type=int, default=None, help='Override config num_samples')
args = parser.parse_args()
config = OmegaConf.load(args.config_path)
if args.checkpoint_path:
config.generator_ckpt = args.checkpoint_path
if args.lora_ckpt:
config.lora_ckpt = args.lora_ckpt
if args.data_path:
config.data_path = args.data_path
if args.output_folder:
config.output_folder = args.output_folder
if args.use_ema:
config.use_ema = True
if args.seed is not None:
config.seed = args.seed
if args.num_samples is not None:
config.num_samples = args.num_samples
if 'LOCAL_RANK' in os.environ:
os.environ['NCCL_CROSS_NIC'] = '1'
os.environ['NCCL_DEBUG'] = os.environ.get('NCCL_DEBUG', 'INFO')
os.environ['NCCL_TIMEOUT'] = os.environ.get('NCCL_TIMEOUT', '1800')
local_rank = int(os.environ['LOCAL_RANK'])
world_size = int(os.environ.get('WORLD_SIZE', '1'))
rank = int(os.environ.get('RANK', str(local_rank)))
torch.cuda.set_device(local_rank)
device = torch.device(f'cuda:{local_rank}')
if not dist.is_initialized():
dist.init_process_group(backend='nccl', rank=rank, world_size=world_size, timeout=torch.distributed.constants.default_pg_timeout)
set_seed(config.seed + local_rank)
config.distributed = True
if rank == 0:
print(f'[Rank {rank}] Initialized distributed processing on device {device}')
else:
local_rank = 0
rank = 0
device = torch.device('cuda')
set_seed(config.seed)
config.distributed = False
print(f'Single GPU mode on device {device}')
print(f'Free VRAM {get_cuda_free_memory_gb(device)} GB')
low_memory = get_cuda_free_memory_gb(device) < 40
torch.set_grad_enabled(False)
pipeline = CausalInferencePipeline(config, device=device)
if config.generator_ckpt:
state_dict = torch.load(config.generator_ckpt, map_location='cpu')
if 'generator' in state_dict or 'generator_ema' in state_dict:
if config.use_ema and 'generator_ema' in state_dict:
raw_gen_state_dict = state_dict['generator_ema']
if 'generator' in state_dict:
enc_keys = {k: v for k, v in state_dict['generator'].items() if 'query_memory_encoder' in k}
if enc_keys:
raw_gen_state_dict = dict(raw_gen_state_dict)
raw_gen_state_dict.update(enc_keys)
else:
raw_gen_state_dict = state_dict.get('generator', state_dict.get('generator_ema'))
elif 'model' in state_dict:
raw_gen_state_dict = state_dict['model']
else:
raise ValueError(f'Generator state dict not found in {config.generator_ckpt}')
def _clean_key(name: str) -> str:
return name.replace('_fsdp_wrapped_module.', '')
cleaned_state_dict = {_clean_key(k): v for k, v in raw_gen_state_dict.items()}
missing, unexpected = pipeline.generator.load_state_dict(cleaned_state_dict, strict=False)
if local_rank == 0:
enc_loaded = sum((1 for k in cleaned_state_dict if 'query_memory_encoder' in k))
if len(missing) > 0:
print(f'[Warning] {len(missing)} parameters missing: {missing[:8]} ...')
if len(unexpected) > 0:
print(f'[Warning] {len(unexpected)} unexpected parameters: {unexpected[:8]} ...')
pipeline.is_lora_enabled = False
if getattr(config, 'adapter', None):
from utils.lora_utils import configure_lora_for_model
import peft
if local_rank == 0:
print(f'LoRA enabled with config: {config.adapter}')
print('Applying LoRA to generator (inference)...')
pipeline.generator.model = configure_lora_for_model(pipeline.generator.model, model_name='generator', lora_config=config.adapter, is_main_process=local_rank == 0)
lora_ckpt_path = getattr(config, 'lora_ckpt', None)
if lora_ckpt_path:
if local_rank == 0:
print(f'Loading LoRA checkpoint from {lora_ckpt_path}')
lora_checkpoint = torch.load(lora_ckpt_path, map_location='cpu')
if isinstance(lora_checkpoint, dict) and 'generator_lora' in lora_checkpoint:
peft.set_peft_model_state_dict(pipeline.generator.model, lora_checkpoint['generator_lora'])
else:
peft.set_peft_model_state_dict(pipeline.generator.model, lora_checkpoint)
if local_rank == 0:
print('LoRA weights loaded for generator')
if isinstance(lora_checkpoint, dict) and 'query_memory_encoder' in lora_checkpoint:
inner = pipeline.generator.model
if inner.query_memory_encoder is not None:
inner.query_memory_encoder.load_state_dict(lora_checkpoint['query_memory_encoder'], strict=False)
elif local_rank == 0:
print('No LoRA checkpoint specified; using base weights with LoRA adapters initialized')
pipeline.is_lora_enabled = True
pipeline = pipeline.to(dtype=torch.bfloat16)
if low_memory:
DynamicSwapInstaller.install_model(pipeline.text_encoder, device=device)
pipeline.generator.to(device=device)
pipeline.vae.to(device=device)
extended_prompt_path = config.data_path
dataset = TextDataset(prompt_path=config.data_path, extended_prompt_path=extended_prompt_path)
num_prompts = len(dataset)
print(f'Number of prompts: {num_prompts}')
if dist.is_initialized():
sampler = DistributedSampler(dataset, shuffle=False, drop_last=True)
else:
sampler = SequentialSampler(dataset)
dataloader = DataLoader(dataset, batch_size=1, sampler=sampler, num_workers=0, drop_last=False)
if local_rank == 0:
os.makedirs(config.output_folder, exist_ok=True)
if dist.is_initialized():
dist.barrier()
manifest = {}
for i, batch_data in tqdm(enumerate(dataloader), disable=local_rank != 0):
idx = batch_data['idx'].item()
if isinstance(batch_data, dict):
batch = batch_data
elif isinstance(batch_data, list):
batch = batch_data[0]
prompt = batch['prompts'][0]
extended_prompt = batch['extended_prompts'][0] if 'extended_prompts' in batch else None
index_str = f'{idx:05d}'
manifest[index_str] = prompt
first_output = os.path.join(config.output_folder, f'{index_str}-0.mp4')
if idx < num_prompts and os.path.exists(first_output):
print(f'Video already exists: {first_output}, skipping')
continue
if extended_prompt is not None:
prompts = [extended_prompt] * config.num_samples
else:
prompts = [prompt] * config.num_samples
sampled_noise = torch.randn([config.num_samples, config.num_output_frames, 16, 60, 104], device=device, dtype=torch.bfloat16)
video, latents = pipeline.inference(noise=sampled_noise, text_prompts=prompts, return_latents=True, low_memory=low_memory, profile=False)
current_video = rearrange(video, 'b t c h w -> b t h w c').cpu()
video = 255.0 * current_video
pipeline.vae.model.clear_cache()
if idx < num_prompts:
for sample_idx in range(config.num_samples):
output_path = os.path.join(config.output_folder, f'{index_str}-{sample_idx}.mp4')
write_video(output_path, video[sample_idx], fps=16)
if config.inference_iter != -1 and i >= config.inference_iter:
break
if dist.is_initialized():
rank_manifest_path = os.path.join(config.output_folder, f'.manifest_rank{rank}.json')
with open(rank_manifest_path, 'w', encoding='utf-8') as f:
json.dump(manifest, f, indent=2, ensure_ascii=False)
dist.barrier()
if local_rank == 0:
merged = {}
manifest_path = os.path.join(config.output_folder, 'manifest.json')
if os.path.exists(manifest_path):
with open(manifest_path) as f:
merged = json.load(f)
for rfile in sorted(glob.glob(os.path.join(config.output_folder, '.manifest_rank*.json'))):
with open(rfile) as f:
merged.update(json.load(f))
os.remove(rfile)
with open(manifest_path, 'w', encoding='utf-8') as f:
json.dump(merged, f, indent=2, ensure_ascii=False)
else:
manifest_path = os.path.join(config.output_folder, 'manifest.json')
if os.path.exists(manifest_path):
with open(manifest_path) as f:
existing = json.load(f)
existing.update(manifest)
manifest = existing
with open(manifest_path, 'w', encoding='utf-8') as f:
json.dump(manifest, f, indent=2, ensure_ascii=False)
if dist.is_initialized():
dist.destroy_process_group()