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

Function train

deita_dataset/train.py:365–405  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

363 return dict(train_dataset=train_dataset, eval_dataset=eval_dataset)
364
365def train():
366 global local_rank
367
368 parser = transformers.HfArgumentParser(
369 (ModelArguments, DataArguments, TrainingArguments)
370 )
371 model_args, data_args, training_args = parser.parse_args_into_dataclasses()
372 training_args.do_eval = False
373 local_rank = training_args.local_rank
374 model = transformers.AutoModelForCausalLM.from_pretrained(
375 model_args.model_name_or_path,
376 cache_dir=training_args.cache_dir,
377 use_flash_attention_2 = True
378 )
379 model.config.use_cache = False
380 tokenizer = transformers.AutoTokenizer.from_pretrained(
381 model_args.model_name_or_path,
382 cache_dir=training_args.cache_dir,
383 model_max_length=training_args.model_max_length,
384 padding_side="right",
385 use_fast=False,
386 )
387 tokenizer.pad_token = tokenizer.unk_token
388
389 if "mistral" in model_args.model_name_or_path.lower():
390 rank0_print("Mistral with Left Padding Side")
391 tokenizer.padding_side = "left"
392
393 data_module = make_supervised_data_module(tokenizer=tokenizer, data_args=data_args, mask_user = training_args.mask_user)
394
395 trainer = Trainer(
396 model=model, tokenizer=tokenizer, args=training_args, **data_module
397 )
398
399 if list(pathlib.Path(training_args.output_dir).glob("checkpoint-*")):
400 trainer.train(resume_from_checkpoint=True)
401 else:
402 trainer.train()
403 trainer.save_state()
404
405 trainer.save_model(output_dir = training_args.output_dir)
406
407
408if __name__ == "__main__":

Callers 1

train.pyFile · 0.70

Calls 2

rank0_printFunction · 0.85

Tested by

no test coverage detected