CharlesCNorton
Image-level person classification on EUPE-ViT-B features with no free parameters
e8b8483 | """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) | |
| def dims(self): | |
| return self._pos.tolist(), self._neg.tolist() | |
| 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 | |
| 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}') | |