abprasadhuggingface/Hindi-BPE-Encoder-Decoder
0
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()) 