Premchan369/Q-TensorFormer
2185
1"""2Data loading and preprocessing.3 4Supported datasets:5 - WikiText-2 (char-level and word-level)6 - WikiText-1037 - Custom text files8 - Synthetic random data (debugging)9 10Tokenization: character-level by default. Simple, deterministic, no external deps.11"""12 13import torch14from torch.utils.data import Dataset, DataLoader15from typing import Optional, Tuple, Dict16from collections import Counter17 18 19class CharTokenizer:20 """Character-level tokenizer. Vocabulary built from data."""21 22 def __init__(self, min_freq: int = 1):23 self.min_freq = min_freq24 self.char_to_idx: Dict[str, int] = {}25 self.idx_to_char: Dict[int, str] = {}26 self.vocab_size = 027 self.special_tokens = {28 "<pad>": 0,29 "<bos>": 1,30 "<eos>": 2,31 "<unk>": 3,32 }33 34 def fit(self, texts: list[str]):35 """Build vocabulary from texts."""36 char_counts = Counter()37 for text in texts:38 char_counts.update(text)39 40 # Special tokens first41 self.char_to_idx = dict(self.special_tokens)42 # Freq-filtered chars43 idx = len(self.special_tokens)44 for char, count in char_counts.most_common():45 if count >= self.min_freq:46 self.char_to_idx[char] = idx47 idx += 148 49 self.idx_to_char = {v: k for k, v in self.char_to_idx.items()}50 self.vocab_size = len(self.char_to_idx)51 52 def encode(self, text: str, add_bos: bool = True,53 add_eos: bool = True, max_len: int = None) -> list[int]:54 """Convert text to token indices."""55 tokens = []56 if add_bos:57 tokens.append(self.special_tokens["<bos>"])58 for ch in text:59 tokens.append(self.char_to_idx.get(ch, self.special_tokens["<unk>"]))60 if add_eos:61 tokens.append(self.special_tokens["<eos>"])62 if max_len is not None:63 if len(tokens) > max_len:64 tokens = tokens[:max_len]65 else:66 tokens.extend([self.special_tokens["<pad>"]] * (max_len - len(tokens)))67 return tokens68 69 def decode(self, indices: list[int], skip_special: bool = True) -> str:70 """Convert indices back to text."""71 chars = []72 for idx in indices:73 ch = self.idx_to_char.get(idx, "?")74 if skip_special and idx in self.special_tokens.values():75 continue76 chars.append(ch)77 return "".join(chars)78 79 def save(self, path: str):80 torch.save({81 "char_to_idx": self.char_to_idx,82 "idx_to_char": self.idx_to_char,83 "vocab_size": self.vocab_size,84 "special_tokens": self.special_tokens,85 }, path)86 87 @classmethod88 def load(cls, path: str) -> "CharTokenizer":89 data = torch.load(path)90 tok = cls()91 tok.char_to_idx = data["char_to_idx"]92 tok.idx_to_char = data["idx_to_char"]93 tok.vocab_size = data["vocab_size"]94 tok.special_tokens = data["special_tokens"]95 return tok96 97 98class TextDataset(Dataset):99 """100 Causal language modeling dataset.101 102 Splits text into overlapping sequences of length seq_len.103 Target = input shifted by 1 (next-token prediction).104 """105 106 def __init__(self, texts: list[str], tokenizer: CharTokenizer,107 seq_len: int = 128, stride: int = None):108 self.seq_len = seq_len109 self.stride = stride or seq_len // 2110 111 # Tokenize all texts112 all_tokens = []113 for text in texts:114 all_tokens.extend(tokenizer.encode(text, add_bos=False, add_eos=True))115 self.tokens = torch.tensor(all_tokens, dtype=torch.long)116 117 # Compute valid starting positions118 self.n_samples = max(0, (len(self.tokens) - seq_len - 1) // self.stride + 1)119 120 def __len__(self):121 return self.n_samples122 123 def __getitem__(self, idx):124 start = idx * self.stride125 end = start + self.seq_len126 x = self.tokens[start:end]127 y = self.tokens[start + 1:end + 1]128 assert len(x) == len(y) == self.seq_len, f"len={len(x)} at idx={idx}"129 return x, y130 131 132def load_wikitext2(tokenizer: CharTokenizer = None,133 seq_len: int = 128,134 batch_size: int = 16) -> Tuple[DataLoader, DataLoader, DataLoader, CharTokenizer]:135 """136 Load WikiText-2 with char-level tokenization.137 138 Returns:139 train_loader, val_loader, test_loader, tokenizer140 """141 try:142 from datasets import load_dataset143 except ImportError:144 raise ImportError("pip install datasets")145 146 ds = load_dataset("wikitext", "wikitext-2-raw-v1")147 148 # Filter empty lines149 train_texts = [t for t in ds["train"]["text"] if t.strip()]150 val_texts = [t for t in ds["validation"]["text"] if t.strip()]151 test_texts = [t for t in ds["test"]["text"] if t.strip()]152 153 if tokenizer is None:154 tokenizer = CharTokenizer()155 tokenizer.fit(train_texts)156 157 train_ds = TextDataset(train_texts, tokenizer, seq_len)158 val_ds = TextDataset(val_texts, tokenizer, seq_len)159 test_ds = TextDataset(test_texts, tokenizer, seq_len)160 161 train_loader = DataLoader(train_ds, batch_size=batch_size, shuffle=True,162 num_workers=0, drop_last=True)163 val_loader = DataLoader(val_ds, batch_size=batch_size, shuffle=False, num_workers=0)164 test_loader = DataLoader(test_ds, batch_size=batch_size, shuffle=False, num_workers=0)165 166 return train_loader, val_loader, test_loader, tokenizer167 168 169def load_synthetic_data(vocab_size: int = 5000, seq_len: int = 128,170 n_samples: int = 2000, batch_size: int = 16):171 """Synthetic random data for debugging."""172 class _SynthDataset(Dataset):173 def __init__(self, n, vocab, slen):174 self.data = torch.randint(1, vocab, (n, slen + 1))175 def __len__(self):176 return len(self.data)177 def __getitem__(self, i):178 return self.data[i, :-1], self.data[i, 1:]179 ds = _SynthDataset(n_samples, vocab_size, seq_len)180 return DataLoader(ds, batch_size=batch_size, shuffle=True, num_workers=0)181 