| 75 | |
| 76 | |
| 77 | def get_sample_reverse(attn_bank, inject_steps, single_layers, double_layers): |
| 78 | @torch.no_grad() |
| 79 | def sample_reverse(model, x, sigmas, extra_args=None, callback=None, disable=None): |
| 80 | if inject_steps > attn_bank["save_steps"]: |
| 81 | raise ValueError( |
| 82 | f'You must save at least as many steps as you want to inject. save_steps: {attn_bank["save_steps"]}, inject_steps: {inject_steps}' |
| 83 | ) |
| 84 | |
| 85 | extra_args = {} if extra_args is None else extra_args |
| 86 | |
| 87 | model_options = extra_args.get("model_options", {}) |
| 88 | model_options = {**model_options} |
| 89 | transformer_options = model_options.get("transformer_options", {}) |
| 90 | transformer_options = {**transformer_options} |
| 91 | model_options["transformer_options"] = transformer_options |
| 92 | extra_args["model_options"] = model_options |
| 93 | |
| 94 | N = len(sigmas) - 1 |
| 95 | s_in = x.new_ones([x.shape[0]]) |
| 96 | for i in trange(N, disable=disable): |
| 97 | sigma = sigmas[i] |
| 98 | sigma_prev = sigmas[i + 1] |
| 99 | |
| 100 | transformer_options["rfedit"] = { |
| 101 | "step": i, |
| 102 | "process": "reverse" if i < inject_steps else None, |
| 103 | "pred": "first", |
| 104 | "bank": attn_bank, |
| 105 | "single_layers": single_layers, |
| 106 | "double_layers": double_layers, |
| 107 | } |
| 108 | |
| 109 | pred = model(x, s_in * sigma, **extra_args) |
| 110 | |
| 111 | transformer_options["rfedit"] = { |
| 112 | "step": i, |
| 113 | "process": "reverse" if i < inject_steps else None, |
| 114 | "pred": "mid", |
| 115 | "bank": attn_bank, |
| 116 | "single_layers": single_layers, |
| 117 | "double_layers": double_layers, |
| 118 | } |
| 119 | |
| 120 | img_mid = x + (sigma_prev - sigma) / 2 * pred |
| 121 | sigma_mid = sigma + (sigma_prev - sigma) / 2 |
| 122 | pred_mid = model(img_mid, s_in * sigma_mid, **extra_args) |
| 123 | |
| 124 | first_order = (pred_mid - pred) / ((sigma_prev - sigma) / 2) |
| 125 | x = ( |
| 126 | x |
| 127 | + (sigma_prev - sigma) * pred |
| 128 | + 0.5 * (sigma_prev - sigma) ** 2 * first_order |
| 129 | ) |
| 130 | |
| 131 | if callback is not None: |
| 132 | callback( |
| 133 | { |
| 134 | "x": x, |