Add safetensors support and fix model switching

This commit is contained in:
asagi4
2024-12-08 15:16:49 +02:00
parent 55d26c10c0
commit ca6864c91d
3 changed files with 53 additions and 13 deletions
+2
View File
@@ -14,3 +14,5 @@ You can also run it on the CPU, though appears to change the output for some rea
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.
You can use `convert_to_safetensors.py` to convert the pre-trained models into safetensors files (with fixed keys)
+16 -13
View File
@@ -1,4 +1,5 @@
import torch
import safetensors.torch
import torch.nn as nn
from torch.nn import functional as F
from timm import create_model
@@ -142,25 +143,27 @@ class NPNet(nn.Module):
self.pretrained_path = pretrained_path
self.unet_embedding = NoiseTransformer(resolution=128)
self.unet_svd = SVDNoiseUnet(resolution=128)
self.alpha = torch.nn.Parameter(torch.empty(1))
self.beta = torch.nn.Parameter(torch.empty(1))
if self.model_id == "DiT":
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=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"]
if ".pth" in pretrained_path:
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 = torch.nn.Parameter(sd["alpha"])
self.beta = torch.nn.Parameter(sd["beta"])
else:
sd = safetensors.torch.load_file(self.pretrained_path)
# safetensors-converted weights with fixed keys
self.load_state_dict(sd)
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
@@ -224,8 +227,8 @@ class NPNetGoldenNoise:
return r
def doit(self, noise, prompt, model_path, model_type, device):
if self.npnet is None:
print("Loading NPNet")
if self.npnet is None or self.npnet.pretrained_path != model_path:
print("Loading NPNet from", model_path)
self.npnet = NPNet(model_type, model_path, device=device)
self.npnet.to(device)
self.noise = noise
+35
View File
@@ -0,0 +1,35 @@
#!/usr/bin/env python3
import sys
import torch
from safetensors.torch import save_file
from pathlib import Path
files = sys.argv[1:]
for f in files:
f = Path(f)
if f.suffix in [".pth"]:
print("Converting", f)
fn = f.with_suffix(".safetensors")
if fn.exists():
print(f"{fn} exists, skipping...")
continue
print(f"Loading {f}...")
try:
model = torch.load(f, weights_only=True, map_location="cpu")
weights = {}
for k in "unet_embedding", "unet_svd", "embeeding":
subdict = model.pop(k)
if k == "embeeding":
k = "text_embedding"
for sk in subdict:
weights[f"{k}.{sk}"] = subdict[sk]
weights["alpha"] = model["alpha"]
weights["beta"] = model["beta"]
print(f"Saving {fn}...")
save_file(weights, fn)
del model
del weights
except Exception as ex:
print(f"ERROR converting {f}: {ex}")
print("Done!")