diff --git a/README.md b/README.md index 7ef6199..3ddb66d 100644 --- a/README.md +++ b/README.md @@ -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) diff --git a/__init__.py b/__init__.py index 41bf4fd..597cfab 100644 --- a/__init__.py +++ b/__init__.py @@ -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 diff --git a/convert_to_safetensors.py b/convert_to_safetensors.py new file mode 100644 index 0000000..ccf3aaa --- /dev/null +++ b/convert_to_safetensors.py @@ -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!")