MCPcopy Create free account
hub / github.com/InternScience/SciReason / change_accelerator

Function change_accelerator

opencompass/utils/run.py:236–354  ·  view source on GitHub ↗
(models, accelerator)

Source from the content-addressed store, hash-verified

234
235
236def change_accelerator(models, accelerator):
237 models = models.copy()
238 logger = get_logger()
239 model_accels = []
240 for model in models:
241 logger.info(f'Transforming {model["abbr"]} to {accelerator}')
242 # change HuggingFace model to VLLM or LMDeploy
243 if model['type'] in [HuggingFace, HuggingFaceCausalLM, HuggingFaceChatGLM3, f'{HuggingFaceBaseModel.__module__}.{HuggingFaceBaseModel.__name__}']:
244 gen_args = dict()
245 if model.get('generation_kwargs') is not None:
246 generation_kwargs = model['generation_kwargs'].copy()
247 gen_args['temperature'] = generation_kwargs.get('temperature', 0.001)
248 gen_args['top_k'] = generation_kwargs.get('top_k', 1)
249 gen_args['top_p'] = generation_kwargs.get('top_p', 0.9)
250 gen_args['stop_token_ids'] = generation_kwargs.get('eos_token_id', None)
251 generation_kwargs['stop_token_ids'] = generation_kwargs.get('eos_token_id', None)
252 generation_kwargs.pop('eos_token_id') if 'eos_token_id' in generation_kwargs else None
253 else:
254 # if generation_kwargs is not provided, set default values
255 generation_kwargs = dict()
256 gen_args['temperature'] = 0.0
257 gen_args['top_k'] = 1
258 gen_args['top_p'] = 0.9
259 gen_args['stop_token_ids'] = None
260
261 if accelerator == 'lmdeploy':
262 logger.info(f'Transforming {model["abbr"]} to {accelerator}')
263 mod = TurboMindModelwithChatTemplate
264 acc_model = dict(
265 type=f'{mod.__module__}.{mod.__name__}',
266 abbr=model['abbr'].replace('hf', 'lmdeploy') if '-hf' in model['abbr'] else model['abbr'] + '-lmdeploy',
267 path=model['path'],
268 engine_config=dict(session_len=model['max_seq_len'],
269 max_batch_size=model['batch_size'],
270 tp=model['run_cfg']['num_gpus']),
271 gen_config=dict(top_k=gen_args['top_k'],
272 temperature=gen_args['temperature'],
273 top_p=gen_args['top_p'],
274 max_new_tokens=model['max_out_len'],
275 stop_words=gen_args['stop_token_ids']),
276 max_out_len=model['max_out_len'],
277 max_seq_len=model['max_seq_len'],
278 batch_size=model['batch_size'],
279 run_cfg=model['run_cfg'],
280 )
281 for item in ['meta_template']:
282 if model.get(item) is not None:
283 acc_model[item] = model[item]
284 elif accelerator == 'vllm':
285 model_kwargs = dict(tensor_parallel_size=model['run_cfg']['num_gpus'], max_model_len=model.get('max_seq_len', None))
286 model_kwargs.update(model.get('model_kwargs'))
287 logger.info(f'Transforming {model["abbr"]} to {accelerator}')
288
289 acc_model = dict(
290 type=f'{VLLM.__module__}.{VLLM.__name__}',
291 abbr=model['abbr'].replace('hf', 'vllm') if '-hf' in model['abbr'] else model['abbr'] + '-vllm',
292 path=model['path'],
293 model_kwargs=model_kwargs,

Callers 1

get_config_from_argFunction · 0.85

Calls 4

get_loggerFunction · 0.90
getMethod · 0.80
replaceMethod · 0.80
updateMethod · 0.80

Tested by

no test coverage detected