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: 14,726 Bytes
3217f9a 7198674 3217f9a | 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 | """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, []
|