Function
loss_fn
(inputs,outputs,loss_fn,z_mean,z_log_var,num_features = 784)
Source from the content-addressed store, hash-verified
| 99 | return z |
| 100 | |
| 101 | def loss_fn(inputs,outputs,loss_fn,z_mean,z_log_var,num_features = 784): |
| 102 | reconstruction_loss = loss_fn(outputs,inputs) |
| 103 | reconstruction_loss = reconstruction_loss * num_features |
| 104 | |
| 105 | #计算KL散度损失值 |
| 106 | kl_loss = 1 + z_log_var - torch.square(z_mean) - torch.exp(z_log_var) |
| 107 | kl_loss = -0.5 * torch.sum(kl_loss,dim = -1) |
| 108 | vae_loss = torch.mean(reconstruction_loss + kl_loss) |
| 109 | |
| 110 | return vae_loss |
Tested by
no test coverage detected