| |
| """ |
| 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) |
| |
| |
| self.tokenizer = AutoTokenizer.from_pretrained(str(self.model_dir)) |
| |
| |
| self.model = AutoModel.from_pretrained( |
| "Ganaraj/rgveda-embedding-gemma" |
| ) |
| self.model.eval() |
| self.model = self.model.to('cpu') |
| |
| |
| 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 = [] |
| |
| |
| for i in range(0, len(texts), batch_size): |
| batch_texts = texts[i:i+batch_size] |
| |
| |
| inputs = self.tokenizer( |
| batch_texts, |
| padding=True, |
| truncation=True, |
| max_length=2048, |
| return_tensors="pt" |
| ) |
| |
| |
| device = next(self.model.parameters()).device |
| inputs = {k: v.to(device) for k, v in inputs.items()} |
| |
| |
| with torch.no_grad(): |
| outputs = self.model(**inputs) |
| token_embeddings = outputs.last_hidden_state |
| |
| |
| pooled = self.mean_pooling(token_embeddings, inputs['attention_mask']) |
| |
| |
| pooled_np = pooled.cpu().numpy() |
| |
| |
| dense1_out = pooled_np @ self.dense1_weight.T |
| |
| |
| dense2_out = dense1_out @ self.dense2_weight.T |
| |
| |
| 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) |
|
|
|
|
| |
| if __name__ == "__main__": |
| |
| model = RgvedaEmbeddingInference(".") |
| |
| |
| prefixes = { |
| "query": "task: search result | query: ", |
| "document": "title: none | text: ", |
| } |
| |
| query = prefixes["query"] + "वृष्टि-विद्युत्-सदृशं दैविकं आगमनम्" |
| documents = [ |
| prefixes["document"] + "असामि हि प्रयज्यवः कण्वं दद प्रचेतसः", |
| prefixes["document"] + "उत द्वार उशतीर् वि श्रयन्ताम् उत देवाṁ उशत आ वहेह", |
| prefixes["document"] + "प्राग्नये बृहते यज्ञियाय ऋतस्य वृष्णे असुराय मन्म", |
| ] |
| |
| |
| query_embedding = model.encode(query) |
| doc_embeddings = model.encode(documents) |
| |
| |
| 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]}...") |
|
|