Extract model saving logic to separate class

This commit is contained in:
filipstrand 2024-09-17 07:13:59 +02:00
parent a30dc4d5cc
commit 38c747b680
4 changed files with 59 additions and 82 deletions

View File

@ -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")

View File

@ -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)

View File

@ -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

View File

@ -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