From 3416edc63df9e352f669856c8d5d1351418f596a Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=88=98=E9=9B=AA=E5=B3=B0?= Date: Fri, 21 Mar 2025 20:25:26 +0800 Subject: [PATCH] change models path to models/LatentSync. use local vae file. use ComfyUI/temp dir. --- __init__.py | 2 +- configs/vae/config.json | 29 ++ fixed_nodes.py | 497 +++++++++++++++++++++++ latentsync/pipelines/lipsync_pipeline.py | 11 +- latentsync/utils/util.py | 2 +- scripts/inference.py | 33 +- utils/diffusers_convert.py | 88 ++++ 7 files changed, 646 insertions(+), 16 deletions(-) create mode 100644 configs/vae/config.json create mode 100644 fixed_nodes.py create mode 100644 utils/diffusers_convert.py diff --git a/__init__.py b/__init__.py index d721463..c644700 100644 --- a/__init__.py +++ b/__init__.py @@ -1,3 +1,3 @@ -from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS +from .fixed_nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS __all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS'] \ No newline at end of file diff --git a/configs/vae/config.json b/configs/vae/config.json new file mode 100644 index 0000000..0db2671 --- /dev/null +++ b/configs/vae/config.json @@ -0,0 +1,29 @@ +{ + "_class_name": "AutoencoderKL", + "_diffusers_version": "0.4.2", + "act_fn": "silu", + "block_out_channels": [ + 128, + 256, + 512, + 512 + ], + "down_block_types": [ + "DownEncoderBlock2D", + "DownEncoderBlock2D", + "DownEncoderBlock2D", + "DownEncoderBlock2D" + ], + "in_channels": 3, + "latent_channels": 4, + "layers_per_block": 2, + "norm_num_groups": 32, + "out_channels": 3, + "sample_size": 256, + "up_block_types": [ + "UpDecoderBlock2D", + "UpDecoderBlock2D", + "UpDecoderBlock2D", + "UpDecoderBlock2D" + ] +} diff --git a/fixed_nodes.py b/fixed_nodes.py new file mode 100644 index 0000000..224fe38 --- /dev/null +++ b/fixed_nodes.py @@ -0,0 +1,497 @@ +import argparse +import importlib.machinery +import importlib.util +import math +import os +import platform +import random +import shutil +import subprocess +import sys + +import torch +import torchaudio +from omegaconf import OmegaConf +from .scripts import inference as inference_module +import folder_paths +import torchvision.io as io + + + +def check_ffmpeg(): + try: + if platform.system() == "Windows": + # Check if ffmpeg exists in PATH + ffmpeg_path = shutil.which("ffmpeg.exe") + if ffmpeg_path is None: + # Look for ffmpeg in common locations + possible_paths = [ + os.path.join(os.environ.get("ProgramFiles", "C:\\Program Files"), "ffmpeg", "bin"), + os.path.join(os.environ.get("ProgramFiles(x86)", "C:\\Program Files (x86)"), "ffmpeg", "bin"), + os.path.join(os.path.dirname(os.path.abspath(__file__)), "ffmpeg", "bin"), + ] + for path in possible_paths: + if os.path.exists(os.path.join(path, "ffmpeg.exe")): + # Add to PATH + os.environ["PATH"] = path + os.pathsep + os.environ.get("PATH", "") + return True + print("FFmpeg not found. Please install FFmpeg and add it to PATH") + return False + return True + else: + subprocess.run(["ffmpeg", "-version"], capture_output=True, check=True) + return True + except (subprocess.CalledProcessError, FileNotFoundError): + print("FFmpeg not found. Please install FFmpeg") + return False + +def check_and_install_dependencies(): + if not check_ffmpeg(): + raise RuntimeError("FFmpeg is required but not found") + + required_packages = [ + 'omegaconf', + 'pytorch_lightning', + 'transformers', + 'accelerate', + 'huggingface_hub', + 'einops', + 'diffusers' + ] + + def is_package_installed(package_name): + return importlib.util.find_spec(package_name) is not None + + def install_package(package): + python_exe = sys.executable + try: + subprocess.check_call([python_exe, '-m', 'pip', 'install', package], + stdout=subprocess.PIPE, + stderr=subprocess.PIPE) + print(f"Successfully installed {package}") + except subprocess.CalledProcessError as e: + print(f"Error installing {package}: {str(e)}") + raise RuntimeError(f"Failed to install required package: {package}") + + for package in required_packages: + if not is_package_installed(package): + print(f"Installing required package: {package}") + try: + install_package(package) + except Exception as e: + print(f"Warning: Failed to install {package}: {str(e)}") + raise + +def normalize_path(path): + """Normalize path to handle spaces and special characters""" + return os.path.normpath(path).replace('\\', '/') + +def get_ext_dir(subpath=None, mkdir=True): + """Get extension directory path, optionally with a subpath""" + # Get the directory containing this script + dir = folder_paths.models_dir + + if subpath is not None: + dir = os.path.join(dir, subpath) + + if mkdir and not os.path.exists(dir): + os.makedirs(dir, exist_ok=True) + + return dir + +def setup_models(): + """Setup and pre-download all required models.""" + # Use our global temp directory + global modles_dir + + # Existing setup logic for LatentSync models + modles_dir = get_ext_dir(subpath="LatentSync") + ckpt_dir = os.path.join(modles_dir, "checkpoints") + whisper_dir = os.path.join(ckpt_dir, "whisper") + os.makedirs(ckpt_dir, exist_ok=True) + os.makedirs(whisper_dir, exist_ok=True) + + # Create a temp_downloads directory in our system temp + temp_downloads = os.path.join(modles_dir, "downloads") + os.makedirs(temp_downloads, exist_ok=True) + + unet_path = os.path.join(ckpt_dir, "latentsync_unet.pt") + whisper_path = os.path.join(whisper_dir, "tiny.pt") + + if not (os.path.exists(unet_path) and os.path.exists(whisper_path)): + print("Downloading required model checkpoints... This may take a while.") + try: + from huggingface_hub import snapshot_download + snapshot_download(repo_id="ByteDance/LatentSync-1.5", + allow_patterns=["latentsync_unet.pt", "whisper/tiny.pt"], + local_dir=ckpt_dir, + local_dir_use_symlinks=False, + cache_dir=temp_downloads) + print("Model checkpoints downloaded successfully!") + except Exception as e: + print(f"Error downloading models: {str(e)}") + print("\nPlease download models manually:") + print("1. Visit: https://huggingface.co/ByteDance/LatentSync-1.5") + print("2. Download: latentsync_unet.pt and whisper/tiny.pt") + print(f"3. Place them in: {ckpt_dir}") + print(f" with whisper/tiny.pt in: {whisper_dir}") + raise RuntimeError("Model download failed. See instructions above.") + +class LatentSyncNode: + def __init__(self): + check_and_install_dependencies() + setup_models() + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "images": ("IMAGE",), + "audio": ("AUDIO", ), + "seed": ("INT", {"default": 1247}), + "lips_expression": ("FLOAT", {"default": 1.5, "min": 1.0, "max": 3.0, "step": 0.1}), + "inference_steps": ("INT", {"default": 20, "min": 1, "max": 999, "step": 1}), + "vae_name": (folder_paths.get_filename_list("vae"), {"tooltip": "sd1.5 vae model"}), + }, + } + + CATEGORY = "LatentSyncNode" + + RETURN_TYPES = ("IMAGE", "AUDIO") + RETURN_NAMES = ("images", "audio") + FUNCTION = "inference" + + def process_batch(self, batch, use_mixed_precision=False): + with torch.cuda.amp.autocast(enabled=use_mixed_precision): + processed_batch = batch.float() / 255.0 + if len(processed_batch.shape) == 3: + processed_batch = processed_batch.unsqueeze(0) + if processed_batch.shape[0] == 3: + processed_batch = processed_batch.permute(1, 2, 0) + if processed_batch.shape[-1] == 4: + processed_batch = processed_batch[..., :3] + return processed_batch + + def inference(self, images, audio, seed, vae_name, lips_expression=1.5, inference_steps=20): + # Get GPU capabilities and memory + device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') + BATCH_SIZE = 4 + use_mixed_precision = False + if torch.cuda.is_available(): + gpu_mem = torch.cuda.get_device_properties(0).total_memory + # Convert to GB + gpu_mem_gb = gpu_mem / (1024 ** 3) + + # Dynamically adjust batch size based on GPU memory + if gpu_mem_gb > 20: # High-end GPUs + BATCH_SIZE = 32 + enable_tf32 = True + use_mixed_precision = True + elif gpu_mem_gb > 8: # Mid-range GPUs + BATCH_SIZE = 16 + enable_tf32 = False + use_mixed_precision = True + else: # Lower-end GPUs + BATCH_SIZE = 8 + enable_tf32 = False + use_mixed_precision = False + + # Set performance options based on GPU capability + torch.backends.cudnn.benchmark = True + if enable_tf32: + torch.backends.cuda.matmul.allow_tf32 = True + torch.backends.cudnn.allow_tf32 = True + + # Clear GPU cache before processing + torch.cuda.empty_cache() + torch.cuda.set_per_process_memory_fraction(0.8) + + # Create a run-specific subdirectory in our temp directory + run_id = ''.join(random.choice("abcdefghijklmnopqrstuvwxyz") for _ in range(5)) + temp_dir = os.path.join(folder_paths.get_temp_directory(), f"run_{run_id}") + os.makedirs(temp_dir, exist_ok=True) + + temp_video_path = None + output_video_path = None + audio_path = None + + try: + # Create temporary file paths in our system temp directory + temp_video_path = os.path.join(temp_dir, f"temp_{run_id}.mp4") + output_video_path = os.path.join(temp_dir, f"latentsync_{run_id}_out.mp4") + audio_path = os.path.join(temp_dir, f"latentsync_{run_id}_audio.wav") + + # Get the extension directory + cur_dir = os.path.dirname(os.path.abspath(__file__)) + + # Process input frames + if isinstance(images, list): + frames = torch.stack(images).to(device) + else: + frames = images.to(device) + frames = (frames * 255).byte() + + if len(frames.shape) == 3: + frames = frames.unsqueeze(0) + + # Process audio with device awareness + waveform = audio["waveform"].to(device) + sample_rate = audio["sample_rate"] + if waveform.dim() == 3: + waveform = waveform.squeeze(0) + + if sample_rate != 16000: + new_sample_rate = 16000 + resampler = torchaudio.transforms.Resample( + orig_freq=sample_rate, + new_freq=new_sample_rate + ).to(device) + waveform_16k = resampler(waveform) + waveform, sample_rate = waveform_16k, new_sample_rate + + # Package resampled audio + resampled_audio = { + "waveform": waveform.unsqueeze(0), + "sample_rate": sample_rate + } + + # Move waveform to CPU for saving + waveform_cpu = waveform.cpu() + torchaudio.save(audio_path, waveform_cpu, sample_rate) + + # Move frames to CPU for saving to video + frames_cpu = frames.cpu() + try: + io.write_video(temp_video_path, frames_cpu, fps=25, video_codec='h264') + except TypeError: + import av + container = av.open(temp_video_path, mode='w') + stream = container.add_stream('h264', rate=25) + stream.width = frames_cpu.shape[2] + stream.height = frames_cpu.shape[1] + + for frame in frames_cpu: + frame = av.VideoFrame.from_ndarray(frame.numpy(), format='rgb24') + packet = stream.encode(frame) + container.mux(packet) + + packet = stream.encode(None) + container.mux(packet) + container.close() + + # Define paths to required files and configs + # inference_script_path = os.path.join(cur_dir, "scripts", "inference.py") + config_path = os.path.join(cur_dir, "configs", "unet", "stage2.yaml") + scheduler_config_path = os.path.join(cur_dir, "configs") + ckpt_path = os.path.join(modles_dir, "checkpoints", "latentsync_unet.pt") + whisper_ckpt_path = os.path.join(modles_dir, "checkpoints", "whisper", "tiny.pt") + + # Create config and args + config = OmegaConf.load(config_path) + + # Set the correct mask image path + mask_image_path = os.path.join(cur_dir, "latentsync", "utils", "mask.png") + # Make sure the mask image exists + if not os.path.exists(mask_image_path): + # Try to find it in the utils directory directly + alt_mask_path = os.path.join(cur_dir, "utils", "mask.png") + if os.path.exists(alt_mask_path): + mask_image_path = alt_mask_path + else: + print(f"Warning: Could not find mask image at expected locations") + + # Set mask path in config + if hasattr(config, "data") and hasattr(config.data, "mask_image_path"): + config.data.mask_image_path = mask_image_path + + vae_path = folder_paths.get_full_path_or_raise("vae", vae_name) + + args = argparse.Namespace( + unet_config_path=config_path, + inference_ckpt_path=ckpt_path, + video_path=temp_video_path, + audio_path=audio_path, + video_out_path=output_video_path, + seed=seed, + inference_steps=inference_steps, + guidance_scale=lips_expression, # Using lips_expression for the guidance_scale + scheduler_config_path=scheduler_config_path, + whisper_ckpt_path=whisper_ckpt_path, + device=device, + batch_size=BATCH_SIZE, + use_mixed_precision=use_mixed_precision, + temp_dir=temp_dir, + mask_image_path=mask_image_path, + vae_model_path=vae_path + ) + + # Set PYTHONPATH to include our directories + # package_root = os.path.dirname(cur_dir) + # if package_root not in sys.path: + # sys.path.insert(0, package_root) + # if cur_dir not in sys.path: + # sys.path.insert(0, cur_dir) + + # Clean GPU cache before inference + if torch.cuda.is_available(): + torch.cuda.empty_cache() + + # Import the inference module + # inference_module = import_inference_script(inference_script_path) + + # Monkey patch any temp directory functions in the inference module + if hasattr(inference_module, 'get_temp_dir'): + inference_module.get_temp_dir = lambda *args, **kwargs: temp_dir + + # Create subdirectories that the inference module might expect + inference_temp = os.path.join(temp_dir, "temp") + os.makedirs(inference_temp, exist_ok=True) + + # Run inference + inference_module.main(config, args) + + # Clean GPU cache after inference + if torch.cuda.is_available(): + torch.cuda.empty_cache() + + # Verify output file exists + if not os.path.exists(output_video_path): + raise FileNotFoundError(f"Output video not found at: {output_video_path}") + + # Read the processed video - ensure it's loaded as CPU tensor + processed_frames = io.read_video(output_video_path, pts_unit='sec')[0] + processed_frames = processed_frames.float() / 255.0 + + # Ensure audio is on CPU before returning + if torch.cuda.is_available(): + if hasattr(resampled_audio["waveform"], 'device') and resampled_audio["waveform"].device.type == 'cuda': + resampled_audio["waveform"] = resampled_audio["waveform"].cpu() + if hasattr(processed_frames, 'device') and processed_frames.device.type == 'cuda': + processed_frames = processed_frames.cpu() + + return processed_frames, resampled_audio + + except Exception as e: + print(f"Error during inference: {str(e)}") + import traceback + traceback.print_exc() + raise + + finally: + # Clean up temporary files individually + for path in [temp_video_path, output_video_path, audio_path]: + if path and os.path.exists(path): + try: + os.remove(path) + print(f"Removed temporary file: {path}") + except Exception as e: + print(f"Failed to remove {path}: {str(e)}") + + # Remove temporary run directory + if temp_dir and os.path.exists(temp_dir): + try: + shutil.rmtree(temp_dir, ignore_errors=True) + print(f"Removed run temporary directory: {temp_dir}") + except Exception as e: + print(f"Failed to remove temp run directory: {str(e)}") + + # Final GPU cache cleanup + if torch.cuda.is_available(): + torch.cuda.empty_cache() + +class VideoLengthAdjuster: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "images": ("IMAGE",), + "audio": ("AUDIO",), + "mode": (["normal", "pingpong", "loop_to_audio"], {"default": "normal"}), + "fps": ("FLOAT", {"default": 25.0, "min": 1.0, "max": 120.0}), + "silent_padding_sec": ("FLOAT", {"default": 0.5, "min": 0.1, "max": 3.0, "step": 0.1}), + } + } + + CATEGORY = "LatentSyncNode" + RETURN_TYPES = ("IMAGE", "AUDIO") + RETURN_NAMES = ("images", "audio") + FUNCTION = "adjust" + + def adjust(self, images, audio, mode, fps=25.0, silent_padding_sec=0.5): + waveform = audio["waveform"].squeeze(0) + sample_rate = int(audio["sample_rate"]) + original_frames = [images[i] for i in range(images.shape[0])] if isinstance(images, torch.Tensor) else images.copy() + + if mode == "normal": + # Bypass video frames exactly + video_duration = len(original_frames) / fps + required_samples = int(video_duration * sample_rate) + + # Adjust audio to match video duration + if waveform.shape[1] >= required_samples: + adjusted_audio = waveform[:, :required_samples] # Trim audio + else: + silence = torch.zeros((waveform.shape[0], required_samples - waveform.shape[1]), dtype=waveform.dtype) + adjusted_audio = torch.cat([waveform, silence], dim=1) # Pad audio + + return ( + torch.stack(original_frames), + {"waveform": adjusted_audio.unsqueeze(0), "sample_rate": sample_rate} + ) + + elif mode == "pingpong": + video_duration = len(original_frames) / fps + audio_duration = waveform.shape[1] / sample_rate + if audio_duration <= video_duration: + required_samples = int(video_duration * sample_rate) + silence = torch.zeros((waveform.shape[0], required_samples - waveform.shape[1]), dtype=waveform.dtype) + adjusted_audio = torch.cat([waveform, silence], dim=1) + + return ( + torch.stack(original_frames), + {"waveform": adjusted_audio.unsqueeze(0), "sample_rate": sample_rate} + ) + + else: + silence_samples = math.ceil(silent_padding_sec * sample_rate) + silence = torch.zeros((waveform.shape[0], silence_samples), dtype=waveform.dtype) + padded_audio = torch.cat([waveform, silence], dim=1) + total_duration = (waveform.shape[1] + silence_samples) / sample_rate + target_frames = math.ceil(total_duration * fps) + reversed_frames = original_frames[::-1][1:-1] # Remove endpoints + frames = original_frames + reversed_frames + while len(frames) < target_frames: + frames += frames[:target_frames - len(frames)] + return ( + torch.stack(frames[:target_frames]), + {"waveform": padded_audio.unsqueeze(0), "sample_rate": sample_rate} + ) + + elif mode == "loop_to_audio": + # Add silent padding then simple loop + silence_samples = math.ceil(silent_padding_sec * sample_rate) + silence = torch.zeros((waveform.shape[0], silence_samples), dtype=waveform.dtype) + padded_audio = torch.cat([waveform, silence], dim=1) + total_duration = (waveform.shape[1] + silence_samples) / sample_rate + target_frames = math.ceil(total_duration * fps) + + frames = original_frames.copy() + while len(frames) < target_frames: + frames += original_frames[:target_frames - len(frames)] + + return ( + torch.stack(frames[:target_frames]), + {"waveform": padded_audio.unsqueeze(0), "sample_rate": sample_rate} + ) + +# Node Mappings for ComfyUI +NODE_CLASS_MAPPINGS = { + "LatentSyncNode": LatentSyncNode, + "VideoLengthAdjuster": VideoLengthAdjuster, +} + +# Display Names for ComfyUI +NODE_DISPLAY_NAME_MAPPINGS = { + "LatentSyncNode": "LatentSync1.5 Node", + "VideoLengthAdjuster": "Video Length Adjuster", +} diff --git a/latentsync/pipelines/lipsync_pipeline.py b/latentsync/pipelines/lipsync_pipeline.py index 446f2ae..28ba3b8 100644 --- a/latentsync/pipelines/lipsync_pipeline.py +++ b/latentsync/pipelines/lipsync_pipeline.py @@ -116,7 +116,7 @@ class LipsyncPipeline(DiffusionPipeline): self.vae_scale_factor = 2 ** (len(self.vae.config.block_out_channels) - 1) - self.set_progress_bar_config(desc="Steps") + # self.set_progress_bar_config(desc="Steps") def enable_vae_slicing(self): self.vae.enable_slicing() @@ -313,7 +313,7 @@ class LipsyncPipeline(DiffusionPipeline): device = self._execution_device mask_image = load_fixed_mask(height, mask_image_path) self.image_processor = ImageProcessor(height, mask=mask, device="cuda", mask_image=mask_image) - self.set_progress_bar_config(desc=f"Sample frames: {num_frames}") + # self.set_progress_bar_config(desc=f"Sample frames: {num_frames}") # 1. Default height and width to unet height = height or self.denoising_unet.config.sample_size * self.vae_scale_factor @@ -361,7 +361,7 @@ class LipsyncPipeline(DiffusionPipeline): generator, ) - for i in tqdm.tqdm(range(num_inferences), desc="Doing inference..."): + for i in tqdm.tqdm(range(num_inferences), position=0, desc="Doing inference..."): if self.denoising_unet.add_audio_layer: audio_embeds = torch.stack(whisper_chunks[i * num_frames : (i + 1) * num_frames]) audio_embeds = audio_embeds.to(device, dtype=weight_dtype) @@ -399,7 +399,8 @@ class LipsyncPipeline(DiffusionPipeline): # 9. Denoising loop num_warmup_steps = len(timesteps) - num_inference_steps * self.scheduler.order - with self.progress_bar(total=num_inference_steps) as progress_bar: + # with self.progress_bar(total=num_inference_steps) as progress_bar: + with tqdm.tqdm(total=num_inference_steps, position=1, desc=f"Sample {num_frames} frames", leave=False) as progress_bar: for j, t in enumerate(timesteps): # expand the latents if we are doing classifier free guidance denoising_unet_input = torch.cat([latents] * 2) if do_classifier_free_guidance else latents @@ -451,7 +452,7 @@ class LipsyncPipeline(DiffusionPipeline): if is_train: self.denoising_unet.train() - temp_dir = "temp" + temp_dir = os.path.join(os.path.dirname(video_path), "temp") if os.path.exists(temp_dir): shutil.rmtree(temp_dir) os.makedirs(temp_dir, exist_ok=True) diff --git a/latentsync/utils/util.py b/latentsync/utils/util.py index 68b620c..d161f81 100644 --- a/latentsync/utils/util.py +++ b/latentsync/utils/util.py @@ -44,7 +44,7 @@ def read_json(filepath: str): def read_video(video_path: str, change_fps=True, use_decord=True): if change_fps: - temp_dir = "temp" + temp_dir = os.path.join(os.path.dirname(video_path), "temp") if os.path.exists(temp_dir): shutil.rmtree(temp_dir) os.makedirs(temp_dir, exist_ok=True) diff --git a/scripts/inference.py b/scripts/inference.py index 8ca2b62..16fee62 100644 --- a/scripts/inference.py +++ b/scripts/inference.py @@ -17,11 +17,13 @@ import os from omegaconf import OmegaConf import torch from diffusers import AutoencoderKL, DDIMScheduler -from latentsync.models.unet import UNet3DConditionModel -from latentsync.pipelines.lipsync_pipeline import LipsyncPipeline +from ..latentsync.models.unet import UNet3DConditionModel +from ..latentsync.pipelines.lipsync_pipeline import LipsyncPipeline from accelerate.utils import set_seed -from latentsync.whisper.audio2feature import Audio2Feature - +from ..latentsync.whisper.audio2feature import Audio2Feature +import json +import safetensors.torch +from ..utils.diffusers_convert import convert_vae_state_dict_to_hf def main(config, args): if not os.path.exists(args.video_path): @@ -40,7 +42,7 @@ def main(config, args): # Use relative path for scheduler configuration current_dir = os.path.dirname(os.path.abspath(__file__)) scheduler_path = os.path.join(current_dir, "..", "configs", "scheduler") - + # Check if scheduler directory exists if not os.path.exists(scheduler_path): print(f"Creating scheduler directory at {scheduler_path}") @@ -65,7 +67,6 @@ def main(config, args): "skip_prk_steps": True } - import json with open(scheduler_config_file, 'w') as f: json.dump(scheduler_config, f, indent=2) @@ -89,11 +90,12 @@ def main(config, args): skip_prk_steps=True ) + model_dir = os.path.dirname(args.inference_ckpt_path) # Use relative paths for whisper models as well if config.model.cross_attention_dim == 768: - whisper_model_path = os.path.join(current_dir, "..", "checkpoints", "whisper", "small.pt") + whisper_model_path = os.path.join(model_dir, "..", "checkpoints", "whisper", "small.pt") elif config.model.cross_attention_dim == 384: - whisper_model_path = os.path.join(current_dir, "..", "checkpoints", "whisper", "tiny.pt") + whisper_model_path = os.path.join(model_dir, "..", "checkpoints", "whisper", "tiny.pt") else: raise NotImplementedError("cross_attention_dim must be 768 or 384") @@ -104,7 +106,19 @@ def main(config, args): audio_feat_length=config.data.audio_feat_length, ) - vae = AutoencoderKL.from_pretrained("stabilityai/sd-vae-ft-mse", torch_dtype=dtype) + if args.vae_model_path and len(args.vae_model_path) > 0: + # 根据模型配置初始化 + vae_config_path = os.path.join(current_dir, "..", "configs", "vae", "config.json") + with open(vae_config_path, "r", encoding="utf-8") as f: + vae_config_dict = json.load(f) + vae = AutoencoderKL.from_config(vae_config_dict) + # 加载单文件权重 + state_dict = safetensors.torch.load_file(args.vae_model_path) + vae.load_state_dict(convert_vae_state_dict_to_hf(state_dict), strict=False) + vae.to(dtype=dtype) + else: + vae = AutoencoderKL.from_pretrained("stabilityai/sd-vae-ft-mse", torch_dtype=dtype) + vae.config.scaling_factor = 0.18215 vae.config.shift_factor = 0 @@ -155,6 +169,7 @@ if __name__ == "__main__": parser.add_argument("--inference_steps", type=int, default=20) parser.add_argument("--guidance_scale", type=float, default=1.0) parser.add_argument("--seed", type=int, default=1247) + parser.add_argument("--vae_model_path", type=str, default="") args = parser.parse_args() config = OmegaConf.load(args.unet_config_path) diff --git a/utils/diffusers_convert.py b/utils/diffusers_convert.py new file mode 100644 index 0000000..7469898 --- /dev/null +++ b/utils/diffusers_convert.py @@ -0,0 +1,88 @@ +import re +import torch +import logging + +# conversion code from https://github.com/huggingface/diffusers/blob/main/scripts/convert_diffusers_to_original_stable_diffusion.py + +# ================# +# VAE Conversion # +# ================# + +vae_conversion_map = [ + # (stable-diffusion, HF Diffusers) + ("nin_shortcut", "conv_shortcut"), + ("norm_out", "conv_norm_out"), + ("mid.attn_1.", "mid_block.attentions.0."), +] + +for i in range(4): + # down_blocks have two resnets + for j in range(2): + hf_down_prefix = f"encoder.down_blocks.{i}.resnets.{j}." + sd_down_prefix = f"encoder.down.{i}.block.{j}." + vae_conversion_map.append((sd_down_prefix, hf_down_prefix)) + + if i < 3: + hf_downsample_prefix = f"down_blocks.{i}.downsamplers.0." + sd_downsample_prefix = f"down.{i}.downsample." + vae_conversion_map.append((sd_downsample_prefix, hf_downsample_prefix)) + + hf_upsample_prefix = f"up_blocks.{i}.upsamplers.0." + sd_upsample_prefix = f"up.{3 - i}.upsample." + vae_conversion_map.append((sd_upsample_prefix, hf_upsample_prefix)) + + # up_blocks have three resnets + # also, up blocks in hf are numbered in reverse from sd + for j in range(3): + hf_up_prefix = f"decoder.up_blocks.{i}.resnets.{j}." + sd_up_prefix = f"decoder.up.{3 - i}.block.{j}." + vae_conversion_map.append((sd_up_prefix, hf_up_prefix)) + +# this part accounts for mid blocks in both the encoder and the decoder +for i in range(2): + hf_mid_res_prefix = f"mid_block.resnets.{i}." + sd_mid_res_prefix = f"mid.block_{i + 1}." + vae_conversion_map.append((sd_mid_res_prefix, hf_mid_res_prefix)) + +vae_conversion_map_attn = [ + # (stable-diffusion, HF Diffusers) + ("norm.", "group_norm."), + ("q.", "query."), + ("k.", "key."), + ("v.", "value."), + ("q.", "to_q."), + ("k.", "to_k."), + ("v.", "to_v."), + ("proj_out.", "to_out.0."), + ("proj_out.", "proj_attn."), +] + + +def reshape_weight_for_sd(w): + # convert SD conv2d weights to HF linear weights + w.view(w.shape[0], w.shape[1]) + + +def convert_vae_state_dict_to_hf(vae_state_dict): + mapping = {k: k for k in vae_state_dict.keys()} + conv3d = False + for k, v in mapping.items(): + for sd_part, hf_part in vae_conversion_map: + v = v.replace(sd_part, hf_part) + # if v.endswith(".conv.weight"): + # if not conv3d and vae_state_dict[k].ndim == 5: + # conv3d = True + mapping[k] = v + for k, v in mapping.items(): + if "attentions" in k: + for sd_part, hf_part in vae_conversion_map_attn: + v = v.replace(sd_part, hf_part) + mapping[k] = v + new_state_dict = {v: vae_state_dict[k] for k, v in mapping.items()} + weights_to_convert = ["to_q", "to_k", "to_v", "to_out.0", "proj_attn", "value", "key", "query"] + for k, v in new_state_dict.items(): + for weight_name in weights_to_convert: + if f"mid_block.attentions.0.{weight_name}.weight" in k: + logging.debug(f"Reshaping {k} for HF format") + new_state_dict[k] = reshape_weight_for_sd(v) + return new_state_dict