From ca6864c91d0696aa6c46304a4a14ee6f6ab1d4e5 Mon Sep 17 00:00:00 2001 From: asagi4 <130366179+asagi4@users.noreply.github.com> Date: Sun, 8 Dec 2024 15:16:30 +0200 Subject: [PATCH] Add safetensors support and fix model switching --- README.md | 2 ++ __init__.py | 29 ++++++++++++++++------------- convert_to_safetensors.py | 35 +++++++++++++++++++++++++++++++++++ 3 files changed, 53 insertions(+), 13 deletions(-) create mode 100644 convert_to_safetensors.py 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!")