Update nodes_rf_inversion.py

This commit is contained in:
kijai
2024-12-11 19:24:04 +02:00
parent 5b349848d4
commit c1500cb04e
+6 -3
View File
@@ -20,7 +20,7 @@ def generate_eta_values(
):
assert start_step < end_step and start_step >= 0 and end_step <= len(timesteps), "Invalid start_step and end_step"
# timesteps are monotonically decreasing, from 1.0 to 0.0
print("eta timesteps", timesteps)
eta_values = [0.0] * (len(timesteps) - 1)
if eta_trend == 'constant':
@@ -36,7 +36,7 @@ def generate_eta_values(
eta_values[i] = eta * (timesteps[i] - timesteps[end_step - 1]) / total_time
else:
raise NotImplementedError(f"Unsupported eta_trend: {eta_trend}")
print("eta_values", eta_values)
return eta_values
class HyVideoEmptyTextEmbeds:
@@ -160,7 +160,7 @@ class HyVideoInverseSampler:
int(height) // pipeline.vae_scale_factor,
int(width) // pipeline.vae_scale_factor,
)
noise = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
noise = randn_tensor(shape, generator=generator, device=device, dtype=torch.float32)
frames_needed = noise.shape[1]
current_frames = latents.shape[1]
@@ -221,10 +221,13 @@ class HyVideoInverseSampler:
)["x"]
sigma = t / 1000.0
sigma_prev = t_prev / 1000.0
latents = latents.to(torch.float32)
noise_pred = noise_pred.to(torch.float32)
target_noise_velocity = (noise - latents) / (1.0 - sigma)
interpolated_velocity = gamma * target_noise_velocity + (1 - gamma) * noise_pred
latents = latents + (sigma_prev - sigma) * interpolated_velocity
latents = latents.to(torch.bfloat16)
# compute the previous noisy sample x_t -> x_t-1
#latents = pipeline.scheduler.step(noise_pred, t, latents, return_dict=False)[0]