Add safetensors support and fix model switching
This commit is contained in:
@@ -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
@@ -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
|
||||
|
||||
@@ -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!")
|
||||
Reference in New Issue
Block a user