Spaces:
Runtime error
Runtime error
| from dotenv import load_dotenv | |
| load_dotenv() | |
| from transformers import AutoTokenizer | |
| from datasets import load_dataset, load_from_disk, concatenate_datasets | |
| from concurrent.futures import ThreadPoolExecutor, as_completed | |
| from tqdm import tqdm | |
| import time | |
| import os | |
| # Load a pre-trained tokenizer | |
| tokenizer = AutoTokenizer.from_pretrained("Qwen/Qwen3-8B") | |
| tokenizer.pad_token = tokenizer.eos_token | |
| # Define all dataset loading tasks | |
| def load_wikitext_103_train(): | |
| ds = load_dataset("Salesforce/wikitext", "wikitext-103-raw-v1", split="train") | |
| ds = ds.filter(lambda x: len(x["text"].strip()) > 0) | |
| return ds.select_columns(["text"]) | |
| def load_wikitext_103_val_test(): | |
| parts = [] | |
| for split in ["validation", "test"]: | |
| ds = load_dataset("Salesforce/wikitext", "wikitext-103-raw-v1", split=split) | |
| ds = ds.filter(lambda x: len(x["text"].strip()) > 0) | |
| parts.append(ds.select_columns(["text"])) | |
| return concatenate_datasets(parts) | |
| def load_wikitext_2(): | |
| parts = [] | |
| for split in ["train", "validation", "test"]: | |
| ds = load_dataset("Salesforce/wikitext", "wikitext-2-raw-v1", split=split) | |
| ds = ds.filter(lambda x: len(x["text"].strip()) > 0) | |
| parts.append(ds.select_columns(["text"])) | |
| return concatenate_datasets(parts) | |
| def load_english_dict(): | |
| ds = load_dataset("npvinHnivqn/EnglishDictionary", split="train") | |
| ds = ds.map(lambda x: {"text": f"{x['word']}: {x['definition']}"}) | |
| return ds.select_columns(["text"]) | |
| def load_wordnet(): | |
| ds = load_dataset("marksverdhei/wordnet-definitions-en-2021", split="train") | |
| ds = ds.map(lambda x: {"text": f"{x['Word']}: {x['Definition']}. Example: {x['Example']}"}) | |
| return ds.select_columns(["text"]) | |
| def load_ag_news(): | |
| ds = load_dataset("fancyzhx/ag_news", split="train") | |
| return ds.select_columns(["text"]) | |
| def load_imdb(): | |
| ds = load_dataset("stanfordnlp/imdb", split="train") | |
| return ds.select_columns(["text"]) | |
| def load_rotten_tomatoes(): | |
| ds = load_dataset("cornell-movie-review-data/rotten_tomatoes", split="train") | |
| return ds.select_columns(["text"]) | |
| def load_cnn_dailymail(): | |
| ds = load_dataset("abisee/cnn_dailymail", "3.0.0", split="train") | |
| ds = ds.rename_column("article", "text") | |
| return ds.select_columns(["text"]) | |
| def load_yelp(): | |
| ds = load_dataset("Yelp/yelp_review_full", split="train") | |
| return ds.select_columns(["text"]) | |
| def load_urban_dictionary(): | |
| ds = load_dataset("daspartho/urban_dictionary", split="train") | |
| ds = ds.map(lambda x: {"text": f"{x['word']}: {x['definition']}. Example: {x['example']}"}) | |
| return ds.select_columns(["text"]) | |
| def load_slang(): | |
| ds = load_dataset("LM-Lexicon/Slang", split="train") | |
| ds = ds.map(lambda x: {"text": f"{x['term']}: {x['definition']}. Example: {x['context']}"}) | |
| return ds.select_columns(["text"]) | |
| def load_genz_slang(): | |
| ds = load_dataset("MLBtrio/genz-slang-dataset", split="train") | |
| ds = ds.map(lambda x: {"text": f"{x['Slang']}: {x['Description']}. Example: {x['Example']}"}) | |
| return ds.select_columns(["text"]) | |
| # All tasks with labels | |
| tasks = [ | |
| ("WikiText-103 (train)", load_wikitext_103_train), | |
| ("WikiText-103 (val+test)", load_wikitext_103_val_test), | |
| ("WikiText-2", load_wikitext_2), | |
| ("English Dictionary", load_english_dict), | |
| ("WordNet Definitions", load_wordnet), | |
| ("AG News", load_ag_news), | |
| ("IMDB", load_imdb), | |
| ("Rotten Tomatoes", load_rotten_tomatoes), | |
| ("CNN/DailyMail", load_cnn_dailymail), | |
| ("Yelp Reviews", load_yelp), | |
| ("Urban Dictionary", load_urban_dictionary), | |
| ("LM-Lexicon Slang", load_slang), | |
| ("Gen Z Slang", load_genz_slang), | |
| ] | |
| # Directory where datasets are cached locally | |
| data_dir = os.path.join(os.path.dirname(__file__), "data") | |
| os.makedirs(data_dir, exist_ok=True) | |
| def safe_name_for(label): | |
| return label.lower().replace(" ", "_").replace("/", "_").replace("(", "").replace(")", "").replace("+", "_") | |
| print(f"Loading {len(tasks)} datasets...\n") | |
| # Create one tqdm bar per dataset, each on its own line | |
| bars = [] | |
| for i, (label, _) in enumerate(tasks): | |
| bar = tqdm(total=1, desc=f" {i+1:>2}/{len(tasks)} {label:<25}", position=i, leave=True, | |
| bar_format="{desc} {bar} {postfix}") | |
| bar.set_postfix_str("pending...") | |
| bars.append(bar) | |
| results = {} | |
| def run_task(index, label, fn): | |
| # Check if dataset already exists on disk | |
| save_path = os.path.join(data_dir, safe_name_for(label)) | |
| if os.path.isdir(save_path): | |
| try: | |
| ds = load_from_disk(save_path) | |
| bars[index].update(1) | |
| bars[index].set_postfix_str(f"✓ {len(ds)} examples (cached)") | |
| return label, ds | |
| except Exception: | |
| pass # cache corrupt, re-download | |
| # Download from HuggingFace | |
| bars[index].set_postfix_str("downloading...") | |
| max_retries = 5 | |
| for attempt in range(max_retries): | |
| try: | |
| ds = fn() | |
| # Save to disk for next time | |
| ds.save_to_disk(save_path) | |
| bars[index].update(1) | |
| bars[index].set_postfix_str(f"✓ {len(ds)} examples") | |
| return label, ds | |
| except Exception as e: | |
| if "429" in str(e) or "rate limit" in str(e).lower(): | |
| wait_time = 60 * (attempt + 1) | |
| bars[index].set_postfix_str(f"rate limited, retry in {wait_time}s...") | |
| time.sleep(wait_time) | |
| else: | |
| bars[index].set_postfix_str(f"error, retry {attempt+1}/{max_retries}...") | |
| time.sleep(10 * (attempt + 1)) | |
| if attempt == max_retries - 1: | |
| bars[index].set_postfix_str(f"✗ FAILED: {e}") | |
| raise | |
| with ThreadPoolExecutor(max_workers=3) as executor: | |
| futures = {executor.submit(run_task, i, label, fn): label for i, (label, fn) in enumerate(tasks)} | |
| for future in as_completed(futures): | |
| label, ds = future.result() | |
| results[label] = ds | |
| # Close all bars and move cursor below them | |
| for bar in bars: | |
| bar.close() | |
| print(f"\nAll {len(tasks)} datasets loaded!") | |
| # Collect in original order | |
| datasets_list = [results[label] for label, _ in tasks] | |
| # Combine all datasets | |
| print("\nCombining datasets...") | |
| combined_dataset = concatenate_datasets(datasets_list) | |
| print(f"Total examples: {len(combined_dataset)}") | |
| # Tokenize dataset | |
| def tokenize_fn(examples): | |
| return tokenizer(examples["text"], truncation=True, padding="max_length", max_length=128) | |
| print("Tokenizing...") | |
| tokenized_dataset = combined_dataset.map(tokenize_fn, batched=True) | |
| print("Done!") | |