fix loss visualizer, and other small stuff

This commit is contained in:
kijai
2024-08-19 02:51:19 +03:00
parent fdcb5a8686
commit 2cd3c8e170
4 changed files with 29 additions and 22 deletions
+2 -1
View File
@@ -577,7 +577,8 @@ class FluxTrainer:
def training_loop(break_at_steps, epoch):
global optimizer_hooked_count
steps_done = 0
accelerator.print(f"\nepoch {epoch+1}/{num_train_epochs}")
#accelerator.print(f"\nepoch {epoch+1}/{num_train_epochs}")
progress_bar.set_description(f"Epoch {epoch + 1}/{num_train_epochs} - steps")
current_epoch.value = epoch + 1
for m in training_models:
+6 -3
View File
@@ -94,8 +94,8 @@ DEFAULT_EPOCH_NAME = "epoch"
DEFAULT_LAST_OUTPUT_NAME = "last"
DEFAULT_STEP_NAME = "at"
STEP_STATE_NAME = "{}-step{:08d}-state"
STEP_FILE_NAME = "{}-step{:08d}"
STEP_STATE_NAME = "{}-step{:05d}-state"
STEP_FILE_NAME = "{}-step{:05d}"
STEP_DIFFUSERS_DIR_NAME = "{}-step{:08d}"
# region dataset
@@ -5820,16 +5820,19 @@ class collator_class:
class LossRecorder:
def __init__(self):
self.loss_list: List[float] = []
self.global_loss_list: List[float] = []
self.loss_total: float = 0.0
def add(self, *, epoch: int, step: int, loss: float) -> None:
def add(self, *, epoch: int, step: int, global_step: int, loss: float) -> None:
if epoch == 0:
self.loss_list.append(loss)
self.global_loss_list.append(loss)
else:
while len(self.loss_list) <= step:
self.loss_list.append(0.0)
self.loss_total -= self.loss_list[step]
self.loss_list[step] = loss
self.global_loss_list.append(loss)
self.loss_total += loss
@property
+16 -12
View File
@@ -171,7 +171,7 @@ class InitFluxLoRATraining:
"discrete_flow_shift": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "for the Euler Discrete Scheduler, default is 3.0"}),
"highvram": ("BOOLEAN", {"default": False, "tooltip": "memory mode"}),
"fp8_base": ("BOOLEAN", {"default": True, "tooltip": "use fp8 for base model"}),
"training_dtype": (["fp32", "fp16", "bf16"], {"default": "fp32", "tooltip": "the actual dtype training uses"}),
"gradient_dtype": (["fp32", "fp16", "bf16"], {"default": "fp32", "tooltip": "the actual dtype training uses"}),
"save_dtype": (["fp32", "fp16", "bf16", "fp8_e4m3fn"], {"default": "bf16", "tooltip": "the dtype to save checkpoints as"}),
"attention_mode": (["sdpa", "xformers", "disabled"], {"default": "sdpa", "tooltip": "memory efficient attention mode"}),
"sample_prompts": ("STRING", {"multiline": True, "default": "illustration of a kitten | photograph of a turtle", "tooltip": "validation sample prompts, for multiple prompts, separate by `|`"}),
@@ -183,7 +183,8 @@ class InitFluxLoRATraining:
FUNCTION = "init_training"
CATEGORY = "FluxTrainer"
def init_training(self, flux_models, dataset_settings, optimizer_settings, sample_prompts, output_name, attention_mode, training_dtype, save_dtype, **kwargs,):
def init_training(self, flux_models, dataset_settings, optimizer_settings, sample_prompts, output_name, attention_mode,
gradient_dtype, save_dtype, **kwargs,):
mm.soft_empty_cache()
dataset = dataset_settings["dataset"]
@@ -254,11 +255,11 @@ class InitFluxLoRATraining:
}
config_dict.update(attention_settings.get(attention_mode, {}))
training_dtype_settings = {
"fp16": {"full_fp16": True, "full_bf16": False},
"bf16": {"full_bf16": True, "full_fp16": False}
gradient_dtype_settings = {
"fp16": {"full_fp16": True, "full_bf16": False, "mixed_precision": "fp16"},
"bf16": {"full_bf16": True, "full_fp16": False, "mixed_precision": "bf16"}
}
config_dict.update(training_dtype_settings.get(training_dtype, {}))
config_dict.update(gradient_dtype_settings.get(gradient_dtype, {}))
if optimizer_settings["optimizer_type"] == "adafactor":
config_dict["optimizer_args"] = [
@@ -623,26 +624,29 @@ class VisualizeLoss:
def INPUT_TYPES(s):
return {"required": {
"network_trainer": ("NETWORKTRAINER",),
"plot_style": (['default', 'classic', 'dark_background', 'ggplot', 'seaborn', 'seaborn-dark', 'seaborn-darkgrid', 'fast', 'grayscale'],),
},
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("plot",)
RETURN_TYPES = ("IMAGE", "FLOAT",)
RETURN_NAMES = ("plot", "loss_list",)
FUNCTION = "draw"
CATEGORY = "FluxTrainer"
def draw(self, network_trainer):
def draw(self, network_trainer, plot_style):
import matplotlib.pyplot as plt
import io
from PIL import Image
# Example list of loss values
loss_values = network_trainer["network_trainer"].loss_recorder.loss_list
loss_values = network_trainer["network_trainer"].loss_recorder.global_loss_list
plt.style.use(plot_style)
# Create a plot
fig, ax = plt.subplots()
ax.plot(loss_values, label='Training Loss')
ax.set_xlabel('Epoch')
ax.set_xlabel('Step')
ax.set_ylabel('Loss')
ax.set_title('Training Loss Over Time')
ax.legend()
@@ -661,7 +665,7 @@ class VisualizeLoss:
image_tensor = transforms.ToTensor()(image)
image_tensor = image_tensor.unsqueeze(0).permute(0, 2, 3, 1).cpu().float()
return image_tensor,
return image_tensor, loss_values,
class FluxKohyaInferenceSampler:
@classmethod
+5 -6
View File
@@ -966,7 +966,7 @@ class NetworkTrainer:
init_kwargs=init_kwargs,
)
loss_recorder = train_util.LossRecorder()
self.loss_recorder = train_util.LossRecorder()
del train_dataset_group
pbar.update(1)
@@ -1040,10 +1040,9 @@ class NetworkTrainer:
self.metadata = metadata
self.optimizer = optimizer
self.lr_scheduler = lr_scheduler
self.loss_recorder = loss_recorder
self.save_model = save_model
self.remove_model = remove_model
progress_bar = tqdm(range(args.max_train_steps - initial_step), smoothing=0, disable=False, desc="steps")
def training_loop(break_at_steps, epoch):
steps_done = 0
@@ -1178,8 +1177,8 @@ class NetworkTrainer:
self.global_step += 1
current_loss = loss.detach().item()
loss_recorder.add(epoch=epoch, step=step, loss=current_loss)
avr_loss: float = loss_recorder.moving_average
self.loss_recorder.add(epoch=epoch, step=step, global_step=self.global_step, loss=current_loss)
avr_loss: float = self.loss_recorder.moving_average
logs = {"avr_loss": avr_loss} # , "lr": lr_scheduler.get_last_lr()[0]}
progress_bar.set_postfix(**logs)
@@ -1197,7 +1196,7 @@ class NetworkTrainer:
steps_done += 1
if args.logging_dir is not None:
logs = {"loss/epoch": loss_recorder.moving_average}
logs = {"loss/epoch": self.loss_recorder.moving_average}
accelerator.log(logs, step=epoch + 1)
return steps_done