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
)
| 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, |
nothing calls this directly
no outgoing calls
no test coverage detected