MCPcopy Create free account
hub / github.com/Lightricks/ComfyUI-LTXVideo / get_sample_reverse

Function get_sample_reverse

tricks/nodes/rf_edit_sampler_nodes.py:77–144  ·  view source on GitHub ↗
(attn_bank, inject_steps, single_layers, double_layers)

Source from the content-addressed store, hash-verified

75
76
77def 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,

Callers 1

buildMethod · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected