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
| """MoR-ViT: Mixture-of-Recursions on a pretrained DINOv3 ViT stem. | |
| Maps the llm-pipeline MoR recipe to a vision transformer: | |
| entry block -> pretrained DINOv3 stem (unique weights) | |
| shared core -> ``RecursiveAttentionBlock`` stack reused at every recursion | |
| depth router -> per-TOKEN ``TokenRouter`` (expert choice over patch tokens) | |
| recursion emb -> learned per-recursion embedding added before the core | |
| token freeze -> a patch token that stops routing keeps its current state | |
| Routing is per patch token (adaptive compute over the image), not per slice: | |
| every study keeps the full recursion budget, but only the "hard" patch tokens | |
| are recursed deeply. Slice features are mean-pooled over patch tokens, then | |
| weighted by a ``SliceAttentionPool`` into a study vector. | |
| ``features()`` exposes the frozen feature extractor used by linear probing | |
| without running the classification head. | |
| """ | |
| from __future__ import annotations | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| try: | |
| import timm | |
| HAS_TIMM = True | |
| except ImportError: # pragma: no cover - optional dependency | |
| HAS_TIMM = False | |
| class TokenScore(nn.Module): | |
| """Per-token continuation score: token embedding -> scalar logit. | |
| ``slice_gain`` (zero-initialised) couples the slice-level abnormality | |
| score into the patch decision; the forward is exactly legacy at init. | |
| """ | |
| def __init__(self, dim: int, hidden: int, init_bias: float = 0.0): | |
| super().__init__() | |
| self.in_proj = nn.Linear(dim, hidden) | |
| self.out_proj = nn.Linear(hidden, 1) | |
| self.slice_gain = nn.Parameter(torch.zeros(1)) | |
| with torch.no_grad(): | |
| self.out_proj.bias.fill_(init_bias) | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| return self.out_proj(F.silu(self.in_proj(x))).squeeze(-1) | |
| class SliceScore(nn.Module): | |
| """Per-slice continuation score: pooled patch feature -> scalar logit. | |
| Gives the patch router a slice-level view, so routing stays consistent | |
| across the slices of a study (a finding shows on several adjacent slices). | |
| """ | |
| def __init__(self, dim: int, hidden: int, init_bias: float = 0.0): | |
| super().__init__() | |
| self.in_proj = nn.Linear(dim, hidden) | |
| self.out_proj = nn.Linear(hidden, 1) | |
| with torch.no_grad(): | |
| self.out_proj.bias.fill_(init_bias) | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| return self.out_proj(F.silu(self.in_proj(x))).squeeze(-1) | |
| class TokenRouter(nn.Module): | |
| """Expert-choice depth router over patch tokens. | |
| One ``TokenScore`` head per recursion, plus a learned recursion embedding | |
| so the router can condition on how deep a token already is. The router | |
| input can be spatially conditioned (stem patch position embedding) and | |
| slice-conditioned (``SliceScore`` coupled through a zero-init gain). | |
| Router biases initialise from the per-recursion capacity so the right | |
| fraction of tokens continues from step 0. | |
| """ | |
| def __init__( | |
| self, | |
| dim: int, | |
| hidden: int, | |
| n_recursions: int, | |
| capacities: list[float], | |
| init_bias: float = 0.0, | |
| warmup_steps: int = 0, | |
| init_from_capacity: bool = True, | |
| ): | |
| super().__init__() | |
| self.n_recursions = n_recursions | |
| self.capacities = list(capacities) | |
| self.warmup_steps = warmup_steps | |
| self.heads = nn.ModuleList( | |
| [TokenScore(dim, hidden, init_bias) for _ in range(n_recursions)] | |
| ) | |
| self.rec_emb = nn.Parameter(torch.zeros(n_recursions, dim)) | |
| self.slice_score = SliceScore(dim, hidden) | |
| if init_from_capacity: | |
| for r, head in enumerate(self.heads): | |
| cap = max(1e-3, min(1.0 - 1e-3, self.capacities[min(r, len(self.capacities) - 1)])) | |
| with torch.no_grad(): | |
| head.out_proj.bias.fill_(float(torch.logit(torch.tensor(cap)))) | |
| def capacity(self, r: int, step: int) -> float: | |
| """Token fraction kept at recursion ``r``, ramping from 1.0 during warmup.""" | |
| target = self.capacities[r] | |
| if self.warmup_steps > 0 and step < self.warmup_steps: | |
| t = step / self.warmup_steps | |
| return 1.0 - (1.0 - target) * t | |
| return target | |
| def forward( | |
| self, | |
| x: torch.Tensor, | |
| r: int, | |
| patch_pos: torch.Tensor | None = None, | |
| slice_feat: torch.Tensor | None = None, | |
| ) -> torch.Tensor: | |
| h = x + self.rec_emb[r] | |
| if patch_pos is not None: | |
| h = h + patch_pos | |
| score = self.heads[r](h) | |
| if slice_feat is not None: | |
| score = score + self.heads[r].slice_gain * self.slice_score(slice_feat).unsqueeze(1) | |
| return score | |
| class RecursiveAttentionBlock(nn.Module): | |
| """Shared pre-norm attention block reused at every recursion. | |
| Operates on the active token group only (the caller gathers/scatters), so | |
| attention is plain unmasked full attention - no -inf masks, no NaN paths. | |
| """ | |
| def __init__(self, dim: int, n_heads: int = 6, mlp_ratio: float = 4.0, dropout: float = 0.0): | |
| super().__init__() | |
| self.n_heads = n_heads | |
| self.head_dim = dim // n_heads | |
| self.q = nn.Linear(dim, dim) | |
| self.k = nn.Linear(dim, dim) | |
| self.v = nn.Linear(dim, dim) | |
| self.proj = nn.Linear(dim, dim) | |
| self.norm1 = nn.LayerNorm(dim) | |
| self.norm2 = nn.LayerNorm(dim) | |
| self.mlp = nn.Sequential( | |
| nn.Linear(dim, int(dim * mlp_ratio)), | |
| nn.GELU(), | |
| nn.Linear(int(dim * mlp_ratio), dim), | |
| ) | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| S, T, D = x.shape | |
| H = self.n_heads | |
| xn = self.norm1(x) | |
| q = self.q(xn).reshape(S, T, H, self.head_dim).transpose(1, 2) | |
| k = self.k(xn).reshape(S, T, H, self.head_dim).transpose(1, 2) | |
| v = self.v(xn).reshape(S, T, H, self.head_dim).transpose(1, 2) | |
| attn = (q @ k.transpose(-1, -2) / (self.head_dim ** 0.5)).softmax(dim=-1) | |
| out = (attn @ v).transpose(1, 2).reshape(S, T, D) | |
| x = x + self.proj(out) | |
| x = x + self.mlp(self.norm2(x)) | |
| return x | |
| class SliceAttentionPool(nn.Module): | |
| """Weight slices by learned relevance before aggregating into a study vector.""" | |
| def __init__(self, dim: int): | |
| super().__init__() | |
| self.query = nn.Parameter(torch.randn(dim)) | |
| self.scale = dim ** -0.5 | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| w = F.softmax(x @ self.query * self.scale, dim=0) | |
| return (w.unsqueeze(1) * x).sum(0) | |
| class MoRViT(nn.Module): | |
| """Pretrained DINOv3 stem + MoR recursive core with per-token routing.""" | |
| def __init__( | |
| self, | |
| *, | |
| stem_name: str = "vit_small_patch16_dinov3", | |
| stem: nn.Module | None = None, | |
| pretrained: bool = True, | |
| freeze_stem: bool = True, | |
| use_mor: bool = True, | |
| dense: bool = False, | |
| n_recursions: int = 3, | |
| capacities: tuple[float, ...] = (1.0, 2.0 / 3.0, 1.0 / 3.0), | |
| core_blocks: int = 2, | |
| router_hidden: int = 128, | |
| n_heads: int = 6, | |
| n_classes: int = 12, | |
| router_warmup_steps: int = 0, | |
| router_init_bias: float = 0.0, | |
| ): | |
| super().__init__() | |
| if stem is None: | |
| if not HAS_TIMM: | |
| raise ImportError("timm is required for a pretrained stem") | |
| stem = timm.create_model(stem_name, pretrained=pretrained, num_classes=0) | |
| self.n_recursions = n_recursions | |
| self.capacities = list(capacities) | |
| assert len(self.capacities) == n_recursions | |
| self.use_mor = use_mor | |
| self.dense = dense | |
| self.stem = stem | |
| stem_dim = self.stem.embed_dim | |
| self.core_dim = stem_dim | |
| if stem_dim != 384: | |
| self.input_proj = nn.Linear(stem_dim, 384) | |
| self.core_dim = 384 | |
| else: | |
| self.input_proj = None | |
| dim = self.core_dim | |
| self.n_prefix = int(getattr(self.stem, "num_prefix_tokens", 1)) | |
| pe = getattr(self.stem, "pos_embed", None) | |
| if isinstance(pe, torch.Tensor) and pe.dim() == 3 and pe.size(1) > self.n_prefix: | |
| patch_pos = pe[:, self.n_prefix:].detach().clone() # [1, P, stem_dim] | |
| if self.input_proj is not None: | |
| with torch.no_grad(): | |
| patch_pos = self.input_proj(patch_pos) | |
| self.register_buffer("patch_pos", patch_pos, persistent=False) | |
| else: | |
| self.patch_pos = None | |
| if freeze_stem: | |
| for p in self.stem.parameters(): | |
| p.requires_grad_(False) | |
| self._soft_route_maps: list[torch.Tensor] = [] | |
| self.core = nn.ModuleList( | |
| [RecursiveAttentionBlock(dim, n_heads=n_heads) for _ in range(core_blocks)] | |
| ) | |
| self.router = ( | |
| TokenRouter( | |
| dim, | |
| router_hidden, | |
| n_recursions, | |
| self.capacities, | |
| init_bias=router_init_bias, | |
| warmup_steps=router_warmup_steps, | |
| ) | |
| if not dense | |
| else None | |
| ) | |
| self.rec_emb = nn.Parameter(torch.zeros(n_recursions, dim)) | |
| self.exit_norm = nn.LayerNorm(dim) | |
| self.pool = SliceAttentionPool(dim) | |
| self.head = nn.Sequential( | |
| nn.Linear(dim, dim * 2), | |
| nn.GELU(), | |
| nn.Linear(dim * 2, n_classes), | |
| ) | |
| def _run_core(self, tokens: torch.Tensor, r: int) -> torch.Tensor: | |
| x = tokens + self.rec_emb[r] | |
| for block in self.core: | |
| x = block(x) | |
| return x | |
| def _moR(self, tokens: torch.Tensor, step: int = 0) -> tuple[torch.Tensor, list[float]]: | |
| """Run the recursion loop over a batch of token sequences. | |
| Dense mode runs the shared core over every token at every recursion | |
| (the same-budget control for adaptive routing). Adaptive mode gathers | |
| the top-k patch tokens (by router score) plus the prefix tokens, | |
| processes them through the core, and scatters them back; non-selected | |
| patch tokens keep their previous state. | |
| The returned tokens are a soft mixture of the per-recursion states | |
| weighted by the router's continuation probabilities, so the router | |
| receives real gradient (hard top-k alone is non-differentiable). The | |
| soft routing maps are stored for the optional smoothness loss. | |
| """ | |
| if self.dense: | |
| for r in range(self.n_recursions): | |
| tokens = self._run_core(tokens, r) | |
| return tokens, [1.0] * self.n_recursions | |
| S, T, D = tokens.shape | |
| n_pref = self.n_prefix | |
| P = T - n_pref | |
| stats: list[float] = [] | |
| self._soft_route_maps = [] | |
| num = tokens.clone() | |
| den = torch.ones(S, T, device=tokens.device) | |
| w_pat = torch.ones(S, P, device=tokens.device) | |
| pref_ones = torch.ones(S, n_pref, device=tokens.device) | |
| pe = ( | |
| self.patch_pos | |
| if (self.patch_pos is not None and self.patch_pos.shape[1] == P) | |
| else None | |
| ) | |
| for r in range(self.n_recursions - 1): | |
| cap = self.router.capacity(r, step) | |
| pat = tokens[:, n_pref:] | |
| pat_scores = self.router(pat, r, patch_pos=pe, slice_feat=pat.mean(1)) | |
| p = torch.sigmoid(pat_scores) | |
| self._soft_route_maps.append(p) | |
| k = max(1, min(P, int(round(P * cap)))) | |
| sel = pat_scores.topk(k, dim=1).indices # [S, k] | |
| pref = torch.arange(n_pref, device=tokens.device).expand(S, n_pref) | |
| idx = torch.cat([pref, sel + n_pref], dim=1) # [S, M] | |
| idx3 = idx.unsqueeze(-1).expand(S, idx.size(1), D) | |
| stats.append((n_pref + k) / T) | |
| gathered = tokens.gather(1, idx3) | |
| updated = self._run_core(gathered, r) | |
| tokens = tokens.scatter(1, idx3, updated) | |
| w_pat = w_pat * p | |
| w = torch.cat([pref_ones, w_pat], dim=1) | |
| num = num + tokens * w.unsqueeze(-1) | |
| den = den + w | |
| tokens = self._run_core(tokens, self.n_recursions - 1) | |
| stats.append(1.0) | |
| num = num + tokens * w.unsqueeze(-1) | |
| den = den + w | |
| mixed = num / den.unsqueeze(-1) | |
| return mixed, stats | |
| def routing_smoothness_loss(self) -> torch.Tensor | float: | |
| """Total-variation penalty on the soft routing maps (spatial coherence). | |
| Knee findings occupy contiguous regions; a scattered routing map is a | |
| bug signal. Add ``weight * this`` to the training loss. | |
| """ | |
| if not self._soft_route_maps: | |
| return 0.0 | |
| total = None | |
| for p in self._soft_route_maps: # [S, P] | |
| P = p.size(1) | |
| h = int(round(P ** 0.5)) | |
| if h * h != P: | |
| continue | |
| g = p.reshape(p.size(0), h, h) | |
| tv = (g[:, 1:, :] - g[:, :-1, :]).abs().mean() + ( | |
| g[:, :, 1:] - g[:, :, :-1] | |
| ).abs().mean() | |
| total = tv if total is None else total + tv | |
| if total is None: | |
| return 0.0 | |
| return total / len(self._soft_route_maps) | |
| def features( | |
| self, x: torch.Tensor, step: int = 0, pool: str = "patchmean" | |
| ) -> torch.Tensor: | |
| """Per-slice features (no head): [S, D] from [S, 3, H, W] slices.""" | |
| tokens = self.stem.forward_features(x) # [S, T, D] | |
| return self.features_from_tokens(tokens, step, pool) | |
| def features_from_tokens( | |
| self, tokens: torch.Tensor, step: int = 0, pool: str = "patchmean" | |
| ) -> torch.Tensor: | |
| """Features from cached stem tokens (feat cache path): [S, D].""" | |
| tokens = tokens.float() # _moR internals are f32; bf16 cache upcasts exactly | |
| if self.input_proj is not None: | |
| with torch.autocast("cuda", enabled=False): | |
| tokens = self.input_proj(tokens).float() # keep f32; autocast would make it bf16 | |
| if self.use_mor or self.dense: | |
| tokens, _ = self._moR(tokens, step) | |
| tokens = self.exit_norm(tokens) | |
| if pool == "cls": | |
| return tokens[:, 0] | |
| return tokens[:, self.n_prefix :].mean(1) | |
| def forward( | |
| self, x: torch.Tensor, step: int = 0 | |
| ) -> tuple[torch.Tensor, list[float]]: | |
| """Forward one study: [S, 3, H, W] slices -> (logits [n_classes], route_stats).""" | |
| slice_feats = self.features(x, step) # [S, D] | |
| study = self.pool(slice_feats) # [D] | |
| logits = self.head(study) | |
| return logits, [] | |