predict.py
1,094 bytes
| 1 | import argparse |
|---|---|
| 2 | |
| 3 | import timm |
| 4 | import torch |
| 5 | from PIL import Image |
| 6 | |
| 7 | from .data import build_transforms |
| 8 | from .model import load_checkpoint |
| 9 | |
| 10 | |
| 11 | def main(): |
| 12 | ap = argparse.ArgumentParser(description="Ennustab pildi riigi") |
| 13 | ap.add_argument("--checkpoint", default="runs/clip_vit_b16/best.pt") |
| 14 | ap.add_argument("--image", required=True) |
| 15 | ap.add_argument("--topk", type=int, default=5) |
| 16 | args = ap.parse_args() |
| 17 | |
| 18 | device = torch.device("cuda" if torch.cuda.is_available() else "cpu") |
| 19 | model, classes = load_checkpoint(args.checkpoint, device) |
| 20 | data_cfg = timm.data.resolve_model_data_config(model.backbone) |
| 21 | tfm = build_transforms(data_cfg["mean"], data_cfg["std"], |
| 22 | data_cfg["input_size"][-1], train=False) |
| 23 | |
| 24 | img = Image.open(args.image).convert("RGB") |
| 25 | with torch.no_grad(): |
| 26 | probs = model(tfm(img).unsqueeze(0).to(device)).softmax(1)[0] |
| 27 | top = probs.topk(min(args.topk, len(classes))) |
| 28 | for p, i in zip(top.values.tolist(), top.indices.tolist()): |
| 29 | print(f"{classes[i]:<30s} {p:6.1%}") |
| 30 | |
| 31 | |
| 32 | if __name__ == "__main__": |
| 33 | main() |
| 34 | |