Extract model saving logic to separate class
This commit is contained in:
parent
a30dc4d5cc
commit
38c747b680
@ -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")
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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
|
||||
|
||||
50
src/mflux/weights/model_saver.py
Normal file
50
src/mflux/weights/model_saver.py
Normal 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
|
||||
Loading…
Reference in New Issue
Block a user