Make dataset and collator for supervised fine-tuning.
(
tokenizer: transformers.PreTrainedTokenizer, data_args
)
| 223 | |
| 224 | |
| 225 | def make_supervised_data_module( |
| 226 | tokenizer: transformers.PreTrainedTokenizer, data_args |
| 227 | ) -> Dict: |
| 228 | """Make dataset and collator for supervised fine-tuning.""" |
| 229 | dataset_cls = ( |
| 230 | LazySupervisedDataset if data_args.lazy_preprocess else SupervisedDataset |
| 231 | ) |
| 232 | rank0_print("Loading data...") |
| 233 | raw_data = json.load(open(data_args.data_path, "r")) |
| 234 | if data_args.eval_data_path is not None: |
| 235 | train_raw_data = raw_data |
| 236 | eval_raw_data = json.load(open(data_args.eval_data_path, "r")) |
| 237 | else: |
| 238 | # Split train/test |
| 239 | perm = np.random.permutation(len(raw_data)) |
| 240 | split = int(len(perm) * 0.98) |
| 241 | train_indices = perm[:split] |
| 242 | eval_indices = perm[split:] |
| 243 | train_raw_data = [raw_data[i] for i in train_indices] |
| 244 | eval_raw_data = [raw_data[i] for i in eval_indices] |
| 245 | rank0_print(f"#train {len(train_raw_data)}, #eval {len(eval_raw_data)}") |
| 246 | train_dataset = dataset_cls(train_raw_data, tokenizer=tokenizer, template=data_args.conv_template) |
| 247 | eval_dataset = dataset_cls(eval_raw_data, tokenizer=tokenizer, template=data_args.conv_template) |
| 248 | return dict(train_dataset=train_dataset, eval_dataset=eval_dataset) |
| 249 | |
| 250 | |
| 251 | def train(): |
no test coverage detected