mirror of
https://git.datalinker.icu/comfyanonymous/ComfyUI
synced 2026-09-04 01:17:18 +08:00
LoRA training QoL improvements
This commit is contained in:
parent
7a13f74220
commit
28fc7abc19
@ -342,6 +342,13 @@ class TrainLoraNode:
|
|||||||
["bf16", "fp32"],
|
["bf16", "fp32"],
|
||||||
{"default": "bf16", "tooltip": "The dtype to use for lora."},
|
{"default": "bf16", "tooltip": "The dtype to use for lora."},
|
||||||
),
|
),
|
||||||
|
"gradient_checkpointing": (
|
||||||
|
IO.BOOLEAN,
|
||||||
|
{
|
||||||
|
"default": True,
|
||||||
|
"tooltip": "Use gradient checkpointing to reduce memory usage at the cost of speed)",
|
||||||
|
},
|
||||||
|
),
|
||||||
"existing_lora": (
|
"existing_lora": (
|
||||||
folder_paths.get_filename_list("loras") + ["[None]"],
|
folder_paths.get_filename_list("loras") + ["[None]"],
|
||||||
{
|
{
|
||||||
@ -372,9 +379,11 @@ class TrainLoraNode:
|
|||||||
seed,
|
seed,
|
||||||
training_dtype,
|
training_dtype,
|
||||||
lora_dtype,
|
lora_dtype,
|
||||||
|
gradient_checkpointing,
|
||||||
existing_lora,
|
existing_lora,
|
||||||
):
|
):
|
||||||
mp = model.clone()
|
mp = model.clone()
|
||||||
|
device = comfy.model_management.get_torch_device()
|
||||||
dtype = node_helpers.string_to_torch_dtype(training_dtype)
|
dtype = node_helpers.string_to_torch_dtype(training_dtype)
|
||||||
lora_dtype = node_helpers.string_to_torch_dtype(lora_dtype)
|
lora_dtype = node_helpers.string_to_torch_dtype(lora_dtype)
|
||||||
mp.set_model_compute_dtype(dtype)
|
mp.set_model_compute_dtype(dtype)
|
||||||
@ -384,8 +393,9 @@ class TrainLoraNode:
|
|||||||
|
|
||||||
with torch.inference_mode(False):
|
with torch.inference_mode(False):
|
||||||
lora_sd = {}
|
lora_sd = {}
|
||||||
generator = torch.Generator()
|
old_cpu_rng_state = torch.get_rng_state()
|
||||||
generator.manual_seed(seed)
|
old_device_rng_state = torch.cuda.get_rng_state(device)
|
||||||
|
torch.manual_seed(seed)
|
||||||
|
|
||||||
# Load existing LoRA weights if provided
|
# Load existing LoRA weights if provided
|
||||||
existing_weights = {}
|
existing_weights = {}
|
||||||
@ -472,8 +482,12 @@ class TrainLoraNode:
|
|||||||
criterion = torch.nn.SmoothL1Loss()
|
criterion = torch.nn.SmoothL1Loss()
|
||||||
|
|
||||||
# setup models
|
# setup models
|
||||||
for m in find_all_highest_child_module_with_forward(mp.model.diffusion_model):
|
if gradient_checkpointing:
|
||||||
patch(m)
|
modules_to_patch = find_all_highest_child_module_with_forward(mp.model.diffusion_model)
|
||||||
|
for m in modules_to_patch:
|
||||||
|
patch(m)
|
||||||
|
logging.info(f"Added gradient checkpoints to {len(modules_to_patch)} modules")
|
||||||
|
|
||||||
comfy.model_management.load_models_gpu([mp], memory_required=1e20, force_full_load=True)
|
comfy.model_management.load_models_gpu([mp], memory_required=1e20, force_full_load=True)
|
||||||
|
|
||||||
# Setup sampler and guider like in test script
|
# Setup sampler and guider like in test script
|
||||||
@ -493,6 +507,9 @@ class TrainLoraNode:
|
|||||||
# Training loop
|
# Training loop
|
||||||
torch.cuda.empty_cache()
|
torch.cuda.empty_cache()
|
||||||
try:
|
try:
|
||||||
|
if comfy.utils.PROGRESS_BAR_ENABLED:
|
||||||
|
ui_pbar = comfy.utils.ProgressBar(steps)
|
||||||
|
|
||||||
for step in (pbar:=tqdm.trange(steps, desc="Training LoRA", smoothing=0.01, disable=not comfy.utils.PROGRESS_BAR_ENABLED)):
|
for step in (pbar:=tqdm.trange(steps, desc="Training LoRA", smoothing=0.01, disable=not comfy.utils.PROGRESS_BAR_ENABLED)):
|
||||||
# Generate random sigma
|
# Generate random sigma
|
||||||
sigma = mp.model.model_sampling.percent_to_sigma(
|
sigma = mp.model.model_sampling.percent_to_sigma(
|
||||||
@ -506,6 +523,7 @@ class TrainLoraNode:
|
|||||||
ss.sample(
|
ss.sample(
|
||||||
noise, guider, train_sampler, sigma, {"samples": latents[indices].clone()}
|
noise, guider, train_sampler, sigma, {"samples": latents[indices].clone()}
|
||||||
)
|
)
|
||||||
|
ui_pbar.update(1)
|
||||||
finally:
|
finally:
|
||||||
for m in mp.model.modules():
|
for m in mp.model.modules():
|
||||||
unpatch(m)
|
unpatch(m)
|
||||||
@ -518,6 +536,9 @@ class TrainLoraNode:
|
|||||||
for param in lora_sd:
|
for param in lora_sd:
|
||||||
lora_sd[param] = lora_sd[param].to(lora_dtype)
|
lora_sd[param] = lora_sd[param].to(lora_dtype)
|
||||||
|
|
||||||
|
torch.set_rng_state(old_cpu_rng_state)
|
||||||
|
torch.cuda.set_rng_state(old_device_rng_state, device)
|
||||||
|
|
||||||
return (mp, lora_sd, loss_map, steps + existing_steps)
|
return (mp, lora_sd, loss_map, steps + existing_steps)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user