MCPcopy Create free account
hub / github.com/OpenBitSys/BitDistiller / __init__

Method __init__

test/gsm8k/test.py:109–131  ·  view source on GitHub ↗
(self, data_path: str, tokenizer: transformers.PreTrainedTokenizer)

Source from the content-addressed store, hash-verified

107 """Dataset for supervised fine-tuning."""
108
109 def __init__(self, data_path: str, tokenizer: transformers.PreTrainedTokenizer):
110 super(SupervisedDataset, self).__init__()
111
112 # dataset_for_eval = load_dataset(data_path)['train']
113
114
115 with open(data_path, 'r') as f:
116 dataset_for_eval = f.readlines()
117
118 dataset_for_eval = [json.loads(item.strip()) for item in dataset_for_eval]
119 try:
120 sources = [PROMPT_DICT["prompt_no_input"].format_map(item)for item in dataset_for_eval]
121 except:
122 sources = [PROMPT_DICT["prompt_no_input_v2"].format_map(item)for item in dataset_for_eval]
123 try:
124 targets = [item['answer'] for item in dataset_for_eval]
125 except:
126 targets = [item['response'] for item in dataset_for_eval]
127
128 data_dict = preprocess(sources, targets, tokenizer)
129
130 self.input_ids = data_dict["input_ids"]
131 self.labels = data_dict["labels"]
132
133 def __len__(self):
134 return len(self.input_ids)

Callers

nothing calls this directly

Calls 1

preprocessFunction · 0.70

Tested by

no test coverage detected