Rebased on main (by filipstrand)

This commit is contained in:
Alessandro 2025-01-27 23:12:17 +01:00 committed by filipstrand
parent cb93e1aaf6
commit 401da9fd75
3 changed files with 123 additions and 22 deletions

View File

@ -70,12 +70,13 @@ class LoRALayers:
blocks: list[JointTransformerBlock] | list[SingleTransformerBlock],
block_prefix: str,
) -> dict:
start = block_spec.block_range.start
end = block_spec.block_range.end
block_indices = block_spec.block_range.get_blocks()
lora_layers = {}
for i in range(start, end):
block = blocks[i]
for idx in block_indices:
if idx >= len(blocks):
raise IndexError(f"Indice {idx} over range")
block = blocks[idx]
for layer_type in block_spec.layer_types:
original_layer = LoRALayers._get_nested_attr(block, layer_type)
is_list = isinstance(original_layer, list)
@ -84,7 +85,7 @@ class LoRALayers:
linear=original_layer[0] if is_list else original_layer,
r=block_spec.lora_rank,
)
layer_path = f"{block_prefix}.{i}.{layer_type}"
layer_path = f"{block_prefix}.{idx}.{layer_type}"
lora_layers[layer_path] = [lora_layer] if is_list else lora_layer

View File

@ -3,7 +3,7 @@ import json
import os
from dataclasses import asdict, dataclass
from pathlib import Path
from typing import List
from typing import List, Optional
from mflux.dreambooth.state.zip_util import ZipUtil
@ -56,8 +56,16 @@ class StatisticsSpec:
@dataclass
class BlockRange:
start: int
end: int
start: Optional[int] = None
end: Optional[int] = None
indices: Optional[List[int]] = None
def get_blocks(self) -> List[int]:
if self.indices:
return self.indices
if self.start is not None and self.end is not None:
return list(range(self.start, self.end - 1))
raise ValueError("Either start and end or indices must be provided.")
@dataclass

View File

@ -1,3 +1,6 @@
import time
from datetime import timedelta
import matplotlib.pyplot as plt
from mflux.dreambooth.state.training_spec import TrainingSpec
@ -5,8 +8,10 @@ from mflux.dreambooth.state.training_state import TrainingState
class Plotter:
start_time = time.time() # Class variable to track time
@staticmethod
def update_loss_plot(training_spec: TrainingSpec, training_state: TrainingState) -> None:
def update_loss_plot(training_spec: TrainingSpec, training_state: TrainingState, target_loss: float = 0.3) -> None:
plt.style.use("bmh")
# Create figure with 16:9 aspect ratio
@ -17,37 +22,124 @@ class Plotter:
plt.plot(stats.steps, stats.losses, "b-", linewidth=2, label="Validation Loss", zorder=1) # Line
plt.plot(stats.steps, stats.losses, "bo", markersize=6, label="_nolegend_", zorder=2) # Points
# Customize the plot
# Find significant points - convert to Python float
max_loss = float(max(stats.losses))
max_step = int(stats.steps[stats.losses.index(max_loss)])
min_loss = float(min(stats.losses))
min_step = int(stats.steps[stats.losses.index(min_loss)])
# Function to position annotations avoiding overlaps
def get_annotation_position(x, y, is_max=True, y_limits=None):
vertical_offset = 0.10
horizontal_offset = 0.10 # Horizontal offset to move to the left
if y_limits is None:
y_limits = plt.ylim()
if is_max:
if y + vertical_offset > y_limits[1]:
return float(x - horizontal_offset), float(y)
return float(x), float(y + vertical_offset)
else:
if y - vertical_offset < y_limits[0]:
return float(x - horizontal_offset), float(y)
return float(x), float(y - vertical_offset)
# Get the limits of the y-axis
y_limits = (0, 1.5) if max_loss <= 1.5 else (0, max_loss + max_loss * 0.2)
# Annotate for maximum and minimum loss
pos = get_annotation_position(max_step, max_loss, is_max=True, y_limits=y_limits)
plt.annotate(
f"Max: {max_loss:.3f}\nStep: {max_step}",
xy=(max_step, max_loss),
xytext=pos,
bbox=dict(facecolor="red", alpha=0.5),
ha="center",
arrowprops=dict(arrowstyle="->"),
zorder=5,
)
plt.plot(max_step, max_loss, "ro", markersize=8, zorder=4)
pos = get_annotation_position(min_step, min_loss, is_max=False, y_limits=y_limits)
plt.annotate(
f"Min: {min_loss:.3f}\nStep: {min_step}",
xy=(min_step, min_loss),
xytext=pos,
bbox=dict(facecolor="green", alpha=0.5),
ha="center",
arrowprops=dict(arrowstyle="->"),
zorder=5,
)
plt.plot(min_step, min_loss, "go", markersize=8, zorder=4)
# Line for target loss
plt.axhline(y=float(target_loss), color="red", linestyle="--", linewidth=0.5, alpha=0.7)
# Calculate elapsed time
elapsed_time = time.time() - Plotter.start_time
elapsed_str = timedelta(seconds=int(elapsed_time))
# Initialize counter and sum of losses
loss_counter = 0
loss_sum = 0
total_step = int(training_state.iterator.total_number_of_steps())
for loss in stats.losses:
loss_counter += 1
loss_sum += loss
if loss_counter > 0:
avg_loss = loss_sum / loss_counter
legend_text = [
f"Elapsed Time: {elapsed_str}",
f"Img Dim {training_spec.width}x{training_spec.height}",
f"Total steps: {int(max(stats.steps))} / {total_step}",
f"Last Loss: {float(stats.losses[-1]):.4f}",
f"Lower Loss: {min_loss:.4f} (Step {min_step})",
f"Higher loss: {max_loss:.4f} (Step {max_step})",
f"Avg Loss: {avg_loss:.2f}",
f"Ideal Goal (lower of): {target_loss}",
]
plt.text(
0.98,
0.98,
"\n".join(legend_text),
transform=plt.gca().transAxes,
bbox=dict(facecolor="white", alpha=0.8, edgecolor="gray"),
verticalalignment="top",
horizontalalignment="right",
fontsize=8,
)
plt.title("Validation Loss Over Time", fontsize=16, pad=20)
plt.xlabel("Steps", fontsize=12)
plt.ylabel("Loss", fontsize=12)
# Set integer grid
plt.grid(True, linestyle="--", alpha=0.7)
ax = plt.gca()
ax.xaxis.set_major_locator(plt.MaxNLocator(integer=True))
# Dynamic x-axis limit with 20% padding
plt.subplots_adjust(right=0.85)
# Dinamiyc padding for x axes
max_x = max(stats.steps)
padding = max_x * 0.2
initial_padding = 0.4
final_padding = 0.01
# formula for padding and centering graphos, it will progressivaly tent to final_padding.
padding_limit = initial_padding - (initial_padding - final_padding) * (max_x / total_step)
padding = max_x * padding_limit
plt.xlim(0, max_x + padding)
# Dynamic y-axis limit with 20% padding
max_y = float(max(stats.losses))
padding = max_y * 0.2
plt.ylim(0, max_y + padding)
plt.ylim(y_limits)
plt.legend(fontsize=12)
# Add margins for better visibility
plt.margins(x=0.02)
# Tight layout to prevent label cutoff
plt.tight_layout()
# Save to desktop with high PPI
path = training_state.get_current_loss_plot_path(training_spec)
plt.savefig(path, format="pdf", dpi=300, bbox_inches="tight")
# Close the figure to free memory
plt.close()