(beta, t)
| 58 | |
| 59 | |
| 60 | def compute_alpha(beta, t): |
| 61 | beta = torch.cat([torch.zeros(1).to(beta.device), beta], dim=0) |
| 62 | a = (1 - beta).cumprod(dim=0).index_select(0, t + 1).view(-1, 1, 1) |
| 63 | return a |
| 64 | |
| 65 | |
| 66 | def p_xt(xt, noise, t, next_t, beta, eta=0): |