MCPcopy Create free account
hub / github.com/NVlabs/CTG / p_sample_loop

Method p_sample_loop

tbsim/models/scenediffuser.py:1490–1565  ·  view source on GitHub ↗

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')

Source from the content-addressed store, hash-verified

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 ======')

Callers 1

conditional_sampleMethod · 0.95

Calls 9

p_sampleMethod · 0.95
descale_trajMethod · 0.95
scale_trajMethod · 0.95
ProgressClass · 0.90
SilentClass · 0.90
apply_constraintsFunction · 0.90
closeMethod · 0.80
updateMethod · 0.45

Tested by

no test coverage detected