-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathrjs_training.py
More file actions
executable file
·62 lines (51 loc) · 2.23 KB
/
Copy pathrjs_training.py
File metadata and controls
executable file
·62 lines (51 loc) · 2.23 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
from arguments import TrainingArguments, CustomTrainingArguments
from transformers import HfArgumentParser
from utils import print_object_on_main_process, print_rank_0, getDataset, loadTokenizerAndModel
from typing import Dict
import os
import json
from collator import rjs_data_collator
from utils import load_data_from_paths, load_jsonl_data, read_json_or_jsonl_data
from datasets import Dataset
from transformers import Trainer
def main():
parser = HfArgumentParser(CustomTrainingArguments)
args = parser.parse_args_into_dataclasses()[0]
print_rank_0(args)
print_object_on_main_process("Arguments", args)
print_rank_0("Loading data>>>>>>>>>>>>>>>>>>>>>>>>>>>")
# data_list = load_jsonl_data(args.data_path)
# train_dataset = Dataset.from_list(data_list)
full_data = []
for data_path in args.train_data_path:
data_list = read_json_or_jsonl_data(data_path)
full_data.extend(data_list)
train_dataset = Dataset.from_list(full_data)
# eval_dataset = getDataset(args, type='eval')
print_object_on_main_process("training set", train_dataset, split_line_color="green", object_color="cyan")
# print_object_on_main_process("evaluation set", eval_dataset, split_line_color="green", object_color="cyan")
tokenizer, model = loadTokenizerAndModel(args)
print_object_on_main_process("tokenizer", tokenizer, split_line_color="green", object_color="cyan")
print_object_on_main_process("model", model, split_line_color="green", object_color="cyan")
print_rank_0("Using rejection sampling data collator")
data_collator = rjs_data_collator(tokenizer, args)
compute_metrics = None
trainer = Trainer(
model=model,
tokenizer=tokenizer,
args=args,
train_dataset=train_dataset,
# eval_dataset=eval_dataset,
data_collator=data_collator,
compute_metrics=compute_metrics
)
if args.do_train:
train_result = trainer.train()
metrics = train_result.metrics
if args.save_training_states:
trainer.save_state()
trainer.save_model(output_dir=args.output_dir)
trainer.log_metrics("train", metrics)
trainer.save_metrics("train", metrics)
if __name__ == '__main__':
main()