Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 3 additions & 1 deletion diffsynth/pipelines/flux2_image.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down
99 changes: 99 additions & 0 deletions tests/test_flux2_tokenizer_selection.py
Original file line number Diff line number Diff line change
@@ -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()