diff --git a/documentation/OPTIONS.es.md b/documentation/OPTIONS.es.md index 6e11c8040..51f0ea1d2 100644 --- a/documentation/OPTIONS.es.md +++ b/documentation/OPTIONS.es.md @@ -253,6 +253,13 @@ Donde `foo` es tu entorno de configuración; o simplemente usa `config/config.js - **Qué**: Ruta al modelo Gemma preentrenado o su identificador en . - **Por qué**: Al entrenar modelos basados en Gemma (por ejemplo LTX-2, Sana o Lumina2), puedes apuntar a un checkpoint Gemma compartido sin cambiar la ruta del modelo base de difusión. +### `--qwen_text_encoder_model_name_or_path` + +- **Qué**: Ruta a un codificador de texto Qwen preentrenado o su identificador en . +- **Predeterminado**: `None` (usa la fuente del codificador de texto Qwen definida por el modelo seleccionado). +- **Por qué**: Úsalo para compartir o reemplazar el codificador de texto Qwen en familias de modelos basadas en Qwen sin editar la caché de Hugging Face. +- **Notas**: Se aplica a familias de modelos con un solo codificador de texto Qwen. Si una familia define varios codificadores Qwen, la opción se ignora y SimpleTuner registra una advertencia. + ### `--max_grounding_entities` - Numero maximo de entidades de grounding por imagen para anotaciones espaciales estilo GLIGEN. Por defecto: 0 (deshabilitado). Valores tipicos: 4-16. @@ -1748,6 +1755,7 @@ usage: train.py [-h] --model_family [--pretrained_unet_subfolder PRETRAINED_UNET_SUBFOLDER] [--pretrained_t5_model_name_or_path PRETRAINED_T5_MODEL_NAME_OR_PATH] [--pretrained_gemma_model_name_or_path PRETRAINED_GEMMA_MODEL_NAME_OR_PATH] + [--qwen_text_encoder_model_name_or_path QWEN_TEXT_ENCODER_MODEL_NAME_OR_PATH] [--revision REVISION] [--variant VARIANT] [--base_model_default_dtype {bf16,fp32}] [--unet_attention_slice [UNET_ATTENTION_SLICE]] @@ -2079,6 +2087,8 @@ options: Path to pretrained T5 model --pretrained_gemma_model_name_or_path PRETRAINED_GEMMA_MODEL_NAME_OR_PATH Path to pretrained Gemma model + --qwen_text_encoder_model_name_or_path QWEN_TEXT_ENCODER_MODEL_NAME_OR_PATH + Path to pretrained Qwen text encoder model --revision REVISION Git branch/tag/commit for model version --variant VARIANT Model variant (e.g., fp16, bf16) --base_model_default_dtype {bf16,fp32} diff --git a/documentation/OPTIONS.hi.md b/documentation/OPTIONS.hi.md index f4b121b93..499ab72c1 100644 --- a/documentation/OPTIONS.hi.md +++ b/documentation/OPTIONS.hi.md @@ -253,6 +253,13 @@ simpletuner configure config/foo/config.json - **What**: pretrained Gemma model का path या से उसका identifier. - **Why**: Gemma‑based models (जैसे LTX-2, Sana, Lumina2) ट्रेन करते समय आप base diffusion model path बदले बिना Gemma weights का source specify कर सकते हैं। +### `--qwen_text_encoder_model_name_or_path` + +- **What**: pretrained Qwen text encoder model का path या से उसका identifier. +- **Default**: `None` (selected model में defined Qwen text encoder source उपयोग होता है). +- **Why**: Qwen-based model families में Qwen text encoder को share या replace करने के लिए इसका उपयोग करें, बिना Hugging Face cache edit किए। +- **Notes**: यह उन model families पर लागू होता है जिनमें एक Qwen text encoder है। अगर कोई model family multiple Qwen encoders define करती है, तो option ignore होता है और SimpleTuner warning log करता है। + ### `--max_grounding_entities` - GLIGEN-style spatial annotations के लिए प्रति image grounding entities की अधिकतम संख्या। Default: 0 (disabled)। सामान्य मान: 4-16। @@ -1746,6 +1753,7 @@ usage: train.py [-h] --model_family [--pretrained_unet_subfolder PRETRAINED_UNET_SUBFOLDER] [--pretrained_t5_model_name_or_path PRETRAINED_T5_MODEL_NAME_OR_PATH] [--pretrained_gemma_model_name_or_path PRETRAINED_GEMMA_MODEL_NAME_OR_PATH] + [--qwen_text_encoder_model_name_or_path QWEN_TEXT_ENCODER_MODEL_NAME_OR_PATH] [--revision REVISION] [--variant VARIANT] [--base_model_default_dtype {bf16,fp32}] [--unet_attention_slice [UNET_ATTENTION_SLICE]] @@ -2077,6 +2085,8 @@ options: Path to pretrained T5 model --pretrained_gemma_model_name_or_path PRETRAINED_GEMMA_MODEL_NAME_OR_PATH Path to pretrained Gemma model + --qwen_text_encoder_model_name_or_path QWEN_TEXT_ENCODER_MODEL_NAME_OR_PATH + Path to pretrained Qwen text encoder model --revision REVISION Git branch/tag/commit for model version --variant VARIANT Model variant (e.g., fp16, bf16) --base_model_default_dtype {bf16,fp32} diff --git a/documentation/OPTIONS.ja.md b/documentation/OPTIONS.ja.md index b7d251c62..95b1eedcb 100644 --- a/documentation/OPTIONS.ja.md +++ b/documentation/OPTIONS.ja.md @@ -254,6 +254,13 @@ simpletuner configure config/foo/config.json - **内容**: 事前学習済み Gemma モデルのパス、または の識別子。 - **理由**: Gemma 系モデル(例: LTX-2、Sana、Lumina2)を学習する際、ベース拡散モデルのパスを変えずに Gemma 重みの参照先を指定できます。 +### `--qwen_text_encoder_model_name_or_path` + +- **内容**: 事前学習済み Qwen テキストエンコーダーモデルのパス、または の識別子。 +- **既定**: `None`(選択したモデルが定義する Qwen テキストエンコーダーの参照元を使用します) +- **理由**: Hugging Face キャッシュを編集せずに、Qwen 系モデルファミリーの Qwen テキストエンコーダーを共有または置き換えるために使用します。 +- **注記**: Qwen テキストエンコーダーが 1 つのモデルファミリーに適用されます。複数の Qwen エンコーダーを定義するモデルファミリーでは、このオプションは無視され、SimpleTuner が警告を記録します。 + ### `--max_grounding_entities` - GLIGEN スタイルの空間アノテーション用に、画像あたりのグラウンディングエンティティの最大数を指定します。デフォルト: 0(無効)。一般的な値: 4-16。 @@ -1749,6 +1756,7 @@ usage: train.py [-h] --model_family [--pretrained_unet_subfolder PRETRAINED_UNET_SUBFOLDER] [--pretrained_t5_model_name_or_path PRETRAINED_T5_MODEL_NAME_OR_PATH] [--pretrained_gemma_model_name_or_path PRETRAINED_GEMMA_MODEL_NAME_OR_PATH] + [--qwen_text_encoder_model_name_or_path QWEN_TEXT_ENCODER_MODEL_NAME_OR_PATH] [--revision REVISION] [--variant VARIANT] [--base_model_default_dtype {bf16,fp32}] [--unet_attention_slice [UNET_ATTENTION_SLICE]] @@ -2079,6 +2087,8 @@ options: Path to pretrained T5 model --pretrained_gemma_model_name_or_path PRETRAINED_GEMMA_MODEL_NAME_OR_PATH Path to pretrained Gemma model + --qwen_text_encoder_model_name_or_path QWEN_TEXT_ENCODER_MODEL_NAME_OR_PATH + Path to pretrained Qwen text encoder model --revision REVISION Git branch/tag/commit for model version --variant VARIANT Model variant (e.g., fp16, bf16) --base_model_default_dtype {bf16,fp32} diff --git a/documentation/OPTIONS.md b/documentation/OPTIONS.md index ca83f86c3..5a9e03f13 100644 --- a/documentation/OPTIONS.md +++ b/documentation/OPTIONS.md @@ -253,6 +253,13 @@ Where `foo` is your config environment - or just use `config/config.json` if you - **What**: Path to the pretrained Gemma model or its identifier from . - **Why**: When training Gemma-based models (for example LTX-2, Sana, or Lumina2), you can point at a shared Gemma checkpoint without changing the base diffusion model path. +### `--qwen_text_encoder_model_name_or_path` + +- **What**: Path to a pretrained Qwen text encoder model or its identifier from . +- **Default**: `None` (use the Qwen text encoder source defined by the selected model). +- **Why**: Use this to share or replace the Qwen text encoder used by Qwen-based model families without editing the Hugging Face cache. +- **Notes**: This applies to model families with one Qwen text encoder. If a model family defines multiple Qwen text encoders, the option is ignored and SimpleTuner logs a warning. + ### `--max_grounding_entities` - **What**: Maximum number of grounding entities per image for GLIGEN-style spatial annotations. @@ -1752,6 +1759,7 @@ usage: train.py [-h] --model_family [--pretrained_unet_subfolder PRETRAINED_UNET_SUBFOLDER] [--pretrained_t5_model_name_or_path PRETRAINED_T5_MODEL_NAME_OR_PATH] [--pretrained_gemma_model_name_or_path PRETRAINED_GEMMA_MODEL_NAME_OR_PATH] + [--qwen_text_encoder_model_name_or_path QWEN_TEXT_ENCODER_MODEL_NAME_OR_PATH] [--revision REVISION] [--variant VARIANT] [--base_model_default_dtype {bf16,fp32}] [--unet_attention_slice [UNET_ATTENTION_SLICE]] @@ -2083,6 +2091,8 @@ options: Path to pretrained T5 model --pretrained_gemma_model_name_or_path PRETRAINED_GEMMA_MODEL_NAME_OR_PATH Path to pretrained Gemma model + --qwen_text_encoder_model_name_or_path QWEN_TEXT_ENCODER_MODEL_NAME_OR_PATH + Path to pretrained Qwen text encoder model --revision REVISION Git branch/tag/commit for model version --variant VARIANT Model variant (e.g., fp16, bf16) --base_model_default_dtype {bf16,fp32} diff --git a/documentation/OPTIONS.pt-BR.md b/documentation/OPTIONS.pt-BR.md index bb934b7e0..ae3e2023d 100644 --- a/documentation/OPTIONS.pt-BR.md +++ b/documentation/OPTIONS.pt-BR.md @@ -253,6 +253,13 @@ Onde `foo` e seu ambiente de config — ou use `config/config.json` se nao estiv - **O que**: Caminho para o modelo Gemma pre-treinado ou seu identificador em . - **Por que**: Ao treinar modelos baseados em Gemma (por exemplo LTX-2, Sana ou Lumina2), voce pode apontar para um checkpoint Gemma compartilhado sem mudar o caminho do modelo base de difusao. +### `--qwen_text_encoder_model_name_or_path` + +- **O que**: Caminho para um encoder de texto Qwen pre-treinado ou seu identificador em . +- **Padrao**: `None` (usa a origem do encoder de texto Qwen definida pelo modelo selecionado). +- **Por que**: Use para compartilhar ou substituir o encoder de texto Qwen em familias de modelos baseadas em Qwen sem editar o cache do Hugging Face. +- **Notas**: Aplica-se a familias de modelos com um unico encoder de texto Qwen. Se uma familia definir varios encoders Qwen, a opcao e ignorada e o SimpleTuner registra um aviso. + ### `--max_grounding_entities` - Numero maximo de entidades de grounding por imagem para anotacoes espaciais no estilo GLIGEN. Padrao: 0 (desabilitado). Valores tipicos: 4-16. @@ -1744,6 +1751,7 @@ usage: train.py [-h] --model_family [--pretrained_unet_subfolder PRETRAINED_UNET_SUBFOLDER] [--pretrained_t5_model_name_or_path PRETRAINED_T5_MODEL_NAME_OR_PATH] [--pretrained_gemma_model_name_or_path PRETRAINED_GEMMA_MODEL_NAME_OR_PATH] + [--qwen_text_encoder_model_name_or_path QWEN_TEXT_ENCODER_MODEL_NAME_OR_PATH] [--revision REVISION] [--variant VARIANT] [--base_model_default_dtype {bf16,fp32}] [--unet_attention_slice [UNET_ATTENTION_SLICE]] @@ -2074,6 +2082,8 @@ options: Path to pretrained T5 model --pretrained_gemma_model_name_or_path PRETRAINED_GEMMA_MODEL_NAME_OR_PATH Path to pretrained Gemma model + --qwen_text_encoder_model_name_or_path QWEN_TEXT_ENCODER_MODEL_NAME_OR_PATH + Path to pretrained Qwen text encoder model --revision REVISION Git branch/tag/commit for model version --variant VARIANT Model variant (e.g., fp16, bf16) --base_model_default_dtype {bf16,fp32} diff --git a/documentation/OPTIONS.zh.md b/documentation/OPTIONS.zh.md index c64afa903..5db8f63fe 100644 --- a/documentation/OPTIONS.zh.md +++ b/documentation/OPTIONS.zh.md @@ -254,6 +254,13 @@ simpletuner configure config/foo/config.json - **内容**:预训练 Gemma 模型路径或 上的标识符。 - **原因**:训练 Gemma 系模型(例如 LTX-2、Sana、Lumina2)时,可单独指定 Gemma 权重来源,而无需更换基础扩散模型路径。 +### `--qwen_text_encoder_model_name_or_path` + +- **内容**:预训练 Qwen 文本编码器模型路径,或 上的标识符。 +- **默认**:`None`(使用所选模型定义的 Qwen 文本编码器来源)。 +- **原因**:用于在 Qwen 系模型家族中共享或替换 Qwen 文本编码器,而无需编辑 Hugging Face 缓存。 +- **说明**:此选项适用于只有一个 Qwen 文本编码器的模型家族。如果某个模型家族定义了多个 Qwen 编码器,该选项会被忽略,SimpleTuner 会记录警告。 + ### `--max_grounding_entities` - 每张图像用于 GLIGEN 风格空间标注的最大 grounding 实体数。默认值:0(禁用)。典型值:4-16。 @@ -1751,6 +1758,7 @@ usage: train.py [-h] --model_family [--pretrained_unet_subfolder PRETRAINED_UNET_SUBFOLDER] [--pretrained_t5_model_name_or_path PRETRAINED_T5_MODEL_NAME_OR_PATH] [--pretrained_gemma_model_name_or_path PRETRAINED_GEMMA_MODEL_NAME_OR_PATH] + [--qwen_text_encoder_model_name_or_path QWEN_TEXT_ENCODER_MODEL_NAME_OR_PATH] [--revision REVISION] [--variant VARIANT] [--base_model_default_dtype {bf16,fp32}] [--unet_attention_slice [UNET_ATTENTION_SLICE]] @@ -2081,6 +2089,8 @@ options: Path to pretrained T5 model --pretrained_gemma_model_name_or_path PRETRAINED_GEMMA_MODEL_NAME_OR_PATH Path to pretrained Gemma model + --qwen_text_encoder_model_name_or_path QWEN_TEXT_ENCODER_MODEL_NAME_OR_PATH + Path to pretrained Qwen text encoder model --revision REVISION Git branch/tag/commit for model version --variant VARIANT Model variant (e.g., fp16, bf16) --base_model_default_dtype {bf16,fp32} diff --git a/documentation/quickstart/FLUX2.es.md b/documentation/quickstart/FLUX2.es.md index 31f2841bf..dc8b27847 100644 --- a/documentation/quickstart/FLUX2.es.md +++ b/documentation/quickstart/FLUX2.es.md @@ -25,7 +25,7 @@ Para seleccionar una variante, configura `model_flavour` en tu configuración: } ``` -> **Importante**: Para `klein-4b` y `klein-9b`, deja `pretrained_text_encoder_model_name_or_path` sin definir a menos que realmente quieras reemplazar el codificador Qwen3 incluido. Si configuras ese campo, anulas el valor predeterminado de Klein y puedes provocar la descarga de otro codificador de texto. +> **Importante**: Para `klein-4b` y `klein-9b`, deja `qwen_text_encoder_model_name_or_path` sin definir para usar el codificador Qwen3 incluido. Configuralo solo cuando reemplaces el codificador Qwen3 incluido por otra fuente compatible con Qwen. ## Resumen del modelo diff --git a/documentation/quickstart/FLUX2.hi.md b/documentation/quickstart/FLUX2.hi.md index b3c07dc30..5e30432c7 100644 --- a/documentation/quickstart/FLUX2.hi.md +++ b/documentation/quickstart/FLUX2.hi.md @@ -25,7 +25,7 @@ FLUX.2 तीन वेरिएंट में आता है: } ``` -> **महत्वपूर्ण**: `klein-4b` और `klein-9b` के लिए `pretrained_text_encoder_model_name_or_path` को unset छोड़ें, जब तक कि आप bundled Qwen3 text encoder को जानबूझकर बदलना न चाहते हों। इस field को सेट करने पर Klein का default override हो जाता है और किसी दूसरे text encoder का download शुरू हो सकता है। +> **महत्वपूर्ण**: `klein-4b` और `klein-9b` के लिए bundled Qwen3 text encoder उपयोग करने के लिए `qwen_text_encoder_model_name_or_path` को unset छोड़ें। इसे केवल तब सेट करें जब bundled Qwen3 text encoder को किसी दूसरे Qwen-compatible source से बदलना हो। ## मॉडल ओवरव्यू diff --git a/documentation/quickstart/FLUX2.ja.md b/documentation/quickstart/FLUX2.ja.md index 5831ac7c2..5359f2fb3 100644 --- a/documentation/quickstart/FLUX2.ja.md +++ b/documentation/quickstart/FLUX2.ja.md @@ -25,7 +25,7 @@ FLUX.2は3つのバリアントがあります: } ``` -> **重要**: `klein-4b` と `klein-9b` では、同梱のQwen3テキストエンコーダーを意図的に置き換えたい場合を除き、`pretrained_text_encoder_model_name_or_path` は設定しないでください。この項目を設定するとKleinのデフォルトを上書きし、別のテキストエンコーダーのダウンロードが発生することがあります。 +> **重要**: `klein-4b` と `klein-9b` では、同梱の Qwen3 テキストエンコーダーを使用するため `qwen_text_encoder_model_name_or_path` は未設定のままにします。同梱の Qwen3 テキストエンコーダーを別の Qwen 互換ソースに置き換える場合のみ設定してください。 ## モデル概要 diff --git a/documentation/quickstart/FLUX2.md b/documentation/quickstart/FLUX2.md index 5b96d49c5..a71ed561b 100644 --- a/documentation/quickstart/FLUX2.md +++ b/documentation/quickstart/FLUX2.md @@ -25,7 +25,7 @@ To select a variant, set `model_flavour` in your config: } ``` -> **Important**: For `klein-4b` and `klein-9b`, leave `pretrained_text_encoder_model_name_or_path` unset unless you intentionally want to replace the bundled Qwen3 text encoder. Setting that field overrides the Klein default and can trigger downloads of a different text encoder. +> **Important**: For `klein-4b` and `klein-9b`, leave `qwen_text_encoder_model_name_or_path` unset to use the bundled Qwen3 text encoder. Set it only when replacing the bundled Qwen3 text encoder with another Qwen-compatible source. ## Model Overview diff --git a/documentation/quickstart/FLUX2.pt-BR.md b/documentation/quickstart/FLUX2.pt-BR.md index 2d7719f1a..8fac19e0b 100644 --- a/documentation/quickstart/FLUX2.pt-BR.md +++ b/documentation/quickstart/FLUX2.pt-BR.md @@ -25,7 +25,7 @@ Para selecionar uma variante, defina `model_flavour` na sua configuração: } ``` -> **Importante**: Para `klein-4b` e `klein-9b`, deixe `pretrained_text_encoder_model_name_or_path` sem definir, a menos que você realmente queira substituir o encoder Qwen3 incluído. Ao definir esse campo, você sobrescreve o padrão do Klein e pode disparar o download de outro encoder de texto. +> **Importante**: Para `klein-4b` e `klein-9b`, deixe `qwen_text_encoder_model_name_or_path` sem definir para usar o encoder Qwen3 incluído. Defina-o apenas ao substituir o encoder Qwen3 incluído por outra fonte compatível com Qwen. ## Visão geral do modelo diff --git a/documentation/quickstart/FLUX2.zh.md b/documentation/quickstart/FLUX2.zh.md index 027002850..7f17fa61c 100644 --- a/documentation/quickstart/FLUX2.zh.md +++ b/documentation/quickstart/FLUX2.zh.md @@ -25,7 +25,7 @@ FLUX.2 有三个变体: } ``` -> **重要**:对于 `klein-4b` 和 `klein-9b`,除非你明确想替换内置的 Qwen3 文本编码器,否则不要设置 `pretrained_text_encoder_model_name_or_path`。设置这个字段会覆盖 Klein 的默认行为,并可能触发下载其他文本编码器。 +> **重要**:对于 `klein-4b` 和 `klein-9b`,请将 `qwen_text_encoder_model_name_or_path` 留空以使用内置的 Qwen3 文本编码器。仅在需要用另一个兼容 Qwen 的来源替换内置 Qwen3 文本编码器时设置它。 ## 模型概述 diff --git a/simpletuner/helpers/configuration/env_file.py b/simpletuner/helpers/configuration/env_file.py index 5faa513cd..635bb0a1c 100644 --- a/simpletuner/helpers/configuration/env_file.py +++ b/simpletuner/helpers/configuration/env_file.py @@ -32,6 +32,7 @@ "MODEL_TYPE": "--model_type", "MODEL_NAME": "--pretrained_model_name_or_path", "MODEL_FAMILY": "--model_family", + "QWEN_TEXT_ENCODER_MODEL_NAME_OR_PATH": "--qwen_text_encoder_model_name_or_path", "TRAIN_BATCH_SIZE": "--train_batch_size", "USE_GRADIENT_CHECKPOINTING": "--gradient_checkpointing", "ENABLE_CHUNKED_FEED_FORWARD": "--enable_chunked_feed_forward", diff --git a/simpletuner/helpers/models/ace_step/model.py b/simpletuner/helpers/models/ace_step/model.py index cae5dd1fe..c2c4362b6 100644 --- a/simpletuner/helpers/models/ace_step/model.py +++ b/simpletuner/helpers/models/ace_step/model.py @@ -378,9 +378,13 @@ def _resolve_v15_layout(self, base_path: Optional[str] = None) -> Optional[Dict[ if available_variants: variant_dir = available_variants[0] + qwen_text_encoder_path = self._get_optional_config_model_path("qwen_text_encoder_model_name_or_path") tokenizer_dir = shared_root / self.V15_SHARED_TEXT_ENCODER_SUBFOLDER vae_dir = shared_root / self.V15_SHARED_VAE_SUBFOLDER - if variant_dir is None or not tokenizer_dir.is_dir() or not vae_dir.is_dir(): + if variant_dir is None or not vae_dir.is_dir(): + self._v15_layout = None + return None + if not qwen_text_encoder_path and not tokenizer_dir.is_dir(): self._v15_layout = None return None @@ -404,7 +408,7 @@ def _resolve_v15_layout(self, base_path: Optional[str] = None) -> Optional[Dict[ self._v15_layout = { "root_path": str(shared_root), "variant_path": str(variant_dir), - "tokenizer_path": str(tokenizer_dir), + "tokenizer_path": qwen_text_encoder_path or str(tokenizer_dir), "vae_path": str(vae_dir), "silence_latent_path": str(silence_path), } diff --git a/simpletuner/helpers/models/anima/model.py b/simpletuner/helpers/models/anima/model.py index 3cbf8544d..80aff1c62 100644 --- a/simpletuner/helpers/models/anima/model.py +++ b/simpletuner/helpers/models/anima/model.py @@ -283,20 +283,21 @@ def load_model(self, move_to_device: bool = True): return super().load_model(move_to_device=move_to_device) def _prompt_tokenizer_sources(self) -> tuple[str, str]: + qwen_tokenizer_source = self._get_optional_config_model_path("qwen_text_encoder_model_name_or_path") model_path = getattr(self.config, "pretrained_model_name_or_path", None) if isinstance(model_path, str): model_dir = Path(model_path) qwen_dir = model_dir / "tokenizer" t5_dir = model_dir / "t5_tokenizer" if qwen_dir.is_dir() and t5_dir.is_dir(): - return str(qwen_dir), str(t5_dir) + return qwen_tokenizer_source or str(qwen_dir), str(t5_dir) qwen_dir = model_dir / "prompt_tokenizer_qwen" t5_dir = model_dir / "prompt_tokenizer_t5" if qwen_dir.is_dir() and t5_dir.is_dir(): - return str(qwen_dir), str(t5_dir) + return qwen_tokenizer_source or str(qwen_dir), str(t5_dir) if self._uses_diffusers_repo_layout(): - return f"{model_path}::tokenizer", f"{model_path}::t5_tokenizer" - return _QWEN_TOKENIZER_SOURCE, _T5_TOKENIZER_SOURCE + return qwen_tokenizer_source or f"{model_path}::tokenizer", f"{model_path}::t5_tokenizer" + return qwen_tokenizer_source or _QWEN_TOKENIZER_SOURCE, _T5_TOKENIZER_SOURCE def load_text_tokenizer(self): qwen_source, t5_source = self._prompt_tokenizer_sources() @@ -313,8 +314,10 @@ def load_text_tokenizer(self): def load_text_encoder(self, move_to_device: bool = True): if self.text_encoders is not None and len(self.text_encoders) > 0: return + qwen_text_encoder_path = self._get_optional_config_model_path("qwen_text_encoder_model_name_or_path") model_path = ( - getattr(self.config, "pretrained_text_encoder_model_name_or_path", None) + qwen_text_encoder_path + or getattr(self.config, "pretrained_text_encoder_model_name_or_path", None) or self.config.pretrained_model_name_or_path ) revision = getattr(self.config, "text_encoder_revision", None) or getattr(self.config, "revision", None) @@ -324,6 +327,34 @@ def load_text_encoder(self, move_to_device: bool = True): execution_device=self.accelerator.device.type, ) load_device = self.accelerator.device.type if move_to_device else "cpu" + if qwen_text_encoder_path: + if os.path.isfile(model_path): + text_encoder = load_text_encoder_single_file( + file_path=model_path, + device=load_device, + dtype=dtype, + ) + else: + load_kwargs = { + "pretrained_model_name_or_path": model_path, + "revision": revision, + "torch_dtype": dtype, + "local_files_only": bool(getattr(self.config, "local_files_only", False)), + "cache_dir": getattr(self.config, "cache_dir", None), + "force_download": bool(getattr(self.config, "force_download", False)), + } + self._add_hf_token_kwarg(load_kwargs) + text_encoder = Qwen3Model.from_pretrained(**load_kwargs) + text_encoder.eval().requires_grad_(False) + text_encoder.to(device=load_device, dtype=dtype) + self.text_encoders = [text_encoder] + self.text_encoder = text_encoder + self.text_encoder_1 = text_encoder + if not move_to_device: + text_encoder.to("cpu") + if getattr(self, "prompt_tokenizer", None) is None: + self.load_text_tokenizer() + return if self._uses_diffusers_repo_layout(model_path, component_subfolder="text_encoder"): load_kwargs = { "pretrained_model_name_or_path": model_path, diff --git a/simpletuner/helpers/models/boogu_image/model.py b/simpletuner/helpers/models/boogu_image/model.py index 4b789e0c6..473489878 100644 --- a/simpletuner/helpers/models/boogu_image/model.py +++ b/simpletuner/helpers/models/boogu_image/model.py @@ -16,8 +16,6 @@ get_sdnq_presets, get_torchao_presets, ) -from simpletuner.helpers.models.common import ImageModelFoundation, ModelTypes, PipelineTypes, PredictionTypes -from simpletuner.helpers.models.flux.model import Flux from simpletuner.helpers.models.boogu_image.pipeline import BooguImagePipeline from simpletuner.helpers.models.boogu_image.pipeline_edit import BooguImageEditPipeline from simpletuner.helpers.models.boogu_image.pipeline_img2img import BooguImageImg2ImgPipeline @@ -26,6 +24,8 @@ FlowMatchEulerDiscreteScheduler, ) from simpletuner.helpers.models.boogu_image.transformer import BooguImageTransformer2DModel +from simpletuner.helpers.models.common import ImageModelFoundation, ModelTypes, PipelineTypes, PredictionTypes +from simpletuner.helpers.models.flux.model import Flux from simpletuner.helpers.models.registry import ModelRegistry from simpletuner.helpers.training.deepspeed import deepspeed_zero_init_disabled_context_manager @@ -188,8 +188,8 @@ def text_embed_cache_key(self): def _load_processor_for_pipeline(self): if self.processor is not None: return self.processor - processor_path = getattr(self.config, "processor_pretrained_model_name_or_path", None) or self._model_config_path() - processor_subfolder = getattr(self.config, "processor_subfolder", self.PROCESSOR_SUBFOLDER) + processor_path = self._resolve_qwen_processor_path(self._model_config_path()) + processor_subfolder = self._resolve_qwen_processor_subfolder(self.PROCESSOR_SUBFOLDER) processor_kwargs = { "pretrained_model_name_or_path": processor_path, "subfolder": processor_subfolder, diff --git a/simpletuner/helpers/models/common.py b/simpletuner/helpers/models/common.py index 5253afd7d..df365cd7b 100644 --- a/simpletuner/helpers/models/common.py +++ b/simpletuner/helpers/models/common.py @@ -2517,6 +2517,84 @@ def _is_gemma_component(component_cls) -> bool: component_name = getattr(component_cls, "__name__", "") return "gemma" in component_name.lower() + @staticmethod + def _component_value_mentions_qwen(value) -> bool: + if value is None: + return False + if isinstance(value, str): + return "qwen" in value.lower() + component_name = getattr(value, "__name__", "") + return "qwen" in component_name.lower() + + @staticmethod + def _optional_model_path(value) -> Optional[str]: + if value is None or value is False: + return None + if isinstance(value, os.PathLike): + return os.fspath(value) + if isinstance(value, str): + value = value.strip() + return value or None + return None + + def _get_optional_config_model_path(self, name: str) -> Optional[str]: + return self._optional_model_path(getattr(self.config, name, None)) + + def _is_qwen_text_encoder_config(self, text_encoder_config: dict) -> bool: + return any( + self._component_value_mentions_qwen(text_encoder_config.get(key)) + for key in ("name", "model", "tokenizer", "path") + ) + + def _qwen_text_encoder_config_count(self) -> int: + text_encoder_configuration = getattr(self, "TEXT_ENCODER_CONFIGURATION", None) or {} + return sum( + 1 + for text_encoder_config in text_encoder_configuration.values() + if self._is_qwen_text_encoder_config(text_encoder_config) + ) + + def _warn_qwen_text_encoder_override_ignored(self, qwen_count: int) -> None: + if getattr(self, "_qwen_text_encoder_override_warning_emitted", False): + return + logger.warning( + "Ignoring qwen_text_encoder_model_name_or_path for %s because this model defines %s Qwen text encoders.", + self.NAME, + qwen_count, + ) + self._qwen_text_encoder_override_warning_emitted = True + + def _uses_qwen_text_encoder_override(self, text_encoder_config: dict) -> bool: + qwen_path = self._get_optional_config_model_path("qwen_text_encoder_model_name_or_path") + if not qwen_path or not self._is_qwen_text_encoder_config(text_encoder_config): + return False + qwen_count = self._qwen_text_encoder_config_count() + if qwen_count == 1: + return True + if qwen_count > 1: + self._warn_qwen_text_encoder_override_ignored(qwen_count) + return False + + def _resolve_text_encoder_subfolder(self, text_encoder_config: dict, key: str, default=None): + if self._uses_qwen_text_encoder_override(text_encoder_config): + return None + return text_encoder_config.get(key, default) + + def _resolve_qwen_processor_path(self, default_path: str) -> str: + return ( + self._get_optional_config_model_path("processor_pretrained_model_name_or_path") + or self._get_optional_config_model_path("qwen_text_encoder_model_name_or_path") + or default_path + ) + + def _resolve_qwen_processor_subfolder(self, default_subfolder: Optional[str]) -> Optional[str]: + processor_subfolder = getattr(self.config, "processor_subfolder", None) + if isinstance(processor_subfolder, str): + return processor_subfolder + if self._get_optional_config_model_path("qwen_text_encoder_model_name_or_path"): + return None + return default_subfolder + def _resolve_text_encoder_path(self, text_encoder_config: dict) -> str: text_encoder_path = get_model_config_path(self.config.model_family, self.config.pretrained_model_name_or_path) config_path = text_encoder_config.get("path", None) @@ -2525,6 +2603,9 @@ def _resolve_text_encoder_path(self, text_encoder_config: dict) -> str: gemma_path = getattr(self.config, "pretrained_gemma_model_name_or_path", None) if gemma_path and self._is_gemma_component(text_encoder_config.get("model")): text_encoder_path = gemma_path + qwen_path = self._get_optional_config_model_path("qwen_text_encoder_model_name_or_path") + if qwen_path and self._uses_qwen_text_encoder_override(text_encoder_config): + text_encoder_path = qwen_path return text_encoder_path def load_text_tokenizer(self): @@ -2551,7 +2632,11 @@ def load_text_tokenizer(self): for attr_name, text_encoder_config in self.TEXT_ENCODER_CONFIGURATION.items(): tokenizer_idx += 1 tokenizer_cls = text_encoder_config.get("tokenizer") - tokenizer_kwargs["subfolder"] = text_encoder_config.get("tokenizer_subfolder", "tokenizer") + tokenizer_kwargs["subfolder"] = self._resolve_text_encoder_subfolder( + text_encoder_config, + "tokenizer_subfolder", + "tokenizer", + ) tokenizer_kwargs["use_fast"] = text_encoder_config.get("use_fast", False) tokenizer_kwargs["pretrained_model_name_or_path"] = self._resolve_text_encoder_path(text_encoder_config) logger.info(f"Loading tokenizer {tokenizer_idx}: {tokenizer_cls.__name__} with args: {tokenizer_kwargs}") @@ -2673,7 +2758,12 @@ def load_text_encoder(self, move_to_device: bool = True): "pretrained_model_name_or_path": text_encoder_path, "variant": self.config.variant, "revision": self.config.revision, - "subfolder": text_encoder_config.get("subfolder", "text_encoder") or "", + "subfolder": self._resolve_text_encoder_subfolder( + text_encoder_config, + "subfolder", + "text_encoder", + ) + or "", **extra_kwargs, } accelerator = getattr(self, "accelerator", None) diff --git a/simpletuner/helpers/models/flux2/model.py b/simpletuner/helpers/models/flux2/model.py index 0f8d0447a..b470dc690 100644 --- a/simpletuner/helpers/models/flux2/model.py +++ b/simpletuner/helpers/models/flux2/model.py @@ -455,7 +455,9 @@ def _load_text_encoder_qwen3(self, move_to_device: bool = True): # For Klein models, text encoder is bundled in the model repo under "text_encoder" subfolder # and tokenizer is in a separate "tokenizer" subfolder model_path = self.config.pretrained_model_name_or_path - text_encoder_path = getattr(self.config, "pretrained_text_encoder_model_name_or_path", None) + text_encoder_path = self._get_optional_config_model_path("qwen_text_encoder_model_name_or_path") or getattr( + self.config, "pretrained_text_encoder_model_name_or_path", None + ) if text_encoder_path is None: text_encoder_path = model_path text_encoder_subfolder = "text_encoder" diff --git a/simpletuner/helpers/models/hunyuanvideo/model.py b/simpletuner/helpers/models/hunyuanvideo/model.py index 524059ec1..472d78d64 100644 --- a/simpletuner/helpers/models/hunyuanvideo/model.py +++ b/simpletuner/helpers/models/hunyuanvideo/model.py @@ -282,7 +282,11 @@ def load_text_encoder(self, move_to_device: bool = True): Load the Qwen2.5 VL text encoder and ByT5 glyph encoder. """ device = self.accelerator.device if move_to_device else torch.device("cpu") - qwen_path = getattr(self.config, "hunyuan_text_encoder_path", None) or self.TEXT_ENCODER_REPO + qwen_path = ( + self._get_optional_config_model_path("qwen_text_encoder_model_name_or_path") + or getattr(self.config, "hunyuan_text_encoder_path", None) + or self.TEXT_ENCODER_REPO + ) logger.info(f"Loading HunyuanVideo text encoder from {qwen_path}") tokenizer = Qwen2Tokenizer.from_pretrained(qwen_path) diff --git a/simpletuner/helpers/models/ideogram/model.py b/simpletuner/helpers/models/ideogram/model.py index b3407f41d..9d585a711 100644 --- a/simpletuner/helpers/models/ideogram/model.py +++ b/simpletuner/helpers/models/ideogram/model.py @@ -103,12 +103,13 @@ def load_vae(self, move_to_device: bool = True): def load_text_encoder(self, move_to_device: bool = True): repo_id = getattr(self.config, "pretrained_model_name_or_path", None) or self.HUGGINGFACE_PATHS["fp8"] pipe_config = Ideogram4PipelineConfig(weights_repo=repo_id) + qwen_repo_id = self._get_optional_config_model_path("qwen_text_encoder_model_name_or_path") tokenizer, text_encoder = _load_qwen3_vl( - repo_id, + qwen_repo_id or repo_id, self.accelerator.device, self.config.weight_dtype, - tokenizer_subfolder=pipe_config.tokenizer_subfolder, - text_encoder_subfolder=pipe_config.text_encoder_subfolder, + tokenizer_subfolder=None if qwen_repo_id else pipe_config.tokenizer_subfolder, + text_encoder_subfolder=None if qwen_repo_id else pipe_config.text_encoder_subfolder, ) self.tokenizers = [tokenizer] self.text_encoders = [text_encoder] diff --git a/simpletuner/helpers/models/krea2/model.py b/simpletuner/helpers/models/krea2/model.py index 4d4e803b9..5d5c86b27 100644 --- a/simpletuner/helpers/models/krea2/model.py +++ b/simpletuner/helpers/models/krea2/model.py @@ -178,8 +178,8 @@ def _load_processor_for_pipeline(self): if self.processor is not None: return self.processor - processor_path = getattr(self.config, "processor_pretrained_model_name_or_path", None) or self.PROCESSOR_PATH - processor_subfolder = getattr(self.config, "processor_subfolder", self.PROCESSOR_SUBFOLDER) + processor_path = self._resolve_qwen_processor_path(self.PROCESSOR_PATH) + processor_subfolder = self._resolve_qwen_processor_subfolder(self.PROCESSOR_SUBFOLDER) processor_revision = getattr(self.config, "processor_revision", getattr(self.config, "revision", None)) processor_kwargs = {"pretrained_model_name_or_path": processor_path} diff --git a/simpletuner/helpers/models/longcat_image/model.py b/simpletuner/helpers/models/longcat_image/model.py index 86eb54f04..41671375d 100644 --- a/simpletuner/helpers/models/longcat_image/model.py +++ b/simpletuner/helpers/models/longcat_image/model.py @@ -211,8 +211,15 @@ def _load_text_processor_for_pipeline(self): text_processor = getattr(self, "text_processor", None) if text_processor is not None: return text_processor - model_path = get_model_config_path(self.config.model_family, self.config.pretrained_model_name_or_path) - text_processor = AutoProcessor.from_pretrained(model_path, subfolder="text_processor") + qwen_text_encoder_path = self._get_optional_config_model_path("qwen_text_encoder_model_name_or_path") + model_path = qwen_text_encoder_path or get_model_config_path( + self.config.model_family, + self.config.pretrained_model_name_or_path, + ) + processor_kwargs = {"pretrained_model_name_or_path": model_path} + if not qwen_text_encoder_path: + processor_kwargs["subfolder"] = "text_processor" + text_processor = AutoProcessor.from_pretrained(**processor_kwargs) self.text_processor = text_processor return text_processor diff --git a/simpletuner/helpers/models/mageflow/model.py b/simpletuner/helpers/models/mageflow/model.py index c8c99ed72..ea9283bc7 100644 --- a/simpletuner/helpers/models/mageflow/model.py +++ b/simpletuner/helpers/models/mageflow/model.py @@ -338,8 +338,8 @@ def _select_crepa_hidden_states(self, prepared_batch: dict, hidden_states_buffer def _load_processor_for_pipeline(self): if self.processor is not None: return self.processor - processor_path = getattr(self.config, "processor_pretrained_model_name_or_path", None) or self._model_config_path() - processor_subfolder = getattr(self.config, "processor_subfolder", self.PROCESSOR_SUBFOLDER) + processor_path = self._resolve_qwen_processor_path(self._model_config_path()) + processor_subfolder = self._resolve_qwen_processor_subfolder(self.PROCESSOR_SUBFOLDER) self.processor = self.PROCESSOR_CLASS.from_pretrained( processor_path, subfolder=processor_subfolder, diff --git a/simpletuner/helpers/models/qwen_image/model.py b/simpletuner/helpers/models/qwen_image/model.py index 9b4a09179..f36e2f0eb 100644 --- a/simpletuner/helpers/models/qwen_image/model.py +++ b/simpletuner/helpers/models/qwen_image/model.py @@ -245,8 +245,8 @@ def _load_processor_for_pipeline(self): if processor_cls is None: return None - processor_path = getattr(self.config, "processor_pretrained_model_name_or_path", None) or self._model_config_path() - processor_subfolder = getattr(self.config, "processor_subfolder", self.PROCESSOR_SUBFOLDER) + processor_path = self._resolve_qwen_processor_path(self._model_config_path()) + processor_subfolder = self._resolve_qwen_processor_subfolder(self.PROCESSOR_SUBFOLDER) processor_revision = getattr(self.config, "processor_revision", getattr(self.config, "revision", None)) processor_kwargs = {"pretrained_model_name_or_path": processor_path} diff --git a/simpletuner/simpletuner_sdk/server/services/field_registry/sections/model.py b/simpletuner/simpletuner_sdk/server/services/field_registry/sections/model.py index b9151c28c..82f4baa6a 100644 --- a/simpletuner/simpletuner_sdk/server/services/field_registry/sections/model.py +++ b/simpletuner/simpletuner_sdk/server/services/field_registry/sections/model.py @@ -1131,6 +1131,26 @@ def _quant_label(value: str) -> str: ) ) + # Qwen Text Encoder Model Path + registry._add_field( + ConfigField( + name="qwen_text_encoder_model_name_or_path", + arg_name="--qwen_text_encoder_model_name_or_path", + ui_label="Qwen Model Path", + field_type=FieldType.TEXT, + tab="model", + section="model_config", + subsection="advanced_paths", + default_value=None, + placeholder="path/to/qwen", + help_text="Path to pretrained Qwen text encoder model", + tooltip="HuggingFace model ID or local path for the Qwen text encoder component.", + importance=ImportanceLevel.ADVANCED, + order=30, + documentation="OPTIONS.md#--qwen_text_encoder_model_name_or_path", + ) + ) + # Revision registry._add_field( ConfigField( @@ -1146,7 +1166,7 @@ def _quant_label(value: str) -> str: help_text="Git branch/tag/commit for model version", tooltip="Specific version of the model to load from HuggingFace. Useful for reproducible training.", importance=ImportanceLevel.ADVANCED, - order=30, + order=31, ) ) @@ -1165,7 +1185,7 @@ def _quant_label(value: str) -> str: help_text="Model variant (e.g., fp16, bf16)", tooltip="Specific variant of the model to load, such as precision variants.", importance=ImportanceLevel.ADVANCED, - order=31, + order=32, ) ) @@ -1187,7 +1207,7 @@ def _quant_label(value: str) -> str: help_text="Default precision for quantized base model weights", tooltip="Precision for non-quantized weights in quantized models. BF16 recommended for stability.", importance=ImportanceLevel.ADVANCED, - order=32, + order=33, ) ) @@ -1208,7 +1228,7 @@ def _quant_label(value: str) -> str: tooltip="Experimental feature for memory savings. May impact training quality. Only available for UNet-based architectures.", importance=ImportanceLevel.EXPERIMENTAL, model_specific=["sd15", "sd20", "sdxl", "deepfloyd"], - order=33, + order=34, ) ) diff --git a/tests/test_ace_step_model.py b/tests/test_ace_step_model.py index 4ccbbf89e..a96b807ec 100644 --- a/tests/test_ace_step_model.py +++ b/tests/test_ace_step_model.py @@ -24,6 +24,7 @@ def setUp(self): self.config.pretrained_model_name_or_path = "dummy_path" self.config.pretrained_transformer_model_name_or_path = None self.config.pretrained_transformer_subfolder = None + self.config.qwen_text_encoder_model_name_or_path = None self.config.model_flavour = "base" self.config.controlnet = False self.config.peft_lora_target_modules = None @@ -185,6 +186,21 @@ def test_resolve_v15_layout_uses_requested_variant(self): self.assertEqual(layout["tokenizer_path"], str(root / "Qwen3-Embedding-0.6B")) self.assertEqual(layout["vae_path"], str(root / "vae")) + def test_resolve_v15_layout_uses_qwen_text_encoder_override(self): + self.config.model_flavour = "v15-base" + self.config.qwen_text_encoder_model_name_or_path = "custom/qwen" + with TemporaryDirectory() as tmpdir: + root = Path(tmpdir) + (root / "vae").mkdir() + (root / "acestep-v15-base").mkdir() + torch.save(torch.zeros(1, 64, 4), root / "acestep-v15-base" / "silence_latent.pt") + + layout = self.model._resolve_v15_layout(str(root)) + + self.assertIsNotNone(layout) + self.assertEqual(layout["tokenizer_path"], "custom/qwen") + self.assertEqual(layout["vae_path"], str(root / "vae")) + def test_resolve_v15_layout_caches_negative_result_for_same_base_path(self): with TemporaryDirectory() as tmpdir: root = Path(tmpdir) diff --git a/tests/test_acestep_lora_targets.py b/tests/test_acestep_lora_targets.py index 3edb8171b..ad547b438 100644 --- a/tests/test_acestep_lora_targets.py +++ b/tests/test_acestep_lora_targets.py @@ -20,6 +20,7 @@ def setUp(self): self.config.pretrained_model_name_or_path = "dummy_path" self.config.pretrained_transformer_model_name_or_path = None self.config.pretrained_transformer_subfolder = None + self.config.qwen_text_encoder_model_name_or_path = None self.config.model_family = "ace_step" self.config.model_flavour = "base" self.config.peft_lora_target_modules = None diff --git a/tests/test_hunyuanvideo_model.py b/tests/test_hunyuanvideo_model.py index c9d356bca..dfb435d10 100644 --- a/tests/test_hunyuanvideo_model.py +++ b/tests/test_hunyuanvideo_model.py @@ -76,6 +76,53 @@ def test_load_text_encoder_registers_both_hunyuan_encoders_for_device_management self.assertIs(model.text_encoder_2, byt5_model) self.assertIs(model.get_text_encoder(1), byt5_model) + def test_load_text_encoder_prefers_qwen_text_encoder_override(self): + model = HunyuanVideo.__new__(HunyuanVideo) + model.accelerator = SimpleNamespace(device=torch.device("cuda:0")) + model.config = SimpleNamespace( + qwen_text_encoder_model_name_or_path="custom/qwen", + hunyuan_text_encoder_path="legacy/qwen", + glyph_byt5_repo="glyph/repo", + glyph_byt5_fallback_repo="glyph/fallback", + ) + model._ramtorch_text_encoders_requested = MagicMock(return_value=False) + model._ramtorch_text_encoder_percent = MagicMock(return_value=1.0) + model._apply_ramtorch_layers = MagicMock() + + qwen_tokenizer = MagicMock() + byt5_tokenizer = MagicMock() + text_encoder = MagicMock() + text_encoder.to.return_value = text_encoder + byt5_model = MagicMock() + byt5_model.to.return_value = byt5_model + + with ( + patch( + "simpletuner.helpers.models.hunyuanvideo.model.Qwen2Tokenizer.from_pretrained", + return_value=qwen_tokenizer, + ) as mock_qwen_tokenizer, + patch( + "simpletuner.helpers.models.hunyuanvideo.model.Qwen2_5_VLTextModel.from_pretrained", + return_value=text_encoder, + ) as mock_qwen_text_encoder, + patch( + "simpletuner.helpers.models.hunyuanvideo.model.ByT5Tokenizer.from_pretrained", + return_value=byt5_tokenizer, + ), + patch( + "simpletuner.helpers.models.hunyuanvideo.model.T5EncoderModel.from_pretrained", + return_value=byt5_model, + ), + patch( + "simpletuner.helpers.models.hunyuanvideo.model.hf_hub_download", + side_effect=RuntimeError("no glyph checkpoint"), + ), + ): + model.load_text_encoder(move_to_device=True) + + mock_qwen_tokenizer.assert_called_once_with("custom/qwen") + mock_qwen_text_encoder.assert_called_once_with("custom/qwen", torch_dtype=torch.bfloat16) + def test_model_supports_crepa_self_flow(self): model = HunyuanVideo.__new__(HunyuanVideo) self.assertTrue(model.supports_crepa_self_flow()) diff --git a/tests/test_qwen_text_encoder_override.py b/tests/test_qwen_text_encoder_override.py new file mode 100644 index 000000000..1e4e423fc --- /dev/null +++ b/tests/test_qwen_text_encoder_override.py @@ -0,0 +1,155 @@ +import unittest +from types import SimpleNamespace +from unittest.mock import patch + +from simpletuner.helpers.models.common import ImageModelFoundation + + +class QwenTextModel: + pass + + +class QwenTokenizer: + pass + + +class ClipTextModel: + pass + + +class ClipTokenizer: + pass + + +class DummyQwenFoundation(ImageModelFoundation): + NAME = "Dummy Qwen" + + def model_predict(self, prepared_batch, custom_timesteps: list = None): + raise NotImplementedError + + def _encode_prompts(self, prompts: list, is_negative_prompt: bool = False): + raise NotImplementedError + + def convert_text_embed_for_pipeline(self, text_embedding): + raise NotImplementedError + + def convert_negative_text_embed_for_pipeline(self, text_embedding): + raise NotImplementedError + + +class QwenTextEncoderOverrideTests(unittest.TestCase): + def _model(self, text_encoder_configuration): + model = object.__new__(DummyQwenFoundation) + model.config = SimpleNamespace( + model_family="qwen_image", + pretrained_model_name_or_path="base/model", + qwen_text_encoder_model_name_or_path="custom/qwen", + ) + model.TEXT_ENCODER_CONFIGURATION = text_encoder_configuration + return model + + def test_single_qwen_encoder_uses_qwen_override_and_clears_component_subfolders(self): + qwen_config = { + "name": "Qwen2.5-VL", + "tokenizer": QwenTokenizer, + "tokenizer_subfolder": "tokenizer", + "model": QwenTextModel, + "subfolder": "text_encoder", + } + clip_config = { + "name": "CLIP-L/14", + "tokenizer": ClipTokenizer, + "tokenizer_subfolder": "tokenizer_2", + "model": ClipTextModel, + "subfolder": "text_encoder_2", + } + model = self._model( + { + "text_encoder": qwen_config, + "text_encoder_2": clip_config, + } + ) + + self.assertEqual(model._resolve_text_encoder_path(qwen_config), "custom/qwen") + self.assertIsNone(model._resolve_text_encoder_subfolder(qwen_config, "subfolder", "text_encoder")) + self.assertIsNone(model._resolve_text_encoder_subfolder(qwen_config, "tokenizer_subfolder", "tokenizer")) + + self.assertEqual(model._resolve_text_encoder_path(clip_config), "base/model") + self.assertEqual( + model._resolve_text_encoder_subfolder(clip_config, "subfolder", "text_encoder"), + "text_encoder_2", + ) + + def test_multiple_qwen_encoders_ignore_override_and_warn(self): + first_qwen_config = { + "name": "Qwen3-A", + "tokenizer": QwenTokenizer, + "model": QwenTextModel, + "subfolder": "text_encoder", + } + second_qwen_config = { + "name": "Qwen3-B", + "tokenizer": QwenTokenizer, + "model": QwenTextModel, + "subfolder": "text_encoder_2", + } + model = self._model( + { + "text_encoder": first_qwen_config, + "text_encoder_2": second_qwen_config, + } + ) + + with patch("simpletuner.helpers.models.common.logger.warning") as warning: + self.assertEqual(model._resolve_text_encoder_path(first_qwen_config), "base/model") + + warning.assert_called_once() + self.assertIn("Ignoring qwen_text_encoder_model_name_or_path", warning.call_args.args[0]) + self.assertEqual(warning.call_args.args[2], 2) + self.assertEqual( + model._resolve_text_encoder_subfolder(first_qwen_config, "subfolder", "text_encoder"), + "text_encoder", + ) + + def test_webui_field_and_cli_parser_include_qwen_override(self): + from simpletuner.helpers.configuration.cmd_args import get_argument_parser + from simpletuner.simpletuner_sdk.server.services.field_registry.registry import FieldRegistry + + registry = FieldRegistry() + field = registry.get_field("qwen_text_encoder_model_name_or_path") + + self.assertIsNotNone(field) + self.assertEqual(field.arg_name, "--qwen_text_encoder_model_name_or_path") + self.assertEqual(field.documentation, "OPTIONS.md#--qwen_text_encoder_model_name_or_path") + + parser = get_argument_parser() + args = parser.parse_args( + [ + "--model_family", + "krea2", + "--output_dir", + "/tmp/simpletuner-test", + "--model_type", + "lora", + "--optimizer", + "adamw_bf16", + "--data_backend_config", + "/tmp/backend.json", + "--qwen_text_encoder_model_name_or_path", + "custom/qwen", + ] + ) + + self.assertEqual(args.qwen_text_encoder_model_name_or_path, "custom/qwen") + + def test_env_mapping_includes_qwen_override(self): + from simpletuner.helpers.configuration.env_file import env_to_args_map + + self.assertEqual( + env_to_args_map["QWEN_TEXT_ENCODER_MODEL_NAME_OR_PATH"], + "--qwen_text_encoder_model_name_or_path", + ) + + +if __name__ == "__main__": + unittest.main()