""" BatikAI — Inference Module Load model and predict batik motif from an image. """ import json import torch import torch.nn as nn import torch.nn.functional as F import torchvision.transforms as transforms from torchvision import models from PIL import Image import numpy as np from batik_info import get_batik_info MODEL_DIR = "./model" def load_model(model_dir: str = MODEL_DIR): """Load trained model and config.""" # Load config with open(f"{model_dir}/config.json") as f: config = json.load(f) class_names = config["class_names"] num_classes = config["num_classes"] img_size = config.get("img_size", 224) # Detect device if torch.cuda.is_available(): device = torch.device("cuda") elif hasattr(torch.backends, "mps") and torch.backends.mps.is_available(): device = torch.device("mps") else: device = torch.device("cpu") # Build model model = models.resnet50(weights=None) num_features = model.fc.in_features model.fc = nn.Sequential( nn.Dropout(0.4), nn.Linear(num_features, 512), nn.ReLU(), nn.Dropout(0.2), nn.Linear(512, num_classes) ) model.load_state_dict( torch.load(f"{model_dir}/best_model.pth", map_location=device, weights_only=False) ) model = model.to(device) model.eval() # Transform for inference transform = transforms.Compose([ transforms.Resize((img_size, img_size)), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), ]) return model, device, transform, class_names, config def predict(image: Image.Image, model, device, transform, class_names, top_k: int = 5): """ Predict batik motif from a PIL Image. Returns: predictions: list of (class_name, confidence) sorted by confidence desc info: cultural info dict for the top prediction """ if image is None: return [], {} # Ensure RGB image = image.convert("RGB") # Transform tensor = transform(image).unsqueeze(0).to(device) with torch.no_grad(): outputs = model(tensor) probs = F.softmax(outputs, dim=1)[0].cpu().numpy() # Top-K predictions top_indices = np.argsort(probs)[::-1][:top_k] predictions = [(class_names[i], float(probs[i])) for i in top_indices] # Cultural info for top prediction top_class = predictions[0][0] info = get_batik_info(top_class) return predictions, info if __name__ == "__main__": import sys import argparse parser = argparse.ArgumentParser(description="BatikAI single-image inference") parser.add_argument("--image", type=str, required=True, help="Path to input image") args = parser.parse_args() print("Loading model...") model, device, transform, class_names, config = load_model() print(f"Model loaded. Classes: {class_names}\n") img = Image.open(args.image) preds, info = predict(img, model, device, transform, class_names) print("Top-5 Predictions:") for i, (cls, conf) in enumerate(preds): bar = "█" * int(conf * 30) print(f" {i+1}. {cls:<20} {conf*100:5.1f}% {bar}") print(f"\nCultural Info for '{preds[0][0]}':") print(f" Origin : {info['origin']}") print(f" Meaning : {info['meaning']}") print(f" Usage : {info['usage']}") print(f" Fun Fact: {info['fun_fact']}")