chest-xray-classifier / src /inference.py
zaidshabbir's picture
Upload 9 files
81d1ded verified
Raw
History Blame Contribute Delete
3.82 kB
"""
inference.py – Standalone predictor class.
Used by app.py and can be imported independently by any downstream service.
Example:
from src.inference import ChestXRayPredictor
predictor = ChestXRayPredictor("configs/config.yaml")
result = predictor.predict(Image.open("xray.jpg"))
# {
# "label": "PNEUMONIA",
# "confidence": 0.94,
# "probabilities": {"NORMAL": 0.06, "PNEUMONIA": 0.94}
# }
"""
import logging
import sys
from pathlib import Path
from typing import Dict, Union
import torch
import torch.nn.functional as F
import yaml
from PIL import Image
from torchvision import transforms
sys.path.insert(0, str(Path(__file__).parent))
from dataset import IMAGENET_MEAN, IMAGENET_STD, LABEL_NAMES
from model import build_model
logger = logging.getLogger(__name__)
class ChestXRayPredictor:
"""
Self-contained inference wrapper.
The class loads the model once at initialisation and exposes a single
`predict()` method, keeping the preprocessing pipeline inside the class
to ensure consistency between training and inference.
"""
def __init__(self, config_path: str = "configs/config.yaml", checkpoint_path: str = None):
"""
Args:
config_path: Path to config.yaml.
checkpoint_path: Path to .pth checkpoint. If None, reads from
cfg['inference']['model_path'].
"""
with open(config_path) as f:
self.cfg = yaml.safe_load(f)
self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
image_size = self.cfg["data"]["image_size"]
# Preprocessing (identical to val/test transforms used during training)
self.transform = transforms.Compose([
transforms.Resize((image_size, image_size)),
transforms.ToTensor(),
transforms.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),
])
# Resolve checkpoint path
if checkpoint_path is None:
checkpoint_path = self.cfg["inference"]["model_path"]
ckpt_path = Path(checkpoint_path)
if not ckpt_path.exists():
raise FileNotFoundError(
f"Model checkpoint not found: {ckpt_path}\n"
"Train a model first with: python src/train.py --model mobilenet_v2"
)
logger.info(f"Loading checkpoint: {ckpt_path}")
ckpt = torch.load(ckpt_path, map_location=self.device, weights_only=False)
self.model = build_model(self.cfg).to(self.device)
self.model.load_state_dict(ckpt["model_state"])
self.model.eval()
logger.info(
f"Model ready | epoch={ckpt['epoch']}, val_acc={ckpt['val_acc']:.4f}"
)
@torch.no_grad()
def predict(self, image: Union[Image.Image, str, Path]) -> Dict:
"""
Run inference on a single image.
Args:
image: PIL Image, or a path (str / Path) to an image file.
Returns:
dict with:
label – predicted class name ('NORMAL' or 'PNEUMONIA')
confidence – probability of the predicted class (float)
probabilities – {class_name: probability} for all classes
"""
if isinstance(image, (str, Path)):
image = Image.open(image)
image = image.convert("RGB")
tensor = self.transform(image).unsqueeze(0).to(self.device)
logits = self.model(tensor)
probs = F.softmax(logits, dim=1).squeeze(0).cpu().numpy()
label_idx = int(probs.argmax())
return {
"label": LABEL_NAMES[label_idx],
"confidence": float(probs[label_idx]),
"probabilities": {LABEL_NAMES[i]: float(p) for i, p in enumerate(probs)},
}