bsbarkur's picture
Upload folder using huggingface_hub
9ee4de4 verified
Raw
History Blame Contribute Delete
5.06 kB
#!/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]}...")