[feature] remove dependency to download hfsummary manually

This commit is contained in:
theblackcat102
2023-01-01 08:22:30 +00:00
parent 4b7f1f25a1
commit 8b1553642f
3 changed files with 34 additions and 34 deletions
@@ -17,3 +17,4 @@ from dataset import load_dataset
from torch.utils.data import Dataset
class
+28 -30
View File
@@ -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
@@ -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: