Image Classification
LiteRT
LiteRT
ONNX
English
vision
botany
western-australia
dinov3
mixture-of-experts
adaround
fp8
int8
android
biodiversity
flora
Instructions to use thenukegun10x/PLantDetect-WA with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- LiteRT
How to use thenukegun10x/PLantDetect-WA with LiteRT:
# No code snippets available yet for this library. # To use this model, check the repository files and the library's documentation. # Want to help? PRs adding snippets are welcome at: # https://github.com/huggingface/huggingface.js
- Notebooks
- Google Colab
- Kaggle
File size: 47,260 Bytes
d499940 54a4bbb d499940 54a4bbb d499940 54a4bbb d499940 cdc1714 d499940 cdc1714 d499940 54a4bbb cdc1714 d499940 253e323 d499940 cdc1714 d499940 54a4bbb b799883 54a4bbb b799883 54a4bbb b799883 54a4bbb b799883 54a4bbb d499940 cdc1714 d499940 cdc1714 d499940 cdc1714 d499940 54a4bbb d499940 54a4bbb d499940 54a4bbb d499940 54a4bbb d499940 54a4bbb d499940 54a4bbb d499940 54a4bbb b799883 54a4bbb b799883 54a4bbb d499940 54a4bbb d499940 54a4bbb d499940 54a4bbb d499940 253e323 d499940 253e323 54a4bbb d499940 54a4bbb d499940 54a4bbb d499940 253e323 d499940 253e323 d499940 54a4bbb d499940 54a4bbb d499940 54a4bbb d499940 54a4bbb d499940 54a4bbb d499940 54a4bbb d499940 cdc1714 d499940 cdc1714 d499940 cdc1714 54a4bbb cdc1714 54a4bbb cdc1714 d499940 cdc1714 54a4bbb d499940 cdc1714 54a4bbb d499940 54a4bbb d499940 54a4bbb d499940 54a4bbb d499940 54a4bbb d499940 54a4bbb d499940 | 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 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 395 396 397 398 399 400 401 402 403 404 405 406 407 408 409 410 411 412 413 414 415 416 417 418 419 420 421 422 423 424 425 426 427 428 429 430 431 432 433 434 435 436 437 438 439 440 441 442 443 444 445 446 447 448 449 450 451 452 453 454 455 456 457 458 459 460 461 462 463 464 465 466 467 468 469 470 471 472 473 474 475 476 477 478 479 480 481 482 483 484 485 486 487 488 489 490 491 492 493 494 495 496 497 498 499 500 501 502 503 504 505 506 507 508 509 510 511 512 513 514 515 516 517 518 519 520 521 522 523 524 525 526 527 528 529 530 531 532 533 534 535 536 537 538 539 540 541 542 543 544 545 546 547 548 549 550 551 552 553 554 555 556 557 558 559 560 561 562 563 564 565 566 567 568 569 570 571 572 573 574 575 576 577 578 579 580 581 582 583 584 585 586 587 588 589 590 591 592 593 594 595 596 597 598 599 600 601 602 603 604 605 606 607 608 609 610 611 612 613 614 615 616 617 618 619 620 621 622 623 624 625 626 627 628 629 630 631 632 633 634 635 636 637 638 639 640 641 642 643 644 645 646 647 648 649 650 651 652 653 654 655 656 657 658 659 660 661 662 663 664 665 666 667 668 669 670 671 672 673 674 675 676 677 678 679 680 681 682 683 684 685 686 687 688 689 690 691 692 693 694 695 696 697 698 699 700 701 702 703 704 705 706 707 708 709 710 711 712 713 714 715 716 717 718 719 720 721 722 723 724 725 726 727 728 729 730 731 732 733 734 735 736 737 738 739 740 741 742 743 744 745 746 747 748 749 750 751 752 753 754 755 756 757 758 759 760 761 762 763 764 765 766 767 768 769 770 771 772 773 774 775 776 777 778 779 780 781 782 783 784 785 786 787 788 789 790 791 792 793 794 795 796 797 798 799 800 801 802 803 804 805 806 807 808 809 810 811 812 813 814 815 816 817 818 819 820 821 822 823 824 825 826 827 828 829 830 831 832 833 834 835 836 837 838 839 840 841 842 843 844 845 846 847 848 849 850 851 852 853 854 855 856 857 858 859 860 861 862 863 864 865 866 867 868 869 870 871 872 873 874 875 876 | #!/usr/bin/env python3
"""
WA Plant Identifier — Seamless CLI for DSLR / Phone images
Built with DINOv3 — see LICENSE.md
Usage:
python plant_cli.py identify photo.jpg --topk 5 --text "yellow puff" --lat -31.95 --lon 115.86
python plant_cli.py identify img1.jpg img2.jpg img3.jpg --text "blue flower" # 4-view mean-logits 99.2%
python plant_cli.py batch --dir DCIM --pattern "*.jpg" --recursive --text "yellow" --out results.csv
python plant_cli.py info
python plant_cli.py build-centroids --out species_centroids.npz [--dummy | --manifest data/wa_plants_200k/manifest_for_train.csv]
python plant_cli.py test-unknown known.jpg unknown.jpg --centroids species_centroids.npz
python plant_cli.py bench --split val
Handles: DSLR JPEG (45MP draft768), phone HEIC/JPG, EXIF orientation, 1-4 views, optional text + geo.
Models: plant_phase3b_otherblue.pt (89.21% single) + plant_phase5_fusion.pt (92% with text) auto-selected.
Unknown: centroid cosine (512-D) + MSP + energy ensemble, cached 999x512 ~2MB.
"""
import argparse, json, sys, time, os
from pathlib import Path
import torch
import torch.nn.functional as F
from PIL import Image, ExifTags
import pandas as pd
import numpy as np
BASE = Path(__file__).resolve().parent
sys.path.insert(0, str(BASE))
from src.models.plant_vit import PlantViT
import torchvision.transforms as T
from src.data.plant import IMAGENET_MEAN, IMAGENET_STD
MANIFEST = BASE/"data/wa_plants_200k/manifest_for_train.csv"
SPECIES_JSON = BASE/"species_labels.json"
CKPT_VISUAL = BASE/"data/plant_phase3b_otherblue.pt"
# Public HF release fallbacks (root *.safetensors) - used when private data/ ckpts missing
PUBLIC_CKPTS = [
BASE/"PlantDetect-FP8-AdaRound.safetensors",
BASE/"PlantDetect-BF16.safetensors",
BASE/"PlantDetect-Dense-FP8-AdaRound.safetensors",
BASE/"PlantDetect-Dense-BF16.safetensors",
BASE/"PlantDetect-4View-FP8-AdaRound.safetensors",
BASE/"PlantDetect-4View-BF16.safetensors",
BASE/"PlantDetect-Dense-4View-FP8-AdaRound.safetensors",
BASE/"PlantDetect-Dense-4View-BF16.safetensors",
]
CKPT_FUSION = BASE/"data/plant_phase5_fusion.pt"
E5_DIR = BASE/"data/plant_phase5_e5"
CKPT_OLD = BASE/"data/plant_phase2_200k.pt"
# Unknown detection: centroid cache (few MB: 999*512*4=2MB fp32, ~1MB fp16)
CENTROIDS_PATH = BASE/"species_centroids.npz"
DEFAULT_COS_THR = 0.50 # for placeholder green-centered centroids; recalibrate after true build
DEFAULT_CONF_THR = 0.60
DEFAULT_MARGIN_THR = 0.15
_CENTROIDS_CACHE = None
_CENTROIDS_META = None
def _resolve_ckpt(preferred=CKPT_VISUAL):
if preferred and preferred.exists():
return preferred
for cand in PUBLIC_CKPTS:
if cand.exists():
return cand
return preferred
# fix PIL large image
Image.MAX_IMAGE_PIXELS = 300_000_000
def log(msg): print(msg, flush=True)
def get_transform(img_size=336):
return T.Compose([T.Resize(int(img_size*1.14)), T.CenterCrop(img_size), T.ToTensor(), T.Normalize(IMAGENET_MEAN, IMAGENET_STD)])
def get_vit_heatmap(visual_model, image_path, device="cuda", img_size=336):
"""Feature heatmap where the ViT looks - per-patch logits for top class, 21x21 -> 336."""
try:
from pathlib import Path as _P
im = open_image_pil(image_path)
tf = get_transform(img_size)
x = tf(im).unsqueeze(0).to(device)
core = visual_model.core
with torch.no_grad():
with torch.autocast("cuda", dtype=torch.bfloat16, enabled=device=="cuda"):
tokens = core.stem.forward_features(x) # [1, T, D]
if core.input_proj is not None:
tokens = core.input_proj(tokens.float())
else:
tokens = tokens.float()
for b in core.blocks:
tokens, _, _, _ = b(tokens)
tokens = core.exit_norm(tokens) # [1, T, D]
n_prefix = int(getattr(core, "n_prefix", 1))
patch_tokens = tokens[:, n_prefix:, :] # [1, P, D]
P = patch_tokens.shape[1]
h = w = int(round(P ** 0.5))
pt = patch_tokens.squeeze(0) # [P, D]
logits_patch = core.head(pt) # [P, 999]
top_idx = int(logits_patch.mean(0).argmax())
heat = logits_patch[:, top_idx]
heat = heat - heat.min()
heat = heat / (heat.max() - heat.min() + 1e-8)
heat_np = heat.float().cpu().numpy().reshape(h, w)
heat_img = Image.fromarray((heat_np * 255).astype(np.uint8)).resize((img_size, img_size), Image.BILINEAR)
heat_arr = np.array(heat_img).astype(np.float32) / 255.0
return heat_arr, top_idx, heat_np
except Exception as e:
log(f"heatmap failed {e}")
return None, None, None
def load_species_map(manifest=MANIFEST):
# Public release: species_labels.json (999 WA species) - preferred
if SPECIES_JSON.exists():
data=json.loads(SPECIES_JSON.read_text(encoding="utf-8"))
if "idx_to_species" in data:
i2s={int(k):v for k,v in data["idx_to_species"].items()}
spp=[i2s[i] for i in sorted(i2s)]
s2i={s:i for i,s in i2s.items()}
return s2i, i2s, spp
spp=data.get("species", [])
s2i={s:i for i,s in enumerate(spp)}
i2s={i:s for s,i in s2i.items()}
return s2i, i2s, spp
# Fallback: private manifest CSV (training)
df=pd.read_csv(manifest)
spp=sorted(df[df["status"]=="downloaded"]["species"].unique())
s2i={s:i for i,s in enumerate(spp)}
i2s={i:s for s,i in s2i.items()}
return s2i, i2s, spp
# --- Centroid unknown detection (caching ~2MB) ---
def load_centroids(path=CENTROIDS_PATH):
global _CENTROIDS_CACHE, _CENTROIDS_META
if _CENTROIDS_CACHE is not None:
return _CENTROIDS_CACHE, _CENTROIDS_META
if not path.exists():
return None, None
try:
data=np.load(str(path), allow_pickle=True)
if "centroids" in data:
cents=data["centroids"] # [C,512]
meta={}
for k in ["thr_cos","thr_conf","species","n_classes"]:
if k in data:
meta[k]=data[k]
# handle species as array
if "species" in meta and isinstance(meta["species"], np.ndarray):
meta["species"]=meta["species"].tolist()
_CENTROIDS_CACHE=cents
_CENTROIDS_META=meta
log(f"loaded centroids {cents.shape} from {path.name} thr_cos={meta.get('thr_cos', DEFAULT_COS_THR)}")
return cents, meta
else:
# single array file
arr=data[data.files[0]]
_CENTROIDS_CACHE=arr
_CENTROIDS_META={}
return arr, {}
except Exception as e:
log(f"WARN centroids load failed {e}")
return None, None
def centroid_similarity(emb, centroids):
"""cosine similarity between L2 512-D emb [1,512] and centroids [C,512]"""
if centroids is None:
return 1.0, -1, None
if isinstance(centroids, np.ndarray):
centroids_t=torch.from_numpy(centroids).to(emb.device).float()
else:
centroids_t=centroids
# L2 normalize both
centroids_t=F.normalize(centroids_t.float(), dim=1)
emb_n=F.normalize(emb.float(), dim=1) # [1,512]
sims=(emb_n @ centroids_t.T).squeeze(0) # [C]
max_sim, idx = sims.max(0)
return float(max_sim.item()), int(idx.item()), sims
def detect_unknown(emb, logits, centroids, thr_cos=DEFAULT_COS_THR, thr_conf=DEFAULT_CONF_THR, thr_margin=DEFAULT_MARGIN_THR):
probs=F.softmax(logits.float(), dim=1)
conf=float(probs.max().item())
sorted_probs=probs[0].sort(descending=True).values
margin=float(sorted_probs[0]-sorted_probs[1]) if len(sorted_probs)>1 else 1.0
energy=float(torch.logsumexp(logits.float(), dim=1).item())
# centroid distance
if centroids is not None:
sim_max, sim_idx, sims = centroid_similarity(emb, centroids)
dist = 1 - sim_max
else:
sim_max, sim_idx, sims = 1.0, -1, None
dist = 0.0
# ensemble rule - dummy centroids are placeholder (sim -0.02 for real eucalyptus) so require BOTH low sim and low conf
is_unknown=False
reasons=[]
# 1) far from centroids AND low conf -> unknown. High conf (>0.75) overrides dummy cache false positive (Eucalyptus 98% should stay KNOWN)
if centroids is not None and sim_max < thr_cos and conf < 0.75:
is_unknown=True
reasons.append(f"far from centroids sim {sim_max:.2f} < {thr_cos:.2f} (dist {dist:.2f}) + conf {conf:.2%} < 75%")
# 2) low sim + low conf joint (slightly higher sim but still low conf)
elif centroids is not None and sim_max < thr_cos+0.12 and conf < thr_conf:
is_unknown=True
reasons.append(f"low sim {sim_max:.2f} + low conf {conf:.2%} < {thr_conf:.2%}")
# 3) fallback if no centroids: low conf + small margin
elif centroids is None and conf < thr_conf and margin < thr_margin:
is_unknown=True
reasons.append(f"low conf {conf:.2%} margin {margin:.3f}")
elif centroids is None and conf < 0.35:
is_unknown=True
reasons.append(f"very low conf {conf:.2%}")
# else: high conf keeps KNOWN - prevents dummy -0.02 false UNKNOWN on Eucalyptus 98%
return is_unknown, {"sim_max": sim_max, "sim_idx": sim_idx, "conf": conf, "margin": margin, "energy": energy, "dist": dist, "reasons": "; ".join(reasons)}
def open_image_pil(path):
p=Path(path)
if not p.exists():
raise FileNotFoundError(f"not found {p}")
# use draft for large DSLR to save RAM
im=Image.open(p)
# EXIF orientation
try:
exif = im._getexif()
if exif:
orientation = exif.get(274)
if orientation==3: im=im.rotate(180, expand=True)
elif orientation==6: im=im.rotate(270, expand=True)
elif orientation==8: im=im.rotate(90, expand=True)
except: pass
# draft shrink for large DSLR (>20MP) before decode
try:
if im.size[0]*im.size[1] > 20_000_000:
im.draft("RGB", (768,768))
except: pass
return im.convert("RGB")
def load_visual_model(ckpt=CKPT_VISUAL, n_classes=999, device="cuda"):
# auto-resolve to public safetensors if private ckpt missing
ckpt=_resolve_ckpt(ckpt)
s2i,i2s,spp = load_species_map()
# n_classes from species list (public 999) overrides arg if mismatch
if len(spp)!=n_classes:
n_classes=len(spp)
# Dense vs MoE auto-detect from ckpt name
is_dense = "Dense" in ckpt.name if ckpt else False
use_moe = None if is_dense else True
model=PlantViT(stem_name="vit_base_patch16_dinov3", n_classes=n_classes, use_moe=use_moe, num_ffn=16).to(device)
if ckpt.exists():
if ckpt.suffix==".safetensors":
try:
from safetensors.torch import load_file
try:
sd=load_file(str(ckpt), device=device)
except Exception:
sd=load_file(str(ckpt))
except ImportError:
raise RuntimeError("safetensors required: pip install safetensors")
else:
sd=torch.load(ckpt, map_location=device)
if isinstance(sd, dict) and "model" in sd: sd=sd["model"]
model.load_state_dict(sd, strict=False)
log(f"loaded {ckpt.name} {ckpt.stat().st_size/1e6:.1f}MB")
else:
log(f"WARN ckpt not found {ckpt} using random")
model.eval()
return model, s2i, i2s
# Fusion helpers
class FusionMLP(torch.nn.Module):
def __init__(self, vis_dim=512, text_dim=384, n_classes=999, hidden=512, p_drop=0.0):
super().__init__()
self.fc1=torch.nn.Linear(vis_dim+text_dim, hidden)
self.ln=torch.nn.LayerNorm(hidden)
self.drop=torch.nn.Dropout(p_drop)
self.fc2=torch.nn.Linear(hidden, n_classes)
def forward(self, vis, txt):
x=torch.cat([vis, txt], dim=1)
x=self.fc1(x); x=self.ln(x); x=F.gelu(x); x=self.drop(x)
return self.fc2(x)
NOTICE="The notice period started at 9:45 am on Friday, 12 December 2025"
def clean_desc(txt):
if NOTICE in txt:
idx=txt.find("Habit and leaf form")
if idx!=-1: txt=txt[idx:]
if "WAHerb" in txt and "read-only" in txt:
idx=txt.find("Recent taxonomic")
if idx!=-1: txt=txt[idx+len("Recent taxonomic changes are not currently being reflected in Florabase, herbarium collections, or the census. "):]
return txt.strip()[:1800]
def load_e5_and_fusion(ckpt_fusion=CKPT_FUSION, e5_dir=E5_DIR, device="cuda"):
from transformers import AutoTokenizer, AutoModel
e5_name=str(e5_dir) if (e5_dir/"config.json").exists() else "intfloat/multilingual-e5-small"
log(f"loading e5 {e5_name}")
tok=AutoTokenizer.from_pretrained(e5_name, local_files_only=(e5_dir/"config.json").exists())
emodel=AutoModel.from_pretrained(e5_name, local_files_only=(e5_dir/"config.json").exists()).to(device)
emodel.eval()
# load flora cleaned for prototypes (optional)
flora_path=BASE/"data/wa_plants_200k/florabase_200k_multi.json"
cleaned={}
if flora_path.exists():
import json as js
flora=js.loads(flora_path.read_text(encoding="utf-8"))
cleaned={k:clean_desc(v) for k,v in flora.items()}
# fusion
fusion=FusionMLP(p_drop=0.0).to(device)
if ckpt_fusion.exists():
sd=torch.load(ckpt_fusion, map_location=device)
fusion.load_state_dict(sd)
log(f"loaded fusion {ckpt_fusion.name} {ckpt_fusion.stat().st_size/1e6:.2f}MB")
else:
log(f"WARN fusion ckpt not found {ckpt_fusion}")
fusion.eval()
return tok, emodel, fusion, cleaned
def embed_text(text, tok, emodel, device):
if not text or text.strip()=="":
return None
# e5 expects query/passage prefix
q=f"query: {text.strip()}"
enc=tok([q], padding=True, truncation=True, max_length=512, return_tensors="pt")
enc={k:v.to(device) for k,v in enc.items()}
with torch.no_grad():
with torch.autocast("cuda", dtype=torch.bfloat16, enabled=device=="cuda"):
out=emodel(**enc).last_hidden_state
mask=enc["attention_mask"].unsqueeze(-1).float()
pooled=(out*mask).sum(1)/mask.sum(1).clamp(min=1)
pooled=F.normalize(pooled, dim=-1)
return pooled # [1,384]
def predict_images(image_paths, text=None, geo=None, topk=5, device="cuda", use_fusion_auto=True, img_size=336, return_emb=False, fusion_thr=0.98, fusion_margin=0.70, centroids_path=CENTROIDS_PATH, thr_cos=DEFAULT_COS_THR, thr_conf=DEFAULT_CONF_THR, thr_margin=DEFAULT_MARGIN_THR, use_centroids=True):
"""
image_paths: list Path, 1-4 views -> mean logits
text: optional user text
geo: (lat, lon) optional for prior (currently soft boost, not hard)
fusion_thr: if visual conf>thr and margin>fusion_margin skip fusion to preserve 99% visual
centroids: cached 999x512 for unknown detection (few MB)
"""
# decide model
visual_model, s2i, i2s = load_visual_model(device=device)
tf=get_transform(img_size)
# load centroids once
centroids, cent_meta = load_centroids(centroids_path) if use_centroids else (None, None)
if use_centroids and centroids is None:
log(f"centroids not found {centroids_path} -> using confidence/margin only (run build-centroids)")
thr_cos_eff = thr_cos
else:
thr_cos_eff = float(cent_meta.get("thr_cos", thr_cos)) if cent_meta else thr_cos
# load e5/fusion if needed
tok=emodel=fusion=None
text_emb=None
cleaned={}
if text and use_fusion_auto and CKPT_FUSION.exists() and E5_DIR.exists():
tok, emodel, fusion, cleaned = load_e5_and_fusion(device=device)
text_emb = embed_text(text, tok, emodel, device) # [1,384]
elif text and use_fusion_auto:
log("text provided but fusion not found, using visual only + text ignored for ranking (still show)")
# encode images
logits_list=[]
embeds_list=[]
for p in image_paths:
im=open_image_pil(p)
x=tf(im).unsqueeze(0).to(device)
with torch.no_grad():
with torch.autocast("cuda", dtype=torch.bfloat16, enabled=device=="cuda"):
logits, emb, _ = visual_model(x)
logits_list.append(logits.float())
embeds_list.append(emb.float())
# mean logits for multi-view
if len(logits_list)==1:
logits=logits_list[0]
emb=embeds_list[0]
else:
logits=torch.stack(logits_list).mean(0) # [1,999] mean-logits = log P(s|all) prior
emb=torch.stack(embeds_list).mean(0)
log(f"multi-view {len(image_paths)} mean-logits -> 99.2% expected")
# fusion if text - gated to avoid degrading high-conf visual (99.85%->76% case)
if text_emb is not None:
# compute visual conf/margin before fusion
vis_probs=F.softmax(logits.float(), dim=1)
vis_conf=float(vis_probs.max().item())
sorted_probs=vis_probs[0].sort(descending=True).values
vis_margin=float(sorted_probs[0]-sorted_probs[1]) if len(sorted_probs)>1 else 1.0
if vis_conf > fusion_thr and vis_margin > fusion_margin:
log(f"fusion skipped: vis conf {vis_conf:.2%} margin {vis_margin:.3f} > thr {fusion_thr:.2f}/{fusion_margin:.2f} -> keep visual 99%")
else:
vis_t = emb # [1,512]
txt = text_emb # [1,384]
with torch.no_grad():
with torch.autocast("cuda", dtype=torch.bfloat16, enabled=device=="cuda"):
f_logits = fusion(vis_t, txt)
logits = 0.85*f_logits + 0.15*logits
log(f"fusion used: text \"{text[:60]}\" conf {vis_conf:.2%} margin {vis_margin:.3f} -> 92% mode")
# geo prior soft boost (if provided, boost species with known WA region)
# For now, simple: if lat<-30 (SW) boost MW/N species? Placeholder: no boost, just log
if geo:
lat, lon = geo
log(f"geo {lat:.4f},{lon:.4f} -> soft prior (not hard): SW region boost if lat<-30")
# TODO: load bioregion prior from manifest_gold, for now no change
probs=F.softmax(logits, dim=1).cpu().numpy()[0]
logits_np=logits.cpu().numpy()[0]
topk_idx=logits[0].topk(topk).indices.cpu().numpy()
topk_probs=probs[topk_idx]
results=[]
for rank, (idx, p) in enumerate(zip(topk_idx, topk_probs), start=1):
spp=i2s[idx]
results.append({"rank":rank, "species":spp, "prob":float(p), "logit":float(logits_np[idx]), "idx":int(idx)})
# confidence and margin
conf=float(probs.max())
margin=float(sorted(probs)[-1] - sorted(probs)[-2]) if len(probs)>1 else 0
# unknown detection via centroids
is_unknown, unk_info = detect_unknown(emb, logits, centroids, thr_cos=thr_cos_eff, thr_conf=thr_conf, thr_margin=thr_margin)
# return with meta
meta={"topk":results, "confidence":conf, "margin":margin, "n_views":len(image_paths), "text_used": bool(text_emb is not None), "geo":geo, "emb": emb.cpu().numpy()[0] if return_emb else None, "is_unknown": is_unknown, "unknown_info": unk_info, "centroids_used": centroids is not None, "thr_cos": thr_cos_eff, "thr_conf": thr_conf}
return meta, i2s, probs, logits_np, s2i
def cmd_identify(args):
device="cuda" if torch.cuda.is_available() and not args.cpu else "cpu"
image_paths=[Path(p) for p in args.images]
# check exists
for p in image_paths:
if not p.exists():
log(f"ERROR not found {p}")
sys.exit(1)
if len(image_paths)>4:
log(f"WARN {len(image_paths)} images >4, using first 4 + mean-logits")
image_paths=image_paths[:4]
# geo tuple
geo=None
if args.lat is not None and args.lon is not None:
geo=(args.lat, args.lon)
t0=time.time()
centroids_path = Path(args.centroids) if args.centroids else CENTROIDS_PATH
meta,i2s,probs,logits_np,s2i = predict_images(image_paths, text=args.text, geo=geo, topk=args.topk, device=device, use_fusion_auto=not args.no_fusion, img_size=args.size, fusion_thr=args.fusion_thr, fusion_margin=args.fusion_margin, centroids_path=centroids_path, thr_cos=args.thr_cos, thr_conf=args.thr_conf, thr_margin=args.thr_margin, use_centroids=not args.no_centroids)
dt=time.time()-t0
# pretty print
print("\n" + "="*70)
print(f"WA Plant Identifier — Built with DINOv3 | {len(image_paths)} view(s) | {dt*1000:.0f}ms | {device}")
if args.text:
print(f'Text: "{args.text}" {"(fusion 92% mode)" if meta["text_used"] else "(visual only)"}')
if geo:
print(f"Geo: {geo[0]:.4f},{geo[1]:.4f}")
# unknown banner
if meta["is_unknown"]:
print(f"⚠️ UNKNOWN SPECIES WARNING: {meta['unknown_info']['reasons']}")
print(f" sim_max {meta['unknown_info']['sim_max']:.3f} thr {meta['thr_cos']:.2f} conf {meta['unknown_info']['conf']:.2%} margin {meta['unknown_info']['margin']:.3f} energy {meta['unknown_info']['energy']:.1f}")
print(f" -> Not in 999 WA species or non-plant/out-of-distribution. Treat Top-K as nearest known, not confident ID.")
else:
if meta["centroids_used"]:
# dummy centroids give -0.02 for real eucalyptus, so show raw sim + note high conf override
sim_note = " (high conf overrides low sim - dummy centroids, rebuild for true)" if meta['unknown_info']['sim_max'] < meta['thr_cos'] else ""
print(f"Known species: sim {meta['unknown_info']['sim_max']:.3f} thr {meta['thr_cos']:.2f} conf {meta['unknown_info']['conf']:.2%}{sim_note}")
else:
print(f"Known check (no centroids): conf {meta['unknown_info']['conf']:.2%} margin {meta['unknown_info']['margin']:.3f}")
print("-"*70)
for r in meta["topk"]:
marker="*" if r["rank"]==1 else " "
flag=" ?" if meta["is_unknown"] else ""
print(f"{marker} {r['rank']}. {r['species']:<45} {r['prob']:6.2%} logit {r['logit']:6.2f}{flag}")
print("-"*70)
print(f"Confidence {meta['confidence']:.2%} Margin {meta['margin']:.3f} Top-{args.topk} sum {sum([x['prob'] for x in meta['topk']]):.2%}")
if len(image_paths)>1:
print(f"Multi-view boost: single 89.21% -> 4-view 99.33% (pseudo) / 99.2% (train)")
# process-of-elimination hint
if meta["is_unknown"]:
print("Hint: UNKNOWN -> try flora description + geo, or collect more views. If truly unknown, consider iNaturalist/GBIF search outside 999.")
elif meta["topk"][0]["prob"] < 0.6:
print("Hint: Top-1 <60% -> try another angle/flower/leaf + text e.g. 'yellow puff' + geo for elimination")
# save outputs
out={"images":[str(p) for p in image_paths], "text":args.text, "geo":geo, "topk":meta["topk"], "confidence":meta["confidence"], "is_unknown": meta["is_unknown"], "unknown_info": meta["unknown_info"], "device":device, "time_ms": dt*1000, "model": "fusion 92% with text" if meta["text_used"] else "visual 89.21% plant_phase3b_otherblue", "centroids_used": meta["centroids_used"]}
if args.json:
Path(args.json).write_text(json.dumps(out, indent=2), encoding="utf-8")
print(f"wrote {args.json}")
if args.output:
# csv
pd.DataFrame(meta["topk"]).to_csv(args.output, index=False)
print(f"wrote {args.output}")
if args.vis:
# visualize topk bar + image + feature heatmap (where ViT looks)
try:
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
n=len(image_paths)
# --- generate heatmap for first image (where feature extractor finds features) ---
heat_arr = None
try:
vm, _, _ = load_visual_model(device=device)
heat_arr, hm_idx, _ = get_vit_heatmap(vm, image_paths[0], device=device, img_size=args.size)
except Exception as e:
log(f"heatmap gen failed {e}")
heat_arr = None
# 1x3 layout: image | heatmap overlay | bar (if heatmap available, else 1x2)
if heat_arr is not None:
fig, axes = plt.subplots(1, 3, figsize=(18,5), gridspec_kw={"width_ratios":[1,1,1.2]})
ax_img, ax_heat, ax_bar = axes
# original image
im_raw=open_image_pil(image_paths[0]).resize((args.size,args.size))
ax_img.imshow(im_raw)
title=f"{meta['topk'][0]['species']}\n{meta['topk'][0]['prob']:.1%} conf {meta['confidence']:.1%} {dt*1000:.0f}ms"
if meta["is_unknown"]:
title="UNKNOWN\n"+title
if len(image_paths)>1: title+=f" {n}-view"
ax_img.set_title(title, fontsize=9, color="red" if meta["is_unknown"] else "black")
ax_img.axis("off")
# heatmap overlay (jet)
ax_heat.imshow(im_raw)
ax_heat.imshow(heat_arr, cmap="jet", alpha=0.55, vmin=0, vmax=1)
ax_heat.set_title(f"Feature heatmap\nTop patch {hm_idx} {heat_arr.max():.2f}", fontsize=9, color="red" if meta["is_unknown"] else "black")
ax_heat.axis("off")
else:
fig, axes = plt.subplots(1, 2, figsize=(14,5), gridspec_kw={"width_ratios":[1,1.2]})
ax_img, ax_bar = axes
im=open_image_pil(image_paths[0]).resize((336,336))
ax_img.imshow(im)
title=f"{meta['topk'][0]['species']}\n{meta['topk'][0]['prob']:.1%} conf {meta['confidence']:.1%} {dt*1000:.0f}ms"
if meta["is_unknown"]:
title="UNKNOWN\n"+title
if len(image_paths)>1: title+=f" {n}-view"
ax_img.set_title(title, fontsize=9, color="red" if meta["is_unknown"] else "black")
ax_img.axis("off")
# heat axis not present, will reuse ax_bar below
# bar chart (common)
if heat_arr is not None:
probs_bar=[r["prob"] for r in meta["topk"]]
species=[r["species"] for r in meta["topk"]]
else:
probs_bar=[r["prob"] for r in meta["topk"]]
species=[r["species"] for r in meta["topk"]]
# ax_bar is defined in both branches (for heat case it's third axis)
colors=["red" if meta["is_unknown"] and i==0 else "green" if i==0 else "steelblue" for i in range(len(probs_bar))]
bars=ax_bar.barh(range(len(probs_bar))[::-1], probs_bar[::-1], color=colors[::-1])
ax_bar.set_yticks(range(len(probs_bar))[::-1])
ax_bar.set_yticklabels([f"{r['rank']}. {s[:32]}" for r,s in zip(meta["topk"], species)][::-1], fontsize=7)
ax_bar.set_xlabel("prob")
ax_bar.set_title(f"WA 999 spp Top-{args.topk} {'UNKNOWN' if meta['is_unknown'] else ''}", fontsize=10, color="red" if meta["is_unknown"] else "black")
for p, bar in zip(probs_bar[::-1], bars):
ax_bar.text(p+0.01, bar.get_y()+bar.get_height()/2, f"{p:.1%}", va="center", fontsize=7)
if meta["is_unknown"]:
ax_bar.text(0.5, -1, f"sim {meta['unknown_info']['sim_max']:.2f} < {meta['thr_cos']:.2f}", ha="center", fontsize=8, color="red", transform=ax_bar.transAxes)
plt.tight_layout()
plt.savefig(args.vis, dpi=150)
print(f"wrote {args.vis} {'with heatmap' if heat_arr is not None else ''}")
except Exception as e:
import traceback; traceback.print_exc()
log(f"vis failed {e}")
def cmd_batch(args):
device="cuda" if torch.cuda.is_available() and not args.cpu else "cpu"
dirp=Path(args.dir)
if not dirp.exists():
log(f"ERROR dir not found {dirp}")
sys.exit(1)
pattern=args.pattern
files=list(dirp.rglob(pattern) if args.recursive else dirp.glob(pattern))
# filter images
exts={".jpg",".jpeg",".png",".tif",".tiff",".heic",".webp",".bmp"}
files=[f for f in files if f.suffix.lower() in exts]
# optional DSLR RAW -> ignore
if not files:
log(f"no images found {dirp} {pattern}")
sys.exit(1)
files=sorted(files)[:args.limit] if args.limit else sorted(files)
log(f"batch {len(files)} images {dirp} pattern {pattern} text={args.text}")
# group for 4-view: if filenames share prefix before _ or -? Simple: each image independent unless --group
# For seamless, each image => independent predict, but if --group 4 => chunk 4
rows=[]
centroids_path = Path(args.centroids) if getattr(args, 'centroids', None) else CENTROIDS_PATH
thr_cos = getattr(args, 'thr_cos', DEFAULT_COS_THR)
thr_conf = getattr(args, 'thr_conf', DEFAULT_CONF_THR)
if args.group:
# group by stem prefix or consecutive 4
for i in range(0,len(files),args.group):
chunk=files[i:i+args.group]
meta,i2s,_,_,_ = predict_images(chunk, text=args.text, topk=args.topk, device=device, use_fusion_auto=not args.no_fusion, centroids_path=centroids_path, thr_cos=thr_cos, thr_conf=thr_conf)
for r in meta["topk"]:
rows.append({"group": i//args.group, "images": ";".join([str(c) for c in chunk]), "n_views": len(chunk), "rank": r["rank"], "species": r["species"], "prob": r["prob"], "confidence": meta["confidence"], "is_unknown": meta["is_unknown"], "sim_max": meta["unknown_info"]["sim_max"]})
log(f"group {i//args.group} {chunk[0].name} -> {meta['topk'][0]['species']} {meta['topk'][0]['prob']:.1%} {'UNKNOWN' if meta['is_unknown'] else ''}")
else:
for idx, f in enumerate(files):
meta,i2s,_,_,_ = predict_images([f], text=args.text, topk=args.topk, device=device, use_fusion_auto=not args.no_fusion, centroids_path=centroids_path, thr_cos=thr_cos, thr_conf=thr_conf)
for r in meta["topk"]:
rows.append({"image": str(f), "rank": r["rank"], "species": r["species"], "prob": r["prob"], "confidence": meta["confidence"], "is_unknown": meta["is_unknown"], "sim_max": meta["unknown_info"]["sim_max"]})
if idx%10==0:
log(f"{idx+1}/{len(files)} {f.name} -> {meta['topk'][0]['species']} {meta['topk'][0]['prob']:.1%} {'UNKNOWN' if meta['is_unknown'] else ''} sim {meta['unknown_info']['sim_max']:.2f}")
if args.out:
outp=Path(args.out)
df=pd.DataFrame(rows)
if outp.suffix==".json":
outp.write_text(json.dumps(rows, indent=2), encoding="utf-8")
else:
df.to_csv(outp, index=False)
print(f"wrote {outp} {len(rows)} rows")
else:
# print summary
print(json.dumps(rows[:5], indent=2))
print(f"... {len(rows)} rows, use --out results.csv to save")
def cmd_info(args):
print("WA Plant Identifier — Built with DINOv3")
print("Models:")
# show both private and public candidates
cand_paths = [_resolve_ckpt(CKPT_VISUAL)] + PUBLIC_CKPTS
seen=set()
for p in cand_paths:
if str(p) in seen: continue
seen.add(str(p))
exists=p.exists()
try:
rel=p.relative_to(BASE)
except ValueError:
rel=p
size=f"{p.stat().st_size/1e6:.1f}MB" if exists and p.is_file() else ("dir" if exists else "missing")
print(f" {rel} {size} {'OK' if exists else 'MISSING'}")
for p in [CKPT_FUSION, E5_DIR, MANIFEST, SPECIES_JSON, CENTROIDS_PATH]:
exists=p.exists()
try:
rel=p.relative_to(BASE)
except ValueError:
rel=p
size=f"{p.stat().st_size/1e6:.1f}MB" if exists and p.is_file() else ("dir" if p.exists() and p.is_dir() else "missing")
if p==CENTROIDS_PATH and exists:
try:
data=np.load(str(p), allow_pickle=True)
if "centroids" in data:
cents=data["centroids"]
size=f"{p.stat().st_size/1e6:.1f}MB {cents.shape} ~{cents.nbytes/1e6:.1f}MB cache"
except: pass
print(f" {rel} {size} {'OK' if exists else 'MISSING'}")
try:
s2i,i2s,spp = load_species_map()
print(f"Species {len(spp)} ({len(spp)} loaded) e.g. {spp[0]}, {spp[10] if len(spp)>10 else spp[-1]}")
except Exception as e:
print(f"Species map failed: {e}")
spp=[]
print(f"Device {'cuda '+torch.cuda.get_device_name(0) if torch.cuda.is_available() else 'cpu'} torch {torch.__version__}")
print(f"Benchmark val 27673: visual 89.21% Top-1 96.19% Top-5 -> fusion with text 92.00% 97.51%")
print(f"Multi-view: 1-view 89.04% -> 4-view 99.33% (pseudo) | Public HF: MoE 89.31% / Dense 82.08% single, 99.42%/98.31% 4-view")
print(f"Unknown detection: centroids {CENTROIDS_PATH.name} ~2MB (999x512) | thr_cos {DEFAULT_COS_THR} thr_conf {DEFAULT_CONF_THR} | cached on first load")
print("\nLicensing: DINOv3 Meta commercial grant required (LICENSE.md), text e5 MIT, data GBIF per-image BY/BY-NC")
if _resolve_ckpt(CKPT_VISUAL).exists():
print(f"Active ckpt: {_resolve_ckpt(CKPT_VISUAL).name} ({_resolve_ckpt(CKPT_VISUAL).stat().st_size/1e6:.1f}MB)")
else:
print("Active ckpt: MISSING - add PlantDetect-*.safetensors to project root")
cents, meta = load_centroids(CENTROIDS_PATH)
if cents is not None:
print(f"Centroids: {cents.shape} loaded, thr_cos {meta.get('thr_cos', DEFAULT_COS_THR) if meta else DEFAULT_COS_THR}")
else:
print(f"Centroids: MISSING -> run `python plant_cli.py build-centroids --dummy` for placeholder or --manifest for true")
def cmd_build_centroids(args):
# Build 999x512 centroids ~2MB cache
out = Path(args.out) if args.out else CENTROIDS_PATH
thr_cos = args.thr_cos
# try real manifest
manifest = Path(args.manifest) if args.manifest else MANIFEST
if not args.dummy and manifest.exists():
log(f"building centroids from manifest {manifest} (true per-class mean)...")
# need to load model and iterate dataset
device="cuda" if torch.cuda.is_available() and not args.cpu else "cpu"
model, s2i, i2s = load_visual_model(device=device)
s2i_manifest, i2s_manifest, spp_manifest = load_species_map(manifest)
# but use species_labels.json order for centroids
s2i_lab, i2s_lab, spp_lab = load_species_map()
# map manifest species to lab index
# collect embeddings per class
from src.data.plant import PlantDataset # reuse if needed but manifest is different format
# Instead manual csv reading like PlantDataset but for our manifest
import csv
from collections import defaultdict
rows=[]
with open(manifest, newline="", encoding="utf-8") as f:
for r in csv.DictReader(f):
if r.get("split","train")==args.split and r.get("status","downloaded") in ("downloaded","skip_exists"):
# path handling
p=Path(r.get("path",""))
if not p.exists():
# try gbifID under data/wa_plants
gbif=r.get("gbifID","")
base=BASE/"data"/"wa_plants"
for split in ("train","val"):
cand=base/split/r.get("species","").replace(" ","_").replace("/","_")[:120]/f"{gbif}.jpg"
if cand.exists():
p=cand
break
if p.exists():
rows.append((r["species"], p))
if args.limit and len(rows)>=args.limit:
break
if not rows:
log(f"no rows found in {manifest} for split {args.split}, fallback to dummy")
args.dummy=True
else:
# group by species
from collections import defaultdict
per_species=defaultdict(list)
for spp_name, p in rows:
if spp_name in s2i_lab:
per_species[spp_name].append(p)
tf=get_transform(336)
centroids=np.zeros((len(spp_lab), 512), dtype=np.float32)
counts=np.zeros(len(spp_lab), dtype=int)
model.eval()
for idx, spp_name in enumerate(spp_lab):
paths=per_species.get(spp_name, [])
if not paths:
# no data -> keep random small
centroids[idx]=np.random.randn(512).astype(np.float32)
continue
embs=[]
for p in paths[:args.per_class]:
try:
im=open_image_pil(p)
x=tf(im).unsqueeze(0).to(device)
with torch.no_grad():
with torch.autocast("cuda", dtype=torch.bfloat16, enabled=device=="cuda"):
_, emb, _ = model(x)
embs.append(emb.float().cpu().numpy()[0])
except Exception as e:
continue
if embs:
embs=np.stack(embs)
# L2 normalize then mean then normalize
embs=embs/np.linalg.norm(embs, axis=1, keepdims=True).clip(min=1e-8)
mean=embs.mean(0)
mean=mean/np.linalg.norm(mean).clip(min=1e-8)
centroids[idx]=mean
counts[idx]=len(embs)
else:
centroids[idx]=np.random.randn(512).astype(np.float32)
if idx%100==0:
log(f"{idx+1}/{len(spp_lab)} {spp_name} {len(paths)} imgs -> {counts[idx]} used")
# save
np.savez_compressed(str(out), centroids=centroids, species=np.array(spp_lab), thr_cos=np.array(thr_cos), thr_conf=np.array(DEFAULT_CONF_THR), counts=counts)
log(f"wrote centroids {centroids.shape} {centroids.nbytes/1e6:.1f}MB to {out} (per_class {args.per_class}, {counts.sum()} embeddings)")
return
if args.dummy:
log(f"building DUMMY centroids (placeholder, 2MB) -> run with --manifest for true centroids. Using green-centered + hue mix for demo.")
# Build proxy centroids that cover plant color space ~ green-centered + small noise
# Use green reference embedding as base (plant manifold)
device="cuda" if torch.cuda.is_available() and not args.cpu else "cpu"
model, s2i, i2s = load_visual_model(device=device)
s2i_lab, i2s_lab, spp_lab = load_species_map()
# get green base emb
ref_pil=Image.new("RGB",(336,336),(60,120,60)) # muted green (plant)
tf=get_transform(336)
x=tf(ref_pil).unsqueeze(0).to(device)
with torch.no_grad():
with torch.autocast("cuda", dtype=torch.bfloat16, enabled=device=="cuda"):
_, emb_base, _ = model(x)
emb_base=emb_base.float().cpu().numpy()[0]
emb_base=emb_base/np.linalg.norm(emb_base)
# generate 999 centroids around base with small per-class offset (hue)
np.random.seed(42)
centroids=np.zeros((len(spp_lab),512), dtype=np.float32)
for i in range(len(spp_lab)):
# per-class hue offset: small deterministic vector
noise=np.random.randn(512).astype(np.float32)*0.04
# also add hue-like variation for diversity: species name hash
h=hash(spp_lab[i]) % 1000 / 1000.0
noise[0]+= (h-0.5)*0.02
c=emb_base + noise
c=c/np.linalg.norm(c)
centroids[i]=c
# recalibrate thr: use 0.50 for dummy (as tested: green 0.72 vs noise 0.42)
np.savez_compressed(str(out), centroids=centroids.astype(np.float16), species=np.array(spp_lab), thr_cos=np.array(thr_cos), thr_conf=np.array(DEFAULT_CONF_THR))
# also save fp16 for half size
size_mb=Path(out).stat().st_size/1e6
raw_mb=centroids.nbytes/1e6
log(f"wrote DUMMY centroids {centroids.shape} raw {raw_mb:.1f}MB compressed {size_mb:.1f}MB to {out} (placeholder, rebuild with --manifest for true)")
log(f"NOTE: Dummy centroids only for demo/testing unknown logic. For production, run: python plant_cli.py build-centroids --manifest {MANIFEST} --per-class 10")
return
log(f"ERROR: no manifest {manifest} and --dummy not set. Use --dummy for placeholder or provide --manifest")
def cmd_test_unknown(args):
# Quick test: known vs unknown images with current centroids
centroids_path=Path(args.centroids) if args.centroids else CENTROIDS_PATH
centroids, meta = load_centroids(centroids_path)
thr_cos = args.thr_cos if args.thr_cos else (float(meta.get("thr_cos", DEFAULT_COS_THR)) if meta else DEFAULT_COS_THR)
thr_conf = args.thr_conf if args.thr_conf else DEFAULT_CONF_THR
device="cuda" if torch.cuda.is_available() and not args.cpu else "cpu"
if not args.images or len(args.images)<2:
log("need at least 2 images: known + unknown")
sys.exit(1)
for p in args.images:
path=Path(p)
if not path.exists():
log(f"not found {path}")
continue
meta_res, _, _, _, _ = predict_images([path], device=device, centroids_path=centroids_path, thr_cos=thr_cos, thr_conf=thr_conf)
status="UNKNOWN" if meta_res["is_unknown"] else "KNOWN"
print(f"\n{path.name}: {status} sim {meta_res['unknown_info']['sim_max']:.3f} thr {thr_cos:.2f} conf {meta_res['unknown_info']['conf']:.2%} margin {meta_res['unknown_info']['margin']:.3f} -> {meta_res['topk'][0]['species']} {meta_res['topk'][0]['prob']:.2%}")
if meta_res["is_unknown"]:
print(f" REASON: {meta_res['unknown_info']['reasons']}")
def main():
parser=argparse.ArgumentParser(description="WA Plant Identifier — Built with DINOv3 | DSLR/phone seamless CLI", formatter_class=argparse.RawTextHelpFormatter)
parser.add_argument("--cpu", action="store_true", help="force CPU")
sub=parser.add_subparsers(dest="cmd", required=True)
p_id=sub.add_parser("identify", help="identify 1-4 images (DLSR/phone JPEG)")
p_id.add_argument("images", nargs="+", help="image path(s) 1-4 for multi-view")
p_id.add_argument("--text", type=str, default=None, help='optional user text e.g. "yellow puff flower red loam" for +2.8 pct')
p_id.add_argument("--lat", type=float, default=None, help="latitude for geo prior")
p_id.add_argument("--lon", type=float, default=None, help="longitude for geo prior")
p_id.add_argument("--topk", type=int, default=5, help="Top-K")
p_id.add_argument("--size", type=int, default=336, help="image size")
p_id.add_argument("--no-fusion", action="store_true", help="disable fusion even if text provided")
p_id.add_argument("--fusion-thr", type=float, default=0.98, help="fusion gate conf thr 0.98")
p_id.add_argument("--fusion-margin", type=float, default=0.70, help="fusion gate margin thr 0.70")
p_id.add_argument("--json", type=str, default=None, help="save JSON")
p_id.add_argument("--output", type=str, default=None, help="save CSV")
p_id.add_argument("--vis", type=str, default=None, help="save visualization PNG")
p_id.add_argument("--centroids", type=str, default=None, help="centroids npz path (default species_centroids.npz, ~2MB)")
p_id.add_argument("--thr-cos", type=float, default=DEFAULT_COS_THR, help=f"unknown cos thr {DEFAULT_COS_THR}")
p_id.add_argument("--thr-conf", type=float, default=DEFAULT_CONF_THR, help=f"unknown conf thr {DEFAULT_CONF_THR}")
p_id.add_argument("--thr-margin", type=float, default=DEFAULT_MARGIN_THR, help=f"unknown margin thr {DEFAULT_MARGIN_THR}")
p_id.add_argument("--no-centroids", action="store_true", help="disable centroid unknown check (use conf/margin only)")
p_b=sub.add_parser("batch", help="batch folder of DSLR/phone images")
p_b.add_argument("--dir", type=str, required=True, help="folder with images")
p_b.add_argument("--pattern", type=str, default="*.jpg", help="glob e.g. star.jpg or star.JPG")
p_b.add_argument("--recursive", action="store_true", help="rglob")
p_b.add_argument("--group", type=int, default=0, help="group N images as multi-view (e.g. 4)")
p_b.add_argument("--text", type=str, default=None, help="optional text for all")
p_b.add_argument("--topk", type=int, default=5)
p_b.add_argument("--out", type=str, default=None, help="out.csv or out.json")
p_b.add_argument("--limit", type=int, default=0, help="limit N images for test")
p_b.add_argument("--no-fusion", action="store_true")
p_b.add_argument("--centroids", type=str, default=None, help="centroids path")
p_b.add_argument("--thr-cos", type=float, default=DEFAULT_COS_THR)
p_b.add_argument("--thr-conf", type=float, default=DEFAULT_CONF_THR)
p_bench=sub.add_parser("info", help="show model info")
p_build=sub.add_parser("build-centroids", help="build 999x512 centroid cache ~2MB")
p_build.add_argument("--out", type=str, default=str(CENTROIDS_PATH), help="out npz")
p_build.add_argument("--manifest", type=str, default=str(MANIFEST), help="manifest csv for true centroids")
p_build.add_argument("--split", type=str, default="train", choices=["train","val","gold"])
p_build.add_argument("--per-class", type=int, default=10, help="images per class for mean")
p_build.add_argument("--limit", type=int, default=0, help="limit total rows for quick test")
p_build.add_argument("--thr-cos", type=float, default=DEFAULT_COS_THR, help="store thr")
p_build.add_argument("--dummy", action="store_true", help="create dummy placeholder centroids (2MB) when no manifest")
p_build.add_argument("--cpu", action="store_true", help="force CPU for build")
p_test=sub.add_parser("test-unknown", help="test known vs unknown with centroids")
p_test.add_argument("images", nargs="+", help="known.jpg unknown.jpg ...")
p_test.add_argument("--centroids", type=str, default=None, help="centroids path")
p_test.add_argument("--thr-cos", type=float, default=None)
p_test.add_argument("--thr-conf", type=float, default=None)
p_test.add_argument("--cpu", action="store_true")
p_bench2=sub.add_parser("bench", help="bench on val/gold split")
p_bench2.add_argument("--split", type=str, default="val", choices=["val","gold","train"])
p_bench2.add_argument("--manifest", type=str, default=str(MANIFEST))
args=parser.parse_args()
if args.cmd=="identify":
cmd_identify(args)
elif args.cmd=="batch":
cmd_batch(args)
elif args.cmd=="info":
cmd_info(args)
elif args.cmd=="build-centroids":
cmd_build_centroids(args)
elif args.cmd=="test-unknown":
cmd_test_unknown(args)
elif args.cmd=="bench":
# quick bench via bench_intense
import subprocess
ckpt = _resolve_ckpt(CKPT_VISUAL)
cmd=[sys.executable, "scripts/bench_intense.py", "--ckpt", str(ckpt), "--manifest", args.manifest, "--hardneg", str(BASE/"data/wa_plants_200k/hard_negatives.json"), "--n-classes", "999", "--split", args.split, "--size", "336", "--batch", "32", "--workers", "0", "--out", str(BASE/f"data/bench_cli_{args.split}.json")]
print(" ".join(cmd))
subprocess.run(cmd)
if __name__=="__main__":
main()
|