DL_Empa / inference.py
Daimka's picture
Upload 3 files
607293d verified
Raw
History Blame Contribute Delete
3.68 kB
"""
Inference script for trained models
Usage: python inference.py --dataset emotion --text "I am feeling happy today"
"""
import torch
import argparse
import numpy as np
from model import SimpleRNN
def load_model(checkpoint_path, device):
"""Load trained model from checkpoint"""
checkpoint = torch.load(checkpoint_path, map_location=device)
model_config = checkpoint['model_config']
vocab = checkpoint['vocab']
num_classes = checkpoint['num_classes']
model = SimpleRNN(
vocab_size=len(vocab),
num_classes=num_classes,
**model_config
).to(device)
model.load_state_dict(checkpoint['model_state_dict'])
model.eval()
return model, vocab, num_classes
def preprocess_text(text, vocab, max_length=128):
"""Preprocess text for inference"""
tokens = text.lower().split()
sequence = [vocab.get(token, vocab['<UNK>']) for token in tokens]
if len(sequence) > max_length:
sequence = sequence[:max_length]
else:
sequence = sequence + [vocab['<PAD>']] * (max_length - len(sequence))
return torch.tensor([sequence], dtype=torch.long)
def predict(model, text, vocab, device, max_length=128):
"""Make prediction on a single text"""
model.eval()
# Preprocess
input_tensor = preprocess_text(text, vocab, max_length).to(device)
# Predict
with torch.no_grad():
output = model(input_tensor)
probabilities = torch.softmax(output, dim=1)
predicted_class = torch.argmax(output, dim=1).item()
confidence = probabilities[0][predicted_class].item()
# Convert to numpy array safely
# Always use tolist() first to avoid numpy compatibility issues
probs_tensor = probabilities[0].cpu().detach()
probs_list = probs_tensor.tolist() # Convert to Python list first
probs_array = np.array(probs_list, dtype=np.float32) # Then to numpy array
return predicted_class, confidence, probs_array
def main():
parser = argparse.ArgumentParser(description='Inference with trained RNN model')
parser.add_argument('--dataset', type=str, choices=['emotion', 'ag_news'], required=True,
help='Dataset type: emotion or ag_news')
parser.add_argument('--text', type=str, required=True, help='Text to classify')
parser.add_argument('--checkpoint', type=str, default=None,
help='Path to checkpoint (default: checkpoints/best_model_{dataset}.pt)')
args = parser.parse_args()
# Set device
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
# Load model
checkpoint_path = args.checkpoint or f'checkpoints/best_model_{args.dataset}.pt'
print(f"Loading model from {checkpoint_path}...")
model, vocab, num_classes = load_model(checkpoint_path, device)
# Class labels
if args.dataset == 'emotion':
class_labels = ['anger', 'fear', 'joy', 'love', 'sadness', 'surprise']
else: # ag_news
class_labels = ['World', 'Sports', 'Business', 'Science/Technology']
# Predict
predicted_class, confidence, probabilities = predict(model, args.text, vocab, device)
# Print results
print(f"\n{'='*60}")
print(f"Input Text: {args.text}")
print(f"{'='*60}")
print(f"\nPredicted Class: {class_labels[predicted_class]}")
print(f"Confidence: {confidence:.4f}")
print(f"\nAll Probabilities:")
for i, (label, prob) in enumerate(zip(class_labels, probabilities)):
marker = " <--" if i == predicted_class else ""
print(f" {label}: {prob:.4f}{marker}")
if __name__ == '__main__':
main()