#!/usr/bin/env python3 """ Inference script for rgveda-embedding-gemma. This provides ONNX-like inference using PyTorch model with optimized settings. """ import torch import numpy as np from transformers import AutoTokenizer, AutoModel from pathlib import Path class RgvedaEmbeddingInference: """ Optimized inference for rgveda-embedding-gemma model. Uses PyTorch for transformer, numpy for post-processing. """ def __init__(self, model_dir="."): """Initialize the model.""" print("Loading model...") self.model_dir = Path(model_dir) # Load tokenizer self.tokenizer = AutoTokenizer.from_pretrained(str(self.model_dir)) # Load transformer model self.model = AutoModel.from_pretrained( "Ganaraj/rgveda-embedding-gemma" ) self.model.eval() self.model = self.model.to('cpu') # Or 'cuda' if available # Load dense layer weights weights_dir = self.model_dir / "weights" self.dense1_weight = np.load(weights_dir / "dense1_weight.npy") self.dense2_weight = np.load(weights_dir / "dense2_weight.npy") print(f"Model loaded successfully!") print(f"Device: {next(self.model.parameters()).device}") def mean_pooling(self, token_embeddings, attention_mask): """Mean pooling with attention mask.""" input_mask_expanded = attention_mask.unsqueeze(-1).expand(token_embeddings.size()).float() sum_embeddings = torch.sum(token_embeddings * input_mask_expanded, 1) sum_mask = torch.clamp(input_mask_expanded.sum(1), min=1e-9) return sum_embeddings / sum_mask def encode(self, texts, batch_size=32, show_progress=False): """ Encode texts to embeddings. Args: texts: List of strings or single string batch_size: Batch size for processing show_progress: Show progress bar Returns: embeddings: numpy array of shape (num_texts, 768) """ if isinstance(texts, str): texts = [texts] all_embeddings = [] # Process in batches for i in range(0, len(texts), batch_size): batch_texts = texts[i:i+batch_size] # Tokenize inputs = self.tokenizer( batch_texts, padding=True, truncation=True, max_length=2048, return_tensors="pt" ) # Move to same device as model device = next(self.model.parameters()).device inputs = {k: v.to(device) for k, v in inputs.items()} # Get embeddings with torch.no_grad(): outputs = self.model(**inputs) token_embeddings = outputs.last_hidden_state # Mean pooling pooled = self.mean_pooling(token_embeddings, inputs['attention_mask']) # Convert to numpy for dense layers pooled_np = pooled.cpu().numpy() # Dense layer 1 (768 -> 3072) dense1_out = pooled_np @ self.dense1_weight.T # Dense layer 2 (3072 -> 768) dense2_out = dense1_out @ self.dense2_weight.T # L2 normalization norms = np.linalg.norm(dense2_out, axis=1, keepdims=True) normalized = dense2_out / np.clip(norms, a_min=1e-9, a_max=None) all_embeddings.append(normalized) return np.vstack(all_embeddings) # Example usage if __name__ == "__main__": # Initialize model model = RgvedaEmbeddingInference(".") # Test queries and documents with Devanagari script prefixes = { "query": "task: search result | query: ", "document": "title: none | text: ", } query = prefixes["query"] + "वृष्टि-विद्युत्-सदृशं दैविकं आगमनम्" documents = [ prefixes["document"] + "असामि हि प्रयज्यवः कण्वं दद प्रचेतसः", prefixes["document"] + "उत द्वार उशतीर् वि श्रयन्ताम् उत देवाṁ उशत आ वहेह", prefixes["document"] + "प्राग्नये बृहते यज्ञियाय ऋतस्य वृष्णे असुराय मन्म", ] # Encode query_embedding = model.encode(query) doc_embeddings = model.encode(documents) # Compute similarities similarities = query_embedding @ doc_embeddings.T print("\nQuery:", query) print("\nDocument similarities:") for i, (doc, sim) in enumerate(zip(documents, similarities[0])): print(f" {i+1}. {sim:.4f} - {doc[:60]}...")