1import torch
2import torch.nn as nn
3import torch.nn.functional as F
4from transformers import AutoModel, AutoTokenizer
5from peft import get_peft_model, LoraConfig, TaskType
6from huggingface_hub import hf_hub_download
7
8class DhvaniV6(nn.Module):
9 def __init__(self, cfg):
10 super().__init__()
11 base = AutoModel.from_pretrained(
12 cfg['base_model'], torch_dtype=torch.bfloat16,
13 attn_implementation='eager', trust_remote_code=True
14 )
15 lora_config = LoraConfig(
16 r=cfg['lora_r'], lora_alpha=cfg['lora_alpha'],
17 lora_dropout=cfg['lora_dropout'],
18 target_modules=cfg['lora_targets'],
19 bias='none', task_type=TaskType.FEATURE_EXTRACTION
20 )
21 self.base = get_peft_model(base, lora_config)
22 self.trunk = nn.Sequential(
23 nn.Linear(cfg['hidden_dim'], cfg['trunk_dim']),
24 nn.LayerNorm(cfg['trunk_dim']),
25 nn.GELU(),
26 )
27 self.abhida_head = nn.Sequential(
28 nn.Linear(cfg['trunk_dim'], cfg['subspace_dim']),
29 nn.LayerNorm(cfg['subspace_dim']),
30 )
31 self.vyanjana_head = nn.Sequential(
32 nn.Linear(cfg['trunk_dim'], cfg['subspace_dim']),
33 nn.LayerNorm(cfg['subspace_dim']),
34 )
35 self.surface_head = nn.Sequential(
36 nn.Linear(cfg['trunk_dim'], cfg['subspace_dim']),
37 nn.LayerNorm(cfg['subspace_dim']),
38 )
39
40 @staticmethod
41 def mean_pool(hidden, mask):
42 m = mask.unsqueeze(-1).float()
43 return (hidden * m).sum(1) / m.sum(1).clamp(min=1e-9)
44
45 def encode_tokens(self, input_ids, attention_mask):
46 out = self.base(input_ids=input_ids, attention_mask=attention_mask)
47 pooled = self.mean_pool(out.last_hidden_state.float(), attention_mask)
48 trunk = self.trunk(pooled)
49 return {
50 'abhida': F.normalize(self.abhida_head(trunk), p=2, dim=-1),
51 'vyanjana': F.normalize(self.vyanjana_head(trunk), p=2, dim=-1),
52 'surface': F.normalize(self.surface_head(trunk), p=2, dim=-1),
53 'full': F.normalize(torch.cat([
54 self.abhida_head(trunk),
55 self.vyanjana_head(trunk),
56 self.surface_head(trunk),
57 ], dim=-1), p=2, dim=-1),
58 }
59
60# Load model
61ckpt_path = hf_hub_download(repo_id="rb512/dhvani-v6", filename="v6_best.pt")
62ckpt = torch.load(ckpt_path, map_location="cpu", weights_only=False)
63cfg = ckpt["config"]
64
65tokenizer = AutoTokenizer.from_pretrained(cfg['base_model'], trust_remote_code=True)
66if tokenizer.pad_token is None:
67 tokenizer.pad_token = tokenizer.eos_token
68
69model = DhvaniV6(cfg)
70model.base.load_state_dict(ckpt["lora"])
71model.trunk.load_state_dict(ckpt["trunk"])
72model.abhida_head.load_state_dict(ckpt["abhida_head"])
73model.vyanjana_head.load_state_dict(ckpt["vyanjana_head"])
74model.surface_head.load_state_dict(ckpt["surface_head"])
75model.eval()
76
77# Encode text
78texts = ["Hello world", "Greetings, world"]
79enc = tokenizer(texts, max_length=128, truncation=True, padding='max_length', return_tensors='pt')
80with torch.no_grad():
81 embeddings = model.encode_tokens(enc['input_ids'], enc['attention_mask'])
82
83print("Surface embeddings:", embeddings['surface'].shape) # semantic similarity
84print("Abhida embeddings:", embeddings['abhida'].shape) # meaning-invariant
85print("Vyanjana embeddings:", embeddings['vyanjana'].shape) # expression style
86print("Full embeddings:", embeddings['full'].shape) # concatenated