@@ -124,20 +124,48 @@ def _get_meta_template(meta_template):
124124 return APITemplateParser (meta_template or default_meta_template )
125125
126126
127- def _set_model_kwargs_torch_dtype (model_kwargs ):
127+ def _set_model_kwargs_torch_dtype (model_kwargs , path = None ):
128128 import torch
129- if 'torch_dtype' not in model_kwargs :
130- torch_dtype = torch .float16
129+ from transformers import AutoConfig
130+
131+ # If torch_dtype already exists and is not a string, return directly
132+ if 'torch_dtype' in model_kwargs and not isinstance (model_kwargs ['torch_dtype' ], str ):
133+ return model_kwargs
134+
135+ # Mapping from string to torch data types
136+ dtype_map = {
137+ 'torch.float16' : torch .float16 , 'float16' : torch .float16 ,
138+ 'torch.bfloat16' : torch .bfloat16 , 'bfloat16' : torch .bfloat16 ,
139+ 'torch.float' : torch .float , 'float' : torch .float ,
140+ 'torch.float32' : torch .float32 , 'float32' : torch .float32 ,
141+ 'auto' : 'auto' , 'None' : None
142+ }
143+
144+ # 1. Priority: Use torch_dtype from model_kwargs if available
145+ if 'torch_dtype' in model_kwargs :
146+ torch_dtype = dtype_map .get (model_kwargs ['torch_dtype' ], torch .float16 )
147+
148+ # 2. Secondary: Try to read from model config
149+ elif path is not None :
150+ try :
151+ config = AutoConfig .from_pretrained (path )
152+ if hasattr (config , 'torch_dtype' ):
153+ config_dtype = config .torch_dtype
154+ if isinstance (config_dtype , str ):
155+ torch_dtype = dtype_map .get (config_dtype , torch .float16 )
156+ else :
157+ torch_dtype = config_dtype
158+ else :
159+ torch_dtype = torch .float16
160+ except Exception :
161+ torch_dtype = torch .float16
162+
163+ # 3. Default: Use float16 as fallback
131164 else :
132- torch_dtype = {
133- 'torch.float16' : torch .float16 ,
134- 'torch.bfloat16' : torch .bfloat16 ,
135- 'torch.float' : torch .float ,
136- 'auto' : 'auto' ,
137- 'None' : None ,
138- }.get (model_kwargs ['torch_dtype' ])
139- if torch_dtype is not None :
140- model_kwargs ['torch_dtype' ] = torch_dtype
165+ torch_dtype = torch .float16
166+
167+ # Update model_kwargs with the resolved torch_dtype
168+ model_kwargs ['torch_dtype' ] = torch_dtype
141169 return model_kwargs
142170
143171
@@ -218,12 +246,12 @@ def _load_tokenizer(self, path: Optional[str], kwargs: dict, pad_token_id: Optio
218246 raise ValueError ('pad_token_id is not set for this tokenizer. Please set `pad_token_id={PAD_TOKEN_ID}` in model_cfg.' )
219247
220248 def _load_model (self , path : str , kwargs : dict , peft_path : Optional [str ] = None , peft_kwargs : dict = dict ()):
221- from transformers import AutoModel , AutoModelForCausalLM
249+ from transformers import AutoConfig , AutoModel , AutoModelForCausalLM
222250
223251 DEFAULT_MODEL_KWARGS = dict (device_map = 'auto' , trust_remote_code = True )
224252 model_kwargs = DEFAULT_MODEL_KWARGS
225253 model_kwargs .update (kwargs )
226- model_kwargs = _set_model_kwargs_torch_dtype (model_kwargs )
254+ model_kwargs = _set_model_kwargs_torch_dtype (model_kwargs , path )
227255 self .logger .debug (f'using model_kwargs: { model_kwargs } ' )
228256 if is_npu_available ():
229257 model_kwargs ['device_map' ] = 'npu'
0 commit comments