MCPcopy Create free account
hub / github.com/Hzfinfdu/Diffusion-BERT / qt_reverse

Method qt_reverse

diffusion_condition.py:434–464  ·  view source on GitHub ↗

Get q(x_{t+1} | x_t), the one-step posterior efficiently. Args: qt_plus_1: an array of floats specifying a distribution over p(x_0). t: t in q(x_{t+1} | x_t). return_logits: if True, return the output logits make_one_hot: if True, will convert q0 to fl

(self,
                   qt_plus_1,
                   t,
                   return_logits=False,
                   make_one_hot=False,
                   epsilon=1e-20
                   )

Source from the content-addressed store, hash-verified

432 return (1 - beta) * torch.eye(self.dim) + beta * self._get_mask()
433
434 def qt_reverse(self,
435 qt_plus_1,
436 t,
437 return_logits=False,
438 make_one_hot=False,
439 epsilon=1e-20
440 ):
441 """Get q(x_{t+1} | x_t), the one-step posterior efficiently.
442 Args:
443 qt_plus_1: an array of floats specifying a distribution over p(x_0).
444 t: t in q(x_{t+1} | x_t).
445 return_logits: if True, return the output logits
446 make_one_hot: if True, will convert q0 to floats if needed.
447 epsilon: a small number to normalize logits conversion with, if needed.
448 Returns:
449 q(x_{t+1} | x_t).
450 """
451 if make_one_hot:
452 assert qt_plus_1.dtype == torch.int64
453 qt_plus_1 = torch.nn.functional.one_hot(qt_plus_1, num_classes=self.dim)
454
455 beta = self.schedule(t)
456 qtpls1_at_mask = qt_plus_1[Ellipsis, self.tokenizer.mask_token_id: self.tokenizer.mask_token_id + 1]
457 non_mask_prob0 = (1 - beta) * qt_plus_1[Ellipsis, :self.tokenizer.mask_token_id] + beta * qtpls1_at_mask
458 non_mask_prob1 = (1 - beta) * qt_plus_1[Ellipsis, self.tokenizer.mask_token_id + 1:] + beta * qtpls1_at_mask
459 prob_at_time_t = torch.cat((non_mask_prob0, qtpls1_at_mask, non_mask_prob1), dim=-1)
460
461 if return_logits:
462 return torch.log(prob_at_time_t + epsilon)
463 else:
464 return prob_at_time_t
465
466 def get_qt_given_q0(self,
467 q0,

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected