shape: (5), batch_size, num_samp, num_agents, horizon, self.transition_dim
(self, shape, data_batch, num_samp,
aux_info={},
verbose=True,
return_diffusion=False,
return_guidance_losses=False,
class_free_guide_w=0.0,
apply_guidance=True,
guide_clean=False,
mode='testing')
| 1488 | |
| 1489 | @torch.no_grad() |
| 1490 | def p_sample_loop(self, shape, data_batch, num_samp, |
| 1491 | aux_info={}, |
| 1492 | verbose=True, |
| 1493 | return_diffusion=False, |
| 1494 | return_guidance_losses=False, |
| 1495 | class_free_guide_w=0.0, |
| 1496 | apply_guidance=True, |
| 1497 | guide_clean=False, |
| 1498 | mode='testing'): |
| 1499 | ''' |
| 1500 | shape: (5), batch_size, num_samp, num_agents, horizon, self.transition_dim |
| 1501 | ''' |
| 1502 | # merge B and M to be compatible with guidance loss developed for agent-centric models |
| 1503 | data_batch_for_guidance = {} |
| 1504 | if apply_guidance: |
| 1505 | data_batch_for_guidance = extract_data_batch_for_guidance(data_batch, mode=mode) |
| 1506 | |
| 1507 | device = self.betas.device |
| 1508 | batch_size = shape[0] |
| 1509 | if self.current_perturbation_guidance.current_guidance is not None and not apply_guidance: |
| 1510 | print('DIFFUSER: Note, not using guidance during sampling, only evaluating guidance loss at very end...') |
| 1511 | |
| 1512 | # sample from base distribution |
| 1513 | x = torch.randn(shape, device=device) # (B, N, M, T(+T_hist), transition_dim) |
| 1514 | x = TensorUtils.join_dimensions(x, begin_axis=0, end_axis=2) # (B*N, M, T(+T_hist), transition_dim) |
| 1515 | x_clean_model_out = None |
| 1516 | |
| 1517 | # (B, M, C) -> (B*N, M, C) |
| 1518 | aux_info = TensorUtils.repeat_by_expand_at(aux_info, repeats=num_samp, dim=0) |
| 1519 | if return_diffusion: diffusion = [x] #(1, B*N, M, T, transition_dim) |
| 1520 | progress = Progress(self.n_timesteps) if verbose else Silent() |
| 1521 | |
| 1522 | steps = [i for i in reversed(range(0, self.n_timesteps, self.stride))] |
| 1523 | attn_weights = [] |
| 1524 | for i in steps: |
| 1525 | # (B*N) |
| 1526 | timesteps = torch.full((batch_size*num_samp,), i, device=device, dtype=torch.long) |
| 1527 | x, guide_losses, x_clean_model_out, info = self.p_sample(x, timesteps, data_batch, aux_info=aux_info, num_samp=num_samp, class_free_guide_w=class_free_guide_w, |
| 1528 | apply_guidance=apply_guidance, guide_clean=guide_clean, eval_final_guide_loss=(i == steps[-1]), x_clean_model_out=x_clean_model_out, data_batch_for_guidance=data_batch_for_guidance) |
| 1529 | # apply hard constraints (overwrite waypoints at certain timesteps) |
| 1530 | if self.current_constraints is not None: # and i != steps[-1]: # TODO don't do it for last step? |
| 1531 | # TODO why isn't this working very well? And why is y upside down? |
| 1532 | # apply constraints expects traj in shape (B, N, T, D) and metric space |
| 1533 | x = self.descale_traj(x.reshape((shape[0], shape[1], shape[2], -1))) |
| 1534 | x = apply_constraints(x, data_batch['scene_index'], self.current_constraints) |
| 1535 | x = self.scale_traj(x.reshape((shape[0]*shape[1], shape[2], -1))) |
| 1536 | |
| 1537 | |
| 1538 | progress.update({'t': i}) |
| 1539 | |
| 1540 | if return_diffusion: diffusion.append(x) |
| 1541 | if i == 0: |
| 1542 | attn_weights = info['attn_weights'] |
| 1543 | |
| 1544 | progress.close() |
| 1545 | |
| 1546 | if any(guide_losses): |
| 1547 | print('===== GUIDANCE LOSSES ======') |
no test coverage detected