From 55d26c10c044c88b15e13defc4fbf2c0c58f5153 Mon Sep 17 00:00:00 2001 From: asagi4 <130366179+asagi4@users.noreply.github.com> Date: Sun, 8 Dec 2024 14:12:52 +0200 Subject: [PATCH] Actually fix moving the model between devices. The model produces different output on different devices for some reason, but this matches the behaviour of the original code. --- README.md | 13 +++++++++---- __init__.py | 42 ++++++++++++++++++++++-------------------- 2 files changed, 31 insertions(+), 24 deletions(-) diff --git a/README.md b/README.md index 869e4d1..7ef6199 100644 --- a/README.md +++ b/README.md @@ -1,11 +1,16 @@ -ComfyUI NPNet +# ComfyUI NPNet (Golden Noise) -A very barebones copypaste implementation of https://github.com/xie-lab-ml/Golden-Noise-for-Diffusion-Models +A very barebones mostly-copypaste implementation of https://github.com/xie-lab-ml/Golden-Noise-for-Diffusion-Models +## Requirements +You need the pre-trained weights for your model from https://drive.google.com/drive/folders/1Z0wg4HADhpgrztyT3eWijPbJJN5Y2jQt?usp=drive_link + +## Usage Use with custom sampling and pass in an initial noise from eg. `RandomNoise` and a cond (only the first prompt in the conditioning will be used if multiple exist). -Note: *only* works with 128x128 square noise, apparently. If you pass in other shaped latents, it will reshape the noise into a square before running the noise model, and then reshape the result back to the original resolution. - You can also run it on the CPU, though appears to change the output for some reason. +## Notes +The model works with 128x128 latents, apparently. If you pass in other shaped latents, it will reshape the noise into a square before running the noise model, and then reshape the result back to the original resolution. + If you get an error from the timm module when running this, update your timm package. It may be too old. diff --git a/__init__.py b/__init__.py index 85d1e9d..41bf4fd 100644 --- a/__init__.py +++ b/__init__.py @@ -137,9 +137,8 @@ class NPNet(nn.Module): super().__init__() assert model_id in ["SDXL", "DreamShaper", "DiT"] - - self.model_id = model_id self.device = device + self.model_id = model_id self.pretrained_path = pretrained_path self.unet_embedding = NoiseTransformer(resolution=128) self.unet_svd = SVDNoiseUnet(resolution=128) @@ -147,17 +146,25 @@ class NPNet(nn.Module): self.text_embedding = AdaGroupNorm(1024 * 77, 4, 1, eps=1e-6) else: self.text_embedding = AdaGroupNorm(2048 * 77, 4, 1, eps=1e-6) - sd = torch.load(self.pretrained_path, weights_only=True, map_location=self.device) - self.unet_embedding.load_state_dict(sd['unet_embedding']) - self.unet_svd.load_state_dict(sd['unet_svd']) - self.text_embedding.load_state_dict(sd['embeeding']) - self.alpha = sd['alpha'] - self.beta = sd['beta'] - self.to(torch.float32).to(self.device) + sd = torch.load(self.pretrained_path, weights_only=True, map_location=device) + self.unet_embedding.load_state_dict(sd.pop("unet_embedding")) + self.unet_svd.load_state_dict(sd.pop("unet_svd")) + self.text_embedding.load_state_dict(sd.pop("embeeding")) + self.alpha = sd["alpha"] + self.beta = sd["beta"] + self.to(dtype=torch.float32, device=device) + def to(self, *args, **kwargs): + super().to(*args, **kwargs) + self.unet_embedding.to(*args, **kwargs) + self.unet_svd.to(*args, **kwargs) + self.text_embedding.to(*args, **kwargs) + self.alpha = self.alpha.to(*args, **kwargs) + self.beta = self.beta.to(*args, **kwargs) + self.device = self.alpha.device + return self def forward(self, initial_noise, prompt_embeds): - prompt_embeds = prompt_embeds.float().view(prompt_embeds.shape[0], -1) text_emb = self.text_embedding(initial_noise.float(), prompt_embeds) @@ -189,7 +196,7 @@ class NPNetGoldenNoise: "prompt": ("CONDITIONING",), "model_path": ("STRING", {"default": "/path/to/sdxl.pth"}), "model_type": (["SDXL", "DreamShaper", "DiT"],), - "device": (["cuda", "cpu"],) + "device": (["cuda", "cpu"],), } } @@ -210,15 +217,10 @@ class NPNetGoldenNoise: if cond.shape[1] != 77: print("NPNet can't handle conds >77 tokens, truncating...") cond = cond[:, :77, :] - try: - print("Applying NPNet to noise") - r = self.npnet(init_noise, cond).to("cpu") - if orig_shape[-2] != 128 or orig_shape[-1] != 128: - r = common_upscale(r, orig_shape[-1], orig_shape[-2], "nearest-exact", "disabled") - except Exception as e: - print("Running NPNet failed with error, returning unmodified noise:", e) - return init_noise - print("NPNet ran ok") + print("Applying NPNet to noise") + r = self.npnet(init_noise, cond).to("cpu") + if orig_shape[-2] != 128 or orig_shape[-1] != 128: + r = common_upscale(r, orig_shape[-1], orig_shape[-2], "nearest-exact", "disabled") return r def doit(self, noise, prompt, model_path, model_type, device):