MCPcopy Create free account
hub / github.com/OpenGVLab/EfficientQAT / train

Function train

main_e2e_qp.py:385–531  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

383 return None, False # first training
384
385def train():
386 hfparser = transformers.HfArgumentParser((
387 ModelArguments, DataArguments, TrainingArguments, GenerationArguments
388 ))
389 model_args, data_args, training_args, generation_args, extra_args = \
390 hfparser.parse_args_into_dataclasses(return_remaining_strings=True)
391 training_args.generation_config = transformers.GenerationConfig(**vars(generation_args))
392 args = argparse.Namespace(
393 **vars(model_args), **vars(data_args), **vars(training_args)
394 )
395
396 Path(args.output_dir).mkdir(parents=True, exist_ok=True)
397 logger = utils.create_logger(args.output_dir)
398 logger.info(args)
399
400 checkpoint_dir, completed_training = get_last_checkpoint(args.output_dir)
401 if completed_training:
402 print('Detected that training was already completed!')
403
404 model, tokenizer = get_accelerate_model(args, checkpoint_dir)
405
406 model.config.use_cache = False
407 print('loaded model')
408 set_seed(args.seed)
409
410 data_module = make_data_module(tokenizer=tokenizer, args=args)
411
412
413
414 optimizer_grouped_parameters = []
415 for name, module in model.named_modules():
416 # if isinstance(module, LoraLayer):
417 if isinstance(module, QuantLinear) and not 'head' in name:
418 module.scales.requires_grad = True
419 optimizer_grouped_parameters.append({'params': [p for n, p in model.named_parameters() if 'scale' in n], 'weight_decay': 0.0, 'lr': args.learning_rate})
420 optimizer = AdamW(optimizer_grouped_parameters)
421
422 trainer = Seq2SeqTrainer(
423 model=model,
424 tokenizer=tokenizer,
425 args=training_args,
426 optimizers=(optimizer, None),
427 **{k:v for k,v in data_module.items() if k != 'predict_dataset'},
428 )
429
430 if args.do_ppl_eval:
431 class PPLvalCallback(transformers.TrainerCallback):
432 @torch.no_grad()
433 def on_evaluate(self, args=None, state=None, control=None, model=None, **kwargs):
434 results = test_ppl(trainer.model, trainer.tokenizer, datasets=['wikitext2','c4'],ppl_seqlen=2048)
435 logger.info(results)
436 trainer.log(results)
437
438 trainer.add_callback(PPLvalCallback)
439
440 # Verifying the datatypes and parameter counts before training.
441 print_trainable_parameters(args, model)
442 dtypes = {}

Callers 1

main_e2e_qp.pyFile · 0.70

Calls 4

make_data_moduleFunction · 0.90
get_last_checkpointFunction · 0.85
get_accelerate_modelFunction · 0.85

Tested by

no test coverage detected