67 Commits
Author SHA1 Message Date
mayukhdeb 6e4c86ea0c keep a copy 2024-07-22 11:07:14 -07:00
mayukhdeb 7e30d33dc2 find image filenames recursively in folder + train on big style dataset 2024-07-21 03:56:42 -07:00
mayukhdeb fbd050cb0c cleaner preprocess fn 2024-07-21 03:25:07 -07:00
mayukhdeb 70c8e88195 migrate to huggingface script 2024-07-21 02:05:30 -07:00
mayukhdeb 2f6f0cacd2 ignore stuff 2024-07-21 02:03:42 -07:00
mayukhdeb e8932f30ef start with a different embed string 2024-07-18 09:57:00 -07:00
mayukhdeb b7dd176536 add adamw8bit 2024-07-15 07:48:07 -07:00
mayukhdeb 5a72aca596 useful wandb name 2024-07-15 07:47:34 -07:00
mayukhdeb 85264bd004 wandb 2024-07-15 07:38:27 -07:00
mayukhdeb d46976f3dd sweep params set 2024-07-15 07:33:45 -07:00
mayukhdeb 8ed5996c1f train and inference on same device 2024-07-15 07:32:27 -07:00
mayukhdeb 8247b4b9dd fix OOM during inference (vae.decode) 2024-07-15 07:20:13 -07:00
mayukhdeb dc28fcb988 inference on same device 2024-07-15 03:28:31 -07:00
mayukhdeb 3c0a3be55f save different sh files for each gpu 2024-07-15 03:28:20 -07:00
mayukhdeb 274ec5c90b disable textual inversion if config.ti_lr is None 2024-07-11 12:46:52 -07:00
mayukhdeb fc54948eff more tweaks 2024-07-11 12:39:30 -07:00
mayukhdeb 32a7ae248d prompts for sweep 2024-07-11 12:34:42 -07:00
mayukhdeb 3cbabd2634 sweep dry runs 2024-07-11 12:26:07 -07:00
mayukhdeb 1cdacbc63a remove old todos 2024-07-11 11:16:22 -07:00
mayukhdeb 4668f755ec switch to adamw_8bit for sd3 transformer lora params 2024-07-11 00:15:09 -07:00
mayukhdeb dd093616b5 fix inference prompts bug 2024-07-09 09:09:56 -07:00
mayukhdeb 2c824d6807 more progress 2024-07-09 05:35:21 -07:00
mayukhdeb 022d51f53a inference fixed 2024-07-09 00:25:02 -07:00
mayukhdeb ecacf815ec re-impl textual inversion for first 2 text encoders 2024-07-08 12:45:55 -07:00
mayukhdeb 90d2269572 testing textual inversion training with frozen transformer 2024-07-04 07:40:41 -07:00
mayukhdeb 4955261cae new checkpoint + deterministic inference 2024-07-03 04:28:06 -07:00
mayukhdeb e84932af29 some small changes to text with main_sd3.py 2024-07-03 04:06:21 -07:00
mayukhdeb 115da83e2d cleaner output dir with checkpoints and generated samples in one folder 2024-07-03 03:43:30 -07:00
mayukhdeb 98b78cd6d6 cleanup + save training samples in output dir 2024-07-03 03:11:56 -07:00
mayukhdeb 7c5e0949ba keep changes 2024-07-02 05:19:34 -07:00
mayukhdeb ba2b7532ad more tweaks 2024-07-02 05:16:22 -07:00
mayukhdeb 2c0939a733 run inference less often 2024-07-02 04:25:26 -07:00
mayukhdeb 3b41479aa1 bfloat16 training + inference 2024-07-02 04:12:23 -07:00
mayukhdeb 6e589c76a8 impement some upstream changes and save a sample every 10 train steps 2024-07-02 01:50:14 -07:00
mayukhdeb 95f78d7b91 completely comement out TI for now 2024-07-01 01:56:49 -07:00
mayukhdeb 3c0ebd7d70 apply mask to loss + some hardcoding for banny debugging 2024-06-28 04:38:43 -07:00
mayukhdeb 6350b5344b more progress 2024-06-28 03:22:49 -07:00
mayukhdeb 2c4bd43044 temporarily remove ti grad norms 2024-06-28 03:12:16 -07:00
mayukhdeb b850a1615d watch grad norms 2024-06-28 03:04:40 -07:00
mayukhdeb 6d6dd98bfa dynamic ti lr 2024-06-28 02:03:34 -07:00
mayukhdeb 0444e35729 clip grad norms 2024-06-26 03:01:34 -07:00
mayukhdeb 0748e1d14e better prompt 2024-06-22 02:32:19 -07:00
mayukhdeb db7507c849 handle T5EncoderModel 2024-06-22 01:38:48 -07:00
mayukhdeb b379a28715 update todos 2024-06-22 01:38:12 -07:00
mayukhdeb 54ff8c4977 sd3 concept inference 2024-06-22 01:25:16 -07:00
mayukhdeb 64f28c8589 save TI embeds and lora adapters 2024-06-22 01:24:30 -07:00
mayukhdeb b12ce26fc9 update command 2024-06-20 03:52:30 -07:00
mayukhdeb b3da65dd39 ignore wandb stuff 2024-06-20 03:46:26 -07:00
mayukhdeb f64da5d4f1 update todo 2024-06-20 03:44:42 -07:00
mayukhdeb 5cc09092a4 tweak param 2024-06-20 03:43:44 -07:00
mayukhdeb 3d3d0ee4cb smash more todos 2024-06-20 03:42:32 -07:00
mayukhdeb 4b1efce0e3 compute loss and update weights 2024-06-20 03:28:49 -07:00
mayukhdeb d115c52274 do just forward passes 2024-06-20 03:20:41 -07:00
mayukhdeb 22127c9917 typo 2024-06-20 01:41:28 -07:00
mayukhdeb 4844845d5c more progress 2024-06-20 01:40:59 -07:00
mayukhdeb 3b15250bbc init train dataloader + update todos for training 2024-06-20 00:54:41 -07:00
mayukhdeb fd448e5437 accomodate T5EncoderModel 2024-06-20 00:54:21 -07:00
mayukhdeb a41cce1486 small cleanup 2024-06-20 00:27:51 -07:00
mayukhdeb 12b5960c36 progress bar for latent caching 2024-06-20 00:27:21 -07:00
mayukhdeb 2c2bfcda13 init PreprocessedDataset 2024-06-20 00:27:01 -07:00
mayukhdeb cacd6f2201 count trainable params from model 2024-06-19 23:58:26 -07:00
mayukhdeb 67024b4de6 full or lora finetuning of sd3 transformer 2024-06-19 23:58:16 -07:00
mayukhdeb 6d0d96ec79 more progress on todos 2024-06-19 23:37:31 -07:00
mayukhdeb fb2449a60d init textual inversion token embeds 2024-06-17 07:31:54 -07:00
mayukhdeb 0496ccdbef handle sd3 t5 text encoder 2024-06-17 07:31:33 -07:00
mayukhdeb 044aed9f03 sd3 train script wip 2024-06-17 06:42:16 -07:00
mayukhdeb bdda796c56 ignore notebook checkpoint 2024-06-17 05:21:54 -07:00
11 changed files with 2386 additions and 37 deletions
+7 -2
View File
@@ -1,4 +1,7 @@
data/
sd3_sweep_vis/
sd3_sweep_commands/
.ipynb_checkpoints/
cache
__pycache__
@@ -21,4 +24,6 @@ conditioning_spaces/
training_args_x_*.json
xander_configs/
debug/*
wandb/
sd3_sweep_outputs/
sd3_face_sweep_configs/
+179
View File
@@ -0,0 +1,179 @@
from trainer.utils.json_stuff import save_as_json
import itertools
import copy
import os
import random
random.seed(0)
GPU_IDS = [1,2,3]
wandb_log = True
def divide_list(lst, n):
"""
Divide a list into N equal parts.
Parameters:
lst (list): The list to be divided.
n (int): The number of parts to divide the list into.
Returns:
list of lists: A list containing N sublists, each of which is a part of the original list.
"""
if n <= 0:
raise ValueError("Number of parts must be greater than 0.")
if n > len(lst):
raise ValueError("Number of parts cannot be greater than the length of the list.")
# Calculate the size of each part
k, m = divmod(len(lst), n)
# Create the divided parts
return [lst[i * k + min(i, m):(i + 1) * k + min(i + 1, m)] for i in range(n)]
def generate_sh_file(commands, filename="script.sh"):
"""
Generates a .sh file with each command from the list written on a new line.
:param commands: List of commands to be written to the .sh file.
:param filename: Name of the .sh file to be created. Default is 'script.sh'.
"""
with open(filename, 'w') as file:
for command in commands:
file.write(command + '\n')
print(f"Saved: {filename}")
run_commands_dir = f"./sd3_sweep_commands"
os.system(
f"rm -rf {run_commands_dir} && mkdir -p {run_commands_dir}"
)
config_folder = "./sd3_face_sweep_configs"
os.system(f"rm -rf {config_folder}")
os.system(f"mkdir -p {config_folder}")
sweep_params = {
"unet_learning_rate": [
5e-5,
1e-4,
3e-4,
7e-4,
1e-3,
2e-3,
],
"train_batch_size": [
2,
4,
8,
16
],
"lora_rank": [
2,
4,
6,
8,
],
"ti_lr": [1e-3, None],
"unet_optimizer_type": [
"adamw",
"adamw_8bit",
"prodigy"
],
}
num_total_runs = 1
for key in sweep_params:
num_total_runs *= len(sweep_params[key])
print(f"Num total runs: {num_total_runs}")
default_config = {
"output_dir": "lora_models/sweep",
"sd_model_version": "sd3",
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander_big.zip",
"concept_mode": "face",
"seed": 0,
"resolution": 512,
"train_batch_size": 2,
"n_sample_imgs": 6,
"max_train_steps": 1000,
"token_warmup_steps": 200,
"checkpointing_steps": 1000, ## no need to save any checkpoints
"gradient_accumulation_steps": 2,
"sample_imgs_lora_scale": 0.8,
"n_tokens": 2,
"ti_lr": 0.001,
"remove_ti_token_from_prompts": False,
"text_encoder_lora_optimizer": None,
"text_encoder_lora_lr": 0.5e-4,
"text_encoder_lora_weight_decay": 1e-5,
"text_encoder_lora_rank": 16,
"lora_alpha_multiplier": 1.0,
"lora_rank": 16,
"use_dora": False,
"caption_model": "blip",
"debug": True,
}
keys, values = zip(*sweep_params.items())
combinations = [dict(zip(keys, combination)) for combination in itertools.product(*values)]
all_config_paths = []
for index, c in enumerate(combinations):
config = copy.deepcopy(default_config)
filename = f"{index}"
# override default values with sweep params
for key in c:
"""
instead of editing the train batch size, we simply change the gradient
accumulation value. Which has the same effect.
We will also
"""
if key == "train_batch_size":
config["gradient_accumulation_steps"] = c[key] / config["train_batch_size"]
config["max_train_steps"] = config["max_train_steps"] * config["gradient_accumulation_steps"]
config["checkpointing_steps"] = config["checkpointing_steps"] * config["gradient_accumulation_steps"]
else:
config[key] = c[key]
# print(f"{index} - Setting {key} to {c[key]}")
filename += f"_{key}_{c[key]}"
config_path = os.path.join(
config_folder,
f"{filename}.json"
)
save_as_json(
dictionary_or_list=config,
filename = config_path
)
all_config_paths.append(config_path)
print(f"Saved: {config_path}")
print(f"Total: {index+1} configs")
all_commands = []
for c in all_config_paths:
command = f"python3 main_sd3.py {c}"
if wandb_log:
command = command + " --wandb-log"
all_commands.append(command)
random.shuffle(all_commands)
all_commands_split_by_gpu = divide_list(
lst = all_commands,
n = len(GPU_IDS)
)
for index, gpu_id in enumerate(GPU_IDS):
commands_on_single_gpu = [
f"CUDA_VISIBLE_DEVICES={gpu_id} {x}" for x in all_commands_split_by_gpu[index]
]
generate_sh_file(
commands = commands_on_single_gpu,
filename = os.path.join(
run_commands_dir,
f"run_on_gpu_{gpu_id}.sh"
)
)
+2027
View File
File diff suppressed because it is too large Load Diff
+59
View File
@@ -0,0 +1,59 @@
"""
Pre-trained checkpoint:
https://huggingface.co/stabilityai/stable-diffusion-3-medium-diffusers
"""
import os
import torch
from diffusers import StableDiffusion3Pipeline
import torch
# Load the pretrained model
pipe = StableDiffusion3Pipeline.from_pretrained(
"stabilityai/stable-diffusion-3-medium-diffusers",
torch_dtype=torch.float16,
seed = 0
)
# Load the LoRA weights from file
lora_weights_path = "sd3-xander/checkpoint-1000/pytorch_lora_weights.safetensors"
# Move model to GPU
pipe = pipe.to("cuda")
prompts = [
"This is a picture of a man holding a glass of beer. He is wearing a casual plaid shirt and jeans. The man is holding a frosty glass of golden beer with a thick, foamy head in his right hand, lifting it slightly as if making a toast. The background features wooden tables and chairs, vintage beer signs, and warm ambient lighting",
"A close up shot of a man as a dragon rider with a red sword named Za'roc. His face is clearly visible in the high cinematic shot.",
"A man in 2075, looking for the last drop of water in mars. 4k HDR",
# "A king in Skyrim"
]
for idx, prompt in enumerate(prompts):
image = pipe(
prompt,
negative_prompt="",
num_inference_steps=28,
guidance_scale=7.0,
).images[0]
image.save(
os.path.join(
"./outputs",
f"{idx}_baseline.jpg"
)
)
pipe.load_lora_weights(lora_weights_path, alpha = 8)
for idx, prompt in enumerate(prompts):
image = pipe(
prompt,
negative_prompt="",
num_inference_steps=28,
guidance_scale=7.0,
).images[0]
image.save(
os.path.join(
"./outputs",
f"{idx}.jpg"
)
)
print(f"Done!")
+4 -4
View File
@@ -11,7 +11,7 @@ class TrainingConfig(BaseModel):
concept_mode: Literal["face", "style", "object"]
caption_prefix: str = "" # hardcoding this will inject TOK manually and skip the chatgpt token injection step, not recommended unless you know what you're doing
caption_model: Literal["gpt4-v", "blip"] = "blip"
sd_model_version: Literal["sdxl", "sd15"]
sd_model_version: Literal["sdxl", "sd15", "sd3"]
pretrained_model: dict = None
seed: Union[int, None] = None
resolution: int = 512
@@ -25,14 +25,14 @@ class TrainingConfig(BaseModel):
gradient_accumulation_steps: int = 1
is_lora: bool = True
unet_optimizer_type: Literal["adamw", "prodigy"] = "adamw"
unet_optimizer_type: Literal["adamw", "prodigy", "adamw_8bit"] = "adamw"
unet_lr_warmup_steps: int = None # slowly increase the learning rate of the adamw unet optimizer
unet_lr: float = 1.0e-3
prodigy_d_coef: float = 1.0
unet_prodigy_growth_factor: float = 1.05 # lower values make the lr go up slower (1.01 is for 1k step runs, 1.02 is for 500 step runs)
lora_weight_decay: float = 0.002
ti_lr: float = 1e-3
# if ti_lr is None, then we completely skip textual inversion
ti_lr: Union[float, None] = 1e-3
ti_lr_warmup_steps: int = 20 # slowly ramp up the learning rate to build some momentum
token_warmup_steps: int = 0 # warmup the token embeddings with a pure txt loss
ti_weight_decay: float = 0.0
+2 -1
View File
@@ -6,6 +6,7 @@ import PIL
from PIL import Image
from torch.utils.data import Dataset
from typing import Tuple, Dict, List
from tqdm import tqdm
def prepare_image(
pil_image: PIL.Image.Image, w: int = 512, h: int = 512, pipe=None,
@@ -68,7 +69,7 @@ class PreprocessedDataset(Dataset):
self.masks = []
self.do_cache = True
for idx in range(len(self.data)):
for idx in tqdm(range(len(self.data))):
if len(self.data) < 25:
print(self.captions[idx])
vae_latent, mask = self._process(idx)
+82 -23
View File
@@ -9,6 +9,7 @@ from typing import List, Optional, Dict
from safetensors.torch import save_file, safe_open
import matplotlib.pyplot as plt
from trainer.utils.utils import seed_everything, plot_torch_hist, plot_loss
from transformers import T5EncoderModel
class TokenEmbeddingsHandler:
def __init__(self, text_encoders, tokenizers):
@@ -31,7 +32,10 @@ class TokenEmbeddingsHandler:
continue
# Directly accessing and modifying the original weights tensor
text_encoder.text_model.embeddings.token_embedding.weight.requires_grad_(True)
if isinstance(text_encoder, T5EncoderModel):
text_encoder.encoder.embed_tokens.weight.requires_grad_(True)
else:
text_encoder.text_model.embeddings.token_embedding.weight.requires_grad_(True)
print(f"All embeddings in text_encoder_{idx} are now set to be trainable.")
def get_trainable_embeddings(self):
@@ -49,10 +53,22 @@ class TokenEmbeddingsHandler:
continue
# Ensure indices are a tensor. Use pre-existing dtype and device to match the model's.
indices_tensor = torch.tensor(indices, dtype=torch.long, device=text_encoder.text_model.embeddings.token_embedding.weight.device)
if isinstance(text_encoder, T5EncoderModel):
indices_tensor = torch.tensor(
indices,
dtype=torch.long,
device=text_encoder.encoder.embed_tokens.weight.device
)
# Directly access the embedding weights without detaching
token_embeddings = text_encoder.encoder.embed_tokens.weight[indices_tensor]
else:
indices_tensor = torch.tensor(indices, dtype=torch.long, device=text_encoder.text_model.embeddings.token_embedding.weight.device)
# Directly access the embedding weights without detaching
token_embeddings = text_encoder.text_model.embeddings.token_embedding.weight[indices_tensor]
# Directly access the embedding weights without detaching
token_embeddings = text_encoder.text_model.embeddings.token_embedding.weight[indices_tensor]
embeddings[f'txt_encoder_{idx}'] = token_embeddings
# Get all corresponding tokens for these embeddings
@@ -192,9 +208,19 @@ class TokenEmbeddingsHandler:
self.non_train_ids = all_indices[inu]
# random initialization of new tokens
std_token_embedding = (
text_encoder.text_model.embeddings.token_embedding.weight.data.std(dim=1).mean()
)
"""
handle both T5EncoderModel and other text encoders
T5EncoderModel is present in sd3
"""
if isinstance(text_encoder, T5EncoderModel):
std_token_embedding = (
text_encoder.encoder.embed_tokens.weight.data.std(dim=1).mean()
)
else:
std_token_embedding = (
text_encoder.text_model.embeddings.token_embedding.weight.data.std(dim=1).mean()
)
self.embeddings_settings[f"std_token_embedding_{idx}"] = std_token_embedding
if starting_toks is not None:
@@ -207,14 +233,28 @@ class TokenEmbeddingsHandler:
self.train_ids] = text_encoder.text_model.embeddings.token_embedding.weight.data[self.starting_ids].clone()
else:
std_multiplier = 1.0
init_embeddings = torch.randn(len(self.train_ids), text_encoder.text_model.config.hidden_size).to(device=self.device).to(dtype=self.dtype)
if isinstance(text_encoder, T5EncoderModel):
init_embeddings = torch.randn(len(self.train_ids), text_encoder.config.hidden_size).to(device=self.device).to(dtype=self.dtype)
else:
init_embeddings = torch.randn(len(self.train_ids), text_encoder.text_model.config.hidden_size).to(device=self.device).to(dtype=self.dtype)
current_std = init_embeddings.std(dim=1).mean()
init_embeddings = init_embeddings * std_multiplier * std_token_embedding / current_std
text_encoder.text_model.embeddings.token_embedding.weight.data[self.train_ids] = init_embeddings.clone()
self.embeddings_settings[
f"original_embeddings_{idx}"
] = text_encoder.text_model.embeddings.token_embedding.weight.data.clone()
if isinstance(text_encoder, T5EncoderModel):
text_encoder.encoder.embed_tokens.weight.data[self.train_ids] = init_embeddings.clone()
else:
text_encoder.text_model.embeddings.token_embedding.weight.data[self.train_ids] = init_embeddings.clone()
if isinstance(text_encoder, T5EncoderModel):
self.embeddings_settings[
f"original_embeddings_{idx}"
] = text_encoder.encoder.embed_tokens.weight.data.clone()
else:
self.embeddings_settings[
f"original_embeddings_{idx}"
] = text_encoder.text_model.embeddings.token_embedding.weight.data.clone()
inu = torch.ones((len(tokenizer),), dtype=torch.bool)
inu[self.train_ids] = False
@@ -414,14 +454,27 @@ class TokenEmbeddingsHandler:
for idx, text_encoder in enumerate(self.text_encoders):
if text_encoder is None:
continue
assert text_encoder.text_model.embeddings.token_embedding.weight.data.shape[
0
] == len(self.tokenizers[0]), "Tokenizers should be the same."
new_token_embeddings = (
text_encoder.text_model.embeddings.token_embedding.weight.data[
self.train_ids
]
)
if isinstance(text_encoder, T5EncoderModel):
assert text_encoder.encoder.embed_tokens.weight.data.shape[
0
] == len(self.tokenizers[idx]), "Tokenizers should be the same."
new_token_embeddings = (
text_encoder.encoder.embed_tokens.weight.data[
self.train_ids
]
)
else:
assert text_encoder.text_model.embeddings.token_embedding.weight.data.shape[
0
] == len(self.tokenizers[0]), "Tokenizers should be the same."
new_token_embeddings = (
text_encoder.text_model.embeddings.token_embedding.weight.data[
self.train_ids
]
)
tensors[txt_encoder_keys[idx]] = new_token_embeddings
save_file(tensors, file_path)
@@ -474,9 +527,15 @@ class TokenEmbeddingsHandler:
self.train_ids = tokenizer.convert_tokens_to_ids(self.inserting_toks)
assert self.train_ids is not None, "New tokens could not be converted to IDs."
text_encoder.text_model.embeddings.token_embedding.weight.data[
self.train_ids
] = loaded_embeddings.to(device=self.device).to(dtype=self.dtype)
if isinstance(text_encoder, T5EncoderModel):
text_encoder.encoder.embed_tokens.weight.data[
self.train_ids
] = loaded_embeddings.to(device=self.device).to(dtype=self.dtype)
else:
text_encoder.text_model.embeddings.token_embedding.weight.data[
self.train_ids
] = loaded_embeddings.to(device=self.device).to(dtype=self.dtype)
def load_embeddings(self, file_path: str, txt_encoder_keys = ["clip_l", "clip_g"]):
if not os.path.exists(file_path):
+14 -2
View File
@@ -4,6 +4,7 @@ import matplotlib.pyplot as plt
import torch
from torch.utils._foreach_utils import _group_tensors_by_device_and_dtype, _has_foreach_support
from trainer.inference import get_conditioning_signals
from transformers import T5EncoderModel
def compute_snr(noise_scheduler, timesteps):
"""
@@ -104,7 +105,14 @@ class ConditioningRegularizer:
def __init__(self, config, embedding_handler):
self.config = config
self.embedding_handler = embedding_handler
self.target_norm = 34.5 if config.sd_model_version == 'sdxl' else 27.8
self.target_norms = {
"sdxl": 34.5,
"sd15": 27.8,
"sd3": 34.5
}
print(f'\033[91m[trainer.loss.ConditioningRegularizer] WARNING: Using a magic number: 34.5 for the target norm of sd3. We do not know if this is the ideal value. This might cause bugs or even break training completely.\033[0m')
self.target_norm = self.target_norms[config.sd_model_version]
self.reg_captions = ["a photo of TOK", "TOK", "a photo of TOK next to TOK", "TOK and TOK"]
self.token_replacement = config.token_dict.get("TOK", "TOK") # Fallback to "TOK" if not in dict
@@ -114,7 +122,11 @@ class ConditioningRegularizer:
if tokenizer is None:
idx += 1
continue
pretrained_token_embeddings = text_encoder.text_model.embeddings.token_embedding.weight.data
if isinstance(text_encoder, T5EncoderModel):
pretrained_token_embeddings = text_encoder.encoder.embed_tokens.weight.data
else:
pretrained_token_embeddings = text_encoder.text_model.embeddings.token_embedding.weight.data
self.distribution_regularizers[f'txt_encoder_{idx}'] = DistributionLoss(pretrained_token_embeddings, outdir = self.config.output_dir if config.debug else None)
idx += 1
+3 -1
View File
@@ -14,6 +14,7 @@ SDXL_URL = "https://edenartlab-lfs.s3.amazonaws.com/models/checkpoints/j
SD15_MODEL_CACHE = "./models/juggernaut_reborn.safetensors"
SD15_URL = "https://edenartlab-lfs.s3.amazonaws.com/models/checkpoints/juggernaut_reborn.safetensors"
SD3_MODEL_CACHE = "models/stable-diffusion-3-medium"
#SD15_MODEL_CACHE = "./models/DreamShaper_6.31_BakedVae.safetensors"
#SD15_URL = "https://huggingface.co/Lykon/DreamShaper/resolve/main/DreamShaper_6.31_BakedVae.safetensors"
@@ -23,7 +24,8 @@ SD15_URL = "https://edenartlab-lfs.s3.amazonaws.com/models/checkpoints/j
pretrained_models = {
"sdxl": {"path": SDXL_MODEL_CACHE, "url": SDXL_URL, "version": "sdxl"},
"sd15": {"path": SD15_MODEL_CACHE, "url": SD15_URL, "version": "sd15"}
"sd15": {"path": SD15_MODEL_CACHE, "url": SD15_URL, "version": "sd15"},
"sd3": {"path": SD3_MODEL_CACHE, "url": None, "version": "sd3"}
}
############################################################################################################
+5
View File
@@ -3,6 +3,11 @@ import torch
import prodigyopt
from typing import Iterable
def count_trainable_params(model):
return sum([
x.numel() for x in model.parameters() if x.requires_grad
])
def get_unet_optimizer(
prodigy_d_coef: float,
prodigy_growth_factor: float,
+4 -4
View File
@@ -5,11 +5,11 @@
"concept_mode": "face",
"seed": 0,
"resolution": 512,
"train_batch_size": 4,
"train_batch_size": 7,
"n_sample_imgs": 6,
"max_train_steps": 600,
"max_train_steps": 5000,
"token_warmup_steps": 0,
"checkpointing_steps": 100,
"checkpointing_steps": 50,
"gradient_accumulation_steps": 1,
"sample_imgs_lora_scale": 0.8,
@@ -26,6 +26,6 @@
"lora_alpha_multiplier": 1.0,
"lora_rank": 16,
"use_dora": false,
"caption_model": "gpt4-v",
"caption_model": "blip",
"debug": true
}