CandyEye 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
candyeye/__init__.py ADDED
@@ -0,0 +1,57 @@
1
+ """CandyEye — a lightweight, CPU-first object detector.
2
+
3
+ Public API::
4
+
5
+ from candyeye import CandyEye, train
6
+
7
+ model = CandyEye() # bundled yolo11n architecture
8
+ results = model.train(data="dataset.yaml", epochs=100, pretrained=True)
9
+
10
+ # ...or the function form:
11
+ results = train(data="dataset.yaml", epochs=100)
12
+
13
+ Heavy submodules (torch, torchvision) are imported lazily so ``import
14
+ candyeye`` stays cheap.
15
+ """
16
+ from __future__ import annotations
17
+
18
+ from candyeye.paths import (
19
+ available_configs,
20
+ cache_dir,
21
+ default_config_path,
22
+ default_weights_path,
23
+ download_default_weights,
24
+ resolve_config,
25
+ )
26
+
27
+ __version__ = "0.1.0"
28
+
29
+ __all__ = [
30
+ "CandyEye",
31
+ "CandyEyePredictor",
32
+ "train",
33
+ "predict",
34
+ "resolve_config",
35
+ "default_config_path",
36
+ "default_weights_path",
37
+ "download_default_weights",
38
+ "available_configs",
39
+ "cache_dir",
40
+ "__version__",
41
+ ]
42
+
43
+
44
+ def __getattr__(name):
45
+ if name == "CandyEye":
46
+ from candyeye.core.candyeye import CandyEye
47
+ return CandyEye
48
+ if name == "train":
49
+ from candyeye.training.trainer import train
50
+ return train
51
+ if name == "predict":
52
+ from candyeye.inference.predict import predict
53
+ return predict
54
+ if name == "CandyEyePredictor":
55
+ from candyeye.inference.predict import CandyEyePredictor
56
+ return CandyEyePredictor
57
+ raise AttributeError(f"module 'candyeye' has no attribute {name!r}")
@@ -0,0 +1,37 @@
1
+ # CandyEye — YOLO11n transcription, object detection (nc=20)
2
+ #
3
+ # Entry format: [from, repeats, module, args]
4
+ # from < 0 -> offset into the built layer list (-1 = previous layer)
5
+ # from >= 0 -> absolute layer index
6
+ # from list -> gather several layers (Concat / Detect)
7
+ # Channel args are in BASE units; the builder scales them by width_mult.
8
+ nc: 20
9
+ scale: n
10
+
11
+ backbone:
12
+ - [-1, 1, Conv, [64, 3, 2]] # 0-P1/2
13
+ - [-1, 1, Conv, [128, 3, 2]] # 1-P2/4
14
+ - [-1, 2, C3k2, [256, False, 0.25]] # 2
15
+ - [-1, 1, Conv, [256, 3, 2]] # 3-P3/8
16
+ - [-1, 2, C3k2, [512, False, 0.25]] # 4
17
+ - [-1, 1, Conv, [512, 3, 2]] # 5-P4/16
18
+ - [-1, 2, C3k2, [512, True]] # 6
19
+ - [-1, 1, Conv, [1024, 3, 2]] # 7-P5/32
20
+ - [-1, 2, C3k2, [1024, True]] # 8
21
+ - [-1, 1, SPPF, [1024, 5]] # 9
22
+ - [-1, 2, C2PSA, [1024]] # 10
23
+
24
+ head:
25
+ - [-1, 1, nn.Upsample, [None, 2, "nearest"]] # 11
26
+ - [[-1, 6], 1, Concat, [1]] # 12 cat backbone P4
27
+ - [-1, 1, C3k2, [512, False]] # 13
28
+ - [-1, 1, nn.Upsample, [None, 2, "nearest"]] # 14
29
+ - [[-1, 4], 1, Concat, [1]] # 15 cat backbone P3
30
+ - [-1, 1, C3k2, [256, False]] # 16 P3/8 (small)
31
+ - [-1, 1, Conv, [256, 3, 2]] # 17
32
+ - [[-1, 13], 1, Concat, [1]] # 18 cat head P4
33
+ - [-1, 1, C3k2, [512, False]] # 19 P4/16 (medium)
34
+ - [-1, 1, Conv, [512, 3, 2]] # 20
35
+ - [[-1, 10], 1, Concat, [1]] # 21 cat head P5
36
+ - [-1, 1, C3k2, [1024, True]] # 22 P5/32 (large)
37
+ - [[16, 19, 22], 1, Detect, [nc]] # 23 box+cls, parallel
@@ -0,0 +1,38 @@
1
+ # CandyEye — YOLO11n transcription with the adaptive cross-scale exchange neck.
2
+ #
3
+ # Same backbone/neck as configs/yolo11.yaml; layer 23 is a ScaleExchange that
4
+ # lets P3/P4/P5 exchange gated messages before the Detect head. The gate mode is
5
+ # the last arg: "none" | "static" | "dynamic" (optional second arg = iterations).
6
+ #
7
+ # Entry format: [from, repeats, module, args]
8
+ nc: 20
9
+ scale: n
10
+
11
+ backbone:
12
+ - [-1, 1, Conv, [64, 3, 2]] # 0-P1/2
13
+ - [-1, 1, Conv, [128, 3, 2]] # 1-P2/4
14
+ - [-1, 2, C3k2, [256, False, 0.25]] # 2
15
+ - [-1, 1, Conv, [256, 3, 2]] # 3-P3/8
16
+ - [-1, 2, C3k2, [512, False, 0.25]] # 4
17
+ - [-1, 1, Conv, [512, 3, 2]] # 5-P4/16
18
+ - [-1, 2, C3k2, [512, True]] # 6
19
+ - [-1, 1, Conv, [1024, 3, 2]] # 7-P5/32
20
+ - [-1, 2, C3k2, [1024, True]] # 8
21
+ - [-1, 1, SPPF, [1024, 5]] # 9
22
+ - [-1, 2, C2PSA, [1024]] # 10
23
+
24
+ head:
25
+ - [-1, 1, nn.Upsample, [None, 2, "nearest"]] # 11
26
+ - [[-1, 6], 1, Concat, [1]] # 12 cat backbone P4
27
+ - [-1, 1, C3k2, [512, False]] # 13
28
+ - [-1, 1, nn.Upsample, [None, 2, "nearest"]] # 14
29
+ - [[-1, 4], 1, Concat, [1]] # 15 cat backbone P3
30
+ - [-1, 1, C3k2, [256, False]] # 16 P3/8 (small)
31
+ - [-1, 1, Conv, [256, 3, 2]] # 17
32
+ - [[-1, 13], 1, Concat, [1]] # 18 cat head P4
33
+ - [-1, 1, C3k2, [512, False]] # 19 P4/16 (medium)
34
+ - [-1, 1, Conv, [512, 3, 2]] # 20
35
+ - [[-1, 10], 1, Concat, [1]] # 21 cat head P5
36
+ - [-1, 1, C3k2, [1024, True]] # 22 P5/32 (large)
37
+ - [[16, 19, 22], 1, ScaleExchange, ["dynamic"]] # 23 gated cross-scale exchange
38
+ - [-1, 1, Detect, [nc]] # 24 box+cls, parallel
@@ -0,0 +1,13 @@
1
+ """Public CandyEye model and training entry points, loaded on demand."""
2
+
3
+ __all__ = ["CandyEye", "train"]
4
+
5
+
6
+ def __getattr__(name):
7
+ if name == "CandyEye":
8
+ from candyeye.core.candyeye import CandyEye
9
+ return CandyEye
10
+ if name == "train":
11
+ from candyeye.training.trainer import train
12
+ return train
13
+ raise AttributeError(name)
@@ -0,0 +1,82 @@
1
+ """MobileNetV3-Small backbone with a feature pyramid and detection head."""
2
+ from __future__ import annotations
3
+
4
+ import torch
5
+ from torch import nn
6
+ import torch.nn.functional as F
7
+ from torchvision.models import MobileNet_V3_Small_Weights, mobilenet_v3_small
8
+
9
+ from candyeye.core.modules.blocks import C3k2
10
+ from candyeye.core.modules.conv import Conv
11
+ from candyeye.core.modules.detect import Detect
12
+ from candyeye.core.modules.exchange import ScaleExchange
13
+
14
+
15
+ class MobileNetV3SmallDetector(nn.Module):
16
+ """Small VOC detector using MobileNetV3-Small features at strides 8/16/32.
17
+
18
+ ImageNet initialization is optional because torchvision may need to
19
+ download the checkpoint when it is not in the local torch cache.
20
+
21
+ ``neck="light"`` (default) keeps the compact custom FPN. ``neck="exchange"``
22
+ appends the gated cross-scale exchange neck before the Detect head.
23
+ """
24
+ def __init__(self, nc: int = 20, img_size: int = 128,
25
+ pretrained: bool = False, neck: str = "light",
26
+ exchange_gate: str = "none", exchange_iters: int = 1):
27
+ super().__init__()
28
+ if img_size % 32:
29
+ raise ValueError(f"img_size must be divisible by 32, got {img_size}")
30
+ if neck not in ("light", "exchange"):
31
+ raise ValueError(f"unknown neck {neck!r} (have 'light', 'exchange')")
32
+ weights = MobileNet_V3_Small_Weights.DEFAULT if pretrained else None
33
+ backbone = mobilenet_v3_small(weights=weights)
34
+ self.backbone = backbone.features
35
+
36
+ # MobileNet feature indices 3, 8, and 12 are at /8, /16, and /32.
37
+ self.proj3 = Conv(24, 64, 1)
38
+ self.proj4 = Conv(48, 128, 1)
39
+ self.proj5 = Conv(576, 256, 1)
40
+ self.fuse4 = C3k2(384, 128, n=1, c3k=False)
41
+ self.fuse3 = C3k2(192, 64, n=1, c3k=False)
42
+ self.down4 = Conv(64, 64, 3, 2)
43
+ self.pan4 = C3k2(192, 128, n=1, c3k=False)
44
+ self.down5 = Conv(128, 128, 3, 2)
45
+ self.pan5 = C3k2(384, 256, n=1, c3k=False)
46
+
47
+ self.neck = neck
48
+ self.exchange = (
49
+ ScaleExchange(64, 128, 256, gate=exchange_gate, iters=exchange_iters)
50
+ if neck == "exchange" else None
51
+ )
52
+
53
+ detect = Detect(nc=nc, ch=(64, 128, 256))
54
+ detect.stride = torch.tensor([8., 16., 32.])
55
+ detect.bias_init()
56
+ self.model = nn.ModuleList([detect])
57
+ self.nc = nc
58
+ self.img_size = img_size
59
+ self.register_buffer("stride", torch.tensor([8., 16., 32.]))
60
+
61
+ def forward(self, x: torch.Tensor, *, decode: bool | None = None):
62
+ p3 = p4 = None
63
+ for i, layer in enumerate(self.backbone):
64
+ x = layer(x)
65
+ if i == 3:
66
+ p3 = x
67
+ elif i == 8:
68
+ p4 = x
69
+ elif i == 12:
70
+ p5 = x
71
+
72
+ p3, p4, p5 = self.proj3(p3), self.proj4(p4), self.proj5(p5)
73
+ p4_td = self.fuse4(torch.cat((F.interpolate(p5, scale_factor=2,
74
+ mode="nearest"), p4), 1))
75
+ p3_out = self.fuse3(torch.cat((F.interpolate(p4_td, scale_factor=2,
76
+ mode="nearest"), p3), 1))
77
+ p4_out = self.pan4(torch.cat((self.down4(p3_out), p4_td), 1))
78
+ p5_out = self.pan5(torch.cat((self.down5(p4_out), p5), 1))
79
+ features = [p3_out, p4_out, p5_out]
80
+ if self.exchange is not None:
81
+ features = self.exchange(features)
82
+ return self.model[-1](features, decode=decode)
@@ -0,0 +1,293 @@
1
+ """YAML-driven model builder + forward graph for the YOLO11n detector.
2
+
3
+ The YAML in ``configs/yolo11.yaml`` is the single source of truth for the
4
+ architecture. The builder walks it exactly like the official implementation:
5
+
6
+ - each entry is ``[from, repeats, module, args]``
7
+ - ``from`` < 0 -> offset into the built list (-1 = previous layer)
8
+ - ``from`` >= 0 -> absolute layer index; a list = several layers (Concat / Detect)
9
+ - output channels are tracked alongside the built layers
10
+ - width : ``c2 = make_divisible(min(c2, max_channels) * width_mult, 8)``
11
+ - depth : ``repeats = max(round(repeats * depth_mult), 1)``
12
+
13
+ The result is an ``nn.Sequential`` indexed exactly like the official yolo11n
14
+ (0..23), so state_dict keys land at ``model.0.conv.weight``,
15
+ ``model.23.cv3.1.0.1.conv.weight``, ... and Phase 2 can load the official
16
+ ``yolo11n.pt`` 1:1 after re-initialising the head's class branch for nc=20.
17
+ """
18
+ from __future__ import annotations
19
+
20
+ import ast
21
+ import copy
22
+ from pathlib import Path
23
+
24
+ import torch
25
+ from torch import nn
26
+
27
+ from candyeye.core.functions.layer_utils import make_divisible
28
+ from candyeye.core.modules.blocks import Attention, Bottleneck, C2PSA, C3k, C3k2, PSABlock, SPPF
29
+ from candyeye.core.modules.conv import Conv
30
+ from candyeye.core.modules.detect import Detect
31
+ from candyeye.core.modules.exchange import ScaleExchange
32
+ from candyeye.paths import resolve_config
33
+
34
+
35
+ class Concat(nn.Module):
36
+ """Concatenate several tensors along a dimension (neck merges)."""
37
+
38
+ def __init__(self, dimension: int = 1):
39
+ super().__init__()
40
+ self.d = dimension
41
+
42
+ def forward(self, x: list[torch.Tensor]) -> torch.Tensor:
43
+ return torch.cat(x, self.d)
44
+
45
+
46
+ _MODULE_MAP = {
47
+ "Conv": Conv,
48
+ "Bottleneck": Bottleneck,
49
+ "C3k": C3k,
50
+ "C3k2": C3k2,
51
+ "SPPF": SPPF,
52
+ "Attention": Attention,
53
+ "PSABlock": PSABlock,
54
+ "C2PSA": C2PSA,
55
+ "Concat": Concat,
56
+ "Detect": Detect,
57
+ "ScaleExchange": ScaleExchange,
58
+ "nn.Upsample": nn.Upsample,
59
+ "nn.Conv2d": nn.Conv2d,
60
+ }
61
+
62
+ _CHANNEL_MODULES = {"Conv", "C3k2", "C2PSA", "SPPF", "C3k", "Bottleneck", "Attention", "PSABlock"}
63
+
64
+ # depth_mult, width_mult, max_channels
65
+ _SCALES = {"n": (0.50, 0.25, 1024)}
66
+
67
+
68
+ def _source_channels(src: list[int], cur: int | list[int],
69
+ ch: dict[int, int | list[int]]) -> list[int]:
70
+ """Flatten the input-channel spec of a multi-input layer.
71
+
72
+ A source entry may itself be a channel *list* (a previous multi-output
73
+ layer such as ``ScaleExchange``), in which case its channels are spliced in.
74
+ """
75
+ channels: list[int] = []
76
+ for index in src:
77
+ entry = cur if index == -1 else ch[index]
78
+ if isinstance(entry, (list, tuple)):
79
+ channels.extend(entry)
80
+ else:
81
+ channels.append(entry)
82
+ return channels
83
+
84
+
85
+
86
+ def parse_model(data: dict) -> nn.Sequential:
87
+ """Expand the yaml layer table into an indexed nn.Sequential."""
88
+ nc = data["nc"]
89
+ scale = data.get("scale", "n")
90
+ if scale not in _SCALES:
91
+ raise ValueError(f"scale {scale!r} not supported (have {sorted(_SCALES)})")
92
+ depth, width, max_channels = _SCALES[scale]
93
+
94
+ x = data["backbone"] + data["head"]
95
+ layers, ch, cur = [], {}, 3 # ch[i] = out channels of layer i; cur = last output
96
+
97
+ for i, (f, n, m, args) in enumerate(x):
98
+ args = [nc if a == "nc" else a for a in args]
99
+ for j, a in enumerate(args):
100
+ if isinstance(a, str) and a != "nc":
101
+ try:
102
+ args[j] = ast.literal_eval(a) # "None" -> None, keep "nearest"
103
+ except (ValueError, SyntaxError):
104
+ pass
105
+ n = max(round(n * depth), 1)
106
+
107
+ if m in _CHANNEL_MODULES:
108
+ c1, c2 = cur, args[0]
109
+ if c2 != nc:
110
+ c2 = make_divisible(min(c2, max_channels) * width, 8)
111
+ args = [c1, c2, *args[1:]]
112
+ if m == "C3k2":
113
+ args.insert(2, n) # repeats slot inside the block
114
+ n = 1
115
+ elif m == "nn.Upsample":
116
+ c2 = cur
117
+ elif m == "Concat":
118
+ c2 = sum(cur if x == -1 else ch[x] for x in f)
119
+ elif m == "Detect":
120
+ src = f if isinstance(f, (list, tuple)) else [f]
121
+ args.append(_source_channels(src, cur, ch))
122
+ c2 = nc # not referenced downstream; placeholder
123
+ elif m == "ScaleExchange":
124
+ src = f if isinstance(f, (list, tuple)) else [f]
125
+ in_ch = _source_channels(src, cur, ch)
126
+ gate = str(args[0]) if args else "none"
127
+ iters = int(args[1]) if len(args) > 1 else 1
128
+ args = [*in_ch, gate, iters]
129
+ c2 = list(in_ch) # this layer emits a channel *list*
130
+ else:
131
+ raise ValueError(f"unhandled module {m!r}")
132
+
133
+ module_cls = _MODULE_MAP[m]
134
+ if m == "Detect":
135
+ module = module_cls(nc=args[0], ch=args[1])
136
+ elif n > 1:
137
+ module = nn.Sequential(*(copy.deepcopy(module_cls(*args)) for _ in range(n)))
138
+ else:
139
+ module = module_cls(*args)
140
+ module.type = m
141
+ module.i = i # absolute index
142
+ module.f = f # from-list (int or list of ints)
143
+ module.rep = n # repeats used (rep avoids clobbering SPPF/C2PSA self.n)
144
+
145
+ layers.append(module)
146
+ ch[i] = c2
147
+ cur = c2
148
+ return nn.Sequential(*layers)
149
+
150
+
151
+ class CandyEye(nn.Module):
152
+ """CandyEye detector: YAML-built graph and anchor-free Detect head."""
153
+
154
+ def __init__(self, cfg=None, nc: int | None = None, img_size: int = 128):
155
+ super().__init__()
156
+ resolved = cfg if isinstance(cfg, dict) else resolve_config(cfg)
157
+ data = copy.deepcopy(resolved) if isinstance(resolved, dict) else self._load_yaml(resolved)
158
+ if nc is not None:
159
+ data["nc"] = nc
160
+ self.model = parse_model(data)
161
+ self.nc = self.model[-1].nc
162
+ self.yaml = resolved
163
+ self.img_size = img_size
164
+ if img_size % 32:
165
+ raise ValueError(f"img_size must be divisible by 32, got {img_size}")
166
+ self.stride = self._detect_stride(img_size)
167
+ self.model[-1].stride = self.stride
168
+ self.model[-1].bias_init()
169
+
170
+ def set_classes(self, nc: int):
171
+ """Resize the detection class head, preserving all compatible weights."""
172
+ nc = int(nc)
173
+ if nc <= 0:
174
+ raise ValueError(f"nc must be positive, got {nc}")
175
+ old = self.model[-1]
176
+ if nc == old.nc:
177
+ self.nc = nc
178
+ return self
179
+
180
+ channels = tuple(branch[0].conv.in_channels for branch in old.cv2)
181
+ new = Detect(nc=nc, reg_max=old.reg_max, ch=channels).to(
182
+ device=old.stride.device, dtype=next(old.parameters()).dtype)
183
+ new.stride = old.stride.clone()
184
+ new.bias_init()
185
+ for attribute in ("type", "i", "f", "rep"):
186
+ if hasattr(old, attribute):
187
+ setattr(new, attribute, getattr(old, attribute))
188
+ old_state, new_state = old.state_dict(), new.state_dict()
189
+ with torch.no_grad():
190
+ for key, value in new_state.items():
191
+ if key in old_state and old_state[key].shape == value.shape:
192
+ value.copy_(old_state[key])
193
+ self.model[-1] = new
194
+ self.nc = nc
195
+ return self
196
+
197
+ @staticmethod
198
+ def _load_yaml(cfg: str) -> dict:
199
+ import yaml
200
+
201
+ with open(cfg) as fh:
202
+ data = yaml.safe_load(fh)
203
+ CandyEye._validate_yaml(data)
204
+ return data
205
+
206
+ @staticmethod
207
+ def _validate_yaml(data: dict) -> None:
208
+ path = Path(data.get("path", "?"))
209
+ for section in ("backbone", "head"):
210
+ if section not in data:
211
+ raise ValueError(f"yaml {path} missing section {section!r}")
212
+ for entry in data[section]:
213
+ if not (isinstance(entry, list) and len(entry) == 4):
214
+ raise ValueError(f"bad {section} entry {entry!r} (want [from, repeats, module, args])")
215
+ if "nc" not in data:
216
+ raise ValueError(f"yaml {path} missing nc")
217
+
218
+ def _detect_stride(self, img_size: int) -> torch.Tensor:
219
+ """Feed a blank image and infer P3/P4/P5 strides from output sizes."""
220
+ was_training = self.training
221
+ self.eval() # BatchNorm train mode rejects 1x1 feature maps (img_size=32)
222
+ try:
223
+ with torch.no_grad():
224
+ # decode=False: raw per-level maps regardless of train/eval.
225
+ feats = self(torch.zeros(1, 3, img_size, img_size), decode=False)
226
+ finally:
227
+ if was_training:
228
+ self.train()
229
+ if not isinstance(feats, (list, tuple)):
230
+ return torch.ones(1)
231
+ return torch.tensor([img_size / f.shape[-2] for f in feats])
232
+
233
+ def forward(self, x: torch.Tensor, *, decode: bool | None = None) -> torch.Tensor | list[torch.Tensor]:
234
+ """Run the graph; ``decode=False`` returns raw maps even in eval mode."""
235
+ y = [] # layer-output history, index == absolute layer index
236
+ for m in self.model:
237
+ f = m.f
238
+ if isinstance(f, int):
239
+ xi = x if f == -1 else y[f]
240
+ else:
241
+ xi = [x if j == -1 else y[j] for j in f]
242
+ x = m(xi, decode=decode) if isinstance(m, Detect) else m(xi)
243
+ y.append(x)
244
+ return x
245
+
246
+ def train(self, mode: bool = True, *, data=None, epochs: int = 100,
247
+ imgsz: int | None = None, batch: int = 16, patience: int = 50,
248
+ workers: int = 0, device: str = "cpu",
249
+ project: str | Path = "runs/train", name: str = "exp",
250
+ resume: bool | str | Path = False, optimizer: str = "AdamW",
251
+ lr0: float = 2e-4, weight_decay: float = 5e-4,
252
+ warmup_epochs: float = 3, mosaic: float = .5,
253
+ hsv: bool = True, fliplr: float = .5,
254
+ pretrained: bool | str | Path = False, seed: int = 23,
255
+ exist_ok: bool = False, max_batches: int | None = None,
256
+ threads: int = 4):
257
+ """Set PyTorch mode or launch CandyEye training when `data` is set.
258
+
259
+ Example: ``model.train(data="configs/default.yaml", epochs=100,
260
+ imgsz=128, batch=16, patience=20)``. When called without `data`, this
261
+ retains the standard ``nn.Module.train(mode)`` behavior.
262
+ """
263
+ if data is None:
264
+ return super().train(mode)
265
+ from candyeye.training.trainer import train_model
266
+
267
+ return train_model(
268
+ self, data=data, epochs=epochs, imgsz=imgsz or self.img_size,
269
+ batch=batch, patience=patience, workers=workers, device=device,
270
+ project=project, name=name, resume=resume, optimizer=optimizer,
271
+ lr0=lr0, weight_decay=weight_decay, warmup_epochs=warmup_epochs,
272
+ mosaic=mosaic, hsv=hsv, fliplr=fliplr, pretrained=pretrained,
273
+ seed=seed, exist_ok=exist_ok, max_batches=max_batches,
274
+ threads=threads,
275
+ )
276
+
277
+ def fuse(self):
278
+ """Fold every BatchNorm into its conv (in place).
279
+
280
+ Puts the model in eval mode first (BN running stats are the whole
281
+ point of fusing). After this, state_dict keys match the ONNX
282
+ versions of the model (``model.0.conv.weight`` +
283
+ ``model.0.conv.bias``, no ``.bn.*``), and ``load_fused_from_onnx``
284
+ can copy the official fp32 weights 1:1. Numerical parity with the
285
+ official export is then exact, not "fp16 checkpoint drift" close.
286
+ """
287
+ from candyeye.core.modules.conv import Conv
288
+
289
+ self.eval()
290
+ for m in self.modules():
291
+ if isinstance(m, Conv):
292
+ m.fuse()
293
+ return self
@@ -0,0 +1,157 @@
1
+ """Load the official YOLO11n weights (nc=80) into our model.
2
+
3
+ Reads a *clean* state_dict (plain tensors only — no ultralytics dependency).
4
+ By default it uses the wheel's bundled ``assets/yolo11n.pth``; pass an explicit
5
+ path to override. The loader is shape-checked: any tensor whose shape doesn't
6
+ match is skipped and reported, everything else is copied 1:1 (parameters *and*
7
+ buffers such as BatchNorm running statistics, so eval-mode inference is
8
+ faithful).
9
+
10
+ Two load scenarios:
11
+ - ``CandyEye(yaml, nc=80)`` — exact full load: 0 missing / 0 unexpected.
12
+ - ``CandyEye(yaml, nc=20)`` — our VOC training target: the whole class branch
13
+ ``model.23.cv3.*`` is nc/c3-dependent, so exactly ``EXPECTED_NC20_SKIPPED``
14
+ (51 tensors, generated below from the known structure) are skipped and stay
15
+ random-initialised. The box branch (``cv2`` + ``dfl``) still loads 1:1, so
16
+ box proposals remain the official ones.
17
+ """
18
+ from __future__ import annotations
19
+
20
+ from pathlib import Path
21
+
22
+ import torch
23
+
24
+ from candyeye.paths import default_weights_path
25
+
26
+ ONNX_PATH = "tmp/yolo11n.onnx" # official fp32 export (BN fused)
27
+
28
+ # The official checkpoint was trained with nc=80 (COCO), our training model
29
+ # uses nc=20 (VOC). In Detect, c3 = max(ch[0], min(nc, 100)) pins the class
30
+ # branch's hidden width (80 -> 64) and its final Conv2d in/out (-> 20).
31
+ # Per scale that unmatchable part is:
32
+ # cv3.<s>.0.1 Conv(x -> c3) bn.weight/bias/mean/var + conv.weight
33
+ # cv3.<s>.1.0 DWConv(c3 -> c3) bn.* + conv.weight
34
+ # cv3.<s>.1.1 Conv(c3 -> c3) bn.* + conv.weight
35
+ # cv3.<s>.2 Conv2d(c3 -> nc) weight + bias
36
+ # = 17 tensors/scale x 3 scales = 51. Everything else loads unchanged.
37
+ EXPECTED_NC20_SKIPPED = frozenset(
38
+ [
39
+ *(
40
+ f"model.23.cv3.{s}.{a}.{b}.{comp}"
41
+ for s in range(3)
42
+ for (a, b) in ((0, 1), (1, 0), (1, 1))
43
+ for comp in ("bn.bias", "bn.running_mean", "bn.running_var",
44
+ "bn.weight", "conv.weight")
45
+ ),
46
+ *(
47
+ f"model.23.cv3.{s}.2.{comp}"
48
+ for s in range(3)
49
+ for comp in ("weight", "bias")
50
+ ),
51
+ ]
52
+ )
53
+
54
+
55
+ def load_official_state_dict(path: str | Path | None = None) -> dict[str, torch.Tensor]:
56
+ """Read the clean .pth (tensors only -> safe with torch.load defaults).
57
+
58
+ Defaults to the bundled ``assets/yolo11n.pth`` so callers can load the
59
+ default weights without knowing the package location.
60
+ """
61
+ if path is None:
62
+ path = default_weights_path()
63
+ sd = torch.load(path, map_location="cpu")
64
+ if not (isinstance(sd, dict) and all(isinstance(v, torch.Tensor) for v in sd.values())):
65
+ raise TypeError(f"{path} is not a clean state_dict of tensors "
66
+ f"(re-run scripts/bootstrap_weights.py)")
67
+ return sd
68
+
69
+
70
+ def load_weights(
71
+ model: torch.nn.Module,
72
+ official: dict[str, torch.Tensor],
73
+ verbose: bool = False,
74
+ ) -> dict[str, list]:
75
+ """Shape-checked copy of *official* into *model*; nothing in-place original.
76
+
77
+ Returns {"loaded": [...], "skipped": [(key, our_shape, off_shape)],
78
+ "unexpected": [(key, off_shape)]}. Buffers are copied too.
79
+ """
80
+ ours = model.state_dict()
81
+ loaded, skipped, unexpected = [], [], []
82
+
83
+ for key, value in official.items(): # first pass: classify
84
+ if key not in ours:
85
+ unexpected.append((key, tuple(value.shape)))
86
+ elif ours[key].shape == value.shape:
87
+ loaded.append(key)
88
+ else:
89
+ skipped.append((key, tuple(ours[key].shape), tuple(value.shape)))
90
+
91
+ with torch.no_grad(): # in-place copy onto the model's existing tensors
92
+ for key in loaded:
93
+ ours[key].copy_(official[key])
94
+
95
+ if verbose:
96
+ print(f"loaded {len(loaded)} tensors, skipped {len(skipped)}, "
97
+ f"unexpected {len(unexpected)}")
98
+ for key, a, b in skipped:
99
+ print(f" skip {key}: ours {a} vs official {b}")
100
+ return {"loaded": loaded, "skipped": skipped, "unexpected": unexpected}
101
+
102
+
103
+ # With BN fused away, the nc=20 class branch (c3 = min(nc, 100) -> 64 vs the
104
+ # official 80) differs in 4 tensors per scale x 3 scales = 24. Per scale:
105
+ # cv3.<s>.0.0.1 Conv(c3 -> c3 fixed 64<->80) weight + bias
106
+ # cv3.<s>.0.1.0 DWConv(c3 -> c3, depthwise) weight + bias
107
+ # cv3.<s>.0.1.1 Conv(c3 -> c3) weight + bias
108
+ # cv3.<s>.2 Conv2d(c3 -> nc) weight + bias
109
+ EXPECTED_FUSED_NC20_SKIPPED = frozenset(
110
+ [
111
+ *(
112
+ f"model.23.cv3.{s}.{a}.{b}.{comp}"
113
+ for s in range(3)
114
+ for (a, b) in ((0, 1), (1, 0), (1, 1))
115
+ for comp in ("conv.weight", "conv.bias")
116
+ ),
117
+ *(
118
+ f"model.23.cv3.{s}.2.{comp}"
119
+ for s in range(3)
120
+ for comp in ("weight", "bias")
121
+ ),
122
+ ]
123
+ )
124
+
125
+
126
+ def remap_prefix(official: dict[str, torch.Tensor], source_index: int,
127
+ target_index: int) -> dict[str, torch.Tensor]:
128
+ """Shift ``model.<source_index>.*`` keys onto another layer index.
129
+
130
+ Neck variants that add layers after the original Detect (e.g. the
131
+ ``ScaleExchange`` neck) move Detect from layer 23 to a new index, so the
132
+ official head keys must be re-keyed before loading.
133
+ """
134
+ if source_index == target_index:
135
+ return official
136
+ source = f"model.{source_index}."
137
+ target = f"model.{target_index}."
138
+ return {
139
+ (target + key[len(source):] if key.startswith(source) else key): value
140
+ for key, value in official.items()
141
+ }
142
+
143
+
144
+ def load_fused_from_onnx(path: str = ONNX_PATH) -> dict[str, torch.Tensor]:
145
+ """Read the official ONNX graph initializers as a name -> fp32 Tensor dict.
146
+
147
+ The graph's weights are BN-fused, so this only makes sense for a model
148
+ that has been ``CandyEye(...).fuse()``d — the key set then matches 1:1
149
+ (``model.0.conv.weight`` + ``model.0.conv.bias``, no ``.bn.*``).
150
+ """
151
+ import onnx
152
+
153
+ model = onnx.load(path)
154
+ return {
155
+ init.name: torch.tensor(onnx.numpy_helper.to_array(init), dtype=torch.float32)
156
+ for init in model.graph.initializer
157
+ }
File without changes