abprasadhuggingface/Hindi-BPE-Encoder-Decoder
0
1from download_data import download_hindi_corpus2from preprocessor import load_and_preprocess_data3from bpe import HindiBPE4import statistics5import os6 7def train_and_save_model():8 """Train BPE model and save it to file"""9 # Download the corpus10 corpus_path = download_hindi_corpus()11 12 # Load and preprocess the data13 processed_texts = load_and_preprocess_data(corpus_path)14 15 # Initialize and train BPE16 vocab_size = 500017 bpe = HindiBPE(vocab_size=vocab_size)18 print(f"Training BPE with vocabulary size {vocab_size}...")19 bpe.fit(processed_texts[:1000])20 21 # Save the model22 bpe.save_model()23 print("Model saved successfully!")24 return bpe25 26def load_or_train_model():27 """Load existing model or train new one if not exists"""28 model_path = "data/bpe_model.json"29 if os.path.exists(model_path):30 bpe = HindiBPE()31 bpe.load_model(model_path)32 print("Loaded existing model.")33 return bpe34 else:35 return train_and_save_model()36 37def calculate_stats(original_text: str, indices: list, token_mapping: dict) -> dict:38 """Calculate compression statistics"""39 original_chars = len(original_text)40 encoded_length = len(indices)41 compression_ratio = original_chars / encoded_length if encoded_length > 0 else 042 43 tokens = [token_mapping[idx] for idx in indices]44 token_lengths = [len(token) for token in tokens]45 avg_token_length = statistics.mean(token_lengths) if token_lengths else 046 47 return {48 'original_length': original_chars,49 'encoded_length': encoded_length,50 'compression_ratio': compression_ratio,51 'avg_token_length': avg_token_length,52 'min_token_length': min(token_lengths) if token_lengths else 0,53 'max_token_length': max(token_lengths) if token_lengths else 054 }55 56def main():57 # Load or train model58 bpe = load_or_train_model()59 60 # Get token mapping61 token_mapping = bpe.get_token_mapping()62 63 # Test the encoding and decoding with multiple examples64 test_texts = [65 "नमस्ते भारत",66 "भारतीय संस्कृति विविधता में एकता का प्रतीक है",67 "हिंदी भारत की सबसे अधिक बोली जाने वाली भाषा है"68 ]69 70 print("\nEncoding/Decoding Statistics:")71 print("-" * 50)72 73 all_compression_ratios = []74 all_token_lengths = []75 76 for test_text in test_texts:77 indices = bpe.encode(test_text)78 decoded = bpe.decode(indices)79 stats = calculate_stats(test_text, indices, token_mapping)80 81 print(f"\nTest text: '{test_text}'")82 print(f"Original length: {stats['original_length']} characters")83 print(f"Encoded length: {stats['encoded_length']} tokens")84 print(f"Compression ratio: {stats['compression_ratio']:.2f}x")85 print(f"Average token length: {stats['avg_token_length']:.2f} characters")86 print(f"Token length range: {stats['min_token_length']} to {stats['max_token_length']} characters")87 print(f"Encoded indices: {indices}")88 print(f"Corresponding tokens: {[token_mapping[idx] for idx in indices]}")89 print(f"Decoded text: '{decoded}'")90 print(f"Successful roundtrip: {test_text == decoded}")91 92 all_compression_ratios.append(stats['compression_ratio'])93 all_token_lengths.extend([len(token_mapping[idx]) for idx in indices])94 95 # Print overall statistics96 print("\nOverall Statistics:")97 print("-" * 50)98 print(f"Vocabulary size: {len(bpe.vocab)}")99 print(f"Average compression ratio: {statistics.mean(all_compression_ratios):.2f}x")100 print(f"Average token length: {statistics.mean(all_token_lengths):.2f} characters")101 print(f"Token length range: {min(all_token_lengths)} to {max(all_token_lengths)} characters")102 103if __name__ == "__main__":104 main() 