File size: 2,234 Bytes
e8b8483
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Score images for person presence.

    from infer import PersonDetector
    det = PersonDetector.load('d6')
    present = det.predict('image.jpg')

`predict` returns a bool. `margin` returns the signed difference between the two
sums, which is positive exactly when the answer is yes; it carries no threshold.
"""
import argparse
import sys
from pathlib import Path

import torch

from common import BACKBONE, RES, backbone_pooled, load_image, read_artifact, score
from common.models import load_backbone

HERE = Path(__file__).resolve().parent


class PersonDetector:
    def __init__(self, forward_fn, pos_dims, neg_dims, dev):
        self._forward = forward_fn
        self._dev = dev
        self._pos = torch.tensor(pos_dims, dtype=torch.long, device=dev)
        self._neg = torch.tensor(neg_dims, dtype=torch.long, device=dev)

    @property
    def dims(self):
        return self._pos.tolist(), self._neg.tolist()

    @torch.inference_mode()
    def margin(self, image) -> float:
        pooled = self._forward(load_image(image, RES, self._dev))
        return float(score(pooled, self._pos, self._neg))

    def predict(self, image) -> bool:
        return self.margin(image) > 0.0

    @classmethod
    def load(cls, rule: str = None, backbone_repo: str = BACKBONE, root=None):
        from common import device
        root = Path(root) if root else HERE
        doc = read_artifact(root / 'rules.json')
        rules = doc['rules']
        rule = rule or 'd6'
        if rule not in rules:
            raise ValueError(f'unknown rule {rule!r}; expected one of {sorted(rules)}')
        r = rules[rule]
        dev = device()
        backbone = load_backbone(backbone_repo).to(dev).eval()
        return cls(lambda x: backbone_pooled(backbone, x)[0],
                   r['pos_dims'], r['neg_dims'], dev)


if __name__ == '__main__':
    ap = argparse.ArgumentParser(description=__doc__)
    ap.add_argument('rule')
    ap.add_argument('images', nargs='+')
    args = ap.parse_args()
    det = PersonDetector.load(args.rule)
    pos, neg = det.dims
    print(f'rule {args.rule}: sum{pos} > sum{neg}')
    for path in args.images:
        m = det.margin(path)
        print(f'{path}  margin={m:+.3f}  person={m > 0}')