From 89830707285482bc34b3542fd279eeb4e83382fe Mon Sep 17 00:00:00 2001 From: filipstrand Date: Wed, 4 Sep 2024 19:48:15 +0200 Subject: [PATCH] Code style improvements and some restructuring --- main.py | 4 +- src/flux_1/flux.py | 10 +-- src/flux_1/weights/weight_handler.py | 109 +++++++++++++-------------- 3 files changed, 60 insertions(+), 63 deletions(-) diff --git a/main.py b/main.py index 0f95784..eaf693d 100644 --- a/main.py +++ b/main.py @@ -26,13 +26,11 @@ def main(): parser.add_argument('--apply-lora', type=str, nargs='*', default=[], help='Local safetensors for applying LORA from disk') parser.add_argument('--lora-scales', type=float,nargs='*', default=[1.0], help='Scaling factor to adjust the impact of LoRA weights on the model. A value of 1.0 applies the LoRA weights as they are.') - args = parser.parse_args() if args.path and args.model is None: parser.error("--model must be specified when using --path") - seed = int(time.time()) if args.seed is None else args.seed flux = Flux1( model_config=ModelConfig.from_alias(args.model), quantize_full_weights=args.quantize, @@ -42,7 +40,7 @@ def main(): ) image = flux.generate_image( - seed=seed, + seed=int(time.time()) if args.seed is None else args.seed, prompt=args.prompt, config=Config( num_inference_steps=args.steps, diff --git a/src/flux_1/flux.py b/src/flux_1/flux.py index d76a0e4..7f53a80 100644 --- a/src/flux_1/flux.py +++ b/src/flux_1/flux.py @@ -21,7 +21,6 @@ from flux_1.tokenizer.tokenizer_handler import TokenizerHandler from flux_1.weights.weight_handler import WeightHandler - class Flux1: def __init__( @@ -29,17 +28,17 @@ class Flux1: model_config: ModelConfig, quantize_full_weights: int | None = None, local_path: str | None = None, - lora_files: [str] =[], + lora_files: [str] = [], lora_scales: [float] = [1.0] ): self.model_config = model_config self.quantize_full_weights = quantize_full_weights - self.lora_files = lora_files # Load and initialize the tokenizers from disk, huggingface cache, or download from huggingface tokenizers = TokenizerHandler(model_config.model_name, self.model_config.max_sequence_length, local_path) self.t5_tokenizer = TokenizerT5(tokenizers.t5, max_length=self.model_config.max_sequence_length) self.clip_tokenizer = TokenizerCLIP(tokenizers.clip) + # Initialize the models self.vae = VAE() self.transformer = Transformer(model_config) @@ -47,7 +46,8 @@ class Flux1: self.clip_text_encoder = CLIPEncoder() # Load the weights from disk, huggingface cache, or download from huggingface - weights = WeightHandler(repo_id=model_config.model_name, local_path=local_path,lora_files=lora_files,lora_scales=lora_scales) + weights = WeightHandler(repo_id=model_config.model_name, local_path=local_path, lora_files=lora_files, lora_scales=lora_scales) + # Set the loaded weights if they are not quantized if weights.quantization_level is None: self._set_model_weights(weights) @@ -63,7 +63,7 @@ class Flux1: # If loading previously saved quantized weights, the weights must be set after modules have been quantized if weights.quantization_level is not None: self._set_model_weights(weights) - + def generate_image(self, seed: int, prompt: str, config: Config = Config()) -> PIL.Image.Image: # Create a new runtime config based on the model type and input parameters config = RuntimeConfig(config, self.model_config) diff --git a/src/flux_1/weights/weight_handler.py b/src/flux_1/weights/weight_handler.py index 024228a..41a326f 100644 --- a/src/flux_1/weights/weight_handler.py +++ b/src/flux_1/weights/weight_handler.py @@ -1,102 +1,104 @@ +import logging from pathlib import Path import mlx.core as mx from huggingface_hub import snapshot_download +from mlx.utils import tree_flatten from mlx.utils import tree_unflatten from safetensors import safe_open from flux_1.config.config import Config -from mlx.utils import tree_flatten -import logging log = logging.getLogger(__name__) + class WeightHandler: def __init__( self, repo_id: str | None = None, local_path: str | None = None, - lora_files: [str] =[], - lora_scales: [float] = [1.0] + lora_files=None, + lora_scales=None ): + if lora_files is None: + lora_files = [] + if lora_scales is None: + lora_scales = [1.0] root_path = Path(local_path) if local_path else WeightHandler._download_or_get_cached_weights(repo_id) self.clip_encoder, _ = WeightHandler._clip_encoder(root_path=root_path) self.t5_encoder, _ = WeightHandler._t5_encoder(root_path=root_path) self.vae, _ = WeightHandler._vae(root_path=root_path) self.transformer, self.quantization_level = WeightHandler._transformer(root_path=root_path) - if(lora_files): - if(len(lora_files)< len(lora_scales)): - lora_scales=lora_scales[0:len(lora_files)] - if(len(lora_scales)1.0): + if lora_scale < 0.0 or lora_scale > 1.0: raise Exception(f"Invalid scale {lora_scale} provided for {lora_file}. Valid Range [0.0-1.0] ") try: - lora_transformer,_ = WeightHandler._lora_transformer(lora_file=lora_file) + lora_transformer, _ = WeightHandler._lora_transformer(lora_file=lora_file) if 'transformer' not in lora_transformer: - raise Exception("The key `transformer` is missing in the LoRA safetensors file. Please ensure that the file is correctly formatted and contains the expected keys.") - self._apply_transformer(self.transformer,lora_transformer['transformer'],lora_scale) + raise Exception( + "The key `transformer` is missing in the LoRA safetensors file. Please ensure that the file is correctly formatted and contains the expected keys.") + WeightHandler._apply_transformer(self.transformer, lora_transformer['transformer'], lora_scale) except Exception as e: log.error(f"Error loading the LoRA safetensors file: {e}") - - def _apply_transformer(self,transformer,lora_transformer,lora_scale): + @staticmethod + def _apply_transformer(transformer, lora_transformer, lora_scale): lora_weights = tree_flatten(lora_transformer) - visited={} - - for key,weight in lora_weights: - splits=key.split(".") - target=transformer - visiting=[] + visited = {} + + for key, weight in lora_weights: + splits = key.split(".") + target = transformer + visiting = [] for splitKey in splits: - if isinstance(target,dict) and splitKey in target: - target=target[splitKey] + if isinstance(target, dict) and splitKey in target: + target = target[splitKey] visiting.append(splitKey) - elif isinstance(target,list) and len(target)>0: - if(len(target)< int(splitKey)): - for _ in range(int(splitKey)-len(target)+1): + elif isinstance(target, list) and len(target) > 0: + if len(target) < int(splitKey): + for _ in range(int(splitKey) - len(target) + 1): target.append({}) - - target=target[int(splitKey)] + + target = target[int(splitKey)] visiting.append(splitKey) else: - parentKey=".".join(visiting) - if(parentKey in visited and 'lora_A' in visited[parentKey] and 'lora_B' in visited[parentKey]): + parentKey = ".".join(visiting) + if parentKey in visited and 'lora_A' in visited[parentKey] and 'lora_B' in visited[parentKey]: continue if not splitKey.startswith("lora_"): visiting.append(splitKey) - parentKey=".".join(visiting) - if(splitKey=="net"): - target['net']=list({}) - target=target['net'] - elif (splitKey=="0"): + parentKey = ".".join(visiting) + if splitKey == "net": + target['net'] = list({}) + target = target['net'] + elif splitKey == "0": target.append({}) - target=target[0] + target = target[0] continue - elif (splitKey=="proj"): - target[splitKey]=weight + elif splitKey == "proj": + target[splitKey] = weight if parentKey not in visited: - visited[parentKey]={} + visited[parentKey] = {} continue if parentKey not in visited: - visited[parentKey]={} - visited[parentKey][splitKey]=weight + visited[parentKey] = {} + visited[parentKey][splitKey] = weight if not 'weight' in target: - raise ValueError(f"LoRA weights for layer {parentKey} cannot be loaded into the model.") - if 'lora_A' in visited[parentKey] and 'lora_B' in visited[parentKey]: - lora_a=visited[parentKey]['lora_A'] - lora_b=visited[parentKey]['lora_B'] - transWeight=target['weight'] - weight=transWeight + lora_scale* (lora_b @lora_a) - target['weight']=weight - - - - + raise ValueError(f"LoRA weights for layer {parentKey} cannot be loaded into the model.") + if 'lora_A' in visited[parentKey] and 'lora_B' in visited[parentKey]: + lora_a = visited[parentKey]['lora_A'] + lora_b = visited[parentKey]['lora_B'] + transWeight = target['weight'] + weight = transWeight + lora_scale * (lora_b @ lora_a) + target['weight'] = weight @staticmethod def _lora_transformer(lora_file: Path) -> (dict, int): @@ -117,7 +119,6 @@ class WeightHandler: } return unflatten, quantization_level - @staticmethod def _clip_encoder(root_path: Path) -> (dict, int): weights, quantization_level = WeightHandler._get_weights("text_encoder", root_path) @@ -170,8 +171,6 @@ class WeightHandler: "linear2": block["ff_context"]["net"][2] } return weights, quantization_level - - @staticmethod def _vae(root_path: Path) -> (dict, int):