Skip to content

Commit 9b2bb91

Browse files
liushzsudanl
andauthored
[Fix] Fix default torch dtype loading (#1969)
* Support OlympiadBench Benchmark * Support OlympiadBench Benchmark * Support OlympiadBench Benchmark * update dataset path * Update olmpiadBench * Update olmpiadBench * Update olmpiadBench * Add HLE dataset * Add HLE dataset * Add HLE dataset * Add AIME2025 oss info * Fix torch dtype error --------- Co-authored-by: sudanl <sudanl@foxmail.com>
1 parent 7ce6321 commit 9b2bb91

1 file changed

Lines changed: 42 additions & 14 deletions

File tree

opencompass/models/huggingface_above_v4_33.py

Lines changed: 42 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -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

Comments
 (0)