diff --git a/model/reward/instructor/experimental_dataset.py b/model/reward/instructor/experimental_dataset.py index f705ccf6..85f0c899 100644 --- a/model/reward/instructor/experimental_dataset.py +++ b/model/reward/instructor/experimental_dataset.py @@ -17,3 +17,4 @@ from dataset import load_dataset from torch.utils.data import Dataset +class \ No newline at end of file diff --git a/model/reward/instructor/rank_datasets.py b/model/reward/instructor/rank_datasets.py index 4ba6293c..2f2260c2 100644 --- a/model/reward/instructor/rank_datasets.py +++ b/model/reward/instructor/rank_datasets.py @@ -112,51 +112,49 @@ class HFSummary(Dataset): ''' Human feedback data from OpenAI https://github.com/openai/summarize-from-feedback - - >> azcopy copy "https://openaipublic.blob.core.windows.net/summarize-from-feedback/dataset/*" . --recursive labeling method : pair comparison, 0 or 1 ''' def __init__(self, split='train', - path='summarize-from-feedback/comparisons/*.json', conf_threshold=-1, - max_comparison_per_sample=5) -> None: + max_comparison_per_sample=3) -> None: super().__init__() - assert split in ('train', 'valid1', 'valid2', 'test') + assert split in ('train', 'validation') summaries = {} # using prompt as our index will allows us # to add additional generated prompt later self.index2summary = {} self.max_comparison_per_sample = max_comparison_per_sample - for jsonl_file in glob.glob(path): - with open(jsonl_file, 'r') as f: - for line in f: - data = json.loads(line) - if data['split'] != split: - continue - if 'extra' in data and \ - 'confidence' in data['extra'] and \ - conf_threshold > data['extra']['confidence']: - print('skipping {}'.format(data['info']['id'])) - continue + dataset = load_dataset('Tristan/summarize_from_feedback', 'comparisons')[split] + for data in dataset: + if 'extra' in data and \ + 'confidence' in data['extra'] and \ + data['extra']['confidence'] is not None and \ + conf_threshold > data['extra']['confidence']: + print('skipping {}'.format(data['info']['id'])) + continue - if 'article' in data['info']: - context = data['info']['article'] - elif 'post' in data['info']: - context = data['info']['post'] + if 'article' in data['info'] and \ + data['info']['article'] is not None: + context = data['info']['article'] + elif 'post' in data['info']: + context = data['info']['post'] - if context not in self.index2summary: - self.index2summary[len(self.index2summary)] = context - - if context not in summaries: - summaries[context] = [] + if context is None: + continue - pos, neg = (0, 1) if data['choice'] == 0 else (1, 0) - summaries[context].append(( - data['summaries'][pos]['text'], - data['summaries'][neg]['text'] - )) + if context not in self.index2summary: + self.index2summary[len(self.index2summary)] = context + + if context not in summaries: + summaries[context] = [] + + pos, neg = (0, 1) if data['choice'] == 0 else (1, 0) + summaries[context].append(( + data['summaries'][pos]['text'], + data['summaries'][neg]['text'] + )) self.summaries = summaries diff --git a/model/reward/instructor/tests/test_dataset.py b/model/reward/instructor/tests/test_dataset.py index c452786b..7b432fd3 100644 --- a/model/reward/instructor/tests/test_dataset.py +++ b/model/reward/instructor/tests/test_dataset.py @@ -1,22 +1,23 @@ from transformers import AutoTokenizer from torch.utils.data import DataLoader -from rank_datasets import WebGPT, HFSummary, DataCollatorForMultipleChoice +from rank_datasets import WebGPT, HFSummary, DataCollatorForPairRank def test_hfsummary(): tokenizer = AutoTokenizer.from_pretrained("bigscience/mt0-large") - collate_fn = DataCollatorForMultipleChoice(tokenizer, max_length=200) + collate_fn = DataCollatorForPairRank(tokenizer, max_length=200) dataset = HFSummary() + print(len(dataset)) dataloader = DataLoader(dataset, collate_fn=collate_fn, batch_size=8) for batch in dataloader: - print(batch['input_ids'].shape) + batch['input_ids'].shape def test_webgpt(): tokenizer = AutoTokenizer.from_pretrained("bigscience/mt0-large") - collate_fn = DataCollatorForMultipleChoice(tokenizer, max_length=200) + collate_fn = DataCollatorForPairRank(tokenizer, max_length=200) dataset = WebGPT() dataloader = DataLoader(dataset, collate_fn=collate_fn, batch_size=32) for batch in dataloader: