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
360 changes: 19 additions & 341 deletions modules/ui/BaseModelTabView.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,357 +17,35 @@ 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(
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,
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,
Expand Down
15 changes: 15 additions & 0 deletions modules/ui/ModelTabController.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
Loading