""" MANAS - Model for Awadhi Natural Autoregressive Sequences Interactive inference script. Run: python run.py Requirements: pip install -r requirements.txt Then download the model: python -c "from huggingface_hub import hf_hub_download; hf_hub_download('JayF14/MANAS', 'tulsidas_model.pth', local_dir='.')" """ import torch import torch.nn as nn from torch.nn import functional as F import json import re import sys import os try: from indic_transliteration import sanscript from indic_transliteration.sanscript import transliterate TRANSLITERATION_AVAILABLE = True except ImportError: TRANSLITERATION_AVAILABLE = False print("Note: indic-transliteration not installed. Only Devanagari input will work.") print("Install with: pip install indic-transliteration\n") # Hyperparameters (must match training exactly) block_size = 256 n_embd = 384 n_head = 6 n_layer = 6 dropout = 0.0 # 0 for inference device = 'cuda' if torch.cuda.is_available() else 'cpu' # Load vocabulary from vocab.json VOCAB_FILE = os.path.join(os.path.dirname(__file__), 'vocab.json') MODEL_FILE = os.path.join(os.path.dirname(__file__), 'tulsidas_model.pth') if not os.path.exists(VOCAB_FILE): print("Error: vocab.json not found. Make sure it is in the same folder as this script.") sys.exit(1) if not os.path.exists(MODEL_FILE): print("Error: tulsidas_model.pth not found.") print("Download it with:") print(' python -c "from huggingface_hub import hf_hub_download; hf_hub_download(\'JayF14/MANAS\', \'tulsidas_model.pth\', local_dir=\'.\')"\n') sys.exit(1) with open(VOCAB_FILE, 'r', encoding='utf-8') as f: vocab = json.load(f) vocab_size = vocab['vocab_size'] stoi = vocab['stoi'] itos = {int(k): v for k, v in vocab['itos'].items()} def encode(s): unknown = [c for c in s if c not in stoi] if unknown: raise ValueError(f"Character(s) not in model vocabulary: {unknown}") return [stoi[c] for c in s] def decode(l): return ''.join([itos[i] for i in l]) # Model architecture class Head(nn.Module): def __init__(self, head_size): super().__init__() self.key = nn.Linear(n_embd, head_size, bias=False) self.query = nn.Linear(n_embd, head_size, bias=False) self.value = nn.Linear(n_embd, head_size, bias=False) self.register_buffer('tril', torch.tril(torch.ones(block_size, block_size))) self.dropout = nn.Dropout(dropout) def forward(self, x): B, T, C = x.shape k, q = self.key(x), self.query(x) wei = q @ k.transpose(-2, -1) * (k.shape[-1] ** -0.5) wei = wei.masked_fill(self.tril[:T, :T] == 0, float('-inf')) wei = F.softmax(wei, dim=-1) return self.dropout(wei) @ self.value(x) class MultiHeadAttention(nn.Module): def __init__(self, num_heads, head_size): super().__init__() self.heads = nn.ModuleList([Head(head_size) for _ in range(num_heads)]) self.proj = nn.Linear(n_embd, n_embd) self.dropout = nn.Dropout(dropout) def forward(self, x): return self.dropout(self.proj(torch.cat([h(x) for h in self.heads], dim=-1))) class FeedForward(nn.Module): def __init__(self, n_embd): super().__init__() self.net = nn.Sequential( nn.Linear(n_embd, 4 * n_embd), nn.ReLU(), nn.Linear(4 * n_embd, n_embd), nn.Dropout(dropout), ) def forward(self, x): return self.net(x) class Block(nn.Module): def __init__(self, n_embd, n_head): super().__init__() head_size = n_embd // n_head self.sa = MultiHeadAttention(n_head, head_size) self.ffwd = FeedForward(n_embd) self.ln1 = nn.LayerNorm(n_embd) self.ln2 = nn.LayerNorm(n_embd) def forward(self, x): x = x + self.sa(self.ln1(x)) x = x + self.ffwd(self.ln2(x)) return x class LanguageModel(nn.Module): def __init__(self): super().__init__() self.token_embedding_table = nn.Embedding(vocab_size, n_embd) self.position_embedding_table = nn.Embedding(block_size, n_embd) self.blocks = nn.Sequential(*[Block(n_embd, n_head=n_head) for _ in range(n_layer)]) self.ln_f = nn.LayerNorm(n_embd) self.lm_head = nn.Linear(n_embd, vocab_size) def forward(self, idx, targets=None): B, T = idx.shape x = self.token_embedding_table(idx) + self.position_embedding_table(torch.arange(T, device=device)) logits = self.lm_head(self.ln_f(self.blocks(x))) return logits, None def generate(self, idx, max_new_tokens): for _ in range(max_new_tokens): logits, _ = self(idx[:, -block_size:]) idx_next = torch.multinomial(F.softmax(logits[:, -1, :], dim=-1), num_samples=1) idx = torch.cat((idx, idx_next), dim=1) return idx # Load model print(f"\nUsing device: {device}") print("Loading MANAS model...") model = LanguageModel() model.load_state_dict(torch.load(MODEL_FILE, map_location=device)) model.to(device) model.eval() print("Model loaded successfully!\n") # Common ITRANS overrides for accurate Hindi transliteration OVERRIDES = { "shri ram": "श्री राम", "shri rama": "श्री राम", "jai shri ram": "जय श्री राम", "ram": "राम", "sita": "सीता", "hanuman": "हनुमान", "tulsidas": "तुलसीदास", "ramcharitmanas": "रामचरितमानस", } def to_devanagari(user_input): lower = user_input.lower().strip() if lower in OVERRIDES: return OVERRIDES[lower], True if re.search(r'[a-zA-Z]', user_input) and TRANSLITERATION_AVAILABLE: return transliterate(user_input, sanscript.ITRANS, sanscript.DEVANAGARI), True return user_input, False # Interactive loop print("=" * 55) print(" MANAS - Model for Awadhi Natural Autoregressive Sequences") print("=" * 55) print("Type a Hindi/Awadhi prompt in Devanagari or English.") print("Examples: shri ram | tulsidas | jai shri ram") print("Type 'quit' or press Ctrl+C to exit.") print() print("-" * 55) print(" IF HINDI TEXT LOOKS BROKEN (overlapping/garbled):") print() print(" Windows Terminal (recommended fix):") print(" Settings > Profiles > Appearance > Font face") print(" -> Set to: Nirmala UI or Mangal") print() print(" VS Code terminal:") print(" Settings > terminal.integrated.fontFamily") print(" -> Set to: Nirmala UI") print() print(" Old cmd.exe / PowerShell window:") print(" Right-click title bar > Properties > Font") print(" -> Select: NSimSun (best fallback)") print() print(" Best overall: Use Windows Terminal (free from") print(" Microsoft Store) with Nirmala UI font.") print("-" * 55) print() while True: try: user_input = input("You: ").strip() if not user_input: continue if user_input.lower() in ('quit', 'exit', 'q'): print("जय श्री राम 🙏") break hindi_input, was_translated = to_devanagari(user_input) if was_translated: print(f"[Transliterated] → {hindi_input}") try: context = torch.tensor([encode(hindi_input)], dtype=torch.long, device=device) tokens = model.generate(context, max_new_tokens=400)[0].tolist() print(f"\nMANAS:\n{decode(tokens)}\n") print("-" * 55) except ValueError as e: print(f"[Error] {e}") print("Tip: Type in Devanagari or standard ITRANS (e.g. 'shri ram').\n") except KeyboardInterrupt: print("\nजय श्री राम 🙏") break