@@ -104,17 +104,20 @@ def init_preprocessors(
104104
105105 # Load text encoder and tokenizer.
106106 model_path = str (training_config .model_path )
107- self ._load_text_encoder (model_path )
107+ self ._load_text_encoder (
108+ model_path , training_config
109+ )
108110
109111 # Dummy dataloader for the trainer's outer loop.
110112 self .dataloader = _InfiniteDummyLoader ()
111113 self .start_step = 0
112114
113- def _load_text_encoder (self , model_path : str ) -> None :
114- from transformers import (
115- AutoTokenizer ,
116- UMT5EncoderModel ,
117- )
115+ def _load_text_encoder (
116+ self ,
117+ model_path : str ,
118+ training_config : TrainingConfig ,
119+ ) -> None :
120+ from transformers import AutoTokenizer
118121
119122 logger .info (
120123 "Loading tokenizer from %s" , model_path
@@ -124,17 +127,14 @@ def _load_text_encoder(self, model_path: str) -> None:
124127 )
125128
126129 logger .info (
127- "Loading T5 text encoder from %s" , model_path
130+ "Loading text encoder from %s" , model_path
128131 )
129- dtype = self ._get_training_dtype ()
130- self .text_encoder = UMT5EncoderModel .from_pretrained (
131- model_path ,
132- subfolder = "text_encoder" ,
133- torch_dtype = dtype ,
132+ self .text_encoder = load_module_from_path (
133+ model_path = model_path ,
134+ module_type = "text_encoder" ,
135+ training_config = training_config ,
134136 )
135- self .text_encoder .to (self .device )
136137 self .text_encoder .requires_grad_ (False )
137- self .text_encoder .eval ()
138138
139139 def on_train_start (self ) -> None :
140140 """Skip negative conditioning (handled by method)."""
0 commit comments