diff --git a/pyproject.toml b/pyproject.toml index 5ccf476..3165233 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -84,7 +84,7 @@ mflux-generate-fill = "mflux.generate_fill:main" mflux-generate-depth = "mflux.generate_depth:main" mflux-generate-redux = "mflux.generate_redux:main" mflux-generate-kontext = "mflux.generate_kontext:main" -mflux-generate-qwen-edit-plus = "mflux.generate_qwen_edit_plus:main" +mflux-generate-qwen-edit = "mflux.generate_qwen_edit:main" mflux-concept = "mflux.concept:main" mflux-concept-from-image = "mflux.concept_from_image:main" mflux-save = "mflux.save:main" diff --git a/src/mflux/config/model_config.py b/src/mflux/config/model_config.py index d063055..1d9e51d 100644 --- a/src/mflux/config/model_config.py +++ b/src/mflux/config/model_config.py @@ -91,8 +91,8 @@ class ModelConfig: @staticmethod @lru_cache - def qwen_image_edit_plus() -> "ModelConfig": - return AVAILABLE_MODELS["qwen-image-edit-plus"] + def qwen_image_edit() -> "ModelConfig": + return AVAILABLE_MODELS["qwen-image-edit"] def x_embedder_input_dim(self) -> int: if "Fill" in self.model_name: @@ -307,8 +307,8 @@ AVAILABLE_MODELS = { requires_sigma_shift=None, priority=11, ), - "qwen-image-edit-plus": ModelConfig( - aliases=["qwen-image-edit-plus", "qwen-edit-plus", "qwen-edit-2509"], + "qwen-image-edit": ModelConfig( + aliases=["qwen-image-edit", "qwen-edit", "qwen-edit-plus", "qwen-edit-2509"], model_name="Qwen/Qwen-Image-Edit-2509", base_model=None, controlnet_model=None, @@ -317,6 +317,6 @@ AVAILABLE_MODELS = { max_sequence_length=None, supports_guidance=None, requires_sigma_shift=None, - priority=13, + priority=12, ), } diff --git a/src/mflux/generate_qwen_edit_plus.py b/src/mflux/generate_qwen_edit.py similarity index 90% rename from src/mflux/generate_qwen_edit_plus.py rename to src/mflux/generate_qwen_edit.py index 39b2a59..f208657 100644 --- a/src/mflux/generate_qwen_edit_plus.py +++ b/src/mflux/generate_qwen_edit.py @@ -3,7 +3,7 @@ from pathlib import Path from mflux.callbacks.callback_manager import CallbackManager from mflux.config.config import Config from mflux.error.exceptions import PromptFileReadError, StopImageGenerationException -from mflux.models.qwen.variants.edit.qwen_image_edit_plus import QwenImageEditPlus +from mflux.models.qwen.variants.edit.qwen_image_edit import QwenImageEdit from mflux.ui import defaults as ui_defaults from mflux.ui.cli.parsers import CommandLineParser from mflux.ui.prompt_utils import get_effective_negative_prompt, get_effective_prompt @@ -11,7 +11,7 @@ from mflux.ui.prompt_utils import get_effective_negative_prompt, get_effective_p def main(): # 0. Parse command line arguments - parser = CommandLineParser(description="Generate an image using Qwen Image Edit Plus with image conditioning.") + parser = CommandLineParser(description="Generate an image using Qwen Image Edit with image conditioning.") parser.add_general_arguments() parser.add_model_arguments(require_model_arg=False) parser.add_lora_arguments() @@ -22,7 +22,7 @@ def main(): type=Path, nargs="+", default=None, - help="Local paths to multiple init images (for plus model). If not provided, uses --image-path as single image.", + help="Local paths to multiple init images. If not provided, uses --image-path as single image.", ) parser.add_output_arguments() args = parser.parse_args() @@ -32,7 +32,7 @@ def main(): args.guidance = ui_defaults.GUIDANCE_SCALE_KONTEXT # 1. Load the model - qwen = QwenImageEditPlus( + qwen = QwenImageEdit( quantize=args.quantize, local_path=args.path, lora_paths=args.lora_paths, diff --git a/src/mflux/models/qwen/qwen_edit_initializer.py b/src/mflux/models/qwen/qwen_edit_initializer.py index d792c47..eb019e5 100644 --- a/src/mflux/models/qwen/qwen_edit_initializer.py +++ b/src/mflux/models/qwen/qwen_edit_initializer.py @@ -53,42 +53,36 @@ class QwenImageEditInitializer: root_path = Path(local_path) if local_path else QwenImageEditInitializer._download_vl_processor(repo_id) # 2. Load only the tokenizer (we implement image processing ourselves) - tokenizer_path = root_path / "tokenizer" - if not tokenizer_path.exists(): - tokenizer_path = root_path - - try: - tokenizer = Qwen2TokenizerFast.from_pretrained( - pretrained_model_name_or_path=tokenizer_path, - local_files_only=True, - ) - except OSError: - tokenizer = Qwen2TokenizerFast.from_pretrained(repo_id) + tokenizer = QwenImageEditInitializer._load_tokenizer(root_path, repo_id) # 3. Create our MLX processor with the tokenizer processor = QwenVisionLanguageProcessor(tokenizer=tokenizer) # 4. Initialize vision-language tokenizer wrapper - if model_config is not None: - model_name_lower = model_config.model_name.lower() - use_picture_prefix = ( - "plus" in model_name_lower - or "2509" in model_config.model_name # Check original case for "2509" - or (hasattr(model_config, "use_picture_prefix") and model_config.use_picture_prefix) - ) - else: - # Fallback: check repo_id - use_picture_prefix = "plus" in repo_id.lower() or "2509" in repo_id - qwen_model.qwen_vl_tokenizer = QwenVisionLanguageTokenizer( processor=processor, max_length=1024, - use_picture_prefix=use_picture_prefix, + use_picture_prefix=True, ) # 5. Initialize vision-language encoder (integrated approach like Diffusers) qwen_model.qwen_vl_encoder = QwenVisionLanguageEncoder(encoder=qwen_model.text_encoder.encoder) + @staticmethod + def _load_tokenizer(root_path: Path, repo_id: str) -> Qwen2TokenizerFast: + """Load the tokenizer from local path or download from HuggingFace.""" + tokenizer_path = root_path / "tokenizer" + if not tokenizer_path.exists(): + tokenizer_path = root_path + + try: + return Qwen2TokenizerFast.from_pretrained( + pretrained_model_name_or_path=tokenizer_path, + local_files_only=True, + ) + except OSError: + return Qwen2TokenizerFast.from_pretrained(repo_id) + @staticmethod def _download_vl_processor(repo_id: str) -> Path: return Path( diff --git a/src/mflux/models/qwen/tokenizer/qwen_vision_language_tokenizer.py b/src/mflux/models/qwen/tokenizer/qwen_vision_language_tokenizer.py index 519c5be..7a66aff 100644 --- a/src/mflux/models/qwen/tokenizer/qwen_vision_language_tokenizer.py +++ b/src/mflux/models/qwen/tokenizer/qwen_vision_language_tokenizer.py @@ -59,7 +59,7 @@ class QwenVisionLanguageTokenizer: # Format prompt based on tokenizer mode if self.use_picture_prefix: - # Edit Plus format: Add "Picture N:" prefix for each image + # Edit format: Add "Picture N:" prefix for each image # For multiple images: "Picture 1: ... Picture 2: ... Picture N: ..." img_prompt_template = "Picture {}: <|vision_start|><|image_pad|><|vision_end|>" base_img_prompt = "" diff --git a/src/mflux/models/qwen/variants/edit/__init__.py b/src/mflux/models/qwen/variants/edit/__init__.py index 5180da1..6f56cec 100644 --- a/src/mflux/models/qwen/variants/edit/__init__.py +++ b/src/mflux/models/qwen/variants/edit/__init__.py @@ -1,4 +1,4 @@ # Qwen Image Edit variant -from mflux.models.qwen.variants.edit.qwen_image_edit_plus import QwenImageEditPlus +from mflux.models.qwen.variants.edit.qwen_image_edit import QwenImageEdit -__all__ = ["QwenImageEditPlus"] +__all__ = ["QwenImageEdit"] diff --git a/src/mflux/models/qwen/variants/edit/qwen_image_edit_plus.py b/src/mflux/models/qwen/variants/edit/qwen_image_edit.py similarity index 98% rename from src/mflux/models/qwen/variants/edit/qwen_image_edit_plus.py rename to src/mflux/models/qwen/variants/edit/qwen_image_edit.py index cc510c0..34bdb17 100644 --- a/src/mflux/models/qwen/variants/edit/qwen_image_edit_plus.py +++ b/src/mflux/models/qwen/variants/edit/qwen_image_edit.py @@ -21,7 +21,7 @@ from mflux.post_processing.generated_image import GeneratedImage from mflux.post_processing.image_util import ImageUtil -class QwenImageEditPlus(nn.Module): +class QwenImageEdit(nn.Module): vae: QwenVAE transformer: QwenTransformer text_encoder: QwenTextEncoder @@ -38,7 +38,7 @@ class QwenImageEditPlus(nn.Module): super().__init__() QwenImageEditInitializer.init( qwen_model=self, - model_config=ModelConfig.qwen_image_edit_plus(), + model_config=ModelConfig.qwen_image_edit(), quantize=quantize, local_path=local_path, lora_paths=lora_paths, @@ -70,7 +70,7 @@ class QwenImageEditPlus(nn.Module): ) # 2. Encode the prompt - prompt_embeds, prompt_mask, negative_prompt_embeds, negative_prompt_mask = self._encode_prompts_with_image_plus( + prompt_embeds, prompt_mask, negative_prompt_embeds, negative_prompt_mask = self._encode_prompts_with_images( prompt=prompt, negative_prompt=negative_prompt, image_paths=image_paths, @@ -186,7 +186,7 @@ class QwenImageEditPlus(nn.Module): negative_prompt=negative_prompt, ) - def _encode_prompts_with_image_plus( + def _encode_prompts_with_images( self, prompt: str, negative_prompt: str, diff --git a/tests/image_generation/helpers/image_generation_edit_test_helper.py b/tests/image_generation/helpers/image_generation_edit_test_helper.py index 9b4a448..dc8084d 100644 --- a/tests/image_generation/helpers/image_generation_edit_test_helper.py +++ b/tests/image_generation/helpers/image_generation_edit_test_helper.py @@ -31,9 +31,9 @@ class ImageGeneratorEditTestHelper: reference_image_path = ImageGeneratorEditTestHelper.resolve_path(reference_image_path) output_image_path = ImageGeneratorEditTestHelper.resolve_path(output_image_path) - # For Edit Plus, if image_paths is provided but image_path is not, use first element of image_paths - is_edit_plus = "Plus" in model_class.__name__ - if is_edit_plus and image_paths and not image_path: + # For Qwen Edit, if image_paths is provided but image_path is not, use first element of image_paths + is_qwen_edit = "QwenImageEdit" in model_class.__name__ + if is_qwen_edit and image_paths and not image_path: image_path = image_paths[0] if image_path: @@ -64,9 +64,9 @@ class ImageGeneratorEditTestHelper: "config": Config(**config_kwargs), } - # Edit Plus uses image_paths instead of image_path in config - # Check if it's Edit Plus by checking the class name - if "Plus" in model_class.__name__: + # Qwen Edit uses image_paths instead of image_path in config + # Check if it's Qwen Edit by checking the class name + if "QwenImageEdit" in model_class.__name__: if image_paths: generate_kwargs["image_paths"] = image_paths else: diff --git a/tests/image_generation/test_generate_image_edit.py b/tests/image_generation/test_generate_image_edit.py index 91cb06c..831e59f 100644 --- a/tests/image_generation/test_generate_image_edit.py +++ b/tests/image_generation/test_generate_image_edit.py @@ -1,6 +1,6 @@ from mflux.config.model_config import ModelConfig from mflux.models.flux.variants.kontext.flux_kontext import Flux1Kontext -from mflux.models.qwen.variants.edit import QwenImageEditPlus +from mflux.models.qwen.variants.edit import QwenImageEdit from tests.image_generation.helpers.image_generation_edit_test_helper import ImageGeneratorEditTestHelper @@ -20,12 +20,12 @@ class TestImageGeneratorEdit: image_path="reference_upscaled.png", ) - def test_image_generation_qwen_edit_plus(self): + def test_image_generation_qwen_edit(self): ImageGeneratorEditTestHelper.assert_matches_reference_image( - reference_image_path="reference_qwen_edit_plus.png", - output_image_path="output_qwen_edit_plus.png", - model_class=QwenImageEditPlus, - model_config=ModelConfig.qwen_image_edit_plus(), + reference_image_path="reference_qwen_edit.png", + output_image_path="output_qwen_edit.png", + model_class=QwenImageEdit, + model_config=ModelConfig.qwen_image_edit(), steps=20, seed=4869845, height=384, @@ -36,12 +36,12 @@ class TestImageGeneratorEdit: image_path="reference_upscaled.png", ) - def test_image_generation_qwen_edit_plus_multiple_images(self): + def test_image_generation_qwen_edit_multiple_images(self): ImageGeneratorEditTestHelper.assert_matches_reference_image( - reference_image_path="reference_qwen_edit_plus_multiple_images.png", - output_image_path="output_qwen_edit_plus_multiple_images.png", - model_class=QwenImageEditPlus, - model_config=ModelConfig.qwen_image_edit_plus(), + reference_image_path="reference_qwen_edit_multiple_images.png", + output_image_path="output_qwen_edit_multiple_images.png", + model_class=QwenImageEdit, + model_config=ModelConfig.qwen_image_edit(), steps=20, seed=4869845, height=384, diff --git a/tests/resources/reference_qwen_edit_plus.png b/tests/resources/reference_qwen_edit.png similarity index 100% rename from tests/resources/reference_qwen_edit_plus.png rename to tests/resources/reference_qwen_edit.png diff --git a/tests/resources/reference_qwen_edit_plus_multiple_images.png b/tests/resources/reference_qwen_edit_multiple_images.png similarity index 100% rename from tests/resources/reference_qwen_edit_plus_multiple_images.png rename to tests/resources/reference_qwen_edit_multiple_images.png