change models path to models/LatentSync. use local vae file. use ComfyUI/temp dir.

This commit is contained in:
刘雪峰
2025-03-21 20:25:26 +08:00
parent 26448c3071
commit 3416edc63d
7 changed files with 646 additions and 16 deletions
+1 -1
View File
@@ -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']
+29
View File
@@ -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"
]
}
+497
View File
@@ -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",
}
+6 -5
View File
@@ -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)
+1 -1
View File
@@ -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)
+24 -9
View File
@@ -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)
+88
View File
@@ -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