diff --git a/src/mflux/controlnet/flux_controlnet.py b/src/mflux/controlnet/flux_controlnet.py index 1b8c8c0..4edf228 100644 --- a/src/mflux/controlnet/flux_controlnet.py +++ b/src/mflux/controlnet/flux_controlnet.py @@ -1,10 +1,8 @@ import logging -from pathlib import Path import PIL.Image import mlx.core as mx from mlx import nn -from mlx.utils import tree_flatten from tqdm import tqdm from mflux.config.config import ConfigControlnet @@ -21,6 +19,7 @@ from mflux.post_processing.image_util import ImageUtil from mflux.tokenizer.clip_tokenizer import TokenizerCLIP from mflux.tokenizer.t5_tokenizer import TokenizerT5 from mflux.tokenizer.tokenizer_handler import TokenizerHandler +from mflux.weights.model_saver import ModelSaver from mflux.weights.weight_handler import WeightHandler log = logging.getLogger(__name__) @@ -185,42 +184,6 @@ class Flux1Controlnet: self.t5_text_encoder.update(weights.t5_encoder) self.clip_text_encoder.update(weights.clip_encoder) - def save_model(self, base_path: str): - def _save_tokenizer(tokenizer, subdir: str): - path = Path(base_path) / subdir - path.mkdir(parents=True, exist_ok=True) - tokenizer.save_pretrained(path) - - def _save_weights(model, subdir: str): - path = Path(base_path) / subdir - path.mkdir(parents=True, exist_ok=True) - weights = _split_weights(dict(tree_flatten(model.parameters()))) - for i, weight in enumerate(weights): - mx.save_safetensors(str(path / f"{i}.safetensors"), weight, {"quantization_level": str(self.bits)}) - - def _split_weights(weights: dict, max_file_size_gb: int = 2) -> list: - # Copied from mlx-examples repo - max_file_size_bytes = max_file_size_gb << 30 - shards = [] - shard, shard_size = {}, 0 - for k, v in weights.items(): - if shard_size + v.nbytes > max_file_size_bytes: - shards.append(shard) - shard, shard_size = {}, 0 - shard[k] = v - shard_size += v.nbytes - shards.append(shard) - return shards - - # Save the tokenizers - _save_tokenizer(self.clip_tokenizer.tokenizer, "tokenizer") - _save_tokenizer(self.t5_tokenizer.tokenizer, "tokenizer_2") - - # Save the models - _save_weights(self.vae, "vae") - _save_weights(self.transformer, "transformer") - _save_weights(self.clip_text_encoder, "text_encoder") - _save_weights(self.t5_text_encoder, "text_encoder_2") - _save_weights(self.transformer_controlnet, "transformer_controlnet") - - + def save_model(self, base_path: str) -> None: + ModelSaver.save_model(self, self.bits, base_path) + ModelSaver.save_weights(base_path, self.bits, self.transformer_controlnet, "transformer_controlnet") diff --git a/src/mflux/flux/flux.py b/src/mflux/flux/flux.py index 7cd6524..fd45ee9 100644 --- a/src/mflux/flux/flux.py +++ b/src/mflux/flux/flux.py @@ -1,8 +1,5 @@ -from pathlib import Path - import mlx.core as mx from mlx import nn -from mlx.utils import tree_flatten from tqdm import tqdm from mflux.config.config import Config @@ -17,6 +14,7 @@ from mflux.post_processing.image_util import ImageUtil from mflux.tokenizer.clip_tokenizer import TokenizerCLIP from mflux.tokenizer.t5_tokenizer import TokenizerT5 from mflux.tokenizer.tokenizer_handler import TokenizerHandler +from mflux.weights.model_saver import ModelSaver from mflux.weights.weight_handler import WeightHandler @@ -138,39 +136,5 @@ class Flux1: self.t5_text_encoder.update(weights.t5_encoder) self.clip_text_encoder.update(weights.clip_encoder) - def save_model(self, base_path: str): - def _save_tokenizer(tokenizer, subdir: str): - path = Path(base_path) / subdir - path.mkdir(parents=True, exist_ok=True) - tokenizer.save_pretrained(path) - - def _save_weights(model, subdir: str): - path = Path(base_path) / subdir - path.mkdir(parents=True, exist_ok=True) - weights = _split_weights(dict(tree_flatten(model.parameters()))) - for i, weight in enumerate(weights): - mx.save_safetensors(str(path / f"{i}.safetensors"), weight, {"quantization_level": str(self.bits)}) - - def _split_weights(weights: dict, max_file_size_gb: int = 2) -> list: - # Copied from mlx-examples repo - max_file_size_bytes = max_file_size_gb << 30 - shards = [] - shard, shard_size = {}, 0 - for k, v in weights.items(): - if shard_size + v.nbytes > max_file_size_bytes: - shards.append(shard) - shard, shard_size = {}, 0 - shard[k] = v - shard_size += v.nbytes - shards.append(shard) - return shards - - # Save the tokenizers - _save_tokenizer(self.clip_tokenizer.tokenizer, "tokenizer") - _save_tokenizer(self.t5_tokenizer.tokenizer, "tokenizer_2") - - # Save the models - _save_weights(self.vae, "vae") - _save_weights(self.transformer, "transformer") - _save_weights(self.clip_text_encoder, "text_encoder") - _save_weights(self.t5_text_encoder, "text_encoder_2") + def save_model(self, base_path: str) -> None: + ModelSaver.save_model(self, self.bits, base_path) diff --git a/src/mflux/models/transformer/transformer.py b/src/mflux/models/transformer/transformer.py index 6bbcc53..cd32d83 100644 --- a/src/mflux/models/transformer/transformer.py +++ b/src/mflux/models/transformer/transformer.py @@ -1,7 +1,7 @@ -from typing import Tuple +import math + import mlx.core as mx from mlx import nn -import math from mflux.config.model_config import ModelConfig from mflux.config.runtime_config import RuntimeConfig diff --git a/src/mflux/weights/model_saver.py b/src/mflux/weights/model_saver.py new file mode 100644 index 0000000..0e7d372 --- /dev/null +++ b/src/mflux/weights/model_saver.py @@ -0,0 +1,50 @@ +from pathlib import Path + +import mlx.core as mx +from mlx import nn +from mlx.utils import tree_flatten +from transformers import CLIPTokenizer, T5Tokenizer + + +class ModelSaver: + + @staticmethod + def save_model(model, bits: int, base_path: str): + # Save the tokenizers + ModelSaver._save_tokenizer(base_path, model.clip_tokenizer.tokenizer, "tokenizer") + ModelSaver._save_tokenizer(base_path, model.t5_tokenizer.tokenizer, "tokenizer_2") + + # Save the models + ModelSaver.save_weights(base_path, bits, model.vae, "vae") + ModelSaver.save_weights(base_path, bits, model.transformer, "transformer") + ModelSaver.save_weights(base_path, bits, model.clip_text_encoder, "text_encoder") + ModelSaver.save_weights(base_path, bits, model.t5_text_encoder, "text_encoder_2") + + @staticmethod + def _save_tokenizer(base_path: str, tokenizer: CLIPTokenizer | T5Tokenizer, subdir: str): + path = Path(base_path) / subdir + path.mkdir(parents=True, exist_ok=True) + tokenizer.save_pretrained(path) + + @staticmethod + def save_weights(base_path: str, bits: int, model: nn.Module, subdir: str): + path = Path(base_path) / subdir + path.mkdir(parents=True, exist_ok=True) + weights = ModelSaver._split_weights(base_path, dict(tree_flatten(model.parameters()))) + for i, weight in enumerate(weights): + mx.save_safetensors(str(path / f"{i}.safetensors"), weight, {"quantization_level": str(bits)}) + + @staticmethod + def _split_weights(base_path: str, weights: dict, max_file_size_gb: int = 2) -> list: + # Copied from mlx-examples repo + max_file_size_bytes = max_file_size_gb << 30 + shards = [] + shard, shard_size = {}, 0 + for k, v in weights.items(): + if shard_size + v.nbytes > max_file_size_bytes: + shards.append(shard) + shard, shard_size = {}, 0 + shard[k] = v + shard_size += v.nbytes + shards.append(shard) + return shards