(self, x, data_batch, aux_info={})
| 1056 | return loss, info |
| 1057 | |
| 1058 | def loss(self, x, data_batch, aux_info={}): |
| 1059 | batch_size = len(x) |
| 1060 | |
| 1061 | t = torch.randint(0, self.n_timesteps, (batch_size,), device=x.device).long() |
| 1062 | |
| 1063 | return self.p_losses(x, t, data_batch, aux_info=aux_info) |
| 1064 | |
| 1065 | def unravel_index(index, shape): |
| 1066 | out = [] |
no test coverage detected