Updates of training logic

* use seperated dtype for trainable weight
* force "training module only" before training
* disable gradient after training
* ensure same dtype after training
This commit is contained in:
Kohaku-Blueleaf 2025-06-07 05:09:16 +08:00
parent 9c0cf36cd7
commit 5e43ec9031

View File

@ -8,9 +8,12 @@ import safetensors
import torch import torch
from PIL import Image, ImageDraw, ImageFont from PIL import Image, ImageDraw, ImageFont
from PIL.PngImagePlugin import PngInfo from PIL.PngImagePlugin import PngInfo
import torch.utils.checkpoint
import comfy.samplers import comfy.samplers
import comfy.sd
import comfy.utils import comfy.utils
import comfy.model_management
import comfy_extras.nodes_custom_sampler import comfy_extras.nodes_custom_sampler
import folder_paths import folder_paths
import node_helpers import node_helpers
@ -37,6 +40,7 @@ class TrainSampler(comfy.samplers.Sampler):
) )
# Ensure model is in training mode and computing gradients # Ensure model is in training mode and computing gradients
# x0 pred
denoised = model_wrap(noise, sigmas, **extra_args) denoised = model_wrap(noise, sigmas, **extra_args)
try: try:
loss = self.loss_fn(denoised, latent.clone()) loss = self.loss_fn(denoised, latent.clone())
@ -304,6 +308,10 @@ class TrainLoraNode:
["bf16", "fp32"], ["bf16", "fp32"],
{"default": "bf16", "tooltip": "The dtype to use for training."}, {"default": "bf16", "tooltip": "The dtype to use for training."},
), ),
"lora_dtype": (
["bf16", "fp32"],
{"default": "bf32", "tooltip": "The dtype to use for lora."},
),
"existing_lora": ( "existing_lora": (
folder_paths.get_filename_list("loras") + ["[None]"], folder_paths.get_filename_list("loras") + ["[None]"],
{ {
@ -334,6 +342,7 @@ class TrainLoraNode:
loss_function, loss_function,
seed, seed,
training_dtype, training_dtype,
lora_dtype,
existing_lora, existing_lora,
): ):
num_images = image.shape[0] num_images = image.shape[0]
@ -344,6 +353,7 @@ class TrainLoraNode:
encoded = vae.encode(batch_tensor) encoded = vae.encode(batch_tensor)
mp = model.clone() mp = model.clone()
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)
mp.set_model_compute_dtype(dtype) mp.set_model_compute_dtype(dtype)
with torch.inference_mode(False): with torch.inference_mode(False):
@ -361,46 +371,60 @@ class TrainLoraNode:
if lora_path: if lora_path:
existing_weights = comfy.utils.load_torch_file(lora_path) existing_weights = comfy.utils.load_torch_file(lora_path)
all_weight_adapters = []
for n, m in mp.model.named_modules(): for n, m in mp.model.named_modules():
if hasattr(m, "weight_function"): if hasattr(m, "weight_function"):
if m.weight is not None: if m.weight is not None:
key = "{}.weight".format(n) key = "{}.weight".format(n)
shape = m.weight.shape shape = m.weight.shape
if len(shape) >= 2: if len(shape) >= 2:
existing_adapter = None alpha = float(existing_weights.get(f"{key}.alpha", 1.0))
dora_scale = existing_weights.get(
f"{key}.dora_scale", None
)
for adapter_cls in adapters: for adapter_cls in adapters:
existing_adapter = adapter_cls.load( existing_adapter = adapter_cls.load(
n, existing_weights n, existing_weights, alpha, dora_scale
) )
if existing_adapter is not None: if existing_adapter is not None:
break break
else:
# If no existing adapter found, use LoRA
# We will add algo option in the future
existing_adapter = None
adapter_cls = adapters[0]
if existing_adapter is not None: if existing_adapter is not None:
train_adapter = existing_adapter.to_train() train_adapter = existing_adapter.to_train().to(lora_dtype)
for name, parameter in train_adapter.named_parameters(): for name, parameter in train_adapter.named_parameters():
lora_sd[f"{n}.{name}"] = parameter lora_sd[f"{n}.{name}"] = parameter
else: else:
# Use LoRA with alpha=1.0 by default # Use LoRA with alpha=1.0 by default
train_adapter = adapter_cls[0].create_train( train_adapter = adapter_cls.create_train(
m.weight, rank=rank, alpha=1.0 m.weight, rank=rank, alpha=1.0
) ).to(lora_dtype)
mp.add_weight_wrapper(key, train_adapter) mp.add_weight_wrapper(key, train_adapter)
all_weight_adapters.append(train_adapter)
else: else:
diff = torch.nn.Parameter( diff = torch.nn.Parameter(
torch.zeros( torch.zeros(
m.weight.shape, dtype=dtype, requires_grad=True m.weight.shape, dtype=dtype, requires_grad=True
) )
) )
diff_module = BiasDiff(diff)
mp.add_weight_wrapper(key, BiasDiff(diff)) mp.add_weight_wrapper(key, BiasDiff(diff))
all_weight_adapters.append(diff_module)
lora_sd["{}.diff".format(n)] = diff lora_sd["{}.diff".format(n)] = diff
if hasattr(m, "bias") and m.bias is not None: if hasattr(m, "bias") and m.bias is not None:
key = "{}.bias".format(n) key = "{}.bias".format(n)
bias = torch.nn.Parameter( bias = torch.nn.Parameter(
torch.zeros(m.bias.shape, dtype=dtype, requires_grad=True) torch.zeros(m.bias.shape, dtype=dtype, requires_grad=True)
) )
bias_module = BiasDiff(bias)
lora_sd["{}.diff_b".format(n)] = bias lora_sd["{}.diff_b".format(n)] = bias
mp.add_weight_wrapper(key, BiasDiff(bias)) mp.add_weight_wrapper(key, BiasDiff(bias))
all_weight_adapters.append(bias_module)
if optimizer == "Adam": if optimizer == "Adam":
optimizer = torch.optim.Adam(lora_sd.values(), lr=learning_rate) optimizer = torch.optim.Adam(lora_sd.values(), lr=learning_rate)
@ -433,7 +457,11 @@ class TrainLoraNode:
# yoland: this currently resize to the first image in the dataset # yoland: this currently resize to the first image in the dataset
# setup before training
comfy.model_management.load_models_gpu([mp], memory_required=1e20, force_full_load=True)
# Training loop # Training loop
torch.cuda.empty_cache()
for step in range(steps): for step in range(steps):
# Generate random sigma # Generate random sigma
sigma = mp.model.model_sampling.percent_to_sigma( sigma = mp.model.model_sampling.percent_to_sigma(
@ -446,6 +474,14 @@ class TrainLoraNode:
ss.sample( ss.sample(
noise, guider, train_sampler, sigma, {"samples": encoded.clone()} noise, guider, train_sampler, sigma, {"samples": encoded.clone()}
) )
del ss, train_sampler, optimizer
torch.cuda.empty_cache()
for adapter in all_weight_adapters:
adapter.requires_grad_(False)
for param in lora_sd:
lora_sd[param] = lora_sd[param].to(lora_dtype)
return (mp, lora_sd, loss_map, steps + existing_steps) return (mp, lora_sd, loss_map, steps + existing_steps)
@ -519,7 +555,7 @@ class LossGraphNode:
def plot_loss(self, loss, filename_prefix, prompt=None, extra_pnginfo=None): def plot_loss(self, loss, filename_prefix, prompt=None, extra_pnginfo=None):
loss_values = loss["loss"] loss_values = loss["loss"]
width, height = 500, 300 width, height = 800, 480
margin = 40 margin = 40
img = Image.new( img = Image.new(