- Notifications
You must be signed in to change notification settings - Fork 6
Expand file tree
/
Copy pathtrain.py
More file actions
Latest commit
63 lines (49 loc) · 2.18 KB
/
Copy pathtrain.py
File metadata and controls
63 lines (49 loc) · 2.18 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
importcopy
importlogging
fromdataclassesimportdataclass, field
fromtypingimportDict, Optional, Sequence
importtorch
importtransformers
fromtorch.utils.dataimportDataset
fromtransformersimportTrainer, DataCollatorForLanguageModeling
fromdatasetsimportload_from_disk
@dataclass
classModelArguments:
model_name_or_path: Optional[str] =field(default="facebook/opt-125m")
@dataclass
classDataArguments:
data_path: str=field(default=None, metadata={"help": "Path to the training data."})
@dataclass
classTrainingArguments(transformers.TrainingArguments):
cache_dir: Optional[str] =field(default=None)
optim: str=field(default="adamw_torch")
model_max_length: int=field(
default=512,
metadata={"help": "Maximum sequence length. Sequences will be right padded (and possibly truncated)."},
)
defmake_data_module(tokenizer: transformers.PreTrainedTokenizer, data_args) ->Dict:
"""Make dataset and collator"""
tokenized_datasets=load_from_disk(data_args.data_path)
data_collator=DataCollatorForLanguageModeling(tokenizer=tokenizer, mlm=False, mlm_probability=0.0)
returndict(train_dataset=tokenized_datasets['train'], eval_dataset=None, data_collator=data_collator)
deftrain():
parser=transformers.HfArgumentParser((ModelArguments, DataArguments, TrainingArguments))
model_args, data_args, training_args=parser.parse_args_into_dataclasses()
model=transformers.AutoModelForCausalLM.from_pretrained(
model_args.model_name_or_path,
cache_dir=training_args.cache_dir,
)
tokenizer=transformers.AutoTokenizer.from_pretrained(
model_args.model_name_or_path,
cache_dir=training_args.cache_dir,
model_max_length=training_args.model_max_length,
padding_side="right",
#use_fast=False,
)
data_module=make_data_module(tokenizer=tokenizer, data_args=data_args)
trainer=Trainer(model=model, tokenizer=tokenizer, args=training_args, **data_module)
trainer.train(resume_from_checkpoint=training_args.resume_from_checkpoint)
trainer.save_state()
trainer.save_model(output_dir=training_args.output_dir)
if__name__=="__main__":
train()