diff --git a/__init__.py b/__init__.py index 0022eb4..14b5238 100644 --- a/__init__.py +++ b/__init__.py @@ -17,10 +17,13 @@ def common_ksampler(model, seed, steps, cfg, sampler_name, scheduler, positive, noise_mask = None device = model_management.get_torch_device() - if disable_noise: - noise = torch.zeros(latent_image.size(), dtype=latent_image.dtype, layout=latent_image.layout, device="cpu") + if "noise_sequence" in latent: + noise = latent["noise_sequence"] else: - noise = torch.randn(latent_image.size(), dtype=latent_image.dtype, layout=latent_image.layout, generator=torch.manual_seed(seed), device="cpu") + if disable_noise: + noise = torch.zeros(latent_image.size(), dtype=latent_image.dtype, layout=latent_image.layout, device="cpu") + else: + noise = torch.randn(latent_image.size(), dtype=latent_image.dtype, layout=latent_image.layout, generator=torch.manual_seed(seed), device="cpu") if "noise_mask_sequence" in latent: noise_mask_list = [] @@ -280,7 +283,9 @@ class DdimInversionSequence: model_management.load_model_gpu(model) context = context.to(device) samples = samples.to(device) - s = ddim_inversion(model, ddim_scheduler, samples, steps, context) + s = ddim_inversion(model, ddim_scheduler, samples, steps, context)[-1] + s = rearrange(s.squeeze(0), "c f h w -> f c h w") + s = s.cpu() return (s,) diff --git a/sd.py b/sd.py index 1151fc5..4b7e64f 100644 --- a/sd.py +++ b/sd.py @@ -103,7 +103,8 @@ def load_checkpoint_guess_config(ckpt_path, output_vae=True, output_clip=True, e model = instantiate_from_config(model_config) model = load_model_weights(model, sd, verbose=False, load_state_dict_to=load_state_dict_to) model.model.diffusion_model = convert_unet_checkpoint(sd, OmegaConf.create({"model": model_config})) + if fp16: - model = model.half() + model = model.half() return (ModelPatcher(model), clip, vae) diff --git a/tuneavideo/models/unet.py b/tuneavideo/models/unet.py index d1a216d..20f02b0 100644 --- a/tuneavideo/models/unet.py +++ b/tuneavideo/models/unet.py @@ -297,6 +297,9 @@ class UNet3DConditionModel(ModelMixin, ConfigMixin): sample = rearrange(x.unsqueeze(0), "b f c h w -> b c f h w") + sample = sample.type(self.dtype) + context = context.type(self.dtype) + down_block_additional_residuals = None mid_block_additional_residual = None diff --git a/tuneavideo/util.py b/tuneavideo/util.py index 4a1aeae..ba8df73 100644 --- a/tuneavideo/util.py +++ b/tuneavideo/util.py @@ -22,6 +22,7 @@ def next_step(model_output: Union[torch.FloatTensor, np.ndarray], timestep: int, def get_noise_pred_single(latents, t, context, unet): latents = rearrange(latents.squeeze(0), "c f h w -> f c h w") noise_pred = unet(latents, t.view(1), context=context) + noise_pred = rearrange(noise_pred.unsqueeze(0), "b f c h w -> b c f h w") return noise_pred