diff --git a/src/mflux/controlnet/controlnet_util.py b/src/mflux/controlnet/controlnet_util.py index a66f26f..eaf4af1 100644 --- a/src/mflux/controlnet/controlnet_util.py +++ b/src/mflux/controlnet/controlnet_util.py @@ -18,7 +18,7 @@ class ControlnetUtil: height: int, width: int, controlnet_image_path: str, - ) -> (mx.array, PIL.Image): + ) -> tuple[mx.array, PIL.Image.Image]: from mflux import ImageUtil control_image = ImageUtil.load_image(controlnet_image_path) diff --git a/src/mflux/controlnet/transformer_controlnet.py b/src/mflux/controlnet/transformer_controlnet.py index 8558965..7f69f6a 100644 --- a/src/mflux/controlnet/transformer_controlnet.py +++ b/src/mflux/controlnet/transformer_controlnet.py @@ -37,7 +37,7 @@ class TransformerControlnet(nn.Module): prompt_embeds: mx.array, pooled_prompt_embeds: mx.array, controlnet_condition: mx.array, - ) -> (list[mx.array], list[mx.array]): + ) -> tuple[list[mx.array], list[mx.array]]: # 1. Create embeddings hidden_states = self.x_embedder(hidden_states) + self.controlnet_x_embedder(controlnet_condition) encoder_hidden_states = self.context_embedder(prompt_embeds) diff --git a/src/mflux/dreambooth/dreambooth_initializer.py b/src/mflux/dreambooth/dreambooth_initializer.py index 8b11a51..b560da6 100644 --- a/src/mflux/dreambooth/dreambooth_initializer.py +++ b/src/mflux/dreambooth/dreambooth_initializer.py @@ -16,7 +16,7 @@ class DreamBoothInitializer: def initialize( config_path: str | None, checkpoint_path: str | None, - ) -> (Flux1, RuntimeConfig, TrainingSpec, TrainingState): + ) -> tuple[Flux1, RuntimeConfig, TrainingSpec, TrainingState]: # The training specification describing the details of the training process. It is resolved # differently depending on if training starts from scratch or resumes from checkpoint. training_spec = TrainingSpec.resolve( diff --git a/src/mflux/flux_tools/depth/depth_util.py b/src/mflux/flux_tools/depth/depth_util.py index 0e6595f..521122d 100644 --- a/src/mflux/flux_tools/depth/depth_util.py +++ b/src/mflux/flux_tools/depth/depth_util.py @@ -22,7 +22,7 @@ class DepthUtil: config: RuntimeConfig, image_path: str | Path | None = None, depth_image_path: str | Path | None = None, - ) -> (mx.array, PIL.Image.Image): + ) -> tuple[mx.array, PIL.Image.Image]: # 1. Create the depth map or use existing one depth_image_path, depth_image = DepthUtil.get_or_create_depth_map( depth_pro=depth_pro, diff --git a/src/mflux/flux_tools/redux/flux_redux.py b/src/mflux/flux_tools/redux/flux_redux.py index ea86f61..d507715 100644 --- a/src/mflux/flux_tools/redux/flux_redux.py +++ b/src/mflux/flux_tools/redux/flux_redux.py @@ -163,7 +163,7 @@ class Flux1Redux(nn.Module): image_paths: list[str] | list[Path], image_encoder: SiglipVisionTransformer, image_embedder: ReduxEncoder, - ) -> (mx.array, mx.array): + ) -> tuple[mx.array, mx.array]: # 1. Encode the prompt prompt_embeds_txt, pooled_prompt_embeds = PromptEncoder.encode_prompt( prompt=prompt, diff --git a/src/mflux/flux_tools/redux/weight_handler_redux.py b/src/mflux/flux_tools/redux/weight_handler_redux.py index c861cf3..ee3a1b8 100644 --- a/src/mflux/flux_tools/redux/weight_handler_redux.py +++ b/src/mflux/flux_tools/redux/weight_handler_redux.py @@ -26,11 +26,11 @@ class WeightHandlerRedux: ) # fmt:off @staticmethod - def _load_siglip_weights(root_path: Path) -> (dict, int, str | None): + def _load_siglip_weights(root_path: Path) -> tuple[dict, int, str | None]: weights, _, _ = WeightHandler.get_weights("image_encoder", root_path) return weights, _, _ @staticmethod - def _load_redux_encoder_weights(root_path: Path) -> (dict, int, str | None): + def _load_redux_encoder_weights(root_path: Path) -> tuple[dict, int, str | None]: weights, _, _ = WeightHandler.get_weights("image_embedder", root_path) return weights, _, _ diff --git a/src/mflux/models/depth_pro/depth_pro_model.py b/src/mflux/models/depth_pro/depth_pro_model.py index d463a35..b3e8aa2 100644 --- a/src/mflux/models/depth_pro/depth_pro_model.py +++ b/src/mflux/models/depth_pro/depth_pro_model.py @@ -13,7 +13,7 @@ class DepthProModel(nn.Module): self.decoder = MultiresConvDecoder() self.head = FOVHead() - def __call__(self, x: mx.array) -> (mx.array, mx.array): + def __call__(self, x: mx.array) -> tuple[mx.array, mx.array]: encodings = self.encoder(x) features = self.decoder(encodings) return self.head(features) diff --git a/src/mflux/models/depth_pro/depth_pro_util.py b/src/mflux/models/depth_pro/depth_pro_util.py index eb96d99..1e24d1c 100644 --- a/src/mflux/models/depth_pro/depth_pro_util.py +++ b/src/mflux/models/depth_pro/depth_pro_util.py @@ -8,7 +8,7 @@ import torch.nn.functional as F class DepthProUtil: @staticmethod - def create_pyramid(x: mx.array) -> (mx.array, mx.array, mx.array): + def create_pyramid(x: mx.array) -> tuple[mx.array, mx.array, mx.array]: x0 = x x_np = np.array(x) x_torch = torch.from_numpy(x_np) diff --git a/src/mflux/models/depth_pro/dino_v2/dino_vision_transformer.py b/src/mflux/models/depth_pro/dino_v2/dino_vision_transformer.py index f9172c6..e179698 100644 --- a/src/mflux/models/depth_pro/dino_v2/dino_vision_transformer.py +++ b/src/mflux/models/depth_pro/dino_v2/dino_vision_transformer.py @@ -14,7 +14,7 @@ class DinoVisionTransformer(nn.Module): self.blocks = [TransformerBlock() for i in range(24)] self.norm = nn.LayerNorm(dims=1024, eps=1e-6, bias=True) - def __call__(self, x: mx.array) -> (mx.array, mx.array, mx.array): + def __call__(self, x: mx.array) -> tuple[mx.array, mx.array, mx.array]: backbone_highres_hook0 = None backbone_highres_hook1 = None diff --git a/src/mflux/models/text_encoder/clip_encoder/clip_text_model.py b/src/mflux/models/text_encoder/clip_encoder/clip_text_model.py index 66bfe08..2851073 100644 --- a/src/mflux/models/text_encoder/clip_encoder/clip_text_model.py +++ b/src/mflux/models/text_encoder/clip_encoder/clip_text_model.py @@ -12,7 +12,7 @@ class CLIPTextModel(nn.Module): self.embeddings = CLIPEmbeddings(dims) self.final_layer_norm = nn.LayerNorm(dims=768) - def __call__(self, tokens: mx.array) -> (mx.array, mx.array): + def __call__(self, tokens: mx.array) -> tuple[mx.array, mx.array]: hidden_states = self.embeddings(tokens) causal_attention_mask = CLIPTextModel.create_causal_attention_mask(hidden_states.shape) encoder_outputs = self.encoder(hidden_states, causal_attention_mask) diff --git a/src/mflux/models/text_encoder/prompt_encoder.py b/src/mflux/models/text_encoder/prompt_encoder.py index af70160..812f67f 100644 --- a/src/mflux/models/text_encoder/prompt_encoder.py +++ b/src/mflux/models/text_encoder/prompt_encoder.py @@ -15,7 +15,7 @@ class PromptEncoder: clip_tokenizer: TokenizerCLIP, t5_text_encoder: T5Encoder, clip_text_encoder: CLIPEncoder, - ) -> (mx.array, mx.array): + ) -> tuple[mx.array, mx.array]: # 1. Return prompt encodings if already cached if prompt in prompt_cache: return prompt_cache[prompt] diff --git a/src/mflux/models/transformer/joint_attention.py b/src/mflux/models/transformer/joint_attention.py index 1c71005..a179d04 100644 --- a/src/mflux/models/transformer/joint_attention.py +++ b/src/mflux/models/transformer/joint_attention.py @@ -29,7 +29,7 @@ class JointAttention(nn.Module): hidden_states: mx.array, encoder_hidden_states: mx.array, image_rotary_emb: mx.array, - ) -> (mx.array, mx.array): + ) -> tuple[mx.array, mx.array]: # 1a. Compute Q,K,V for hidden_states query, key, value = AttentionUtils.process_qkv( hidden_states=hidden_states, diff --git a/src/mflux/models/transformer/joint_transformer_block.py b/src/mflux/models/transformer/joint_transformer_block.py index 626372e..f085f26 100644 --- a/src/mflux/models/transformer/joint_transformer_block.py +++ b/src/mflux/models/transformer/joint_transformer_block.py @@ -24,7 +24,7 @@ class JointTransformerBlock(nn.Module): encoder_hidden_states: mx.array, text_embeddings: mx.array, rotary_embeddings: mx.array, - ) -> (mx.array, mx.array): + ) -> tuple[mx.array, mx.array]: # 1a. Compute norm for hidden_states norm_hidden_states, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.norm1( hidden_states=hidden_states, diff --git a/src/mflux/models/transformer/single_transformer_block.py b/src/mflux/models/transformer/single_transformer_block.py index a188b51..2193a0e 100644 --- a/src/mflux/models/transformer/single_transformer_block.py +++ b/src/mflux/models/transformer/single_transformer_block.py @@ -21,7 +21,7 @@ class SingleTransformerBlock(nn.Module): hidden_states: mx.array, text_embeddings: mx.array, rotary_embeddings: mx.array, - ) -> (mx.array, mx.array): + ) -> tuple[mx.array, mx.array]: # 0. Establish residual connection residual = hidden_states diff --git a/src/mflux/weights/weight_handler.py b/src/mflux/weights/weight_handler.py index 7ffaadb..8e91430 100644 --- a/src/mflux/weights/weight_handler.py +++ b/src/mflux/weights/weight_handler.py @@ -65,12 +65,12 @@ class WeightHandler: return len(self.transformer["single_transformer_blocks"]) @staticmethod - def _load_clip_encoder(root_path: Path) -> (dict, int, str | None): + def _load_clip_encoder(root_path: Path) -> tuple[dict, int, str | None]: weights, quantization_level, mflux_version = WeightHandler.get_weights("text_encoder", root_path) return weights, quantization_level, mflux_version @staticmethod - def _load_t5_encoder(root_path: Path) -> (dict, int, str | None): + def _load_t5_encoder(root_path: Path) -> tuple[dict, int, str | None]: weights, quantization_level, mflux_version = WeightHandler.get_weights("text_encoder_2", root_path) # Quantized weights (i.e. ones exported from this project) don't need any post-processing. @@ -97,7 +97,7 @@ class WeightHandler: return weights, quantization_level, mflux_version @staticmethod - def load_transformer(root_path: Path | None = None, lora_path: str | None = None) -> (dict, int, str | None): + def load_transformer(root_path: Path | None = None, lora_path: str | None = None) -> tuple[dict, int, str | None]: weights, quantization_level, mflux_version = WeightHandler.get_weights("transformer", root_path, lora_path) if lora_path: @@ -125,7 +125,7 @@ class WeightHandler: return weights, quantization_level, mflux_version @staticmethod - def _load_vae(root_path: Path) -> (dict, int, str | None): + def _load_vae(root_path: Path) -> tuple[dict, int, str | None]: weights, quantization_level, mflux_version = WeightHandler.get_weights("vae", root_path) # Quantized weights (i.e. ones exported from this project) don't need any post-processing. @@ -146,7 +146,7 @@ class WeightHandler: model_name: str, root_path: Path | None = None, lora_path: str | None = None, - ) -> (dict, int, str | None): + ) -> tuple[dict, int, str | None]: weights = [] quantization_level = None mflux_version = None