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

Function get_sample_forward

tricks/nodes/rectified_sampler_nodes.py:24–95  ·  view source on GitHub ↗
(
    gamma, start_step, end_step, gamma_trend, seed, attn_bank=None, order="first"
)

Source from the content-addressed store, hash-verified

22
23
24def 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

Callers 1

buildMethod · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected