Make dataset and collator for supervised fine-tuning.
(
tokenizer: transformers.PreTrainedTokenizer, data_args, mask_user = True
)
| 333 | |
| 334 | |
| 335 | def make_supervised_data_module( |
| 336 | tokenizer: transformers.PreTrainedTokenizer, data_args, mask_user = True |
| 337 | ) -> Dict: |
| 338 | """Make dataset and collator for supervised fine-tuning.""" |
| 339 | conv_template = data_args.conv_template |
| 340 | dataset_cls = ( |
| 341 | LazySupervisedDataset if data_args.lazy_preprocess else SupervisedDataset |
| 342 | ) |
| 343 | rank0_print("Loading data...") |
| 344 | try: |
| 345 | raw_data = json.load(open(data_args.data_path, "r")) |
| 346 | except FileNotFoundError: |
| 347 | raw_data = load_dataset(data_args.data_path, split = "train") |
| 348 | raw_data = [row for row in raw_data] |
| 349 | |
| 350 | # Split train/eval |
| 351 | np.random.seed(0) |
| 352 | train_raw_data = raw_data |
| 353 | perm = np.random.permutation(len(raw_data)) |
| 354 | split = int(len(perm) * 0.98) |
| 355 | train_indices = perm[:split] |
| 356 | eval_indices = perm[split:] |
| 357 | train_raw_data = [raw_data[i] for i in train_indices] |
| 358 | eval_raw_data = [raw_data[i] for i in eval_indices] |
| 359 | rank0_print(f"#train {len(train_raw_data)}, #eval {len(eval_raw_data)}") |
| 360 | |
| 361 | train_dataset = dataset_cls(train_raw_data, tokenizer=tokenizer, conv_template = conv_template, mask_user = mask_user) |
| 362 | eval_dataset = dataset_cls(eval_raw_data, tokenizer=tokenizer, conv_template = conv_template, mask_user = mask_user) |
| 363 | return dict(train_dataset=train_dataset, eval_dataset=eval_dataset) |
| 364 | |
| 365 | def train(): |
| 366 | global local_rank |