Files
filliptm-ComfyUI_FL-Trainer/WIP/FL_SliderLoraInitWorkspace.py.WIP
T
2024-07-13 05:40:34 -07:00

232 lines
10 KiB
Plaintext

import os
import torch
import json
from safetensors import safe_open
from diffusers import AutoencoderKL, UNet2DConditionModel, DDPMScheduler
from diffusers.models.embeddings import TimestepEmbedding, Timesteps
from transformers import CLIPTextModel, CLIPTokenizer, CLIPConfig
from .FL_SliderLoraCore import SliderLoraSD15Pipeline, SliderLoraSDXLPipeline
import folder_paths
import diffusers
from .FL_train_utils import Utils
class CustomAutoencoder(AutoencoderKL):
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
self.encoder.down_blocks = torch.nn.ModuleList([
torch.nn.ModuleDict({"block": torch.nn.ModuleList([]), "downsample": None})
for _ in range(4)
])
self.decoder.up_blocks = torch.nn.ModuleList([
torch.nn.ModuleDict({"block": torch.nn.ModuleList([]), "upsample": None})
for _ in range(4)
])
self.encoder.mid = torch.nn.Module()
self.decoder.mid = torch.nn.Module()
class FL_SliderLoraInitWorkspace:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model_name": (folder_paths.get_filename_list("checkpoints"),),
"model_type": (["SD1.5", "SDXL"],),
"lora_name": ("STRING", {"default": ""}),
"resolution": ("INT", {"default": 512, "min": 256, "max": 2048, "step": 64}),
"rank": ("INT", {"default": 4, "min": 1, "max": 128, "step": 1}),
"device": (["cuda", "cpu"],),
"seed": ("INT", {"default": 0}),
},
}
RETURN_TYPES = ("FL_SLIDER_LORA_WORKSPACE",)
RETURN_NAMES = ("workspace",)
FUNCTION = "init_workspace"
CATEGORY = "FL_Slider_Lora"
def init_workspace(self, model_name, model_type, lora_name, resolution, rank, device, seed):
model_path = folder_paths.get_full_path("checkpoints", model_name)
vae, unet, scheduler, text_encoder, tokenizer = self.load_model_components(model_path, model_type, device)
# Set up LoRA config
from peft import LoraConfig
lora_config = LoraConfig(
r=rank,
lora_alpha=rank,
target_modules=["q_proj", "v_proj"],
lora_dropout=0.1,
bias="none",
)
# Create workspace directory
workspace_dir = os.path.join(folder_paths.get_output_directory(), "slider_lora_workspaces", lora_name)
os.makedirs(workspace_dir, exist_ok=True)
# Set random seed
torch.manual_seed(seed)
if device == "cuda":
torch.cuda.manual_seed_all(seed)
# Create pipeline
pipeline = self.create_pipeline(model_type, unet, text_encoder, tokenizer, scheduler, vae)
workspace = {
"pipeline": pipeline,
"lora_config": lora_config,
"workspace_dir": workspace_dir,
"model_type": model_type,
"resolution": resolution,
"device": device,
}
return (workspace,)
def load_model_components(self, model_path, model_type, device):
print(f"Diffusers version: {diffusers.__version__}")
config_path = os.path.join(os.path.dirname(__file__), "configs", "models_config")
if model_type == "SD1.5":
config_path = os.path.join(config_path, "stable-diffusion-v1-5")
elif model_type == "SDXL":
config_path = os.path.join(config_path, "stable-diffusion-xl-base-1.0")
else:
raise ValueError(f"Unsupported model type: {model_type}")
if model_path.endswith('.safetensors'):
print(f"Loading .safetensors model: {model_path}")
state_dict = self.load_safetensors(model_path)
# Load configs
with open(os.path.join(config_path, "unet", "config.json"), "r") as f:
unet_config = json.load(f)
print("UNet config:")
for key, value in unet_config.items():
print(f" {key}: {value}")
print("Creating custom UNet...")
print("UNet config keys:", unet_config.keys())
print("UNet block_out_channels:", unet_config.get('block_out_channels'))
print("UNet time_embed_dim:", unet_config.get('time_embed_dim'))
with open(os.path.join(config_path, "text_encoder", "config.json"), "r") as f:
text_encoder_config = json.load(f)
with open(os.path.join(config_path, "vae", "config.json"), "r") as f:
vae_config = json.load(f)
# Initialize models with configs
unet = self.create_custom_unet(unet_config)
print("Custom UNet created successfully.")
text_encoder = CLIPTextModel(config=CLIPConfig(**text_encoder_config))
vae = CustomAutoencoder(**vae_config)
# Load state dicts
unet_state_dict = self.filter_state_dict(state_dict, "model.diffusion_model.")
unet.load_state_dict(unet_state_dict, strict=False)
text_encoder_state_dict = self.filter_state_dict(state_dict, "cond_stage_model.transformer.")
text_encoder_state_dict = {k: v for k, v in text_encoder_state_dict.items() if
k in text_encoder.state_dict()}
text_encoder.load_state_dict(text_encoder_state_dict, strict=False)
vae_state_dict = self.filter_state_dict(state_dict, "first_stage_model.")
vae_state_dict = self.adapt_vae_state_dict(vae_state_dict)
vae.load_state_dict(vae_state_dict, strict=False)
# Move models to device
unet = unet.to(device)
text_encoder = text_encoder.to(device)
vae = vae.to(device)
# Initialize scheduler and tokenizer
scheduler = DDPMScheduler.from_pretrained(os.path.join(config_path, "scheduler"))
tokenizer = CLIPTokenizer.from_pretrained(os.path.join(config_path, "tokenizer"))
else:
print(f"Loading model from directory: {model_path}")
vae = AutoencoderKL.from_pretrained(model_path, subfolder="vae").to(device)
unet = UNet2DConditionModel.from_pretrained(model_path, subfolder="unet").to(device)
scheduler = DDPMScheduler.from_pretrained(model_path, subfolder="scheduler")
text_encoder = CLIPTextModel.from_pretrained(model_path, subfolder="text_encoder").to(device)
tokenizer = CLIPTokenizer.from_pretrained(model_path, subfolder="tokenizer")
return vae, unet, scheduler, text_encoder, tokenizer
def load_safetensors(self, path):
with safe_open(path, framework="pt", device="cpu") as f:
return {key: f.get_tensor(key) for key in f.keys()}
def filter_state_dict(self, state_dict, prefix):
return {k.replace(prefix, ""): v for k, v in state_dict.items() if k.startswith(prefix)}
def create_custom_unet(self, config):
original_unet = UNet2DConditionModel(**config)
class CustomUNetWrapper(torch.nn.Module):
def __init__(self, unet):
super().__init__()
self.unet = unet
self.config = unet.config
# Copy any other necessary attributes from the original UNet
for attr_name in dir(unet):
if not attr_name.startswith('_') and not hasattr(self, attr_name):
setattr(self, attr_name, getattr(unet, attr_name))
def forward(self, sample, timestep, encoder_hidden_states, class_labels=None, return_dict=True):
# Ensure timestep is a 1D tensor
if not torch.is_tensor(timestep):
timestep = torch.tensor([timestep], dtype=torch.long, device=sample.device)
timestep = timestep.to(sample.device).view(-1)
# If timestep is a single value, repeat it to match the batch size
if timestep.shape[0] == 1:
timestep = timestep.repeat(sample.shape[0])
print(f"Debug - CustomUNetWrapper - sample shape: {sample.shape}")
print(f"Debug - CustomUNetWrapper - timestep shape: {timestep.shape}")
print(f"Debug - CustomUNetWrapper - timestep value: {timestep}")
# Call the original UNet's forward method
return self.unet(sample, timestep, encoder_hidden_states, class_labels, return_dict)
return CustomUNetWrapper(original_unet)
def adapt_vae_state_dict(self, state_dict):
new_state_dict = {}
for k, v in state_dict.items():
if k.startswith('encoder.') or k.startswith('decoder.'):
parts = k.split('.')
if parts[1] in ['down', 'up']:
block_idx = int(parts[2])
new_key = f"{parts[0]}.{'down' if parts[1] == 'down' else 'up'}_blocks.{block_idx}"
if parts[3] == 'block':
new_key += f".resnets.{parts[4]}"
elif parts[3] in ['upsample', 'downsample']:
new_key += f".{parts[3]}rs.0"
new_key += '.'.join(parts[5:])
elif parts[1] == 'mid':
new_key = f"{parts[0]}.mid_block"
if parts[2] == 'block':
new_key += f".resnets.{int(parts[3]) - 1}"
elif parts[2] == 'attn':
new_key += ".attentions.0"
new_key += '.'.join(parts[4:])
else:
new_key = k
else:
new_key = k
new_state_dict[new_key] = v
return new_state_dict
def create_pipeline(self, model_type, unet, text_encoder, tokenizer, scheduler, vae):
if model_type == "SD1.5":
return SliderLoraSD15Pipeline(unet, text_encoder, tokenizer, scheduler, vae)
elif model_type == "SDXL":
return SliderLoraSDXLPipeline(unet, text_encoder, tokenizer, scheduler, vae)
else:
raise ValueError(f"Unsupported model type: {model_type}")