Problem
In VarixLoss, the method forward creates an effective_beta that scales the variational loss (vae_loss) and stores it.
effective_beta = self.config.beta * anneal_factor
total_loss = recon_loss + effective_beta * var_loss
return total_loss, {
"recon_loss": recon_loss,
"var_loss": var_loss * effective_beta,
"anneal_factor": torch.tensor(anneal_factor),
"effective_beta_factor": torch.tensor(effective_beta),
}
However, within BaseVisualizer, the static method _make_loss_format() loads the stored loss values and multiplies the variational loss again with beta for plotting
# Now create the DataFrame
loss_df = pd.DataFrame.from_dict(loss_values, orient="index") # type: ignore
# Rest of your code remains the same
if term == "var_loss":
loss_df = loss_df * config.beta
loss_df["Epoch"] = loss_df.index + 1
loss_df["Loss Term"] = term
Effectively, the vae_loss values on the loss curves plot are vae_loss = vae_loss * config.beta * anneal_factor * config.beta.
Solution
Simply delete this part
if term == "var_loss":
loss_df = loss_df * config.beta
in BaseVisualizer._make_loss_format.
Problem
In VarixLoss, the method
forwardcreates aneffective_betathat scales the variational loss (vae_loss) and stores it.However, within BaseVisualizer, the static method
_make_loss_format()loads the stored loss values and multiplies the variational loss again with beta for plottingEffectively, the vae_loss values on the loss curves plot are
vae_loss = vae_loss * config.beta * anneal_factor * config.beta.Solution
Simply delete this part
in BaseVisualizer._make_loss_format.