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

Method conditional_sample

tbsim/models/scenediffuser.py:1569–1574  ·  view source on GitHub ↗
(self, data_batch, horizon=None, num_samp=1, class_free_guide_w=0.0, **kwargs)

Source from the content-addressed store, hash-verified

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

Callers 1

forwardMethod · 0.95

Calls 1

p_sample_loopMethod · 0.95

Tested by

no test coverage detected