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

Function make_supervised_data_module

deita_dataset/train.py:335–363  ·  view source on GitHub ↗

Make dataset and collator for supervised fine-tuning.

(
    tokenizer: transformers.PreTrainedTokenizer, data_args, mask_user = True
)

Source from the content-addressed store, hash-verified

333
334
335def 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
365def train():
366 global local_rank

Callers 1

trainFunction · 0.85

Calls 1

rank0_printFunction · 0.85

Tested by

no test coverage detected