fix loss visualizer, and other small stuff
This commit is contained in:
+2
-1
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user