#!/usr/bin/env python3 """ Quick PALADIM Test Script Run this to test your trained model quickly """ import torch import sys import os sys.path.insert(0, 'd:/CLS/paladim') from paladim import PALADIM from config import PALADIMConfig from transformers import AutoTokenizer def main(): print("=" * 80) print("PALADIM Quick Test") print("=" * 80) # 1. Initialize model print("\nInitializing PALADIM...") config = PALADIMConfig() config.device = 'cpu' # Force CPU (no CUDA) model = PALADIM(config) tokenizer = AutoTokenizer.from_pretrained(config.model_name) print(f"Model initialized!") print(f"Total parameters: {sum(p.numel() for p in model.parameters()):,}") # 2. Load trained model model_path = "d:/CLS/paladim/paladim_20251129_203522.pt" if os.path.exists(model_path): print(f"\nLoading trained model...") checkpoint = torch.load(model_path, map_location="cpu", weights_only=False) if isinstance(checkpoint, dict) and 'model_state_dict' in checkpoint: model.load_state_dict(checkpoint['model_state_dict'], strict=False) if 'epoch' in checkpoint: print(f"Trained for {checkpoint['epoch']} epochs") else: model.load_state_dict(checkpoint, strict=False) model.eval() print(f"Model loaded and ready!") else: print(f"Model file not found, using untrained model") # 3. Test cases test_cases = [ "Patient with hypertension and diabetes, currently on metformin", "45 year old male with high blood pressure and cholesterol", "Patient reports chest pain and shortness of breath", "Elderly patient with heart failure on multiple medications", "Young patient with newly diagnosed type 2 diabetes" ] print(f"\nTesting {len(test_cases)} patient cases:") print("=" * 80) for i, case in enumerate(test_cases, 1): # Tokenize inputs = tokenizer(case, return_tensors='pt', padding=True, truncation=True, max_length=512) # Predict with torch.no_grad(): outputs = model(**inputs) # Get probabilities probs = torch.softmax(outputs.logits, dim=-1) pred_idx = torch.argmax(probs, dim=-1).item() confidence = probs[0][pred_idx].item() # Get top 3 top_3 = torch.topk(probs[0], min(3, probs.shape[-1])) print(f"\nCase {i}: {case[:60]}...") print(f" Predicted drug class: {pred_idx}") print(f" Confidence: {confidence:.2%}") print(f" Top 3: ", end="") for idx, prob in zip(top_3.indices, top_3.values): print(f"#{idx.item()}({prob.item():.1%}) ", end="") print() print("\n" + "=" * 80) print("Testing complete!") print("\nNext steps:") print(" 1. Open test_paladim_quick.ipynb for interactive testing") print(" 2. Add your own patient cases") print(" 3. Check drug_mapping.json for drug names") print("=" * 80) if __name__ == "__main__": main()