x_0: joint gaussian noise X_T x_1: twist gaussian noise X_T
(x_0, x_1, seq, model, b, eta, ctx=None, gen_multi=False)
| 64 | |
| 65 | |
| 66 | def generalized_steps(x_0, x_1, seq, model, b, eta, ctx=None, gen_multi=False): |
| 67 | ''' |
| 68 | x_0: joint gaussian noise X_T |
| 69 | x_1: twist gaussian noise X_T |
| 70 | ''' |
| 71 | with torch.no_grad(): |
| 72 | n = x_0.size(0) |
| 73 | seq_next = [-1] + list(seq[:-1]) |
| 74 | x0_preds = [[],[]] |
| 75 | xs = [[x_0],[x_1]] |
| 76 | for i, j in zip(reversed(seq), reversed(seq_next)): |
| 77 | t = (torch.ones(n) * i).to(x_0.device) |
| 78 | next_t = (torch.ones(n) * j).to(x_0.device) |
| 79 | at = compute_alpha(b, t.long()) |
| 80 | at_next = compute_alpha(b, next_t.long()) |
| 81 | xt_0 = xs[0][-1].to(x_0.device) |
| 82 | xt_1 = xs[1][-1].to(x_0.device) |
| 83 | |
| 84 | et = model(xinj=xt_0,xint=xt_1,t=t.float(),ctx =ctx, gen_multi= gen_multi) # estimated noise for current timestep |
| 85 | x0_t_0 = (xt_0 - et[0] * (1 - at).sqrt()) / at.sqrt() |
| 86 | x0_t_1 = (xt_1 - et[1] * (1 - at).sqrt()) / at.sqrt() |
| 87 | x0_preds[0].append(x0_t_0) # estimated x_0 of current timestep |
| 88 | x0_preds[1].append(x0_t_1) |
| 89 | c1 = ( # signma_t |
| 90 | eta * ((1 - at / at_next) * (1 - at_next) / (1 - at)).sqrt() |
| 91 | ) |
| 92 | c2 = ((1 - at_next) - c1 ** 2).sqrt() |
| 93 | xt_next_0 = at_next.sqrt() * x0_t_0 + c1 * torch.randn_like(x_0) + c2 * et[0] |
| 94 | xt_next_1 = at_next.sqrt() * x0_t_1 + c1 * torch.randn_like(x_1) + c2 * et[1] |
| 95 | xs[0].append(xt_next_0) |
| 96 | xs[1].append(xt_next_1) |
| 97 | return xs, x0_preds |
no test coverage detected