CoolFace
Modelpublic

Premchan369/Q-TensorFormer

sourceHugging Faceapache-2.0updated 9d agoView on Hugging Face
2likes185downloads
data.py181 linesDownload Raw Back to src
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