Added multiplier jitter, min_snr, ability to choose sdxl encoders to use, shuffle generator, and other fun
This commit is contained in:
@@ -3,7 +3,7 @@ import hashlib
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
from typing import TYPE_CHECKING
|
||||
from typing import TYPE_CHECKING, Union
|
||||
import sys
|
||||
from toolkit.paths import SD_SCRIPTS_ROOT
|
||||
|
||||
@@ -32,6 +32,10 @@ SCHEDULER_LINEAR_END = 0.0120
|
||||
SCHEDULER_TIMESTEPS = 1000
|
||||
SCHEDLER_SCHEDULE = "scaled_linear"
|
||||
|
||||
UNET_ATTENTION_TIME_EMBED_DIM = 256 # XL
|
||||
TEXT_ENCODER_2_PROJECTION_DIM = 1280
|
||||
UNET_PROJECTION_CLASS_EMBEDDING_INPUT_DIM = 2816
|
||||
|
||||
|
||||
def get_torch_dtype(dtype_str):
|
||||
# if it is a torch dtype, return it
|
||||
@@ -433,3 +437,183 @@ def addnet_hash_legacy(b):
|
||||
b.seek(0x100000)
|
||||
m.update(b.read(0x10000))
|
||||
return m.hexdigest()[0:8]
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from transformers import CLIPTextModel, CLIPTokenizer, CLIPTextModelWithProjection
|
||||
|
||||
|
||||
def text_tokenize(
|
||||
tokenizer: 'CLIPTokenizer', # 普通ならひとつ、XLならふたつ!
|
||||
prompts: list[str],
|
||||
):
|
||||
return tokenizer(
|
||||
prompts,
|
||||
padding="max_length",
|
||||
max_length=tokenizer.model_max_length,
|
||||
truncation=True,
|
||||
return_tensors="pt",
|
||||
).input_ids
|
||||
|
||||
|
||||
# https://github.com/huggingface/diffusers/blob/78922ed7c7e66c20aa95159c7b7a6057ba7d590d/src/diffusers/pipelines/stable_diffusion_xl/pipeline_stable_diffusion_xl.py#L334-L348
|
||||
def text_encode_xl(
|
||||
text_encoder: Union['CLIPTextModel', 'CLIPTextModelWithProjection'],
|
||||
tokens: torch.FloatTensor,
|
||||
num_images_per_prompt: int = 1,
|
||||
):
|
||||
prompt_embeds = text_encoder(
|
||||
tokens.to(text_encoder.device), output_hidden_states=True
|
||||
)
|
||||
pooled_prompt_embeds = prompt_embeds[0]
|
||||
prompt_embeds = prompt_embeds.hidden_states[-2] # always penultimate layer
|
||||
|
||||
bs_embed, seq_len, _ = prompt_embeds.shape
|
||||
prompt_embeds = prompt_embeds.repeat(1, num_images_per_prompt, 1)
|
||||
prompt_embeds = prompt_embeds.view(bs_embed * num_images_per_prompt, seq_len, -1)
|
||||
|
||||
return prompt_embeds, pooled_prompt_embeds
|
||||
|
||||
|
||||
def encode_prompts_xl(
|
||||
tokenizers: list['CLIPTokenizer'],
|
||||
text_encoders: list[Union['CLIPTextModel', 'CLIPTextModelWithProjection']],
|
||||
prompts: list[str],
|
||||
num_images_per_prompt: int = 1,
|
||||
use_text_encoder_1: bool = True, # sdxl
|
||||
use_text_encoder_2: bool = True # sdxl
|
||||
) -> tuple[torch.FloatTensor, torch.FloatTensor]:
|
||||
# text_encoder and text_encoder_2's penuultimate layer's output
|
||||
text_embeds_list = []
|
||||
pooled_text_embeds = None # always text_encoder_2's pool
|
||||
|
||||
for idx, (tokenizer, text_encoder) in enumerate(zip(tokenizers, text_encoders)):
|
||||
# todo, we are using a blank string to ignore that encoder for now.
|
||||
# find a better way to do this (zeroing?, removing it from the unet?)
|
||||
prompt_list_to_use = prompts
|
||||
if idx == 0 and not use_text_encoder_1:
|
||||
prompt_list_to_use = ["" for _ in prompts]
|
||||
if idx == 1 and not use_text_encoder_2:
|
||||
prompt_list_to_use = ["" for _ in prompts]
|
||||
|
||||
text_tokens_input_ids = text_tokenize(tokenizer, prompt_list_to_use)
|
||||
text_embeds, pooled_text_embeds = text_encode_xl(
|
||||
text_encoder, text_tokens_input_ids, num_images_per_prompt
|
||||
)
|
||||
|
||||
text_embeds_list.append(text_embeds)
|
||||
|
||||
bs_embed = pooled_text_embeds.shape[0]
|
||||
pooled_text_embeds = pooled_text_embeds.repeat(1, num_images_per_prompt).view(
|
||||
bs_embed * num_images_per_prompt, -1
|
||||
)
|
||||
|
||||
return torch.concat(text_embeds_list, dim=-1), pooled_text_embeds
|
||||
|
||||
|
||||
def text_encode(text_encoder: 'CLIPTextModel', tokens):
|
||||
return text_encoder(tokens.to(text_encoder.device))[0]
|
||||
|
||||
|
||||
def encode_prompts(
|
||||
tokenizer: 'CLIPTokenizer',
|
||||
text_encoder: 'CLIPTokenizer',
|
||||
prompts: list[str],
|
||||
):
|
||||
text_tokens = text_tokenize(tokenizer, prompts)
|
||||
text_embeddings = text_encode(text_encoder, text_tokens)
|
||||
|
||||
return text_embeddings
|
||||
|
||||
|
||||
# for XL
|
||||
def get_add_time_ids(
|
||||
height: int,
|
||||
width: int,
|
||||
dynamic_crops: bool = False,
|
||||
dtype: torch.dtype = torch.float32,
|
||||
):
|
||||
if dynamic_crops:
|
||||
# random float scale between 1 and 3
|
||||
random_scale = torch.rand(1).item() * 2 + 1
|
||||
original_size = (int(height * random_scale), int(width * random_scale))
|
||||
# random position
|
||||
crops_coords_top_left = (
|
||||
torch.randint(0, original_size[0] - height, (1,)).item(),
|
||||
torch.randint(0, original_size[1] - width, (1,)).item(),
|
||||
)
|
||||
target_size = (height, width)
|
||||
else:
|
||||
original_size = (height, width)
|
||||
crops_coords_top_left = (0, 0)
|
||||
target_size = (height, width)
|
||||
|
||||
# this is expected as 6
|
||||
add_time_ids = list(original_size + crops_coords_top_left + target_size)
|
||||
|
||||
# this is expected as 2816
|
||||
passed_add_embed_dim = (
|
||||
UNET_ATTENTION_TIME_EMBED_DIM * len(add_time_ids) # 256 * 6
|
||||
+ TEXT_ENCODER_2_PROJECTION_DIM # + 1280
|
||||
)
|
||||
if passed_add_embed_dim != UNET_PROJECTION_CLASS_EMBEDDING_INPUT_DIM:
|
||||
raise ValueError(
|
||||
f"Model expects an added time embedding vector of length {UNET_PROJECTION_CLASS_EMBEDDING_INPUT_DIM}, but a vector of {passed_add_embed_dim} was created. The model has an incorrect config. Please check `unet.config.time_embedding_type` and `text_encoder_2.config.projection_dim`."
|
||||
)
|
||||
|
||||
add_time_ids = torch.tensor([add_time_ids], dtype=dtype)
|
||||
return add_time_ids
|
||||
|
||||
|
||||
def concat_embeddings(
|
||||
unconditional: torch.FloatTensor,
|
||||
conditional: torch.FloatTensor,
|
||||
n_imgs: int,
|
||||
):
|
||||
return torch.cat([unconditional, conditional]).repeat_interleave(n_imgs, dim=0)
|
||||
|
||||
|
||||
def add_all_snr_to_noise_scheduler(noise_scheduler, device):
|
||||
if hasattr(noise_scheduler, "all_snr"):
|
||||
return
|
||||
# compute it
|
||||
with torch.no_grad():
|
||||
alphas_cumprod = noise_scheduler.alphas_cumprod
|
||||
sqrt_alphas_cumprod = torch.sqrt(alphas_cumprod)
|
||||
sqrt_one_minus_alphas_cumprod = torch.sqrt(1.0 - alphas_cumprod)
|
||||
alpha = sqrt_alphas_cumprod
|
||||
sigma = sqrt_one_minus_alphas_cumprod
|
||||
all_snr = (alpha / sigma) ** 2
|
||||
all_snr.requires_grad = False
|
||||
noise_scheduler.all_snr = all_snr.to(device)
|
||||
|
||||
|
||||
def get_all_snr(noise_scheduler, device):
|
||||
if hasattr(noise_scheduler, "all_snr"):
|
||||
return noise_scheduler.all_snr.to(device)
|
||||
# compute it
|
||||
with torch.no_grad():
|
||||
alphas_cumprod = noise_scheduler.alphas_cumprod
|
||||
sqrt_alphas_cumprod = torch.sqrt(alphas_cumprod)
|
||||
sqrt_one_minus_alphas_cumprod = torch.sqrt(1.0 - alphas_cumprod)
|
||||
alpha = sqrt_alphas_cumprod
|
||||
sigma = sqrt_one_minus_alphas_cumprod
|
||||
all_snr = (alpha / sigma) ** 2
|
||||
all_snr.requires_grad = False
|
||||
return all_snr.to(device)
|
||||
|
||||
|
||||
def apply_snr_weight(
|
||||
loss,
|
||||
timesteps,
|
||||
noise_scheduler: Union['DDPMScheduler'],
|
||||
gamma
|
||||
):
|
||||
# will get it form noise scheduler if exist or will calculate it if not
|
||||
all_snr = get_all_snr(noise_scheduler, loss.device)
|
||||
|
||||
snr = torch.stack([all_snr[t] for t in timesteps])
|
||||
gamma_over_snr = torch.div(torch.ones_like(snr) * gamma, snr)
|
||||
snr_weight = torch.minimum(gamma_over_snr, torch.ones_like(gamma_over_snr)).float().to(loss.device) # from paper
|
||||
loss = loss * snr_weight
|
||||
return loss
|
||||
|
||||
Reference in New Issue
Block a user