Make device configurable

This commit is contained in:
asagi4
2024-12-08 01:15:06 +02:00
parent 67ad24bf4b
commit 747ad6e21b
2 changed files with 11 additions and 6 deletions
+6 -3
View File
@@ -2,7 +2,10 @@ ComfyUI NPNet
A very barebones copypaste implementation of https://github.com/xie-lab-ml/Golden-Noise-for-Diffusion-Models
Note: *only* works with 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.
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.
If you get an error from the timm module when running this, update your timm package. It may be too old.
+5 -3
View File
@@ -133,7 +133,7 @@ class NoiseTransformer(nn.Module):
class NPNet(nn.Module):
def __init__(self, model_id, pretrained_path=True, device="cuda") -> None:
def __init__(self, model_id, pretrained_path, device="cuda") -> None:
super().__init__()
assert model_id in ["SDXL", "DreamShaper", "DiT"]
@@ -189,6 +189,7 @@ class NPNetGoldenNoise:
"prompt": ("CONDITIONING",),
"model_path": ("STRING", {"default": "/path/to/sdxl.pth"}),
"model_type": (["SDXL", "DreamShaper", "DiT"],),
"device": (["cuda", "cpu"],)
}
}
@@ -220,10 +221,11 @@ class NPNetGoldenNoise:
print("NPNet ran ok")
return r
def doit(self, noise, prompt, model_path, model_type):
def doit(self, noise, prompt, model_path, model_type, device):
if self.npnet is None:
print("Loading NPNet")
self.npnet = NPNet(model_type, model_path)
self.npnet = NPNet(model_type, model_path, device=device)
self.npnet.to(device)
self.noise = noise
self.cond = prompt[0]