File size: 4,391 Bytes
a6e745e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
"""
Inference script for CAM-AI4G-v1 (Colombian Amazon, AI for Good Lab)

Model: AI4GAmazonRainforest
Input: 224x224 RGB, ImageNet-normalised (applied by PytorchWildlife)
Framework: PyTorch (PytorchWildlife, ResNet)
Classes: 36 Amazon genera
Developer: AI for Good Lab, Microsoft

Ported from AddaxAI's legacy classify_detections.py (pywildlife), with
the double-transform bug fixed. See below.

Author: Peter van Lunteren
"""

from __future__ import annotations

from pathlib import Path

import numpy as np
import supervision as sv
import torch
from PIL import Image, ImageFile
from PytorchWildlife.models import classification as pw_classification

# Don't freak out over truncated images
ImageFile.LOAD_TRUNCATED_IMAGES = True

# The 36 classes are genera, and PytorchWildlife owns the list. It is
# read off the loaded model rather than restated here, so the two can
# never drift apart.


class ModelInference:
    """Colombian Amazon classifier (PytorchWildlife AI4GAmazonRainforest)."""

    def __init__(self, model_dir: Path, model_path: Path) -> None:
        self.model_dir = Path(model_dir)
        self.model_path = Path(model_path)
        self.model = None

    # ------------------------------------------------------------------
    # Required interface
    # ------------------------------------------------------------------

    def check_gpu(self) -> bool:
        try:
            if torch.backends.mps.is_built() and torch.backends.mps.is_available():
                return True
        except Exception:
            pass
        return torch.cuda.is_available()

    def load_model(self) -> None:
        self.model = pw_classification.AI4GAmazonRainforest(
            weights=str(self.model_path)
        )

    def get_crop(
        self, image: Image.Image, bbox: tuple[float, float, float, float]
    ) -> Image.Image:
        """
        Plain crop, no squaring and no padding.

        Kept on supervision's crop_image so it stays identical to the
        legacy adapter, down to how it rounds and how numpy slicing
        handles a box that runs past an edge.
        """
        img_width, img_height = image.size
        left = int(round(bbox[0] * img_width))
        top = int(round(bbox[1] * img_height))
        right = int(round(bbox[2] * img_width)) + left
        bottom = int(round(bbox[3] * img_height)) + top
        return Image.fromarray(
            sv.crop_image(
                np.array(image.convert("RGB")), xyxy=[left, top, right, bottom]
            )
        )

    def get_classification(self, crop: Image.Image) -> list[list]:
        """
        Per-crop inference. Returns [[name, prob], ...] for all classes.

        PytorchWildlife's single_image_classification does the whole
        preprocessing itself: Image.fromarray -> resize 224 -> ToTensor
        -> ImageNet Normalize -> forward. So it takes the plain RGB crop.

        The legacy adapter ran Classification_Inference_Transform first
        and then cast the result back to uint8 before handing it over:

            preprocessed = trans_clf(PIL_crop)                  # float, [-2.1, 2.5]
            preprocessed = preprocessed.permute(1,2,0).numpy().astype(np.uint8)
            model.single_image_classification(preprocessed)     # transforms it AGAIN

        That cast is destructive. The normalised floats sit in about
        [-2.1, 2.5], so uint8 truncates them to {0, 1, 2} and wraps the
        negatives: the 224x224x3 array reaching the model held three
        distinct values, i.e. a near-black image. Measured on a real
        crop, legacy answered `Mazama` at 0.362 with `Unknown` at 0.357
        behind it (a model with no signal, guessing), where the same
        weights on the same crop answer `Bos` at 0.917.

        So this is not a faithful-port question: the legacy path never
        classified the image at all. Passing the crop through once, the
        way PytorchWildlife intends, is the fix.
        """
        assert self.model is not None
        return self.model.single_image_classification(
            np.array(crop.convert("RGB"))
        )["all_confidences"]

    def get_class_names(self) -> dict[str, str]:
        """1-indexed mapping {id: class_name} for the output JSON."""
        assert self.model is not None
        names = self.model.CLASS_NAMES
        return {str(i + 1): names[i] for i in range(len(names))}