Team Ai
Apppublic

abprasadhuggingface/Hindi-BPE-Encoder-Decoder

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
bpe.py236 linesDownload Raw Back to root
1from typing import Dict, List, Tuple, Set2from collections import defaultdict3import re4import json5import os6 7class HindiBPE:8    def __init__(self, vocab_size: int = 5000):9        self.vocab_size = vocab_size10        self.merges: Dict[Tuple[str, str], str] = {}11        self.vocab: Set[str] = set()12        self.reverse_merges: Dict[str, Tuple[str, str]] = {}13        self.token_to_index: Dict[str, int] = {}14        self.index_to_token: Dict[int, str] = {}15        self.UNK_TOKEN = "<UNK>"16        self.min_freq = 217        18    def get_stats(self, words: List[List[str]]) -> Dict[Tuple[str, str], int]:19        """Count frequency of adjacent pairs"""20        pairs = defaultdict(int)21        22        for word in words:23            # Count pairs within each word24            for i in range(len(word) - 1):25                pairs[tuple(word[i:i+2])] += 126            27            # Also consider pairs across word boundaries for better compression28            if len(word) > 2:29                for i in range(1, len(word) - 2):30                    pairs[tuple(word[i:i+3])] += 131                32        return pairs33    34    def merge_vocab(self, words: List[List[str]], pair: Tuple[str, str]) -> List[List[str]]:35        """Merge all occurrences of the most frequent pair"""36        first, second = pair37        new_token = first + second38        39        new_words = []40        for word in words:41            i = 042            new_word = []43            while i < len(word):44                if i < len(word) - 1 and word[i] == first and word[i+1] == second:45                    new_word.append(new_token)46                    i += 247                else:48                    new_word.append(word[i])49                    i += 150            new_words.append(new_word)51            52        return new_words53    54    def fit(self, texts: List[str]) -> None:55        """Learn BPE merges from texts"""56        # Preprocess texts to create larger chunks for better compression57        processed_texts = []58        current_chunk = []59        chunk_size = 100  # Characters per chunk60        61        for text in texts:62            current_chunk.extend(list(text))63            if len(current_chunk) >= chunk_size:64                processed_texts.append(''.join(current_chunk))65                current_chunk = []66        if current_chunk:67            processed_texts.append(''.join(current_chunk))68        69        # Initialize vocabulary with characters70        words = [[char for char in text] for text in processed_texts]71        self.vocab = set(char for text in processed_texts for char in text)72        self.vocab.add(self.UNK_TOKEN)73        74        # Initialize token indices75        self.token_to_index = {self.UNK_TOKEN: 0}76        self.index_to_token = {0: self.UNK_TOKEN}77        78        for idx, token in enumerate(sorted(self.vocab - {self.UNK_TOKEN}), start=1):79            self.token_to_index[token] = idx80            self.index_to_token[idx] = token81        82        next_idx = len(self.vocab)83        84        # Learn merges85        for i in range(self.vocab_size - len(self.vocab)):86            pairs = self.get_stats(words)87            if not pairs:88                break89            90            # Filter pairs by minimum frequency91            frequent_pairs = {pair: freq for pair, freq in pairs.items() 92                            if freq >= self.min_freq}93            if not frequent_pairs:94                break95            96            best_pair = max(frequent_pairs.items(), key=lambda x: x[1])[0]97            words = self.merge_vocab(words, best_pair)98            merged_token = ''.join(best_pair)99            100            # Only add to vocabulary if it provides good compression101            if len(merged_token) < len(best_pair[0]) + len(best_pair[1]):102                self.merges[best_pair] = merged_token103                self.reverse_merges[merged_token] = best_pair104                self.vocab.add(merged_token)105                self.token_to_index[merged_token] = next_idx106                self.index_to_token[next_idx] = merged_token107                next_idx += 1108    109    def encode(self, text: str) -> List[int]:110        """Encode text using learned BPE merges and return indices"""111        if not text:112            return []113        114        # Split text into optimal chunks115        chunks = [text[i:i+50] for i in range(0, len(text), 50)]116        final_tokens = []117        118        for chunk in chunks:119            word = [char for char in chunk]120            121            while True:122                pairs = [(word[i], word[i+1]) 123                        for i in range(len(word)-1)]124                if not pairs:125                    break126                127                mergeable_pairs = [pair for pair in pairs 128                                 if pair in self.merges]129                if not mergeable_pairs:130                    break131                132                # Merge all possible pairs in one pass133                i = 0134                new_word = []135                while i < len(word):136                    if (i < len(word) - 1 and 137                        (word[i], word[i+1]) in self.merges):138                        new_word.append(self.merges[(word[i], word[i+1])])139                        i += 2140                    else:141                        new_word.append(word[i])142                        i += 1143                word = new_word144            145            final_tokens.extend(word)146        147        return [self.token_to_index.get(token, self.token_to_index[self.UNK_TOKEN]) 148                for token in final_tokens]149    150    def decode_token(self, token: str, max_depth: int = 100) -> str:151        """Recursively decode a single token with depth limit"""152        if max_depth <= 0 or token not in self.reverse_merges:153            return token154            155        first, second = self.reverse_merges[token]156        return self.decode_token(first, max_depth - 1) + self.decode_token(second, max_depth - 1)157    158    def decode(self, indices: List[int]) -> str:159        """Decode indices back to text"""160        try:161            tokens = [self.index_to_token.get(idx, self.UNK_TOKEN) 162                     for idx in indices]163            164            result = []165            for token in tokens:166                if token == self.UNK_TOKEN:167                    continue168                decoded = self.decode_token(token)169                result.append(decoded)170            171            return ''.join(result)172            173        except Exception as e:174            print(f"Error during decoding: {e}")175            return ""176    177    def get_token_mapping(self) -> Dict[int, str]:178        """Return the mapping of indices to tokens"""179        return self.index_to_token 180    181    def save_model(self, path: str = "data/bpe_model.json"):182        """Save the BPE model's vocabulary and mappings"""183        model_data = {184            'vocab_size': self.vocab_size,185            'token_to_index': self.token_to_index,186            'index_to_token': {str(k): v for k, v in self.index_to_token.items()},187            'merges': {f"{k[0]}|{k[1]}": v for k, v in self.merges.items()},188            'reverse_merges': {k: f"{v[0]}|{v[1]}" for k, v in self.reverse_merges.items()}189        }190        191        os.makedirs(os.path.dirname(path), exist_ok=True)192        with open(path, 'w', encoding='utf-8') as f:193            json.dump(model_data, f, ensure_ascii=False, indent=2)194    195    def save_encoded_text(self, text: str, indices: List[int], path: str = "data/encoded_texts.json"):196        """Save encoded text with its tokens and mappings"""197        token_mapping = self.get_token_mapping()198        tokens = [token_mapping[idx] for idx in indices]199        200        # Create encoded text entry201        encoded_data = {202            'original_text': text,203            'indices': indices,204            'tokens': tokens,205            'token_mappings': {str(idx): token for idx, token in zip(indices, tokens)}206        }207        208        # Load existing data if file exists209        if os.path.exists(path):210            with open(path, 'r', encoding='utf-8') as f:211                try:212                    all_encoded_data = json.load(f)213                except json.JSONDecodeError:214                    all_encoded_data = {'texts': []}215        else:216            all_encoded_data = {'texts': []}217        218        # Add new encoded text219        all_encoded_data['texts'].append(encoded_data)220        221        # Save updated data222        os.makedirs(os.path.dirname(path), exist_ok=True)223        with open(path, 'w', encoding='utf-8') as f:224            json.dump(all_encoded_data, f, ensure_ascii=False, indent=2)225            226    def load_model(self, path: str = "data/bpe_model.json"):227        """Load the BPE model's vocabulary and mappings"""228        with open(path, 'r', encoding='utf-8') as f:229            model_data = json.load(f)230            231        self.vocab_size = model_data['vocab_size']232        self.token_to_index = model_data['token_to_index']233        self.index_to_token = {int(k): v for k, v in model_data['index_to_token'].items()}234        self.merges = {tuple(k.split('|')): v for k, v in model_data['merges'].items()}235        self.reverse_merges = {k: tuple(v.split('|')) for k, v in model_data['reverse_merges'].items()}236        self.vocab = set(self.token_to_index.keys())