Extract model saving logic to separate class
This commit is contained in:
parent
a30dc4d5cc
commit
38c747b680
@ -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")
|
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@ -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")
|
|
||||||
|
|||||||
@ -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
|
||||||
|
|||||||
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