import argparse import timm import torch from PIL import Image from torchvision import transforms CLASS_NAMES = ["out_of_play", "in_play"] TRANSFORM = transforms.Compose( [ transforms.Resize(224), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]), ] ) def main(): parser = argparse.ArgumentParser() parser.add_argument("image") parser.add_argument("--checkpoint", default="model.pth") args = parser.parse_args() checkpoint = torch.load(args.checkpoint, map_location="cpu", weights_only=True) model = timm.create_model( "convnextv2_base.fcmae_ft_in22k_in1k", pretrained=False, num_classes=2 ) model.load_state_dict(checkpoint["model_state_dict"]) model.eval() image = TRANSFORM(Image.open(args.image).convert("RGB")).unsqueeze(0) with torch.inference_mode(): probabilities = model(image).softmax(dim=-1)[0] index = int(probabilities.argmax()) print({"label": CLASS_NAMES[index], "score": float(probabilities[index])}) if __name__ == "__main__": main()