update loss plots, reduce tok_reg value
This commit is contained in:
@@ -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
@@ -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
@@ -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
@@ -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()
|
||||
Reference in New Issue
Block a user