MANAS / run.py
JayF14's picture
Upload run.py with huggingface_hub
994b3c9 verified
Raw History Blame Contribute Delete
7.71 kB
"""
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