karma689 commited on
Commit
c7f751c
·
verified ·
1 Parent(s): ebdd9bf

Update inference_classifier.py: letterbox inference defaults and docs

Browse files
Files changed (1) hide show
  1. inference_classifier.py +158 -22
inference_classifier.py CHANGED
@@ -1,30 +1,139 @@
1
  #!/usr/bin/env python3
2
- """Run DINOv3 script classifier on image paths (classes from checkpoint ``idx_to_label``)."""
3
 
4
  from __future__ import annotations
5
 
6
  import argparse
7
- import sys
8
  from pathlib import Path
9
 
10
  import torch
 
11
  from PIL import Image
12
- from transformers import AutoImageProcessor
13
 
14
- ROOT = Path(__file__).resolve().parents[1]
15
- if str(ROOT) not in sys.path:
16
- sys.path.insert(0, str(ROOT))
 
17
 
18
- from src.labels import DINOV3_MODEL_ID, MULTICLASS_6_LABELS
19
- from src.models import DINOv3Classifier
20
- from src.transforms import apply_preprocess, normalize_preprocess_mode
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
21
 
22
 
23
  @torch.no_grad()
24
- def predict(model, processor, image_path: Path, device, *, preprocess: str | None, size: int):
 
 
 
 
 
 
 
 
25
  img = Image.open(image_path).convert("RGB")
26
  img = apply_preprocess(img, preprocess, size=size)
27
- pv = processor(images=img, return_tensors="pt")["pixel_values"].to(device)
 
 
 
 
28
  logits = model(pv)
29
  probs = torch.softmax(logits, dim=1).squeeze(0).cpu()
30
  pred = int(probs.argmax())
@@ -32,33 +141,60 @@ def predict(model, processor, image_path: Path, device, *, preprocess: str | Non
32
 
33
 
34
  def main() -> None:
35
- ap = argparse.ArgumentParser(description="DINOv3 script-family inference")
36
- ap.add_argument("--checkpoint", type=Path, required=True)
 
 
 
 
 
 
 
37
  ap.add_argument("--image", type=Path, nargs="+", required=True)
38
- ap.add_argument("--preprocess", default="center_crop", help="none | center_crop")
39
- ap.add_argument("--preprocess-size", type=int, default=224)
 
 
 
 
 
 
 
 
 
40
  ap.add_argument("--model-id", default=DINOV3_MODEL_ID)
41
  args = ap.parse_args()
42
 
 
 
 
 
 
 
43
  ckpt = torch.load(args.checkpoint, map_location="cpu", weights_only=False)
44
- idx_to_label = {int(k): v for k, v in ckpt.get("idx_to_label", {}).items()}
45
- if not idx_to_label:
46
- idx_to_label = {i: lab for i, lab in enumerate(MULTICLASS_6_LABELS)}
47
 
48
  device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
49
- model = DINOv3Classifier(args.model_id, num_classes=len(idx_to_label)).to(device)
50
  model.load_state_dict(ckpt["model_state_dict"])
51
  model.eval()
52
  processor = AutoImageProcessor.from_pretrained(args.model_id)
53
- prep = normalize_preprocess_mode(args.preprocess)
54
 
55
  for path in args.image:
56
  pred, probs = predict(
57
- model, processor, path, device, preprocess=prep, size=args.preprocess_size
 
 
 
 
 
58
  )
59
  name = idx_to_label[pred]
60
  conf = probs[pred]
61
- print(f"{path.name}: {name} ({conf:.3f})")
 
 
62
 
63
 
64
  if __name__ == "__main__":
 
1
  #!/usr/bin/env python3
2
+ """Standalone binary page-orientation inference (copied to Hub as ``inference_classifier.py``)."""
3
 
4
  from __future__ import annotations
5
 
6
  import argparse
7
+ import json
8
  from pathlib import Path
9
 
10
  import torch
11
+ import torch.nn as nn
12
  from PIL import Image
13
+ from transformers import AutoImageProcessor, AutoModel
14
 
15
+ DINOV3_MODEL_ID = "facebook/dinov3-vits16-pretrain-lvd1689m"
16
+ DEFAULT_LABELS = ("non_flipped", "flipped")
17
+ DEFAULT_PREPROCESS = "resize_letterbox"
18
+ DEFAULT_PREPROCESS_SIZE = 448
19
 
20
+
21
+ class DINOv3Classifier(nn.Module):
22
+ def __init__(self, model_id: str, num_classes: int, dropout: float = 0.1):
23
+ super().__init__()
24
+ self.backbone = AutoModel.from_pretrained(model_id)
25
+ hidden = self.backbone.config.hidden_size
26
+ self.head = nn.Sequential(
27
+ nn.LayerNorm(hidden),
28
+ nn.Dropout(dropout),
29
+ nn.Linear(hidden, 128),
30
+ nn.GELU(),
31
+ nn.Dropout(dropout),
32
+ nn.Linear(128, num_classes),
33
+ )
34
+
35
+ def forward(self, pixel_values):
36
+ out = self.backbone(pixel_values=pixel_values)
37
+ cls = out.last_hidden_state[:, 0, :]
38
+ return self.head(cls)
39
+
40
+
41
+ def _resize_short_edge(img: Image.Image, target: int) -> Image.Image:
42
+ w, h = img.size
43
+ if h <= w:
44
+ new_h = target
45
+ new_w = max(1, int(w * target / h))
46
+ else:
47
+ new_w = target
48
+ new_h = max(1, int(h * target / w))
49
+ return img.resize((new_w, new_h), Image.BICUBIC)
50
+
51
+
52
+ def _center_crop(img: Image.Image, size: int = 224) -> Image.Image:
53
+ img = _resize_short_edge(img, size)
54
+ w, h = img.size
55
+ left = max(0, (w - size) // 2)
56
+ top = max(0, (h - size) // 2)
57
+ crop = img.crop((left, top, left + size, top + size))
58
+ if crop.size != (size, size):
59
+ padded = Image.new("RGB", (size, size), (255, 255, 255))
60
+ padded.paste(crop, (0, 0))
61
+ return padded
62
+ return crop
63
+
64
+
65
+ def _letterbox_resize(img: Image.Image, size: int, fill: int = 255) -> Image.Image:
66
+ w, h = img.size
67
+ scale = size / max(w, h)
68
+ nw, nh = round(w * scale), round(h * scale)
69
+ img = img.resize((nw, nh), Image.BILINEAR)
70
+ pad_l = (size - nw) // 2
71
+ pad_t = (size - nh) // 2
72
+ canvas = Image.new("RGB", (size, size), (fill, fill, fill))
73
+ canvas.paste(img, (pad_l, pad_t))
74
+ return canvas
75
+
76
+
77
+ def apply_preprocess(img: Image.Image, mode: str | None, *, size: int = 448) -> Image.Image:
78
+ if not mode or mode == "none":
79
+ return img
80
+ if mode in ("center_crop", "center_crop_whole_page"):
81
+ return _center_crop(img, size)
82
+ if mode == "resize_letterbox":
83
+ return _letterbox_resize(img, size)
84
+ raise ValueError(f"Unknown preprocess mode: {mode!r}")
85
+
86
+
87
+ def processor_skip_resize(mode: str | None) -> bool:
88
+ return mode in ("center_crop", "center_crop_whole_page", "resize_letterbox")
89
+
90
+
91
+ def label_order(ckpt: dict) -> list[str]:
92
+ idx = ckpt.get("idx_to_label") or {}
93
+ if idx:
94
+ return [str(idx[k]) for k in sorted(idx.keys(), key=lambda x: int(x))]
95
+ raw = ckpt.get("label_to_idx") or {}
96
+ if raw:
97
+ return sorted(raw.keys(), key=lambda k: raw[k])
98
+ return list(DEFAULT_LABELS)
99
+
100
+
101
+ def load_model_card_defaults(checkpoint: Path) -> tuple[str, int]:
102
+ card_path = checkpoint.parent / "model_card.json"
103
+ if not card_path.is_file():
104
+ return DEFAULT_PREPROCESS, DEFAULT_PREPROCESS_SIZE
105
+ card = json.loads(card_path.read_text(encoding="utf-8"))
106
+ prep = card.get("preprocess") or {}
107
+ mode = prep.get("test") or prep.get("val") or prep.get("train") or DEFAULT_PREPROCESS
108
+ size = int(prep.get("size") or DEFAULT_PREPROCESS_SIZE)
109
+ return mode, size
110
+
111
+
112
+ def describe_label(name: str) -> str:
113
+ if name == "non_flipped":
114
+ return "upright"
115
+ if name == "flipped":
116
+ return "upside-down (180°)"
117
+ return name
118
 
119
 
120
  @torch.no_grad()
121
+ def predict(
122
+ model,
123
+ processor,
124
+ image_path: Path,
125
+ device,
126
+ *,
127
+ preprocess: str | None,
128
+ size: int,
129
+ ):
130
  img = Image.open(image_path).convert("RGB")
131
  img = apply_preprocess(img, preprocess, size=size)
132
+ pv = processor(
133
+ images=img,
134
+ do_resize=not processor_skip_resize(preprocess),
135
+ return_tensors="pt",
136
+ )["pixel_values"].to(device)
137
  logits = model(pv)
138
  probs = torch.softmax(logits, dim=1).squeeze(0).cpu()
139
  pred = int(probs.argmax())
 
141
 
142
 
143
  def main() -> None:
144
+ ap = argparse.ArgumentParser(
145
+ description="Binary page orientation: non_flipped (upright) vs flipped (180°)."
146
+ )
147
+ ap.add_argument(
148
+ "--checkpoint",
149
+ type=Path,
150
+ default=Path("final_model.pt"),
151
+ help="Weights file (default: final_model.pt in cwd)",
152
+ )
153
  ap.add_argument("--image", type=Path, nargs="+", required=True)
154
+ ap.add_argument(
155
+ "--preprocess",
156
+ default=None,
157
+ help="none | center_crop | resize_letterbox (default: from model_card.json or resize_letterbox)",
158
+ )
159
+ ap.add_argument(
160
+ "--preprocess-size",
161
+ type=int,
162
+ default=None,
163
+ help="PIL preprocess size before DINO processor (default: from model_card.json or 448)",
164
+ )
165
  ap.add_argument("--model-id", default=DINOV3_MODEL_ID)
166
  args = ap.parse_args()
167
 
168
+ card_default_mode, card_default_size = load_model_card_defaults(args.checkpoint)
169
+ preprocess = args.preprocess or card_default_mode
170
+ size = args.preprocess_size if args.preprocess_size is not None else card_default_size
171
+ if preprocess in ("none", ""):
172
+ preprocess = None
173
+
174
  ckpt = torch.load(args.checkpoint, map_location="cpu", weights_only=False)
175
+ classes = label_order(ckpt)
176
+ idx_to_label = {i: lab for i, lab in enumerate(classes)}
 
177
 
178
  device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
179
+ model = DINOv3Classifier(args.model_id, num_classes=len(classes)).to(device)
180
  model.load_state_dict(ckpt["model_state_dict"])
181
  model.eval()
182
  processor = AutoImageProcessor.from_pretrained(args.model_id)
 
183
 
184
  for path in args.image:
185
  pred, probs = predict(
186
+ model,
187
+ processor,
188
+ path,
189
+ device,
190
+ preprocess=preprocess,
191
+ size=size,
192
  )
193
  name = idx_to_label[pred]
194
  conf = probs[pred]
195
+ print(f"{path.name}: {name} ({describe_label(name)}, {conf:.3f})")
196
+ for i, lab in enumerate(classes):
197
+ print(f" {lab:14s} {probs[i]:.3f}")
198
 
199
 
200
  if __name__ == "__main__":