From 3464a05c54707a22ed8e0f079d8bfaa0e1e61c43 Mon Sep 17 00:00:00 2001 From: primorLee Date: Sat, 22 Aug 2026 20:49:25 +0800 Subject: [PATCH] fix(flux2): restore Mistral processor loading --- diffsynth/pipelines/flux2_image.py | 4 +- tests/test_flux2_tokenizer_selection.py | 99 +++++++++++++++++++++++++ 2 files changed, 102 insertions(+), 1 deletion(-) create mode 100644 tests/test_flux2_tokenizer_selection.py diff --git a/diffsynth/pipelines/flux2_image.py b/diffsynth/pipelines/flux2_image.py index 753f6c086..5089168dd 100644 --- a/diffsynth/pipelines/flux2_image.py +++ b/diffsynth/pipelines/flux2_image.py @@ -65,7 +65,9 @@ def from_pretrained( pipe.vae = model_pool.fetch_model("flux2_vae") if tokenizer_config is not None: tokenizer_config.download_if_necessary() - pipe.tokenizer = AutoTokenizer.from_pretrained(tokenizer_config.path) + # Mistral3 needs its multimodal processor for structured chat content. + tokenizer_class = AutoTokenizer if pipe.text_encoder_qwen3 is not None else AutoProcessor + pipe.tokenizer = tokenizer_class.from_pretrained(tokenizer_config.path) # VRAM Management pipe.vram_management_enabled = pipe.check_vram_management_state() diff --git a/tests/test_flux2_tokenizer_selection.py b/tests/test_flux2_tokenizer_selection.py new file mode 100644 index 000000000..b9fe17f02 --- /dev/null +++ b/tests/test_flux2_tokenizer_selection.py @@ -0,0 +1,99 @@ +import unittest +from unittest.mock import patch + +import torch + +from diffsynth.pipelines import flux2_image + + +class _ModelPool: + def __init__(self, text_encoder=None, text_encoder_qwen3=None): + self.models = { + "flux2_text_encoder": text_encoder, + "z_image_text_encoder": text_encoder_qwen3, + "flux2_dit": None, + "flux2_vae": None, + } + + def fetch_model(self, name): + return self.models[name] + + +class _TokenizerConfig: + path = "tokenizer-path" + + def __init__(self): + self.downloaded = False + + def download_if_necessary(self): + self.downloaded = True + + +class Flux2TokenizerSelectionTest(unittest.TestCase): + def _load_pipeline(self, model_pool, tokenizer_config): + with ( + patch.object( + flux2_image.Flux2ImagePipeline, + "download_and_load_models", + return_value=model_pool, + ), + patch.object( + flux2_image.Flux2ImagePipeline, + "check_vram_management_state", + return_value=False, + ), + ): + return flux2_image.Flux2ImagePipeline.from_pretrained( + torch_dtype=torch.float32, + device="cpu", + model_configs=[], + tokenizer_config=tokenizer_config, + ) + + def test_mistral_encoder_loads_multimodal_processor(self): + processor = object() + config = _TokenizerConfig() + pool = _ModelPool(text_encoder=object()) + + with ( + patch.object( + flux2_image.AutoProcessor, + "from_pretrained", + return_value=processor, + ) as load_processor, + patch.object( + flux2_image.AutoTokenizer, "from_pretrained" + ) as load_tokenizer, + ): + pipe = self._load_pipeline(pool, config) + + self.assertTrue(config.downloaded) + self.assertIs(pipe.tokenizer, processor) + load_processor.assert_called_once_with(config.path) + load_tokenizer.assert_not_called() + + def test_qwen3_encoder_keeps_tokenizer_path(self): + tokenizer = object() + config = _TokenizerConfig() + pool = _ModelPool(text_encoder_qwen3=object()) + + with ( + patch.object( + flux2_image.AutoProcessor, "from_pretrained" + ) as load_processor, + patch.object( + flux2_image.AutoTokenizer, + "from_pretrained", + return_value=tokenizer, + ) as load_tokenizer, + ): + pipe = self._load_pipeline(pool, config) + + self.assertTrue(config.downloaded) + self.assertIs(pipe.tokenizer, tokenizer) + load_processor.assert_not_called() + load_tokenizer.assert_called_once_with(config.path) + + +if __name__ == "__main__": + unittest.main()