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 import logging
from pathlib import Path
import PIL.Image import PIL.Image
import mlx.core as mx import mlx.core as mx
from mlx import nn from mlx import nn
from mlx.utils import tree_flatten
from tqdm import tqdm from tqdm import tqdm
from mflux.config.config import ConfigControlnet 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.clip_tokenizer import TokenizerCLIP
from mflux.tokenizer.t5_tokenizer import TokenizerT5 from mflux.tokenizer.t5_tokenizer import TokenizerT5
from mflux.tokenizer.tokenizer_handler import TokenizerHandler from mflux.tokenizer.tokenizer_handler import TokenizerHandler
from mflux.weights.model_saver import ModelSaver
from mflux.weights.weight_handler import WeightHandler from mflux.weights.weight_handler import WeightHandler
log = logging.getLogger(__name__) log = logging.getLogger(__name__)
@ -185,42 +184,6 @@ class Flux1Controlnet:
self.t5_text_encoder.update(weights.t5_encoder) self.t5_text_encoder.update(weights.t5_encoder)
self.clip_text_encoder.update(weights.clip_encoder) self.clip_text_encoder.update(weights.clip_encoder)
def save_model(self, base_path: str): def save_model(self, base_path: str) -> None:
def _save_tokenizer(tokenizer, subdir: str): ModelSaver.save_model(self, self.bits, base_path)
path = Path(base_path) / subdir ModelSaver.save_weights(base_path, self.bits, self.transformer_controlnet, "transformer_controlnet")
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")

View File

@ -1,8 +1,5 @@
from pathlib import Path
import mlx.core as mx import mlx.core as mx
from mlx import nn from mlx import nn
from mlx.utils import tree_flatten
from tqdm import tqdm from tqdm import tqdm
from mflux.config.config import Config 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.clip_tokenizer import TokenizerCLIP
from mflux.tokenizer.t5_tokenizer import TokenizerT5 from mflux.tokenizer.t5_tokenizer import TokenizerT5
from mflux.tokenizer.tokenizer_handler import TokenizerHandler from mflux.tokenizer.tokenizer_handler import TokenizerHandler
from mflux.weights.model_saver import ModelSaver
from mflux.weights.weight_handler import WeightHandler from mflux.weights.weight_handler import WeightHandler
@ -138,39 +136,5 @@ class Flux1:
self.t5_text_encoder.update(weights.t5_encoder) self.t5_text_encoder.update(weights.t5_encoder)
self.clip_text_encoder.update(weights.clip_encoder) self.clip_text_encoder.update(weights.clip_encoder)
def save_model(self, base_path: str): def save_model(self, base_path: str) -> None:
def _save_tokenizer(tokenizer, subdir: str): ModelSaver.save_model(self, self.bits, base_path)
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")

View File

@ -1,7 +1,7 @@
from typing import Tuple import math
import mlx.core as mx import mlx.core as mx
from mlx import nn from mlx import nn
import math
from mflux.config.model_config import ModelConfig from mflux.config.model_config import ModelConfig
from mflux.config.runtime_config import RuntimeConfig 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