| 22 | |
| 23 | |
| 24 | def get_sample_forward( |
| 25 | gamma, start_step, end_step, gamma_trend, seed, attn_bank=None, order="first" |
| 26 | ): |
| 27 | # Controlled Forward ODE (Algorithm 1) |
| 28 | generator = torch.Generator() |
| 29 | generator.manual_seed(seed) |
| 30 | |
| 31 | @torch.no_grad() |
| 32 | def sample_forward(model, y0, sigmas, extra_args=None, callback=None, disable=None): |
| 33 | if attn_bank is not None: |
| 34 | for block_idx in attn_bank["block_map"]: |
| 35 | attn_bank["block_map"][block_idx].clear() |
| 36 | |
| 37 | extra_args = {} if extra_args is None else extra_args |
| 38 | model_options = extra_args.get("model_options", {}) |
| 39 | model_options = {**model_options} |
| 40 | transformer_options = model_options.get("transformer_options", {}) |
| 41 | transformer_options = { |
| 42 | **transformer_options, |
| 43 | "total_steps": len(sigmas) - 1, |
| 44 | "sample_mode": "forward", |
| 45 | "attn_bank": attn_bank, |
| 46 | } |
| 47 | model_options["transformer_options"] = transformer_options |
| 48 | extra_args["model_options"] = model_options |
| 49 | |
| 50 | Y = y0.clone() |
| 51 | y1 = torch.randn(Y.shape, generator=generator).to(y0.device) |
| 52 | N = len(sigmas) - 1 |
| 53 | s_in = y0.new_ones([y0.shape[0]]) |
| 54 | gamma_values = generate_trend_values( |
| 55 | N, start_step, end_step, gamma, gamma_trend |
| 56 | ) |
| 57 | for i in trange(N, disable=disable): |
| 58 | transformer_options["step"] = i |
| 59 | sigma = sigmas[i] |
| 60 | sigma_next = sigmas[i + 1] |
| 61 | t_i = model.inner_model.inner_model.model_sampling.timestep(sigmas[i]) |
| 62 | |
| 63 | conditional_vector_field = (y1 - Y) / (1 - t_i) |
| 64 | |
| 65 | transformer_options["pred_order"] = "first" |
| 66 | pred = model( |
| 67 | Y, s_in * sigmas[i], **extra_args |
| 68 | ) # this implementation takes sigma instead of timestep |
| 69 | |
| 70 | if order == "second": |
| 71 | transformer_options["pred_order"] = "second" |
| 72 | img_mid = Y + (sigma_next - sigma) / 2 * pred |
| 73 | sigma_mid = sigma + (sigma_next - sigma) / 2 |
| 74 | pred_mid = model(img_mid, s_in * sigma_mid, **extra_args) |
| 75 | |
| 76 | first_order = (pred_mid - pred) / ((sigma_next - sigma) / 2) |
| 77 | pred = pred + gamma_values[i] * (conditional_vector_field - pred) |
| 78 | # first_order = first_order + gamma_values[i] * (conditional_vector_field - first_order) |
| 79 | Y = ( |
| 80 | Y |
| 81 | + (sigma_next - sigma) * pred |