| 1567 | |
| 1568 | @torch.no_grad() |
| 1569 | def conditional_sample(self, data_batch, horizon=None, num_samp=1, class_free_guide_w=0.0, **kwargs): |
| 1570 | batch_size, num_agents = data_batch['history_positions'].size()[:2] |
| 1571 | horizon = horizon or self.horizon |
| 1572 | shape = (batch_size, num_samp, num_agents, horizon, self.transition_dim) |
| 1573 | |
| 1574 | return self.p_sample_loop(shape, data_batch, num_samp, class_free_guide_w=class_free_guide_w, **kwargs) |
| 1575 | |
| 1576 | #------------------------------------------ training ------------------------------------------# |
| 1577 | |