srajam696's picture
Upload folder using huggingface_hub
d8ea3f9 verified
Raw
History Blame Contribute Delete
2.55 kB
import json
import gzip
import random
import requests
from transformers import DistilBertTokenizerFast
genre_url_dict = {
'poetry': 'https://mcauleylab.ucsd.edu/public_datasets/gdrive/goodreads/byGenre/goodreads_reviews_poetry.json.gz',
'children': 'https://mcauleylab.ucsd.edu/public_datasets/gdrive/goodreads/byGenre/goodreads_reviews_children.json.gz',
'comics_graphic': 'https://mcauleylab.ucsd.edu/public_datasets/gdrive/goodreads/byGenre/goodreads_reviews_comics_graphic.json.gz',
'fantasy_paranormal': 'https://mcauleylab.ucsd.edu/public_datasets/gdrive/goodreads/byGenre/goodreads_reviews_fantasy_paranormal.json.gz',
'history_biography': 'https://mcauleylab.ucsd.edu/public_datasets/gdrive/goodreads/byGenre/goodreads_reviews_history_biography.json.gz',
'mystery_thriller_crime': 'https://mcauleylab.ucsd.edu/public_datasets/gdrive/goodreads/byGenre/goodreads_reviews_mystery_thriller_crime.json.gz',
'romance': 'https://mcauleylab.ucsd.edu/public_datasets/gdrive/goodreads/byGenre/goodreads_reviews_romance.json.gz',
'young_adult': 'https://mcauleylab.ucsd.edu/public_datasets/gdrive/goodreads/byGenre/goodreads_reviews_young_adult.json.gz'
}
def load_reviews(url, head=10000, sample_size=2000):
reviews = []
response = requests.get(url, stream=True)
with gzip.open(response.raw, 'rt', encoding='utf-8') as file:
for i, line in enumerate(file):
if i >= head: break
reviews.append(json.loads(line)['review_text'])
return random.sample(reviews, min(sample_size, len(reviews)))
def prepare_data(model_name='distilbert-base-cased'):
tokenizer = DistilBertTokenizerFast.from_pretrained(model_name)
train_texts, train_labels, test_texts, test_labels = [], [], [], []
for genre, url in genre_url_dict.items():
reviews = load_reviews(url, head=1000, sample_size=200)
split = int(len(reviews) * 0.8)
train_texts.extend(reviews[:split])
train_labels.extend([genre] * split)
test_texts.extend(reviews[split:])
test_labels.extend([genre] * (len(reviews) - split))
unique_labels = sorted(list(set(train_labels)))
label2id = {label: i for i, label in enumerate(unique_labels)}
train_encodings = tokenizer(train_texts, truncation=True, padding=True, max_length=512)
test_encodings = tokenizer(test_texts, truncation=True, padding=True, max_length=512)
return train_encodings, [label2id[l] for l in train_labels], test_encodings, [label2id[l] for l in test_labels], label2id