()
| 383 | return None, False # first training |
| 384 | |
| 385 | def 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 = {} |
no test coverage detected