File size: 5,063 Bytes
9ee4de4 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 | #!/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]}...")
|