diff --git a/modules/ui/BaseModelTabView.py b/modules/ui/BaseModelTabView.py index 08d873d3b..0757a71c2 100644 --- a/modules/ui/BaseModelTabView.py +++ b/modules/ui/BaseModelTabView.py @@ -17,320 +17,9 @@ def _make_svd_frames(self, parent, row: int): pass def build_content(self, frame, controller, ui_state): - if controller.train_config.model_type.is_stable_diffusion(): # TODO simplify - self.__setup_stable_diffusion_ui(frame, controller, ui_state) - if controller.train_config.model_type.is_stable_diffusion_3(): - self.__setup_stable_diffusion_3_ui(frame, controller, ui_state) - elif controller.train_config.model_type.is_stable_diffusion_xl(): - self.__setup_stable_diffusion_xl_ui(frame, controller, ui_state) - elif controller.train_config.model_type.is_wuerstchen(): - self.__setup_wuerstchen_ui(frame, controller, ui_state) - elif controller.train_config.model_type.is_pixart(): - self.__setup_pixart_alpha_ui(frame, controller, ui_state) - elif controller.train_config.model_type.is_flux_1(): - self.__setup_flux_ui(frame, controller, ui_state) - elif controller.train_config.model_type.is_flux_2(): - self.__setup_flux_2_ui(frame, controller, ui_state) - elif controller.train_config.model_type.is_z_image(): - self.__setup_z_image_ui(frame, controller, ui_state) - elif controller.train_config.model_type.is_chroma(): - self.__setup_chroma_ui(frame, controller, ui_state) - elif controller.train_config.model_type.is_qwen(): - self.__setup_qwen_ui(frame, controller, ui_state) - elif controller.train_config.model_type.is_anima(): - self.__setup_anima_ui(frame, controller, ui_state) - elif controller.train_config.model_type.is_krea2(): - self.__setup_krea2_ui(frame, controller, ui_state) - elif controller.train_config.model_type.is_sana(): - self.__setup_sana_ui(frame, controller, ui_state) - elif controller.train_config.model_type.is_hunyuan_video(): - self.__setup_hunyuan_video_ui(frame, controller, ui_state) - elif controller.train_config.model_type.is_hi_dream(): - self.__setup_hi_dream_ui(frame, controller, ui_state) - elif controller.train_config.model_type.is_ernie(): - self.__setup_ernie_ui(frame, controller, ui_state) - - def __setup_stable_diffusion_ui(self, frame, controller, ui_state): - row = 0 - row = self.__create_base_dtype_components(frame, row, ui_state) - row = self.__create_base_components( - frame, - row, - controller, - ui_state, - has_unet=True, - has_text_encoder=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - controller, - ui_state, - ) - - def __setup_stable_diffusion_3_ui(self, frame, controller, ui_state): - row = 0 - row = self.__create_base_dtype_components(frame, row, ui_state) - row = self.__create_base_components( - frame, - row, - controller, - ui_state, - has_transformer=True, - has_text_encoder_1=True, - has_text_encoder_2=True, - has_text_encoder_3=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - controller, - ui_state, - ) - - def __setup_flux_ui(self, frame, controller, ui_state): - row = 0 - row = self.__create_base_dtype_components(frame, row, ui_state) - row = self.__create_base_components( - frame, - row, - controller, - ui_state, - has_transformer=True, - allow_override_transformer=True, - has_text_encoder_1=True, - has_text_encoder_2=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - controller, - ui_state, - ) - - def __setup_flux_2_ui(self, frame, controller, ui_state): - row = 0 - row = self.__create_base_dtype_components(frame, row, ui_state) - row = self.__create_base_components( - frame, - row, - controller, - ui_state, - has_transformer=True, - allow_override_transformer=True, - has_text_encoder_1=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - controller, - ui_state, - ) - - def __setup_z_image_ui(self, frame, controller, ui_state): - row = 0 - row = self.__create_base_dtype_components(frame, row, ui_state) - row = self.__create_base_components( - frame, - row, - controller, - ui_state, - has_transformer=True, - allow_override_transformer=True, - has_text_encoder_1=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - controller, - ui_state, - ) - - def __setup_ernie_ui(self, frame, controller, ui_state): - row = 0 - row = self.__create_base_dtype_components(frame, row, ui_state) - row = self.__create_base_components( - frame, - row, - controller, - ui_state, - has_transformer=True, - allow_override_transformer=True, - has_text_encoder_1=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - controller, - ui_state, - ) - - def __setup_chroma_ui(self, frame, controller, ui_state): - row = 0 - row = self.__create_base_dtype_components(frame, row, ui_state) - row = self.__create_base_components( - frame, - row, - controller, - ui_state, - has_transformer=True, - allow_override_transformer=True, - has_text_encoder_1=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - controller, - ui_state, - ) - - def __setup_qwen_ui(self, frame, controller, ui_state): - row = 0 - row = self.__create_base_dtype_components(frame, row, ui_state) - row = self.__create_base_components( - frame, - row, - controller, - ui_state, - has_transformer=True, - allow_override_transformer=True, - has_text_encoder_1=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - controller, - ui_state, - ) - - def __setup_anima_ui(self, frame, controller, ui_state): - row = 0 - row = self.__create_base_dtype_components(frame, row, ui_state) - row = self.__create_base_components( - frame, - row, - controller, - ui_state, - has_transformer=True, - allow_override_transformer=True, - has_text_encoder_1=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - controller, - ui_state, - ) - - def __setup_krea2_ui(self, frame, controller, ui_state): - row = 0 - row = self.__create_base_dtype_components(frame, row, ui_state) - row = self.__create_base_components( - frame, - row, - controller, - ui_state, - has_transformer=True, - allow_override_transformer=True, - has_text_encoder_1=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - controller, - ui_state, - ) - - def __setup_stable_diffusion_xl_ui(self, frame, controller, ui_state): - row = 0 - row = self.__create_base_dtype_components(frame, row, ui_state) - row = self.__create_base_components( - frame, - row, - controller, - ui_state, - has_unet=True, - has_text_encoder_1=True, - has_text_encoder_2=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - controller, - ui_state, - ) - - def __setup_wuerstchen_ui(self, frame, controller, ui_state): - row = 0 - row = self.__create_base_dtype_components(frame, row, ui_state) - row = self.__create_base_components( - frame, - row, - controller, - ui_state, - has_prior=True, - allow_override_prior=controller.train_config.model_type.is_stable_cascade(), - has_text_encoder=True, - ) - row = self.__create_effnet_encoder_components(frame, row, ui_state) - row = self.__create_decoder_components(frame, row, ui_state, controller.train_config.model_type.is_wuerstchen_v2()) - row = self.__create_output_components( - frame, - row, - controller, - ui_state, - ) - - def __setup_pixart_alpha_ui(self, frame, controller, ui_state): - row = 0 - row = self.__create_base_dtype_components(frame, row, ui_state) - row = self.__create_base_components( - frame, - row, - controller, - ui_state, - has_transformer=True, - has_text_encoder=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - controller, - ui_state, - ) - - def __setup_sana_ui(self, frame, controller, ui_state): - row = 0 - row = self.__create_base_dtype_components(frame, row, ui_state) - row = self.__create_base_components( - frame, - row, - controller, - ui_state, - has_transformer=True, - has_text_encoder=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - controller, - ui_state, - ) + model_type = controller.train_config.model_type + parts = model_type.model_parts() - def __setup_hunyuan_video_ui(self, frame, controller, ui_state): row = 0 row = self.__create_base_dtype_components(frame, row, ui_state) row = self.__create_base_components( @@ -338,36 +27,25 @@ def __setup_hunyuan_video_ui(self, frame, controller, ui_state): row, controller, ui_state, - has_transformer=True, - allow_override_transformer=True, - has_text_encoder_1=True, - has_text_encoder_2=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - controller, - ui_state, + has_unet="unet" in parts, + has_prior="prior" in parts, + allow_override_prior=model_type.is_stable_cascade(), + has_transformer="transformer" in parts, + allow_override_transformer=controller.supports_override_transformer(), + has_text_encoder=not model_type.has_multiple_text_encoders(), + has_text_encoder_1=model_type.has_multiple_text_encoders(), + has_text_encoder_2="text_encoder_2" in parts, + has_text_encoder_3="text_encoder_3" in parts, + has_text_encoder_4="text_encoder_4" in parts, + allow_override_text_encoder_4="text_encoder_4" in parts, + has_vae="vae" in parts, ) + if "effnet_encoder" in parts: + row = self.__create_effnet_encoder_components(frame, row, ui_state) + if "decoder" in parts: + row = self.__create_decoder_components(frame, row, ui_state, "decoder_text_encoder" in parts) - def __setup_hi_dream_ui(self, frame, controller, ui_state): - row = 0 - row = self.__create_base_dtype_components(frame, row, ui_state) - row = self.__create_base_components( - frame, - row, - controller, - ui_state, - has_transformer=True, - has_text_encoder_1=True, - has_text_encoder_2=True, - has_text_encoder_3=True, - has_text_encoder_4=True, - allow_override_text_encoder_4=True, - has_vae=True, - ) - row = self.__create_output_components( + self.__create_output_components( frame, row, controller, diff --git a/modules/ui/ModelTabController.py b/modules/ui/ModelTabController.py index 7e5522d6e..70835d87a 100644 --- a/modules/ui/ModelTabController.py +++ b/modules/ui/ModelTabController.py @@ -13,6 +13,21 @@ def get_presets(self) -> dict: cls = create.get_model_setup_class(self.train_config.model_type, self.train_config.training_method) return cls.LAYER_PRESETS if cls is not None else {"full": []} + def supports_override_transformer(self) -> bool: + model_type = self.train_config.model_type + # The transformer override path exists only for these architectures; SD3, PixArt, Sana + # and HiDream have a transformer but expose no override field. + return ( + model_type.is_flux() + or model_type.is_z_image() + or model_type.is_ernie() + or model_type.is_chroma() + or model_type.is_qwen() + or model_type.is_anima() + or model_type.is_krea2() + or model_type.is_hunyuan_video() + ) + def get_output_formats(self) -> list[tuple[str, ModelFormat]]: labels = { ModelFormat.SAFETENSORS: "Safetensors", diff --git a/modules/ui/TopBarController.py b/modules/ui/TopBarController.py index 25593a80e..7cdd1b3bb 100644 --- a/modules/ui/TopBarController.py +++ b/modules/ui/TopBarController.py @@ -42,39 +42,13 @@ def get_model_types(self) -> list[tuple[str, ModelType]]: ] def get_training_methods(self, model_type: ModelType) -> list[tuple[str, TrainingMethod]]: - #TODO simplify - if model_type.is_stable_diffusion(): - return [ - ("Fine Tune", TrainingMethod.FINE_TUNE), - ("LoRA", TrainingMethod.LORA), - ("Embedding", TrainingMethod.EMBEDDING), - ("Fine Tune VAE", TrainingMethod.FINE_TUNE_VAE), - ] - elif model_type.is_stable_diffusion_3() \ - or model_type.is_stable_diffusion_xl() \ - or model_type.is_wuerstchen() \ - or model_type.is_pixart() \ - or model_type.is_flux_1() \ - or model_type.is_sana() \ - or model_type.is_hunyuan_video() \ - or model_type.is_hi_dream() \ - or model_type.is_chroma(): - return [ - ("Fine Tune", TrainingMethod.FINE_TUNE), - ("LoRA", TrainingMethod.LORA), - ("Embedding", TrainingMethod.EMBEDDING), - ] - elif model_type.is_qwen() \ - or model_type.is_anima() \ - or model_type.is_krea2() \ - or model_type.is_z_image() \ - or model_type.is_flux_2() \ - or model_type.is_ernie(): - return [ - ("Fine Tune", TrainingMethod.FINE_TUNE), - ("LoRA", TrainingMethod.LORA), - ] - return [] + labels = { + TrainingMethod.FINE_TUNE: "Fine Tune", + TrainingMethod.LORA: "LoRA", + TrainingMethod.EMBEDDING: "Embedding", + TrainingMethod.FINE_TUNE_VAE: "Fine Tune VAE", + } + return [(labels[m], m) for m in model_type.supported_training_methods()] def load_available_config_names(self, dir: str) -> list[tuple[str, str]]: configs = [("", path_util.canonical_join(dir, "#.json"))] diff --git a/modules/util/enum/ModelType.py b/modules/util/enum/ModelType.py index 55377ba3e..5638fd261 100644 --- a/modules/util/enum/ModelType.py +++ b/modules/util/enum/ModelType.py @@ -183,6 +183,24 @@ def is_video_model(self) -> bool: def model_parts(self) -> tuple[str, ...]: return _MODEL_PARTS[self] + def supported_training_methods(self) -> tuple[TrainingMethod, ...]: + if self.is_stable_diffusion(): + return (TrainingMethod.FINE_TUNE, TrainingMethod.LORA, TrainingMethod.EMBEDDING, TrainingMethod.FINE_TUNE_VAE) + if self.is_stable_diffusion_3() \ + or self.is_stable_diffusion_xl() \ + or self.is_wuerstchen() \ + or self.is_pixart() \ + or self.is_flux_1() \ + or self.is_sana() \ + or self.is_hunyuan_video() \ + or self.is_hi_dream() \ + or self.is_chroma(): + return (TrainingMethod.FINE_TUNE, TrainingMethod.LORA, TrainingMethod.EMBEDDING) + if self.is_qwen() or self.is_z_image() or self.is_flux_2() or self.is_ernie() \ + or self.is_anima() or self.is_krea2(): + return (TrainingMethod.FINE_TUNE, TrainingMethod.LORA) + raise ValueError(f"No supported training methods defined for model type {self}") + def denoising_model_part(self) -> str: # the denoising model component (unet / transformer / prior), always listed first in model_parts(). return _MODEL_PARTS[self][0]