update loss plots, reduce tok_reg value

This commit is contained in:
xander
2024-04-24 14:42:05 +02:00
parent 0a72748da6
commit e1b05b1b99
4 changed files with 27 additions and 19 deletions
+2 -1
View File
@@ -300,6 +300,7 @@ def train(
# Compute the loss:
loss = compute_diffusion_loss(config, model_pred, noise, noisy_latent, mask, noise_scheduler, timesteps)
losses['img_loss'].append(loss.item())
if config.l1_penalty > 0.0:
# Compute normalized L1 norm (mean of abs sum) of all lora parameters:
@@ -308,7 +309,7 @@ def train(
if optimizers['textual_inversion'] is not None and optimizers['textual_inversion'].param_groups[0]['lr'] > 0.0:
# Some custom regularization: # TODO test how much these actually help!!
loss, prompt_embeds_norms = condtioning_regularizer.apply_regularization(loss, prompt_embeds_norms, prompt_embeds, pipe)
loss, losses, prompt_embeds_norms = condtioning_regularizer.apply_regularization(loss, losses, prompt_embeds_norms, prompt_embeds, pipe)
losses['tot_loss'].append(loss.item())
loss = loss / config.gradient_accumulation_steps
+1 -1
View File
@@ -33,7 +33,7 @@ class TrainingConfig(BaseModel):
lora_weight_decay: float = 0.002
cond_reg_w: float = 0.0e-5
tok_cond_reg_w: float = 0.0e-5
tok_cov_reg_w: float = 0.005
tok_cov_reg_w: float = 0.001 # regularizes the token covariance matrix wrt pretrained "healthy" tokens
off_ratio_power: float = 0.01 # Pulls the std of the token distribution towards the target std
l1_penalty: float = 0.1 # Makes the unet lora matrix more sparse
noise_offset: float = 0.02 # Noise offset training to improve very dark / very bright images
+4 -3
View File
@@ -113,7 +113,7 @@ class ConditioningRegularizer:
self.distribution_regularizers[f'txt_encoder_{idx}'] = CovarianceLoss(pretrained_token_embeddings)
idx += 1
def apply_regularization(self, loss, prompt_embeds_norms, prompt_embeds, pipe=None):
def apply_regularization(self, loss, losses, prompt_embeds_norms, prompt_embeds, pipe=None):
noise_sigma = 0.0
if noise_sigma > 0.0: # experimental: apply random noise to the conditioning vectors as a form of regularization
prompt_embeds[0,1:-2,:] += torch.randn_like(prompt_embeds[0,2:-2,:]) * noise_sigma
@@ -131,10 +131,11 @@ class ConditioningRegularizer:
if self.config.tok_cov_reg_w > 0.0:
for key, distribution_regularizer in self.distribution_regularizers.items():
reg_loss = distribution_regularizer.compute_covariance_loss(self.embedding_handler.get_trainable_embeddings()[0][key])
print(f"Covariance Regularization loss for {key} token-embeds: {reg_loss.item():.5f}")
#print(f"Covariance Regularization loss for {key} token-embeds: {reg_loss.item():.5f}")
losses['covariance_tok_reg_loss'].append(reg_loss.item())
loss += self.config.tok_cov_reg_w * reg_loss
return loss, prompt_embeds_norms
return loss, losses, prompt_embeds_norms
def _compute_regularization_loss(self, prompt_embeds):
conditioning_norms = prompt_embeds.norm(dim=-1).mean(dim=0)
+20 -14
View File
@@ -196,26 +196,32 @@ def plot_token_stds(token_std_dict, save_path='token_stds.png', target_value_dic
from scipy.signal import savgol_filter
def plot_loss(loss_dict, save_path='losses.png', window_length=31, polyorder=3):
colormap = {'img_loss': 'blue', 'tot_loss': 'green', 'covariance_tok_reg_loss': 'orange'}
plot_smoothed = ['img_loss']
plt.figure()
for key in loss_dict.keys():
losses = loss_dict[key]
smoothed_losses = [0]
if key == 'tot_loss':
continue
losses = np.array(loss_dict[key])
if len(losses) < window_length:
continue
smoothed_losses = savgol_filter(losses, window_length, polyorder)
plt.plot(losses, label=key)
plt.plot(smoothed_losses, label=f'Smoothed {key}', color='red')
# plt.yscale('log') # Uncomment if log scale is desired
# Normalize losses to [0, 1]
normalized_losses = (losses - np.min(losses)) / (np.max(losses) - np.min(losses))
if key in plot_smoothed:
smoothed_losses = savgol_filter(normalized_losses, window_length, polyorder)
plt.plot(smoothed_losses, label=f'Smoothed {key}', color=colormap[key], linestyle='dashed')
else:
plt.plot(normalized_losses, label=key, color = colormap[key], linestyle='solid')
plt.xlabel('Step')
plt.ylabel('Training Loss')
plt.ylim(0, max(0.01, np.max(smoothed_losses)*1.4))
plt.ylabel('Training Losses')
plt.ylim(0, 1.1) # Adjust the y-axis limits for normalized data
plt.legend()
plt.savefig(save_path)
plt.close()
plt.close()