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.
This commit is contained in:
@@ -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.
|
||||
|
||||
+22
-20
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user