import torch from model import ANN model = ANN() model.load_state_dict(torch.load("weights.pth", map_location="cpu")) model.eval() def predict(image_tensor): with torch.no_grad(): return model(image_tensor).argmax(1).item()