Save mflux versions in weight metadata

This commit is contained in:
filipstrand 2025-03-15 10:07:38 +01:00
parent cb09eacb69
commit 696762af2e
3 changed files with 33 additions and 17 deletions

View File

@ -5,6 +5,8 @@ from mlx import nn
from mlx.utils import tree_flatten from mlx.utils import tree_flatten
from transformers import CLIPTokenizer, T5Tokenizer from transformers import CLIPTokenizer, T5Tokenizer
from mflux.post_processing.generated_image import GeneratedImage
class ModelSaver: class ModelSaver:
@staticmethod @staticmethod
@ -34,7 +36,10 @@ class ModelSaver:
mx.save_safetensors( mx.save_safetensors(
str(path / f"{i}.safetensors"), str(path / f"{i}.safetensors"),
weight, weight,
{"quantization_level": str(bits)}, {
"quantization_level": str(bits),
"mflux_version": GeneratedImage.get_version(),
},
) )
@staticmethod @staticmethod

View File

@ -15,6 +15,7 @@ class MetaData:
scale: float | None = None scale: float | None = None
is_lora: bool = False is_lora: bool = False
is_mflux: bool = False is_mflux: bool = False
mflux_version: str | None = None
class WeightHandler: class WeightHandler:
@ -40,10 +41,11 @@ class WeightHandler:
# Load the weights from disk, huggingface cache, or download from huggingface # Load the weights from disk, huggingface cache, or download from huggingface
root_path = Path(local_path) if local_path else WeightHandler._download_or_get_cached_weights(repo_id) root_path = Path(local_path) if local_path else WeightHandler._download_or_get_cached_weights(repo_id)
clip_encoder, _ = WeightHandler._load_clip_encoder(root_path=root_path) clip_encoder, _, _ = WeightHandler._load_clip_encoder(root_path=root_path)
t5_encoder, _ = WeightHandler._load_t5_encoder(root_path=root_path) t5_encoder, _, _ = WeightHandler._load_t5_encoder(root_path=root_path)
vae, _ = WeightHandler._load_vae(root_path=root_path) vae, _, _ = WeightHandler._load_vae(root_path=root_path)
transformer, quantization_level, _ = WeightHandler.load_transformer(root_path=root_path) transformer, quantization_level, mflux_version = WeightHandler.load_transformer(root_path=root_path)
return WeightHandler( return WeightHandler(
clip_encoder=clip_encoder, clip_encoder=clip_encoder,
t5_encoder=t5_encoder, t5_encoder=t5_encoder,
@ -53,7 +55,7 @@ class WeightHandler:
quantization_level=quantization_level, quantization_level=quantization_level,
scale=None, scale=None,
is_lora=False, is_lora=False,
is_mflux=False, mflux_version=mflux_version,
), ),
) )
@ -64,17 +66,17 @@ class WeightHandler:
return len(self.transformer["single_transformer_blocks"]) return len(self.transformer["single_transformer_blocks"])
@staticmethod @staticmethod
def _load_clip_encoder(root_path: Path) -> (dict, int): def _load_clip_encoder(root_path: Path) -> (dict, int, str | None):
weights, quantization_level, _ = WeightHandler._get_weights("text_encoder", root_path) weights, quantization_level, mflux_version = WeightHandler._get_weights("text_encoder", root_path)
return weights, quantization_level return weights, quantization_level, mflux_version
@staticmethod @staticmethod
def _load_t5_encoder(root_path: Path) -> (dict, int): def _load_t5_encoder(root_path: Path) -> (dict, int, str | None):
weights, quantization_level, _ = WeightHandler._get_weights("text_encoder_2", root_path) weights, quantization_level, mflux_version = WeightHandler._get_weights("text_encoder_2", root_path)
# Quantized weights (i.e. ones exported from this project) don't need any post-processing. # Quantized weights (i.e. ones exported from this project) don't need any post-processing.
if quantization_level is not None: if quantization_level is not None:
return weights, quantization_level return weights, quantization_level, mflux_version
# Reshape and process the huggingface weights # Reshape and process the huggingface weights
weights["final_layer_norm"] = weights["encoder"]["final_layer_norm"] weights["final_layer_norm"] = weights["encoder"]["final_layer_norm"]
@ -93,7 +95,7 @@ class WeightHandler:
block["attention"]["SelfAttention"]["relative_attention_bias"] = relative_attention_bias block["attention"]["SelfAttention"]["relative_attention_bias"] = relative_attention_bias
weights.pop("encoder") weights.pop("encoder")
return weights, quantization_level return weights, quantization_level, mflux_version
@staticmethod @staticmethod
def load_transformer(root_path: Path | None = None, lora_path: str | None = None) -> (dict, int, str | None): def load_transformer(root_path: Path | None = None, lora_path: str | None = None) -> (dict, int, str | None):
@ -124,12 +126,12 @@ class WeightHandler:
return weights, quantization_level, mflux_version return weights, quantization_level, mflux_version
@staticmethod @staticmethod
def _load_vae(root_path: Path) -> (dict, int): def _load_vae(root_path: Path) -> (dict, int, str | None):
weights, quantization_level, _ = WeightHandler._get_weights("vae", root_path) weights, quantization_level, mflux_version = WeightHandler._get_weights("vae", root_path)
# Quantized weights (i.e. ones exported from this project) don't need any post-processing. # Quantized weights (i.e. ones exported from this project) don't need any post-processing.
if quantization_level is not None: if quantization_level is not None:
return weights, quantization_level return weights, quantization_level, mflux_version
# Reshape and process the huggingface weights # Reshape and process the huggingface weights
weights["decoder"]["conv_in"] = {"conv2d": weights["decoder"]["conv_in"]} weights["decoder"]["conv_in"] = {"conv2d": weights["decoder"]["conv_in"]}
@ -138,7 +140,7 @@ class WeightHandler:
weights["encoder"]["conv_in"] = {"conv2d": weights["encoder"]["conv_in"]} weights["encoder"]["conv_in"] = {"conv2d": weights["encoder"]["conv_in"]}
weights["encoder"]["conv_out"] = {"conv2d": weights["encoder"]["conv_out"]} weights["encoder"]["conv_out"] = {"conv2d": weights["encoder"]["conv_out"]}
weights["encoder"]["conv_norm_out"] = {"norm": weights["encoder"]["conv_norm_out"]} weights["encoder"]["conv_norm_out"] = {"norm": weights["encoder"]["conv_norm_out"]}
return weights, quantization_level return weights, quantization_level, mflux_version
@staticmethod @staticmethod
def _get_weights( def _get_weights(
@ -156,6 +158,7 @@ class WeightHandler:
weight = list(data[0].items()) weight = list(data[0].items())
if len(data) > 1: if len(data) > 1:
quantization_level = data[1].get("quantization_level") quantization_level = data[1].get("quantization_level")
mflux_version = data[1].get("mflux_version")
weights.extend(weight) weights.extend(weight)
if lora_path and root_path is None: if lora_path and root_path is None:

View File

@ -1,9 +1,12 @@
import os import os
import shutil import shutil
from pathlib import Path
import numpy as np import numpy as np
from mflux import Config, Flux1, ModelConfig from mflux import Config, Flux1, ModelConfig
from mflux.post_processing.generated_image import GeneratedImage
from mflux.weights.weight_handler import WeightHandler
PATH = "tests/4bit/" PATH = "tests/4bit/"
@ -31,6 +34,11 @@ class TestModelSaving:
fluxA.save_model(PATH) fluxA.save_model(PATH)
del fluxA del fluxA
# Verify that the mflux version is correctly saved in the model's metadata
_, quantization_level, mflux_version = WeightHandler._load_vae(root_path=Path(PATH))
assert mflux_version == GeneratedImage.get_version(), "mflux version not correctly saved in metadata" # fmt: off
assert quantization_level == "4", "quantization level not correctly saved in metadata" # fmt: off
# when loading the quantized model (also without specifying bits) # when loading the quantized model (also without specifying bits)
fluxB = Flux1( fluxB = Flux1(
model_config=ModelConfig.dev(), model_config=ModelConfig.dev(),