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-CNN: Mixture-of-Recursions applied to a convolutional vision model. | |
| Ports the LLM MoR idea (llm-pipeline) to 2D medical imaging. The mapping: | |
| entry block -> pretrained CNN stem (unique weights) | |
| shared recursive -> ``RecursiveConvBlock`` stack reused at every recursion | |
| core | |
| depth router -> per-slice ``DepthRouter`` (expert choice over slices) | |
| recursion emb -> learned per-recursion embedding added before the core | |
| slice freeze -> a slice that stops routing keeps its current state | |
| At inference the router spends full recursion depth only on the "hard" slices | |
| (those most likely to be abnormal), so FLOPs scale with content — the lever the | |
| competition's Efficiency Track scores. | |
| """ | |
| from __future__ import annotations | |
| import math | |
| 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 DepthScore(nn.Module): | |
| """Per-slice continuation score: pooled feature -> scalar logit. | |
| Mirrors ``_DepthScore`` in llm-pipeline's MoR router, but over a global | |
| pooled per-slice feature instead of a per-token hidden state. | |
| """ | |
| 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 DepthRouter(nn.Module): | |
| """Expert-choice depth router over slices. | |
| One ``DepthScore`` head per recursion. Each recursion also adds a learned | |
| recursion embedding so the router can condition on how deep a slice | |
| already is (the vision analogue of ``rec_emb`` in the LLM MoR). | |
| """ | |
| def __init__( | |
| self, | |
| dim: int, | |
| hidden: int, | |
| n_recursions: int, | |
| capacities: list[float], | |
| init_bias: float = 0.0, | |
| warmup_steps: int = 0, | |
| ): | |
| super().__init__() | |
| self.n_recursions = n_recursions | |
| self.capacities = capacities | |
| self.warmup_steps = warmup_steps | |
| self.heads = nn.ModuleList( | |
| [DepthScore(dim, hidden, init_bias) for _ in range(n_recursions)] | |
| ) | |
| self.rec_emb = nn.Parameter(torch.zeros(n_recursions, dim)) | |
| def capacity(self, r: int, step: int) -> float: | |
| """Slice 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) -> torch.Tensor: | |
| # x: [S, dim] pooled per-slice features | |
| return self.heads[r](x + self.rec_emb[r]) | |
| class RecursiveConvBlock(nn.Module): | |
| """A shared inverted-residual conv block reused at every recursion. | |
| Depthwise-separable (MobileNet-v2 style) so the recursive core stays cheap | |
| while the stem does the heavy feature extraction. | |
| """ | |
| def __init__(self, dim: int): | |
| super().__init__() | |
| self.norm = nn.LayerNorm(dim) | |
| self.pw1 = nn.Conv2d(dim, dim * 2, 1) | |
| self.dw = nn.Conv2d(dim * 2, dim * 2, 3, padding=1, groups=dim * 2) | |
| self.pw2 = nn.Conv2d(dim * 2, dim, 1) | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| # x: [N, dim, h, w] | |
| identity = x | |
| x = self.norm(x.permute(0, 2, 3, 1)).permute(0, 3, 1, 2) | |
| x = self.pw2(F.gelu(self.dw(F.gelu(self.pw1(x))))) | |
| return identity + 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: | |
| # x: [S, dim] | |
| w = F.softmax(x @ self.query * self.scale, dim=0) | |
| return (w.unsqueeze(1) * x).sum(0) | |
| class ChannelLayerNorm(nn.Module): | |
| """LayerNorm over the channel dim of a [N, C, H, W] feature map.""" | |
| def __init__(self, dim: int): | |
| super().__init__() | |
| self.norm = nn.LayerNorm(dim) | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| return self.norm(x.permute(0, 2, 3, 1)).permute(0, 3, 1, 2) | |
| class MoRCNN(nn.Module): | |
| """Pretrained CNN stem + MoR recursive core + per-slice depth router.""" | |
| def __init__( | |
| self, | |
| *, | |
| stem_name: str = "convnext_tiny", | |
| pretrained: bool = True, | |
| n_recursions: int = 3, | |
| capacities: tuple[float, ...] = (1.0, 2.0 / 3.0, 1.0 / 3.0), | |
| router_hidden: int = 128, | |
| core_blocks: int = 2, | |
| n_classes: int = 12, | |
| router_warmup_steps: int = 0, | |
| router_init_bias: float = 0.0, | |
| ): | |
| super().__init__() | |
| if not HAS_TIMM: | |
| raise ImportError("timm is required for a pretrained stem") | |
| self.n_recursions = n_recursions | |
| self.capacities = list(capacities) | |
| assert len(self.capacities) == n_recursions | |
| self.stem = timm.create_model( | |
| stem_name, pretrained=pretrained, features_only=True, num_classes=0 | |
| ) | |
| out_dim = self.stem.feature_info.channels()[-1] | |
| self.core = nn.Sequential(*[RecursiveConvBlock(out_dim) for _ in range(core_blocks)]) | |
| self.router = DepthRouter( | |
| out_dim, | |
| router_hidden, | |
| n_recursions, | |
| self.capacities, | |
| init_bias=router_init_bias, | |
| warmup_steps=router_warmup_steps, | |
| ) | |
| self.rec_emb = nn.Parameter(torch.zeros(n_recursions, out_dim)) | |
| self.exit = nn.Sequential( | |
| ChannelLayerNorm(out_dim), | |
| nn.Conv2d(out_dim, out_dim, 1), | |
| nn.GELU(), | |
| ) | |
| self.pool = SliceAttentionPool(out_dim) | |
| self.head = nn.Sequential( | |
| nn.Linear(out_dim, out_dim * 2), | |
| nn.GELU(), | |
| nn.Linear(out_dim * 2, n_classes), | |
| ) | |
| def _run_core(self, feat: torch.Tensor, r: int) -> torch.Tensor: | |
| emb = self.rec_emb[r].view(1, -1, 1, 1) | |
| return self.core(feat + emb) | |
| def _global_pool(self, feat: torch.Tensor) -> torch.Tensor: | |
| return feat.mean(dim=(2, 3)) | |
| def forward( | |
| self, x: torch.Tensor, step: int = 0 | |
| ) -> tuple[torch.Tensor, list[float]]: | |
| """Forward one study. | |
| Args: | |
| x: [S, 3, H, W] sampled slices of a single study. | |
| step: current optimizer step, used for router warmup. | |
| Returns: | |
| (logits [n_classes], route_stats) where route_stats is the active | |
| slice fraction at each recursion. | |
| """ | |
| S = x.size(0) | |
| feat = self.stem(x)[-1] # [S, C, h, w] | |
| C = feat.size(1) | |
| pooled = self._global_pool(feat) | |
| stats: list[float] = [] | |
| for r in range(self.n_recursions): | |
| probs = torch.sigmoid(self.router(pooled, r)) | |
| cap = self.router.capacity(r, step) | |
| if r == self.n_recursions - 1: | |
| active = torch.ones(S, dtype=torch.bool, device=x.device) | |
| else: | |
| k = max(1, int(round(S * cap))) | |
| active = torch.zeros(S, dtype=torch.bool, device=x.device) | |
| active[torch.topk(probs, k).indices] = True | |
| stats.append(active.float().mean().item()) | |
| if active.all(): | |
| feat = self._run_core(feat, r) | |
| else: | |
| idx = active.nonzero(as_tuple=False).squeeze(1) | |
| updated = self._run_core(feat[idx], r) | |
| feat = feat.clone() | |
| feat[idx] = updated | |
| pooled = self._global_pool(feat) | |
| slice_feats = self.exit(feat) # [S, C, h, w] | |
| slice_feats = self._global_pool(slice_feats) # [S, C] | |
| study = self.pool(slice_feats) # [C] | |
| logits = self.head(study) # [n_classes] | |
| return logits, stats | |