OSCR

NeuroSeg-MF: robust neuron segmentation in two-photon Ca<sup>2+</sup> imaging using multi-feature fusion and detection-guided SAM.

Code ↔ Paper

11 matches between paragraphs of the paper and lines of its authors' code, computed by the harvester (lexical-v1). Click a colored paragraph or line to see its counterpart.

The 11 matches
  1. [1] § Materials and methods › Framework of NeuroSeg-MF › Multi-feature detection network ↔ ultralytics/nn/tasks.py, lines 1778–1920 · score 0.80 · basic block, GateFusion, A2C2f, RT DETR, AIFI, global
  2. [2] § Materials and methods › Framework of NeuroSeg-MF › Multi-feature detection network ↔ ultralytics/nn/modules/__init__.py, lines 37–71 · score 0.80 · RepConv, basic block, GateFusion, A2C2f, AIFI, decoder
  3. [3] § Materials and methods › Image preprocessing › Data augmentation ↔ ultralytics/data/augment.py, lines 2649–2750 · score 0.76 · horizontal flipping, Random erasing, jitter, augmentation, HSV, probability
  4. [4] § Materials and methods › Framework of NeuroSeg-MF › Multi-feature detection network ↔ ultralytics/nn/modules/__init__.py, lines 37–71 · score 0.74 · RepConv, BasicBlock, GateFusion, A2C2f, AIFI, decoder
  5. [5] § Materials and methods › Framework of NeuroSeg-MF › Detection-guided SAM for neuron segmentation ↔ tools/SAM_box_to_mask.py, lines 1–65 · score 0.74 · candidate masks generated, binary masks, box prompt, segmentation masks, overlays, bounding boxes
  6. [6] § Materials and methods › Framework of NeuroSeg-MF › Multi-feature detection network ↔ ultralytics/nn/tasks.py, lines 1778–1920 · score 0.74 · BasicBlock, GateFusion, A2C2f, RT DETR, AIFI, architecture
  7. [7] § Materials and methods › Framework of NeuroSeg-MF › Detection-guided SAM for neuron segmentation ↔ tools/SAM_box_to_mask.py, lines 1–65 · score 0.64 · generate candidate masks, box prompt, detection box, scores, predicted, neuron
  8. [8] § Materials and methods › Image preprocessing › Pseudo-depth map ↔ tools/generate_depth_twophoton.py, lines 61–145 · score 0.56 · bit grayscale, pseudo depth map
  9. [9] § Materials and methods › Evaluation metrics ↔ ultralytics/utils/coco_metrics.py, lines 282–404 · score 0.51 · IoU threshold, F1 score, recall, precision, metrics, predicted
  10. [10] § Materials and methods › Image preprocessing › Pseudo-depth map ↔ tools/generate_depth_twophoton.py, lines 61–145 · score 0.50 · pseudo depth map, smoothing, V2, filtering, clipping, resized
  11. [11] § Materials and methods › Image preprocessing › Correlation map ↔ tools/generate_corr.py, lines 404–483 · score 0.50 · local neighborhood, correlation map, frames, pixels

Paper

Loaded from Europe PMC by your browser, not stored by OSCR: doi.org · Europe PMC

The paper is loaded when this pane is shown.

The authors' code

Python · 2,035 lines · 82 KB · no license · 2 matches

  1. # Ultralytics 🚀 AGPL-3.0 License - https://ultralytics.com/license
  2. import contextlib
  3. import ast
  4. import pickle
  5. import re
  6. import types
  7. from copy import deepcopy
  8. from pathlib import Path
  9. import torch
  10. import torch.nn as nn
  11. # Extra modules conditional import(defaulttext)
  12. # from ultralytics.nn.extra_modules import *
  13. # from ultralytics.nn.backbone.convnextv2 import *
  14. # from ultralytics.nn.backbone.fasternet import *
  15. # from ultralytics.nn.backbone.efficientViT import *
  16. # from ultralytics.nn.backbone.EfficientFormerV2 import *
  17. # from ultralytics.nn.backbone.VanillaNet import *
  18. # from ultralytics.nn.backbone.revcol import *
  19. # from ultralytics.nn.backbone.lsknet import *
  20. # from ultralytics.nn.backbone.SwinTransformer import *
  21. # from ultralytics.nn.backbone.repvit import *
  22. # from ultralytics.nn.backbone.CSwomTramsformer import *
  23. # from ultralytics.nn.backbone.UniRepLKNet import *
  24. # from ultralytics.nn.backbone.TransNext import *
  25. # from ultralytics.nn.backbone.rmt import *
  26. # from ultralytics.nn.backbone.pkinet import *
  27. # from ultralytics.nn.backbone.mobilenetv4 import *
  28. # from ultralytics.nn.backbone.starnet import *
  29. # from ultralytics.nn.backbone.inceptionnext import *
  30. # from ultralytics.nn.extra_modules.mobileMamba.mobilemamba import *
  31. # from ultralytics.nn.backbone.MambaOut import *
  32. # from ultralytics.nn.backbone.overlock import *
  33. # from ultralytics.nn.backbone.lsnet import *
  34. # except:
  35. # pass
  36. from ultralytics.nn.autobackend import check_class_names
  37. from ultralytics.nn.modules import (
  38. AIFI,
  39. A2C2f,
  40. BasicBlock,
  41. Blocks,
  42. Bottleneck,
  43. C3,
  44. Concat,
  45. Conv,
  46. ConvNormLayer,
  47. DWConv,
  48. GateFusion,
  49. HGBlock,
  50. HGStem,
  51. Index,
  52. RepC3,
  53. RepConv,
  54. RTDETRBottleNeck,
  55. RTDETRDecoder,
  56. get_activation,
  57. )
  58. from ultralytics.utils import DEFAULT_CFG_DICT, DEFAULT_CFG_KEYS, LOGGER, YAML, colorstr, emojis
  59. from ultralytics.utils.checks import check_requirements, check_suffix, check_yaml
  60. from ultralytics.utils.loss import v8DetectionLoss
  61. from ultralytics.utils.ops import make_divisible
  62. from ultralytics.utils.patches import torch_load
  63. from ultralytics.utils.plotting import feature_visualization
  64. from ultralytics.utils.torch_utils import (
  65. fuse_conv_and_bn,
  66. initialize_weights,
  67. intersect_dicts,
  68. model_info,
  69. scale_img,
  70. smart_inference_mode,
  71. time_sync,
  72. )
  73. try:
  74. from ultralytics.nn.mm import MultiModalRouter, MultiModalConfigParser
  75. HookManager = None
  76. MULTIMODAL_AVAILABLE = True
  77. except Exception:
  78. MULTIMODAL_AVAILABLE = False
  79. MultiModalRouter = MultiModalConfigParser = None
  80. HookManager = None
  81. CONTRAST_AVAILABLE = False
  82. DETECT_CLASS: tuple = ()
  83. SEGMENT_CLASS: tuple = ()
  84. POSE_CLASS: tuple = ()
  85. OBB_CLASS: tuple = ()
  86. C3K2_CLASS: tuple = ()
  87. C2PSA_CLASS: tuple = ()
  88. SPPF_CLASS: tuple = ()
  89. NECK_CLASS: tuple = ()
  90. LSCD_AVAILABLE = False
  91. C3K2_EXTRACTION_AVAILABLE = False
  92. SPPF_EXTRACTION_AVAILABLE = False
  93. C2PSA_EXTRACTION_AVAILABLE = False
  94. NECK_EXTRACTION_AVAILABLE = False
  95. class _UnsupportedModule(torch.nn.Module):
  96. """Placeholder for modules removed from the NeuroSeg-MF minimal runtime."""
  97. def __init__(self, *args, **kwargs):
  98. super().__init__()
  99. raise NotImplementedError("This module is not included in the NeuroSeg-MF minimal runtime.")
  100. # Compatibility aliases referenced by unused upstream methods. They are not used by NeuroSeg-MF.
  101. Conv2 = ConvTranspose = DWConvTranspose2d = GhostConv = GhostBottleneck = Focus = BottleneckCSP = _UnsupportedModule
  102. C1 = C2 = C2f = C2fAttn = C2fCIB = C2PSA = C3TR = C3Ghost = C3k2 = C3x = _UnsupportedModule
  103. AConv = ADown = ELAN1 = PSA = SPP = SPPF = SPPELAN = SCDown = _UnsupportedModule
  104. RepNCSPELAN4 = RepVGGDW = ResNetLayer = TorchVision = CBLinear = CBFuse = _UnsupportedModule
  105. Classify = Detect = v8Detect = Segment = Pose = OBB = WorldDetect = YOLOEDetect = YOLOESegment = v10Detect = ImagePoolingAttn = LRPCHead = _UnsupportedModule
  106. MutilScaleEdgeInfoGenetator = ConvEdgeFusion = GetIndexOutput = MCFGatedFusion = _UnsupportedModule
  107. class BaseModel(torch.nn.Module):
  108. """
  109. Base class for all YOLO models in the Ultralytics family.
  110. This class provides common functionality for YOLO models including forward pass handling, model fusion,
  111. information display, and weight loading capabilities.
  112. Attributes:
  113. model (torch.nn.Module): The neural network model.
  114. save (list): List of layer indices to save outputs from.
  115. stride (torch.Tensor): Model stride values.
  116. Methods:
  117. forward: Perform forward pass for training or inference.
  118. predict: Perform inference on input tensor.
  119. fuse: Fuse Conv2d and BatchNorm2d layers for optimization.
  120. info: Print model information.
  121. load: Load weights into the model.
  122. loss: Compute loss for training.
  123. Examples:
  124. Create a BaseModel instance
  125. >>> model = BaseModel()
  126. >>> model.info() # Display model information
  127. """
  128. def forward(self, x, *args, **kwargs):
  129. """
  130. Perform forward pass of the model for either training or inference.
  131. If x is a dict, calculates and returns the loss for training. Otherwise, returns predictions for inference.
  132. Args:
  133. x (torch.Tensor | dict): Input tensor for inference, or dict with image tensor and labels for training.
  134. *args (Any): Variable length argument list.
  135. **kwargs (Any): Arbitrary keyword arguments.
  136. Returns:
  137. (torch.Tensor): Loss if x is a dict (training), or network predictions (inference).
  138. """
  139. if isinstance(x, dict): # for cases of training and validating while training.
  140. return self.loss(x, *args, **kwargs)
  141. return self.predict(x, *args, **kwargs)
  142. def predict(self, x, profile=False, visualize=False, augment=False, embed=None):
  143. """
  144. Perform a forward pass through the network.
  145. Args:
  146. x (torch.Tensor): The input tensor to the model.
  147. profile (bool): Print the computation time of each layer if True.
  148. visualize (bool): Save the feature maps of the model if True.
  149. augment (bool): Augment image during prediction.
  150. embed (list, optional): A list of feature vectors/embeddings to return.
  151. Returns:
  152. (torch.Tensor): The last output of the model.
  153. """
  154. if augment:
  155. return self._predict_augment(x)
  156. return self._predict_once(x, profile, visualize, embed)
  157. def _predict_once(self, x, profile=False, visualize=False, embed=None):
  158. """
  159. Perform a forward pass through the network.
  160. Args:
  161. x (torch.Tensor): The input tensor to the model.
  162. profile (bool): Print the computation time of each layer if True.
  163. visualize (bool): Save the feature maps of the model if True.
  164. embed (list, optional): A list of feature vectors/embeddings to return.
  165. Returns:
  166. (torch.Tensor): The last output of the model.
  167. """
  168. # ===== MULTIMODAL EXTENSION START - multi-modaltext =====
  169. mm_router = None
  170. mm_routing_enabled = False
  171. mm_input_sources = None
  172. # Check if this model has a persistent router (RTDETRDetectionModel)
  173. if hasattr(self, 'mm_router') and self.mm_router is not None:
  174. # Use persistent router from model initialization
  175. mm_router = self.mm_router
  176. mm_routing_enabled, mm_input_sources = mm_router.setup_multimodal_routing(x, profile)
  177. if profile:
  178. LOGGER.info("MultiModal: textrouter")
  179. elif MULTIMODAL_AVAILABLE:
  180. try:
  181. from ultralytics.nn.mm import MultiModalRouter
  182. # Create temporary router for other model types
  183. config_dict = getattr(self, 'yaml', None)
  184. mm_router = MultiModalRouter(config_dict, verbose=profile)
  185. mm_routing_enabled, mm_input_sources = mm_router.setup_multimodal_routing(x, profile)
  186. if profile:
  187. LOGGER.info("MultiModal: textrouter")
  188. except Exception as e:
  189. if profile:
  190. LOGGER.warning(f"MultiModal routing initialization failed: {e}")
  191. # ===== MULTIMODAL EXTENSION END =====
  192. y, dt, embeddings = [], [], [] # outputs
  193. embed = frozenset(embed) if embed is not None else {-1}
  194. max_idx = max(embed)
  195. for m in self.model:
  196. if m.f != -1: # if not from previous layer
  197. x = y[m.f] if isinstance(m.f, int) else [x if j == -1 else y[j] for j in m.f] # from earlier layers
  198. # ===== MULTIMODAL EXTENSION START - multi-modaltext =====
  199. # Apply multimodal routing if enabled and module has MM attributes
  200. if mm_routing_enabled and mm_input_sources and mm_router:
  201. routed_x = mm_router.route_layer_input(x, m, mm_input_sources, profile)
  202. if routed_x is not None:
  203. x = routed_x
  204. # Check for spatial reset requirement
  205. if mm_router and hasattr(m, '_mm_spatial_reset') and m._mm_spatial_reset:
  206. x = mm_router.reset_spatial_input(x, m, mm_input_sources, profile)
  207. # ===== MULTIMODAL EXTENSION END =====
  208. if profile:
  209. self._profile_one_layer(m, x, dt)
  210. x = m(x) # run
  211. y.append(x if m.i in self.save else None) # save output
  212. if visualize:
  213. feature_visualization(x, m.type, m.i, save_dir=visualize)
  214. if m.i in embed:
  215. embeddings.append(torch.nn.functional.adaptive_avg_pool2d(x, (1, 1)).squeeze(-1).squeeze(-1)) # flatten
  216. if m.i == max_idx:
  217. return torch.unbind(torch.cat(embeddings, 1), dim=0)
  218. return x
  219. def _predict_augment(self, x):
  220. """Perform augmentations on input image x and return augmented inference."""
  221. LOGGER.warning(
  222. f"{self.__class__.__name__} does not support 'augment=True' prediction. "
  223. f"Reverting to single-scale prediction."
  224. )
  225. return self._predict_once(x)
  226. def _profile_one_layer(self, m, x, dt):
  227. """
  228. Profile the computation time and FLOPs of a single layer of the model on a given input.
  229. Args:
  230. m (torch.nn.Module): The layer to be profiled.
  231. x (torch.Tensor): The input data to the layer.
  232. dt (list): A list to store the computation time of the layer.
  233. """
  234. try:
  235. import thop
  236. except ImportError:
  237. thop = None # conda support without 'ultralytics-thop' installed
  238. c = m == self.model[-1] and isinstance(x, list) # is final layer list, copy input as inplace fix
  239. flops = thop.profile(m, inputs=[x.copy() if c else x], verbose=False)[0] / 1e9 * 2 if thop else 0 # GFLOPs
  240. t = time_sync()
  241. for _ in range(10):
  242. m(x.copy() if c else x)
  243. dt.append((time_sync() - t) * 100)
  244. if m == self.model[0]:
  245. LOGGER.info(f"{'time (ms)':>10s} {'GFLOPs':>10s} {'params':>10s} module")
  246. LOGGER.info(f"{dt[-1]:10.2f} {flops:10.2f} {m.np:10.0f} {m.type}")
  247. if c:
  248. LOGGER.info(f"{sum(dt):10.2f} {'-':>10s} {'-':>10s} Total")
  249. def fuse(self, verbose=True):
  250. """
  251. Fuse the `Conv2d()` and `BatchNorm2d()` layers of the model into a single layer for improved computation
  252. efficiency.
  253. Returns:
  254. (torch.nn.Module): The fused model is returned.
  255. """
  256. if not self.is_fused():
  257. for m in self.model.modules():
  258. if isinstance(m, (Conv, DWConv)) and hasattr(m, "bn"):
  259. m.conv = fuse_conv_and_bn(m.conv, m.bn) # update conv
  260. delattr(m, "bn") # remove batchnorm
  261. m.forward = m.forward_fuse # update forward
  262. if isinstance(m, RepConv):
  263. m.fuse_convs()
  264. m.forward = m.forward_fuse # update forward
  265. self.info(verbose=verbose)
  266. return self
  267. def is_fused(self, thresh=10):
  268. """
  269. Check if the model has less than a certain threshold of BatchNorm layers.
  270. Args:
  271. thresh (int, optional): The threshold number of BatchNorm layers.
  272. Returns:
  273. (bool): True if the number of BatchNorm layers in the model is less than the threshold, False otherwise.
  274. """
  275. bn = tuple(v for k, v in torch.nn.__dict__.items() if "Norm" in k) # normalization layers, i.e. BatchNorm2d()
  276. return sum(isinstance(v, bn) for v in self.modules()) < thresh # True if < 'thresh' BatchNorm layers in model
  277. def info(self, detailed=False, verbose=True, imgsz=640):
  278. """
  279. Print model information.
  280. Args:
  281. detailed (bool): If True, prints out detailed information about the model.
  282. verbose (bool): If True, prints out the model information.
  283. imgsz (int): The size of the image that the model will be trained on.
  284. """
  285. return model_info(self, detailed=detailed, verbose=verbose, imgsz=imgsz)
  286. def _apply(self, fn):
  287. """
  288. Apply a function to all tensors in the model that are not parameters or registered buffers.
  289. Args:
  290. fn (function): The function to apply to the model.
  291. Returns:
  292. (BaseModel): An updated BaseModel object.
  293. """
  294. self = super()._apply(fn)
  295. m = self.model[-1] # Detect()/Segment()/Pose()/OBB()
  296. heads = DETECT_CLASS + SEGMENT_CLASS + POSE_CLASS + OBB_CLASS
  297. if isinstance(m, heads):
  298. m.stride = fn(m.stride)
  299. m.anchors = fn(m.anchors)
  300. m.strides = fn(m.strides)
  301. return self
  302. def load(self, weights, verbose=True):
  303. """
  304. Load weights into the model.
  305. Args:
  306. weights (dict | torch.nn.Module): The pre-trained weights to be loaded.
  307. verbose (bool, optional): Whether to log the transfer progress.
  308. """
  309. model = weights["model"] if isinstance(weights, dict) else weights # torchvision models are not dicts
  310. csd = model.float().state_dict() # checkpoint state_dict as FP32
  311. updated_csd = intersect_dicts(csd, self.state_dict()) # intersect
  312. self.load_state_dict(updated_csd, strict=False) # load
  313. len_updated_csd = len(updated_csd)
  314. first_conv = "model.0.conv.weight" # hard-coded to yolo models for now
  315. # mostly used to boost multi-channel training
  316. state_dict = self.state_dict()
  317. if first_conv not in updated_csd and first_conv in state_dict:
  318. c1, c2, h, w = state_dict[first_conv].shape
  319. cc1, cc2, ch, cw = csd[first_conv].shape
  320. if ch == h and cw == w:
  321. c1, c2 = min(c1, cc1), min(c2, cc2)
  322. state_dict[first_conv][:c1, :c2] = csd[first_conv][:c1, :c2]
  323. len_updated_csd += 1
  324. if verbose:
  325. LOGGER.info(f"Transferred {len_updated_csd}/{len(self.model.state_dict())} items from pretrained weights")
  326. def loss(self, batch, preds=None):
  327. """
  328. Compute loss.
  329. Args:
  330. batch (dict): Batch to compute loss on.
  331. preds (torch.Tensor | List[torch.Tensor], optional): Predictions.
  332. """
  333. if getattr(self, "criterion", None) is None:
  334. self.criterion = self.init_criterion()
  335. preds = self.forward(batch["img"]) if preds is None else preds
  336. det_loss_vec, det_items = self.criterion(preds, batch)
  337. # Optional contrastive branch (enabled only if hooks exist and compute succeeds)
  338. # Requirements: MULTIMODAL + CONTRAST + registered hooks via 6th field
  339. mm_hm = getattr(self, 'mm_hook_manager', None)
  340. has_hooks = False
  341. if mm_hm is not None:
  342. # Only enable contrast branch when YAML registered hooks (6th-field) exist
  343. try:
  344. has_hooks = bool(mm_hm.has_hooks())
  345. except Exception:
  346. has_hooks = False
  347. use_contrast = (
  348. self.training
  349. and MULTIMODAL_AVAILABLE
  350. and CONTRAST_AVAILABLE
  351. and has_hooks
  352. )
  353. if not use_contrast:
  354. return det_loss_vec, det_items
  355. # Read configs from args with safe defaults
  356. args = getattr(self, 'args', None)
  357. cfg = ContrastConfig(
  358. tau=getattr(args, 'contrast_tau', 0.07) if args is not None else 0.07,
  359. proj_dim=getattr(args, 'contrast_dim', 128) if args is not None else 128,
  360. lambda_weight=getattr(args, 'contrast_lambda', 0.1) if args is not None else 0.1,
  361. max_rois_per_image=getattr(args, 'contrast_max_rois', 64) if args is not None else 64,
  362. share_head=getattr(args, 'contrast_share_head', False) if args is not None else False,
  363. preferred_stages=tuple(getattr(args, 'contrast_stages', ("P4", "P5", "P3"))) if args is not None else ("P4", "P5", "P3"),
  364. )
  365. # Lazy create controller (only when hooks exist)
  366. if getattr(self, 'mm_contrast_controller', None) is None and use_contrast:
  367. # Create on correct device to avoid CPU/CUDA mismatch when forward() is first called
  368. try:
  369. dev = next(self.parameters()).device
  370. except StopIteration:
  371. # Fallback to batch image device if model has no parameters (unlikely)
  372. img = batch.get('img')
  373. dev = img.device if isinstance(img, torch.Tensor) else torch.device('cpu')
  374. self.mm_contrast_controller = ContrastController(cfg).to(dev)
  375. # Collect hooked features (do not pop to allow external visualization; keep bounded by latest write)
  376. hook_buffers = mm_hm.collect(pop=False)
  377. loss_c, stats = self.mm_contrast_controller(hook_buffers, batch)
  378. # Debug: detect non-finite contrastive loss (when enabled and computed)
  379. try:
  380. if loss_c is not None and not torch.isfinite(loss_c):
  381. from ultralytics.utils import LOGGER as _LOGGER
  382. _LOGGER.warning(f"[CL][loss] non-finite loss_c detected: {float(loss_c.detach().cpu())}")
  383. except Exception:
  384. pass
  385. if loss_c is None:
  386. return det_loss_vec, det_items # no valid pairs this step
  387. # Compose outputs: scale contrast by lambda for backprop, but report raw value in items
  388. lambda_w = cfg.lambda_weight
  389. if det_loss_vec.dim() == 0:
  390. # safety: ensure vector form
  391. det_loss_vec = det_loss_vec.unsqueeze(0)
  392. total_vec = torch.cat([det_loss_vec, (loss_c * lambda_w).unsqueeze(0)], dim=0)
  393. total_items = torch.cat([det_items, loss_c.detach().unsqueeze(0)], dim=0)
  394. # Optionally, expose stats via side-effect for loggers (trainer can read from model)
  395. self._contrast_last_stats = stats
  396. return total_vec, total_items
  397. def init_criterion(self):
  398. """Initialize the loss criterion for the BaseModel."""
  399. raise NotImplementedError("compute_loss() needs to be implemented by task heads")
  400. class DetectionModel(BaseModel):
  401. """
  402. YOLO detection model.
  403. This class implements the YOLO detection architecture, handling model initialization, forward pass,
  404. augmented inference, and loss computation for object detection tasks.
  405. Attributes:
  406. yaml (dict): Model configuration dictionary.
  407. model (torch.nn.Sequential): The neural network model.
  408. save (list): List of layer indices to save outputs from.
  409. names (dict): Class names dictionary.
  410. inplace (bool): Whether to use inplace operations.
  411. end2end (bool): Whether the model uses end-to-end detection.
  412. stride (torch.Tensor): Model stride values.
  413. Methods:
  414. __init__: Initialize the YOLO detection model.
  415. _predict_augment: Perform augmented inference.
  416. _descale_pred: De-scale predictions following augmented inference.
  417. _clip_augmented: Clip YOLO augmented inference tails.
  418. init_criterion: Initialize the loss criterion.
  419. Examples:
  420. Initialize a detection model
  421. >>> model = DetectionModel("yolo11n.yaml", ch=3, nc=80)
  422. >>> results = model.predict(image_tensor)
  423. """
  424. def __init__(self, cfg="yolo11n.yaml", ch=3, nc=None, verbose=True):
  425. """
  426. Initialize the YOLO detection model with the given config and parameters.
  427. Args:
  428. cfg (str | dict): Model configuration file path or dictionary.
  429. ch (int): Number of input channels.
  430. nc (int, optional): Number of classes.
  431. verbose (bool): Whether to display model information.
  432. """
  433. super().__init__()
  434. self.yaml = cfg if isinstance(cfg, dict) else yaml_model_load(cfg) # cfg dict
  435. if self.yaml["backbone"][0][2] == "Silence":
  436. LOGGER.warning(
  437. "YOLOv9 `Silence` module is deprecated in favor of torch.nn.Identity. "
  438. "Please delete local *.pt file and re-download the latest model checkpoint."
  439. )
  440. self.yaml["backbone"][0][2] = "nn.Identity"
  441. # Define model
  442. self.yaml["channels"] = ch # save channels
  443. if nc and nc != self.yaml["nc"]:
  444. LOGGER.info(f"Overriding model.yaml nc={self.yaml['nc']} with nc={nc}")
  445. self.yaml["nc"] = nc # override YAML value
  446. self.model, self.save = parse_model(deepcopy(self.yaml), ch=ch, verbose=verbose) # model, savelist
  447. self.names = {i: f"{i}" for i in range(self.yaml["nc"])} # default names dict
  448. self.inplace = self.yaml.get("inplace", True)
  449. self.end2end = getattr(self.model[-1], "end2end", False)
  450. # textmulti-modaltext(textexists)
  451. if hasattr(self.model, 'multimodal_router'):
  452. self.multimodal_router = self.model.multimodal_router
  453. else:
  454. self.multimodal_router = None
  455. # Persist router for runtime ablation/filling so BaseModel forward can reuse it
  456. self.mm_router = self.multimodal_router if self.multimodal_router is not None else None
  457. # textcreate: textinfo()textLazytext
  458. # actualcreatetexttrainingtextget_modeltext(textBaseModel.loss).
  459. # textHookManager(textexists)
  460. self.mm_hook_manager = getattr(self.model, 'mm_hook_manager', None)
  461. # Build strides
  462. m = self.model[-1] # Detect()/Segment()/Pose()/OBB()/...
  463. heads = DETECT_CLASS + SEGMENT_CLASS + POSE_CLASS + OBB_CLASS
  464. if isinstance(m, heads): # includes all Detect/Segment/Pose/OBB subclasses (e.g., LSCD variants)
  465. s = 256 # 2x min stride
  466. m.inplace = self.inplace
  467. def _forward(x):
  468. """Perform a forward pass through the model, handling different Detect subclass types accordingly."""
  469. if self.end2end:
  470. return self.forward(x)["one2many"]
  471. seg_pose_obb = SEGMENT_CLASS + POSE_CLASS + OBB_CLASS
  472. return self.forward(x)[0] if isinstance(m, seg_pose_obb) else self.forward(x)
  473. self.model.eval() # Avoid changing batch statistics until training begins
  474. m.training = True # Setting it to True to properly return strides
  475. m.stride = torch.tensor([s / x.shape[-2] for x in _forward(torch.zeros(1, ch, s, s))]) # forward
  476. self.stride = m.stride
  477. self.model.train() # Set model back to training(default) mode
  478. # textdetectiontext(text)
  479. if hasattr(m, "bias_init") and callable(getattr(m, "bias_init")):
  480. m.bias_init() # only run once
  481. else:
  482. self.stride = torch.Tensor([32]) # default stride for i.e. RTDETR
  483. # Init weights, biases
  484. initialize_weights(self)
  485. if verbose:
  486. self.info()
  487. LOGGER.info("")
  488. def _predict_augment(self, x):
  489. """
  490. Perform augmentations on input image x and return augmented inference and train outputs.
  491. Args:
  492. x (torch.Tensor): Input image tensor.
  493. Returns:
  494. (torch.Tensor): Augmented inference output.
  495. """
  496. if getattr(self, "end2end", False) or self.__class__.__name__ != "DetectionModel":
  497. LOGGER.warning("Model does not support 'augment=True', reverting to single-scale prediction.")
  498. return self._predict_once(x)
  499. img_size = x.shape[-2:] # height, width
  500. s = [1, 0.83, 0.67] # scales
  501. f = [None, 3, None] # flips (2-ud, 3-lr)
  502. y = [] # outputs
  503. for si, fi in zip(s, f):
  504. xi = scale_img(x.flip(fi) if fi else x, si, gs=int(self.stride.max()))
  505. yi = super().predict(xi)[0] # forward
  506. yi = self._descale_pred(yi, fi, si, img_size)
  507. y.append(yi)
  508. y = self._clip_augmented(y) # clip augmented tails
  509. return torch.cat(y, -1), None # augmented inference, train
  510. @staticmethod
  511. def _descale_pred(p, flips, scale, img_size, dim=1):
  512. """
  513. De-scale predictions following augmented inference (inverse operation).
  514. Args:
  515. p (torch.Tensor): Predictions tensor.
  516. flips (int): Flip type (0=none, 2=ud, 3=lr).
  517. scale (float): Scale factor.
  518. img_size (tuple): Original image size (height, width).
  519. dim (int): Dimension to split at.
  520. Returns:
  521. (torch.Tensor): De-scaled predictions.
  522. """
  523. p[:, :4] /= scale # de-scale
  524. x, y, wh, cls = p.split((1, 1, 2, p.shape[dim] - 4), dim)
  525. if flips == 2:
  526. y = img_size[0] - y # de-flip ud
  527. elif flips == 3:
  528. x = img_size[1] - x # de-flip lr
  529. return torch.cat((x, y, wh, cls), dim)
  530. def _clip_augmented(self, y):
  531. """
  532. Clip YOLO augmented inference tails.
  533. Args:
  534. y (List[torch.Tensor]): List of detection tensors.
  535. Returns:
  536. (List[torch.Tensor]): Clipped detection tensors.
  537. """
  538. nl = self.model[-1].nl # number of detection layers (P3-P5)
  539. g = sum(4**x for x in range(nl)) # grid points
  540. e = 1 # exclude layer count
  541. i = (y[0].shape[-1] // g) * sum(4**x for x in range(e)) # indices
  542. y[0] = y[0][..., :-i] # large
  543. i = (y[-1].shape[-1] // g) * sum(4 ** (nl - 1 - x) for x in range(e)) # indices
  544. y[-1] = y[-1][..., i:] # small
  545. return y
  546. def init_criterion(self):
  547. """Initialize the loss criterion for the DetectionModel."""
  548. return E2EDetectLoss(self) if getattr(self, "end2end", False) else v8DetectionLoss(self)
  549. class OBBModel(DetectionModel):
  550. """
  551. YOLO Oriented Bounding Box (OBB) model.
  552. This class extends DetectionModel to handle oriented bounding box detection tasks, providing specialized
  553. loss computation for rotated object detection.
  554. Methods:
  555. __init__: Initialize YOLO OBB model.
  556. init_criterion: Initialize the loss criterion for OBB detection.
  557. Examples:
  558. Initialize an OBB model
  559. >>> model = OBBModel("yolo11n-obb.yaml", ch=3, nc=80)
  560. >>> results = model.predict(image_tensor)
  561. """
  562. def __init__(self, cfg="yolo11n-obb.yaml", ch=3, nc=None, verbose=True):
  563. """
  564. Initialize YOLO OBB model with given config and parameters.
  565. Args:
  566. cfg (str | dict): Model configuration file path or dictionary.
  567. ch (int): Number of input channels.
  568. nc (int, optional): Number of classes.
  569. verbose (bool): Whether to display model information.
  570. """
  571. super().__init__(cfg=cfg, ch=ch, nc=nc, verbose=verbose)
  572. def init_criterion(self):
  573. """Initialize the loss criterion for the model."""
  574. return v8OBBLoss(self)
  575. class SegmentationModel(DetectionModel):
  576. """
  577. YOLO segmentation model.
  578. This class extends DetectionModel to handle instance segmentation tasks, providing specialized
  579. loss computation for pixel-level object detection and segmentation.
  580. Methods:
  581. __init__: Initialize YOLO segmentation model.
  582. init_criterion: Initialize the loss criterion for segmentation.
  583. Examples:
  584. Initialize a segmentation model
  585. >>> model = SegmentationModel("yolo11n-seg.yaml", ch=3, nc=80)
  586. >>> results = model.predict(image_tensor)
  587. """
  588. def __init__(self, cfg="yolo11n-seg.yaml", ch=3, nc=None, verbose=True):
  589. """
  590. Initialize Ultralytics YOLO segmentation model with given config and parameters.
  591. Args:
  592. cfg (str | dict): Model configuration file path or dictionary.
  593. ch (int): Number of input channels.
  594. nc (int, optional): Number of classes.
  595. verbose (bool): Whether to display model information.
  596. """
  597. super().__init__(cfg=cfg, ch=ch, nc=nc, verbose=verbose)
  598. def init_criterion(self):
  599. """Initialize the loss criterion for the SegmentationModel."""
  600. return v8SegmentationLoss(self)
  601. class PoseModel(DetectionModel):
  602. """
  603. YOLO pose model.
  604. This class extends DetectionModel to handle human pose estimation tasks, providing specialized
  605. loss computation for keypoint detection and pose estimation.
  606. Attributes:
  607. kpt_shape (tuple): Shape of keypoints data (num_keypoints, num_dimensions).
  608. Methods:
  609. __init__: Initialize YOLO pose model.
  610. init_criterion: Initialize the loss criterion for pose estimation.
  611. Examples:
  612. Initialize a pose model
  613. >>> model = PoseModel("yolo11n-pose.yaml", ch=3, nc=1, data_kpt_shape=(17, 3))
  614. >>> results = model.predict(image_tensor)
  615. """
  616. def __init__(self, cfg="yolo11n-pose.yaml", ch=3, nc=None, data_kpt_shape=(None, None), verbose=True):
  617. """
  618. Initialize Ultralytics YOLO Pose model.
  619. Args:
  620. cfg (str | dict): Model configuration file path or dictionary.
  621. ch (int): Number of input channels.
  622. nc (int, optional): Number of classes.
  623. data_kpt_shape (tuple): Shape of keypoints data.
  624. verbose (bool): Whether to display model information.
  625. """
  626. if not isinstance(cfg, dict):
  627. cfg = yaml_model_load(cfg) # load model YAML
  628. if any(data_kpt_shape) and list(data_kpt_shape) != list(cfg["kpt_shape"]):
  629. LOGGER.info(f"Overriding model.yaml kpt_shape={cfg['kpt_shape']} with kpt_shape={data_kpt_shape}")
  630. cfg["kpt_shape"] = data_kpt_shape
  631. super().__init__(cfg=cfg, ch=ch, nc=nc, verbose=verbose)
  632. def init_criterion(self):
  633. """Initialize the loss criterion for the PoseModel."""
  634. return v8PoseLoss(self)
  635. class ClassificationModel(BaseModel):
  636. """
  637. YOLO classification model.
  638. This class implements the YOLO classification architecture for image classification tasks,
  639. providing model initialization, configuration, and output reshaping capabilities.
  640. Attributes:
  641. yaml (dict): Model configuration dictionary.
  642. model (torch.nn.Sequential): The neural network model.
  643. stride (torch.Tensor): Model stride values.
  644. names (dict): Class names dictionary.
  645. Methods:
  646. __init__: Initialize ClassificationModel.
  647. _from_yaml: Set model configurations and define architecture.
  648. reshape_outputs: Update model to specified class count.
  649. init_criterion: Initialize the loss criterion.
  650. Examples:
  651. Initialize a classification model
  652. >>> model = ClassificationModel("yolo11n-cls.yaml", ch=3, nc=1000)
  653. >>> results = model.predict(image_tensor)
  654. """
  655. def __init__(self, cfg="yolo11n-cls.yaml", ch=3, nc=None, verbose=True):
  656. """
  657. Initialize ClassificationModel with YAML, channels, number of classes, verbose flag.
  658. Args:
  659. cfg (str | dict): Model configuration file path or dictionary.
  660. ch (int): Number of input channels.
  661. nc (int, optional): Number of classes.
  662. verbose (bool): Whether to display model information.
  663. """
  664. super().__init__()
  665. self._from_yaml(cfg, ch, nc, verbose)
  666. def _from_yaml(self, cfg, ch, nc, verbose):
  667. """
  668. Set Ultralytics YOLO model configurations and define the model architecture.
  669. Args:
  670. cfg (str | dict): Model configuration file path or dictionary.
  671. ch (int): Number of input channels.
  672. nc (int, optional): Number of classes.
  673. verbose (bool): Whether to display model information.
  674. """
  675. self.yaml = cfg if isinstance(cfg, dict) else yaml_model_load(cfg) # cfg dict
  676. # Define model
  677. ch = self.yaml["channels"] = self.yaml.get("channels", ch) # input channels
  678. if nc and nc != self.yaml["nc"]:
  679. LOGGER.info(f"Overriding model.yaml nc={self.yaml['nc']} with nc={nc}")
  680. self.yaml["nc"] = nc # override YAML value
  681. elif not nc and not self.yaml.get("nc", None):
  682. raise ValueError("nc not specified. Must specify nc in model.yaml or function arguments.")
  683. self.model, self.save = parse_model(deepcopy(self.yaml), ch=ch, verbose=verbose) # model, savelist
  684. self.stride = torch.Tensor([1]) # no stride constraints
  685. self.names = {i: f"{i}" for i in range(self.yaml["nc"])} # default names dict
  686. self.info()
  687. @staticmethod
  688. def reshape_outputs(model, nc):
  689. """
  690. Update a TorchVision classification model to class count 'n' if required.
  691. Args:
  692. model (torch.nn.Module): Model to update.
  693. nc (int): New number of classes.
  694. """
  695. name, m = list((model.model if hasattr(model, "model") else model).named_children())[-1] # last module
  696. if isinstance(m, Classify): # YOLO Classify() head
  697. if m.linear.out_features != nc:
  698. m.linear = torch.nn.Linear(m.linear.in_features, nc)
  699. elif isinstance(m, torch.nn.Linear): # ResNet, EfficientNet
  700. if m.out_features != nc:
  701. setattr(model, name, torch.nn.Linear(m.in_features, nc))
  702. elif isinstance(m, torch.nn.Sequential):
  703. types = [type(x) for x in m]
  704. if torch.nn.Linear in types:
  705. i = len(types) - 1 - types[::-1].index(torch.nn.Linear) # last torch.nn.Linear index
  706. if m[i].out_features != nc:
  707. m[i] = torch.nn.Linear(m[i].in_features, nc)
  708. elif torch.nn.Conv2d in types:
  709. i = len(types) - 1 - types[::-1].index(torch.nn.Conv2d) # last torch.nn.Conv2d index
  710. if m[i].out_channels != nc:
  711. m[i] = torch.nn.Conv2d(
  712. m[i].in_channels, nc, m[i].kernel_size, m[i].stride, bias=m[i].bias is not None
  713. )
  714. def init_criterion(self):
  715. """Initialize the loss criterion for the ClassificationModel."""
  716. return v8ClassificationLoss()
  717. class RTDETRDetectionModel(DetectionModel):
  718. """
  719. RTDETR (Real-time DEtection and Tracking using Transformers) Detection Model class.
  720. This class is responsible for constructing the RTDETR architecture, defining loss functions, and facilitating both
  721. the training and inference processes. RTDETR is an object detection and tracking model that extends from the
  722. DetectionModel base class.
  723. Attributes:
  724. nc (int): Number of classes for detection.
  725. criterion (RTDETRDetectionLoss): Loss function for training.
  726. Methods:
  727. __init__: Initialize the RTDETRDetectionModel.
  728. init_criterion: Initialize the loss criterion.
  729. loss: Compute loss for training.
  730. predict: Perform forward pass through the model.
  731. Examples:
  732. Initialize an RTDETR model
  733. >>> model = RTDETRDetectionModel("rtdetr-l.yaml", ch=3, nc=80)
  734. >>> results = model.predict(image_tensor)
  735. """
  736. def __init__(self, cfg="rtdetr-l.yaml", ch=3, nc=None, verbose=True):
  737. """
  738. Initialize the RTDETRDetectionModel.
  739. Args:
  740. cfg (str | dict): Configuration file name or path.
  741. ch (int): Number of input channels.
  742. nc (int, optional): Number of classes.
  743. verbose (bool): Print additional information during initialization.
  744. """
  745. super().__init__(cfg=cfg, ch=ch, nc=nc, verbose=verbose)
  746. # ===== MULTIMODAL EXTENSION START - textMultiModalRouter =====
  747. self.mm_router = None
  748. self.mm_routing_enabled = False
  749. self.mm_input_sources = None
  750. if MULTIMODAL_AVAILABLE:
  751. try:
  752. from ultralytics.nn.mm import MultiModalRouter
  753. # Create persistent router with model configuration
  754. config_dict = getattr(self, 'yaml', None)
  755. self.mm_router = MultiModalRouter(config_dict, verbose=verbose)
  756. if verbose:
  757. LOGGER.info("RTDETRDetectionModel: textMultiModalRoutercreate")
  758. except Exception as e:
  759. if verbose:
  760. LOGGER.warning(f"RTDETRDetectionModel: MultiModalRoutertextfailed: {e}")
  761. # ===== MULTIMODAL EXTENSION END =====
  762. def init_criterion(self):
  763. """Initialize the loss criterion for the RTDETRDetectionModel."""
  764. from ultralytics.models.utils.loss import RTDETRDetectionLoss
  765. return RTDETRDetectionLoss(nc=self.nc, use_vfl=True)
  766. def loss(self, batch, preds=None):
  767. """
  768. Compute the loss for the given batch of data.
  769. Args:
  770. batch (dict): Dictionary containing image and label data.
  771. preds (torch.Tensor, optional): Precomputed model predictions.
  772. Returns:
  773. loss_sum (torch.Tensor): Total loss value.
  774. loss_items (torch.Tensor): Main three losses in a tensor.
  775. """
  776. if not hasattr(self, "criterion"):
  777. self.criterion = self.init_criterion()
  778. img = batch["img"]
  779. # NOTE: preprocess gt_bbox and gt_labels to list.
  780. bs = len(img)
  781. batch_idx = batch["batch_idx"]
  782. gt_groups = [(batch_idx == i).sum().item() for i in range(bs)]
  783. targets = {
  784. "cls": batch["cls"].to(img.device, dtype=torch.long).view(-1),
  785. "bboxes": batch["bboxes"].to(device=img.device),
  786. "batch_idx": batch_idx.to(img.device, dtype=torch.long).view(-1),
  787. "gt_groups": gt_groups,
  788. }
  789. preds = self.predict(img, batch=targets) if preds is None else preds
  790. dec_bboxes, dec_scores, enc_bboxes, enc_scores, dn_meta = preds if self.training else preds[1]
  791. if dn_meta is None:
  792. dn_bboxes, dn_scores = None, None
  793. else:
  794. dn_bboxes, dec_bboxes = torch.split(dec_bboxes, dn_meta["dn_num_split"], dim=2)
  795. dn_scores, dec_scores = torch.split(dec_scores, dn_meta["dn_num_split"], dim=2)
  796. dec_bboxes = torch.cat([enc_bboxes.unsqueeze(0), dec_bboxes]) # (7, bs, 300, 4)
  797. dec_scores = torch.cat([enc_scores.unsqueeze(0), dec_scores])
  798. loss = self.criterion(
  799. (dec_bboxes, dec_scores), targets, dn_bboxes=dn_bboxes, dn_scores=dn_scores, dn_meta=dn_meta
  800. )
  801. # NOTE: There are like 12 losses in RTDETR, backward with all losses but only show the main three losses.
  802. return sum(loss.values()), torch.as_tensor(
  803. [loss[k].detach() for k in ["loss_giou", "loss_class", "loss_bbox"]], device=img.device
  804. )
  805. def predict(self, x, profile=False, visualize=False, batch=None, augment=False, embed=None):
  806. """
  807. Perform a forward pass through the model.
  808. Args:
  809. x (torch.Tensor): The input tensor.
  810. profile (bool): If True, profile the computation time for each layer.
  811. visualize (bool): If True, save feature maps for visualization.
  812. batch (dict, optional): Ground truth data for evaluation.
  813. augment (bool): If True, perform data augmentation during inference.
  814. embed (list, optional): A list of feature vectors/embeddings to return.
  815. Returns:
  816. (torch.Tensor): Model's output tensor.
  817. """
  818. # ===== MULTIMODAL EXTENSION START - multi-modaltext =====
  819. mm_router = self.mm_router
  820. mm_routing_enabled, mm_input_sources = (mm_router.setup_multimodal_routing(x, profile) if mm_router is not None else (False, None))
  821. # ===== MULTIMODAL EXTENSION END =====
  822. y, dt, embeddings = [], [], [] # outputs
  823. embed = frozenset(embed) if embed is not None else {-1}
  824. max_idx = max(embed)
  825. for m in self.model[:-1]: # except the head part
  826. if m.f != -1: # if not from previous layer
  827. x = y[m.f] if isinstance(m.f, int) else [x if j == -1 else y[j] for j in m.f] # from earlier layers
  828. # ===== MULTIMODAL EXTENSION START - multi-modaltext =====
  829. # Apply multimodal routing if enabled and module has MM attributes
  830. if mm_routing_enabled and mm_input_sources and mm_router:
  831. routed_x = mm_router.route_layer_input(x, m, mm_input_sources, profile)
  832. if routed_x is not None:
  833. x = routed_x
  834. # Check for spatial reset requirement
  835. if mm_router and hasattr(m, '_mm_spatial_reset') and m._mm_spatial_reset:
  836. x = mm_router.reset_spatial_input(x, m, mm_input_sources, profile)
  837. # ===== MULTIMODAL EXTENSION END =====
  838. if profile:
  839. self._profile_one_layer(m, x, dt)
  840. x = m(x) # run
  841. y.append(x if m.i in self.save else None) # save output
  842. if visualize:
  843. feature_visualization(x, m.type, m.i, save_dir=visualize)
  844. if m.i in embed:
  845. embeddings.append(torch.nn.functional.adaptive_avg_pool2d(x, (1, 1)).squeeze(-1).squeeze(-1)) # flatten
  846. if m.i == max_idx:
  847. return torch.unbind(torch.cat(embeddings, 1), dim=0)
  848. head = self.model[-1]
  849. x = head([y[j] for j in head.f], batch) # head inference
  850. return x
  851. class WorldModel(DetectionModel):
  852. """
  853. YOLOv8 World Model.
  854. This class implements the YOLOv8 World model for open-vocabulary object detection, supporting text-based
  855. class specification and CLIP model integration for zero-shot detection capabilities.
  856. Attributes:
  857. txt_feats (torch.Tensor): Text feature embeddings for classes.
  858. clip_model (torch.nn.Module): CLIP model for text encoding.
  859. Methods:
  860. __init__: Initialize YOLOv8 world model.
  861. set_classes: Set classes for offline inference.
  862. get_text_pe: Get text positional embeddings.
  863. predict: Perform forward pass with text features.
  864. loss: Compute loss with text features.
  865. Examples:
  866. Initialize a world model
  867. >>> model = WorldModel("yolov8s-world.yaml", ch=3, nc=80)
  868. >>> model.set_classes(["person", "car", "bicycle"])
  869. >>> results = model.predict(image_tensor)
  870. """
  871. def __init__(self, cfg="yolov8s-world.yaml", ch=3, nc=None, verbose=True):
  872. """
  873. Initialize YOLOv8 world model with given config and parameters.
  874. Args:
  875. cfg (str | dict): Model configuration file path or dictionary.
  876. ch (int): Number of input channels.
  877. nc (int, optional): Number of classes.
  878. verbose (bool): Whether to display model information.
  879. """
  880. self.txt_feats = torch.randn(1, nc or 80, 512) # features placeholder
  881. self.clip_model = None # CLIP model placeholder
  882. super().__init__(cfg=cfg, ch=ch, nc=nc, verbose=verbose)
  883. def set_classes(self, text, batch=80, cache_clip_model=True):
  884. """
  885. Set classes in advance so that model could do offline-inference without clip model.
  886. Args:
  887. text (List[str]): List of class names.
  888. batch (int): Batch size for processing text tokens.
  889. cache_clip_model (bool): Whether to cache the CLIP model.
  890. """
  891. self.txt_feats = self.get_text_pe(text, batch=batch, cache_clip_model=cache_clip_model)
  892. self.model[-1].nc = len(text)
  893. def get_text_pe(self, text, batch=80, cache_clip_model=True):
  894. """
  895. Set classes in advance so that model could do offline-inference without clip model.
  896. Args:
  897. text (List[str]): List of class names.
  898. batch (int): Batch size for processing text tokens.
  899. cache_clip_model (bool): Whether to cache the CLIP model.
  900. Returns:
  901. (torch.Tensor): Text positional embeddings.
  902. """
  903. from ultralytics.nn.text_model import build_text_model
  904. device = next(self.model.parameters()).device
  905. if not getattr(self, "clip_model", None) and cache_clip_model:
  906. # For backwards compatibility of models lacking clip_model attribute
  907. self.clip_model = build_text_model("clip:ViT-B/32", device=device)
  908. model = self.clip_model if cache_clip_model else build_text_model("clip:ViT-B/32", device=device)
  909. text_token = model.tokenize(text)
  910. txt_feats = [model.encode_text(token).detach() for token in text_token.split(batch)]
  911. txt_feats = txt_feats[0] if len(txt_feats) == 1 else torch.cat(txt_feats, dim=0)
  912. return txt_feats.reshape(-1, len(text), txt_feats.shape[-1])
  913. def predict(self, x, profile=False, visualize=False, txt_feats=None, augment=False, embed=None):
  914. """
  915. Perform a forward pass through the model.
  916. Args:
  917. x (torch.Tensor): The input tensor.
  918. profile (bool): If True, profile the computation time for each layer.
  919. visualize (bool): If True, save feature maps for visualization.
  920. txt_feats (torch.Tensor, optional): The text features, use it if it's given.
  921. augment (bool): If True, perform data augmentation during inference.
  922. embed (list, optional): A list of feature vectors/embeddings to return.
  923. Returns:
  924. (torch.Tensor): Model's output tensor.
  925. """
  926. txt_feats = (self.txt_feats if txt_feats is None else txt_feats).to(device=x.device, dtype=x.dtype)
  927. if len(txt_feats) != len(x) or self.model[-1].export:
  928. txt_feats = txt_feats.expand(x.shape[0], -1, -1)
  929. ori_txt_feats = txt_feats.clone()
  930. y, dt, embeddings = [], [], [] # outputs
  931. embed = frozenset(embed) if embed is not None else {-1}
  932. max_idx = max(embed)
  933. for m in self.model: # except the head part
  934. if m.f != -1: # if not from previous layer
  935. x = y[m.f] if isinstance(m.f, int) else [x if j == -1 else y[j] for j in m.f] # from earlier layers
  936. if profile:
  937. self._profile_one_layer(m, x, dt)
  938. if isinstance(m, C2fAttn):
  939. x = m(x, txt_feats)
  940. elif isinstance(m, WorldDetect):
  941. x = m(x, ori_txt_feats)
  942. elif isinstance(m, ImagePoolingAttn):
  943. txt_feats = m(x, txt_feats)
  944. else:
  945. x = m(x) # run
  946. y.append(x if m.i in self.save else None) # save output
  947. if visualize:
  948. feature_visualization(x, m.type, m.i, save_dir=visualize)
  949. if m.i in embed:
  950. embeddings.append(torch.nn.functional.adaptive_avg_pool2d(x, (1, 1)).squeeze(-1).squeeze(-1)) # flatten
  951. if m.i == max_idx:
  952. return torch.unbind(torch.cat(embeddings, 1), dim=0)
  953. return x
  954. def loss(self, batch, preds=None):
  955. """
  956. Compute loss.
  957. Args:
  958. batch (dict): Batch to compute loss on.
  959. preds (torch.Tensor | List[torch.Tensor], optional): Predictions.
  960. """
  961. if not hasattr(self, "criterion"):
  962. self.criterion = self.init_criterion()
  963. if preds is None:
  964. preds = self.forward(batch["img"], txt_feats=batch["txt_feats"])
  965. return self.criterion(preds, batch)
  966. class YOLOEModel(DetectionModel):
  967. """
  968. YOLOE detection model.
  969. This class implements the YOLOE architecture for efficient object detection with text and visual prompts,
  970. supporting both prompt-based and prompt-free inference modes.
  971. Attributes:
  972. pe (torch.Tensor): Prompt embeddings for classes.
  973. clip_model (torch.nn.Module): CLIP model for text encoding.
  974. Methods:
  975. __init__: Initialize YOLOE model.
  976. get_text_pe: Get text positional embeddings.
  977. get_visual_pe: Get visual embeddings.
  978. set_vocab: Set vocabulary for prompt-free model.
  979. get_vocab: Get fused vocabulary layer.
  980. set_classes: Set classes for offline inference.
  981. get_cls_pe: Get class positional embeddings.
  982. predict: Perform forward pass with prompts.
  983. loss: Compute loss with prompts.
  984. Examples:
  985. Initialize a YOLOE model
  986. >>> model = YOLOEModel("yoloe-v8s.yaml", ch=3, nc=80)
  987. >>> results = model.predict(image_tensor, tpe=text_embeddings)
  988. """
  989. def __init__(self, cfg="yoloe-v8s.yaml", ch=3, nc=None, verbose=True):
  990. """
  991. Initialize YOLOE model with given config and parameters.
  992. Args:
  993. cfg (str | dict): Model configuration file path or dictionary.
  994. ch (int): Number of input channels.
  995. nc (int, optional): Number of classes.
  996. verbose (bool): Whether to display model information.
  997. """
  998. super().__init__(cfg=cfg, ch=ch, nc=nc, verbose=verbose)
  999. @smart_inference_mode()
  1000. def get_text_pe(self, text, batch=80, cache_clip_model=False, without_reprta=False):
  1001. """
  1002. Set classes in advance so that model could do offline-inference without clip model.
  1003. Args:
  1004. text (List[str]): List of class names.
  1005. batch (int): Batch size for processing text tokens.
  1006. cache_clip_model (bool): Whether to cache the CLIP model.
  1007. without_reprta (bool): Whether to return text embeddings cooperated with reprta module.
  1008. Returns:
  1009. (torch.Tensor): Text positional embeddings.
  1010. """
  1011. from ultralytics.nn.text_model import build_text_model
  1012. device = next(self.model.parameters()).device
  1013. if not getattr(self, "clip_model", None) and cache_clip_model:
  1014. # For backwards compatibility of models lacking clip_model attribute
  1015. self.clip_model = build_text_model("mobileclip:blt", device=device)
  1016. model = self.clip_model if cache_clip_model else build_text_model("mobileclip:blt", device=device)
  1017. text_token = model.tokenize(text)
  1018. txt_feats = [model.encode_text(token).detach() for token in text_token.split(batch)]
  1019. txt_feats = txt_feats[0] if len(txt_feats) == 1 else torch.cat(txt_feats, dim=0)
  1020. txt_feats = txt_feats.reshape(-1, len(text), txt_feats.shape[-1])
  1021. if without_reprta:
  1022. return txt_feats
  1023. assert not self.training
  1024. head = self.model[-1]
  1025. assert isinstance(head, YOLOEDetect)
  1026. return head.get_tpe(txt_feats) # run auxiliary text head
  1027. @smart_inference_mode()
  1028. def get_visual_pe(self, img, visual):
  1029. """
  1030. Get visual embeddings.
  1031. Args:
  1032. img (torch.Tensor): Input image tensor.
  1033. visual (torch.Tensor): Visual features.
  1034. Returns:
  1035. (torch.Tensor): Visual positional embeddings.
  1036. """
  1037. return self(img, vpe=visual, return_vpe=True)
  1038. def set_vocab(self, vocab, names):
  1039. """
  1040. Set vocabulary for the prompt-free model.
  1041. Args:
  1042. vocab (nn.ModuleList): List of vocabulary items.
  1043. names (List[str]): List of class names.
  1044. """
  1045. assert not self.training
  1046. head = self.model[-1]
  1047. assert isinstance(head, YOLOEDetect)
  1048. # Cache anchors for head
  1049. device = next(self.parameters()).device
  1050. self(torch.empty(1, 3, self.args["imgsz"], self.args["imgsz"]).to(device)) # warmup
  1051. # re-parameterization for prompt-free model
  1052. self.model[-1].lrpc = nn.ModuleList(
  1053. LRPCHead(cls, pf[-1], loc[-1], enabled=i != 2)
  1054. for i, (cls, pf, loc) in enumerate(zip(vocab, head.cv3, head.cv2))
  1055. )
  1056. for loc_head, cls_head in zip(head.cv2, head.cv3):
  1057. assert isinstance(loc_head, nn.Sequential)
  1058. assert isinstance(cls_head, nn.Sequential)
  1059. del loc_head[-1]
  1060. del cls_head[-1]
  1061. self.model[-1].nc = len(names)
  1062. self.names = check_class_names(names)
  1063. def get_vocab(self, names):
  1064. """
  1065. Get fused vocabulary layer from the model.
  1066. Args:
  1067. names (list): List of class names.
  1068. Returns:
  1069. (nn.ModuleList): List of vocabulary modules.
  1070. """
  1071. assert not self.training
  1072. head = self.model[-1]
  1073. assert isinstance(head, YOLOEDetect)
  1074. assert not head.is_fused
  1075. tpe = self.get_text_pe(names)
  1076. self.set_classes(names, tpe)
  1077. device = next(self.model.parameters()).device
  1078. head.fuse(self.pe.to(device)) # fuse prompt embeddings to classify head
  1079. vocab = nn.ModuleList()
  1080. for cls_head in head.cv3:
  1081. assert isinstance(cls_head, nn.Sequential)
  1082. vocab.append(cls_head[-1])
  1083. return vocab
  1084. def set_classes(self, names, embeddings):
  1085. """
  1086. Set classes in advance so that model could do offline-inference without clip model.
  1087. Args:
  1088. names (List[str]): List of class names.
  1089. embeddings (torch.Tensor): Embeddings tensor.
  1090. """
  1091. assert not hasattr(self.model[-1], "lrpc"), (
  1092. "Prompt-free model does not support setting classes. Please try with Text/Visual prompt models."
  1093. )
  1094. assert embeddings.ndim == 3
  1095. self.pe = embeddings
  1096. self.model[-1].nc = len(names)
  1097. self.names = check_class_names(names)
  1098. def get_cls_pe(self, tpe, vpe):
  1099. """
  1100. Get class positional embeddings.
  1101. Args:
  1102. tpe (torch.Tensor, optional): Text positional embeddings.
  1103. vpe (torch.Tensor, optional): Visual positional embeddings.
  1104. Returns:
  1105. (torch.Tensor): Class positional embeddings.
  1106. """
  1107. all_pe = []
  1108. if tpe is not None:
  1109. assert tpe.ndim == 3
  1110. all_pe.append(tpe)
  1111. if vpe is not None:
  1112. assert vpe.ndim == 3
  1113. all_pe.append(vpe)
  1114. if not all_pe:
  1115. all_pe.append(getattr(self, "pe", torch.zeros(1, 80, 512)))
  1116. return torch.cat(all_pe, dim=1)
  1117. def predict(
  1118. self, x, profile=False, visualize=False, tpe=None, augment=False, embed=None, vpe=None, return_vpe=False
  1119. ):
  1120. """
  1121. Perform a forward pass through the model.
  1122. Args:
  1123. x (torch.Tensor): The input tensor.
  1124. profile (bool): If True, profile the computation time for each layer.
  1125. visualize (bool): If True, save feature maps for visualization.
  1126. tpe (torch.Tensor, optional): Text positional embeddings.
  1127. augment (bool): If True, perform data augmentation during inference.
  1128. embed (list, optional): A list of feature vectors/embeddings to return.
  1129. vpe (torch.Tensor, optional): Visual positional embeddings.
  1130. return_vpe (bool): If True, return visual positional embeddings.
  1131. Returns:
  1132. (torch.Tensor): Model's output tensor.
  1133. """
  1134. y, dt, embeddings = [], [], [] # outputs
  1135. b = x.shape[0]
  1136. embed = frozenset(embed) if embed is not None else {-1}
  1137. max_idx = max(embed)
  1138. for m in self.model: # except the head part
  1139. if m.f != -1: # if not from previous layer
  1140. x = y[m.f] if isinstance(m.f, int) else [x if j == -1 else y[j] for j in m.f] # from earlier layers
  1141. if profile:
  1142. self._profile_one_layer(m, x, dt)
  1143. if isinstance(m, YOLOEDetect):
  1144. vpe = m.get_vpe(x, vpe) if vpe is not None else None
  1145. if return_vpe:
  1146. assert vpe is not None
  1147. assert not self.training
  1148. return vpe
  1149. cls_pe = self.get_cls_pe(m.get_tpe(tpe), vpe).to(device=x[0].device, dtype=x[0].dtype)
  1150. if cls_pe.shape[0] != b or m.export:
  1151. cls_pe = cls_pe.expand(b, -1, -1)
  1152. x = m(x, cls_pe)
  1153. else:
  1154. x = m(x) # run
  1155. y.append(x if m.i in self.save else None) # save output
  1156. if visualize:
  1157. feature_visualization(x, m.type, m.i, save_dir=visualize)
  1158. if m.i in embed:
  1159. embeddings.append(torch.nn.functional.adaptive_avg_pool2d(x, (1, 1)).squeeze(-1).squeeze(-1)) # flatten
  1160. if m.i == max_idx:
  1161. return torch.unbind(torch.cat(embeddings, 1), dim=0)
  1162. return x
  1163. def loss(self, batch, preds=None):
  1164. """
  1165. Compute loss.
  1166. Args:
  1167. batch (dict): Batch to compute loss on.
  1168. preds (torch.Tensor | List[torch.Tensor], optional): Predictions.
  1169. """
  1170. if not hasattr(self, "criterion"):
  1171. from ultralytics.utils.loss import TVPDetectLoss
  1172. visual_prompt = batch.get("visuals", None) is not None # TODO
  1173. self.criterion = TVPDetectLoss(self) if visual_prompt else self.init_criterion()
  1174. if preds is None:
  1175. preds = self.forward(batch["img"], tpe=batch.get("txt_feats", None), vpe=batch.get("visuals", None))
  1176. return self.criterion(preds, batch)
  1177. class YOLOESegModel(YOLOEModel, SegmentationModel):
  1178. """
  1179. YOLOE segmentation model.
  1180. This class extends YOLOEModel to handle instance segmentation tasks with text and visual prompts,
  1181. providing specialized loss computation for pixel-level object detection and segmentation.
  1182. Methods:
  1183. __init__: Initialize YOLOE segmentation model.
  1184. loss: Compute loss with prompts for segmentation.
  1185. Examples:
  1186. Initialize a YOLOE segmentation model
  1187. >>> model = YOLOESegModel("yoloe-v8s-seg.yaml", ch=3, nc=80)
  1188. >>> results = model.predict(image_tensor, tpe=text_embeddings)
  1189. """
  1190. def __init__(self, cfg="yoloe-v8s-seg.yaml", ch=3, nc=None, verbose=True):
  1191. """
  1192. Initialize YOLOE segmentation model with given config and parameters.
  1193. Args:
  1194. cfg (str | dict): Model configuration file path or dictionary.
  1195. ch (int): Number of input channels.
  1196. nc (int, optional): Number of classes.
  1197. verbose (bool): Whether to display model information.
  1198. """
  1199. super().__init__(cfg=cfg, ch=ch, nc=nc, verbose=verbose)
  1200. def loss(self, batch, preds=None):
  1201. """
  1202. Compute loss.
  1203. Args:
  1204. batch (dict): Batch to compute loss on.
  1205. preds (torch.Tensor | List[torch.Tensor], optional): Predictions.
  1206. """
  1207. if not hasattr(self, "criterion"):
  1208. from ultralytics.utils.loss import TVPSegmentLoss
  1209. visual_prompt = batch.get("visuals", None) is not None # TODO
  1210. self.criterion = TVPSegmentLoss(self) if visual_prompt else self.init_criterion()
  1211. if preds is None:
  1212. preds = self.forward(batch["img"], tpe=batch.get("txt_feats", None), vpe=batch.get("visuals", None))
  1213. return self.criterion(preds, batch)
  1214. class Ensemble(torch.nn.ModuleList):
  1215. """
  1216. Ensemble of models.
  1217. This class allows combining multiple YOLO models into an ensemble for improved performance through
  1218. model averaging or other ensemble techniques.
  1219. Methods:
  1220. __init__: Initialize an ensemble of models.
  1221. forward: Generate predictions from all models in the ensemble.
  1222. Examples:
  1223. Create an ensemble of models
  1224. >>> ensemble = Ensemble()
  1225. >>> ensemble.append(model1)
  1226. >>> ensemble.append(model2)
  1227. >>> results = ensemble(image_tensor)
  1228. """
  1229. def __init__(self):
  1230. """Initialize an ensemble of models."""
  1231. super().__init__()
  1232. def forward(self, x, augment=False, profile=False, visualize=False):
  1233. """
  1234. Generate the YOLO network's final layer.
  1235. Args:
  1236. x (torch.Tensor): Input tensor.
  1237. augment (bool): Whether to augment the input.
  1238. profile (bool): Whether to profile the model.
  1239. visualize (bool): Whether to visualize the features.
  1240. Returns:
  1241. y (torch.Tensor): Concatenated predictions from all models.
  1242. train_out (None): Always None for ensemble inference.
  1243. """
  1244. y = [module(x, augment, profile, visualize)[0] for module in self]
  1245. # y = torch.stack(y).max(0)[0] # max ensemble
  1246. # y = torch.stack(y).mean(0) # mean ensemble
  1247. y = torch.cat(y, 2) # nms ensemble, y shape(B, HW, C)
  1248. return y, None # inference, train output
  1249. # Functions ------------------------------------------------------------------------------------------------------------
  1250. @contextlib.contextmanager
  1251. def temporary_modules(modules=None, attributes=None):
  1252. """
  1253. Context manager for temporarily adding or modifying modules in Python's module cache (`sys.modules`).
  1254. This function can be used to change the module paths during runtime. It's useful when refactoring code,
  1255. where you've moved a module from one location to another, but you still want to support the old import
  1256. paths for backwards compatibility.
  1257. Args:
  1258. modules (dict, optional): A dictionary mapping old module paths to new module paths.
  1259. attributes (dict, optional): A dictionary mapping old module attributes to new module attributes.
  1260. Examples:
  1261. >>> with temporary_modules({"old.module": "new.module"}, {"old.module.attribute": "new.module.attribute"}):
  1262. >>> import old.module # this will now import new.module
  1263. >>> from old.module import attribute # this will now import new.module.attribute
  1264. Note:
  1265. The changes are only in effect inside the context manager and are undone once the context manager exits.
  1266. Be aware that directly manipulating `sys.modules` can lead to unpredictable results, especially in larger
  1267. applications or libraries. Use this function with caution.
  1268. """
  1269. if modules is None:
  1270. modules = {}
  1271. if attributes is None:
  1272. attributes = {}
  1273. import sys
  1274. from importlib import import_module
  1275. try:
  1276. # Set attributes in sys.modules under their old name
  1277. for old, new in attributes.items():
  1278. old_module, old_attr = old.rsplit(".", 1)
  1279. new_module, new_attr = new.rsplit(".", 1)
  1280. setattr(import_module(old_module), old_attr, getattr(import_module(new_module), new_attr))
  1281. # Set modules in sys.modules under their old name
  1282. for old, new in modules.items():
  1283. sys.modules[old] = import_module(new)
  1284. yield
  1285. finally:
  1286. # Remove the temporary module paths
  1287. for old in modules:
  1288. if old in sys.modules:
  1289. del sys.modules[old]
  1290. class SafeClass:
  1291. """A placeholder class to replace unknown classes during unpickling."""
  1292. def __init__(self, *args, **kwargs):
  1293. """Initialize SafeClass instance, ignoring all arguments."""
  1294. pass
  1295. def __call__(self, *args, **kwargs):
  1296. """Run SafeClass instance, ignoring all arguments."""
  1297. pass
  1298. class SafeUnpickler(pickle.Unpickler):
  1299. """Custom Unpickler that replaces unknown classes with SafeClass."""
  1300. def find_class(self, module, name):
  1301. """
  1302. Attempt to find a class, returning SafeClass if not among safe modules.
  1303. Args:
  1304. module (str): Module name.
  1305. name (str): Class name.
  1306. Returns:
  1307. (type): Found class or SafeClass.
  1308. """
  1309. safe_modules = (
  1310. "torch",
  1311. "collections",
  1312. "collections.abc",
  1313. "builtins",
  1314. "math",
  1315. "numpy",
  1316. # Add other modules considered safe
  1317. )
  1318. if module in safe_modules:
  1319. return super().find_class(module, name)
  1320. else:
  1321. return SafeClass
  1322. def torch_safe_load(weight, safe_only=False):
  1323. """
  1324. Attempt to load a PyTorch model with the torch.load() function. If a ModuleNotFoundError is raised, it catches the
  1325. error, logs a warning message, and attempts to install the missing module via the check_requirements() function.
  1326. After installation, the function again attempts to load the model using torch.load().
  1327. Args:
  1328. weight (str): The file path of the PyTorch model.
  1329. safe_only (bool): If True, replace unknown classes with SafeClass during loading.
  1330. Returns:
  1331. ckpt (dict): The loaded model checkpoint.
  1332. file (str): The loaded filename.
  1333. Examples:
  1334. >>> from ultralytics.nn.tasks import torch_safe_load
  1335. >>> ckpt, file = torch_safe_load("path/to/best.pt", safe_only=True)
  1336. """
  1337. from ultralytics.utils.downloads import attempt_download_asset
  1338. check_suffix(file=weight, suffix=".pt")
  1339. file = attempt_download_asset(weight) # search online if missing locally
  1340. try:
  1341. with temporary_modules(
  1342. modules={
  1343. "ultralytics.yolo.utils": "ultralytics.utils",
  1344. "ultralytics.yolo.v8": "ultralytics.models.yolo",
  1345. "ultralytics.yolo.data": "ultralytics.data",
  1346. },
  1347. attributes={
  1348. "ultralytics.nn.modules.block.Silence": "torch.nn.Identity", # YOLOv9e
  1349. "ultralytics.nn.tasks.YOLOv10DetectionModel": "ultralytics.nn.tasks.DetectionModel", # YOLOv10
  1350. "ultralytics.utils.loss.v10DetectLoss": "ultralytics.utils.loss.E2EDetectLoss", # YOLOv10
  1351. },
  1352. ):
  1353. if safe_only:
  1354. # Load via custom pickle module
  1355. safe_pickle = types.ModuleType("safe_pickle")
  1356. safe_pickle.Unpickler = SafeUnpickler
  1357. safe_pickle.load = lambda file_obj: SafeUnpickler(file_obj).load()
  1358. with open(file, "rb") as f:
  1359. ckpt = torch_load(f, pickle_module=safe_pickle)
  1360. else:
  1361. ckpt = torch_load(file, map_location="cpu")
  1362. except ModuleNotFoundError as e: # e.name is missing module name
  1363. if e.name == "models":
  1364. raise TypeError(
  1365. emojis(
  1366. f"ERROR ❌️ {weight} appears to be an Ultralytics YOLOv5 model originally trained "
  1367. f"with https://github.com/ultralytics/yolov5.\nThis model is NOT forwards compatible with "
  1368. f"YOLOv8 at https://github.com/ultralytics/ultralytics."
  1369. f"\nRecommend fixes are to train a new model using the latest 'ultralytics' package or to "
  1370. f"run a command with an official Ultralytics model, i.e. 'yolo predict model=yolo11n.pt'"
  1371. )
  1372. ) from e
  1373. elif e.name == "numpy._core":
  1374. raise ModuleNotFoundError(
  1375. emojis(
  1376. f"ERROR ❌️ {weight} requires numpy>=1.26.1, however numpy=={__import__('numpy').__version__} is installed."
  1377. )
  1378. ) from e
  1379. LOGGER.warning(
  1380. f"{weight} appears to require '{e.name}', which is not in Ultralytics requirements."
  1381. f"\nAutoInstall will run now for '{e.name}' but this feature will be removed in the future."
  1382. f"\nRecommend fixes are to train a new model using the latest 'ultralytics' package or to "
  1383. f"run a command with an official Ultralytics model, i.e. 'yolo predict model=yolo11n.pt'"
  1384. )
  1385. check_requirements(e.name) # install missing module
  1386. ckpt = torch_load(file, map_location="cpu")
  1387. if not isinstance(ckpt, dict):
  1388. # File is likely a YOLO instance saved with i.e. torch.save(model, "saved_model.pt")
  1389. LOGGER.warning(
  1390. f"The file '{weight}' appears to be improperly saved or formatted. "
  1391. f"For optimal results, use model.save('filename.pt') to correctly save YOLO models."
  1392. )
  1393. ckpt = {"model": ckpt.model}
  1394. return ckpt, file
  1395. def attempt_load_weights(weights, device=None, inplace=True, fuse=False):
  1396. """
  1397. Load an ensemble of models weights=[a,b,c] or a single model weights=[a] or weights=a.
  1398. Args:
  1399. weights (str | List[str]): Model weights path(s).
  1400. device (torch.device, optional): Device to load model to.
  1401. inplace (bool): Whether to do inplace operations.
  1402. fuse (bool): Whether to fuse model.
  1403. Returns:
  1404. (torch.nn.Module): Loaded model.
  1405. """
  1406. ensemble = Ensemble()
  1407. for w in weights if isinstance(weights, list) else [weights]:
  1408. ckpt, w = torch_safe_load(w) # load ckpt
  1409. args = {**DEFAULT_CFG_DICT, **ckpt["train_args"]} if "train_args" in ckpt else None # combined args
  1410. model = (ckpt.get("ema") or ckpt["model"]).to(device).float() # FP32 model
  1411. # Model compatibility updates
  1412. model.args = args # attach args to model
  1413. model.pt_path = w # attach *.pt file path to model
  1414. model.task = getattr(model, "task", guess_model_task(model))
  1415. if not hasattr(model, "stride"):
  1416. model.stride = torch.tensor([32.0])
  1417. # Append
  1418. ensemble.append(model.fuse().eval() if fuse and hasattr(model, "fuse") else model.eval()) # model in eval mode
  1419. # Module updates
  1420. for m in ensemble.modules():
  1421. if hasattr(m, "inplace"):
  1422. m.inplace = inplace
  1423. elif isinstance(m, torch.nn.Upsample) and not hasattr(m, "recompute_scale_factor"):
  1424. m.recompute_scale_factor = None # torch 1.11.0 compatibility
  1425. # Return model
  1426. if len(ensemble) == 1:
  1427. return ensemble[-1]
  1428. # Return ensemble
  1429. LOGGER.info(f"Ensemble created with {weights}\n")
  1430. for k in "names", "nc", "yaml":
  1431. setattr(ensemble, k, getattr(ensemble[0], k))
  1432. ensemble.stride = ensemble[int(torch.argmax(torch.tensor([m.stride.max() for m in ensemble])))].stride
  1433. assert all(ensemble[0].nc == m.nc for m in ensemble), f"Models differ in class counts {[m.nc for m in ensemble]}"
  1434. return ensemble
  1435. def attempt_load_one_weight(weight, device=None, inplace=True, fuse=False):
  1436. """
  1437. Load a single model weights.
  1438. Args:
  1439. weight (str): Model weight path.
  1440. device (torch.device, optional): Device to load model to.
  1441. inplace (bool): Whether to do inplace operations.
  1442. fuse (bool): Whether to fuse model.
  1443. Returns:
  1444. model (torch.nn.Module): Loaded model.
  1445. ckpt (dict): Model checkpoint dictionary.
  1446. """
  1447. ckpt, weight = torch_safe_load(weight) # load ckpt
  1448. args = {**DEFAULT_CFG_DICT, **(ckpt.get("train_args", {}))} # combine model and default args, preferring model args
  1449. model = (ckpt.get("ema") or ckpt["model"]).to(device).float() # FP32 model
  1450. # Model compatibility updates
  1451. model.args = {k: v for k, v in args.items() if k in DEFAULT_CFG_KEYS} # attach args to model
  1452. model.pt_path = weight # attach *.pt file path to model
  1453. model.task = getattr(model, "task", guess_model_task(model))
  1454. if not hasattr(model, "stride"):
  1455. model.stride = torch.tensor([32.0])
  1456. model = model.fuse().eval() if fuse and hasattr(model, "fuse") else model.eval() # model in eval mode
  1457. # Module updates
  1458. for m in model.modules():
  1459. if hasattr(m, "inplace"):
  1460. m.inplace = inplace
  1461. elif isinstance(m, torch.nn.Upsample) and not hasattr(m, "recompute_scale_factor"):
  1462. m.recompute_scale_factor = None # torch 1.11.0 compatibility
  1463. # Return model and ckpt
  1464. return model, ckpt
  1465. def _validate_and_fill_dea_args(args, c_left):
  1466. """Validate DEA args strictly and auto-fill only the channel whentextNone.
  1467. text: DEA(channel, kernel_size, p_kernel=None, m_kernel=None, reduction=16)
  1468. text: DEA([None, kernel_size]) text DEA([channel, kernel_size]).
  1469. textautotext, text.
  1470. """
  1471. a = list(args) if isinstance(args, (list, tuple)) else [args]
  1472. # text [channel/None, kernel_size]
  1473. if len(a) == 2:
  1474. ch, ks = a
  1475. ch = c_left if ch is None else ch
  1476. if not isinstance(ch, int) or not isinstance(ks, int) or ks <= 0:
  1477. raise ValueError(
  1478. f"DEA expects [channel(or None), kernel_size] as minimal form, got {args}"
  1479. )
  1480. return [ch, ks, None, None, 16]
  1481. # text: text5
  1482. while len(a) < 5:
  1483. a.append(None)
  1484. ch, ks, pk, mk, rd = a[:5]
  1485. ch = c_left if ch is None else ch
  1486. # text
  1487. if not isinstance(ch, int) or ch <= 0:
  1488. raise ValueError(f"DEA arg[0]=channel must be positive int, got {ch}")
  1489. if not isinstance(ks, int) or ks <= 0:
  1490. raise ValueError(f"DEA arg[1]=kernel_size must be positive int, got {ks}")
  1491. if pk is not None and (not isinstance(pk, (list, tuple)) or len(pk) != 2):
  1492. raise ValueError(f"DEA arg[2]=p_kernel must be 2-list/tuple or None, got {pk}")
  1493. if mk is not None and (not isinstance(mk, (list, tuple)) or len(mk) != 2):
  1494. raise ValueError(f"DEA arg[3]=m_kernel must be 2-list/tuple or None, got {mk}")
  1495. if rd is None:
  1496. rd = 16
  1497. if not isinstance(rd, int) or rd <= 0:
  1498. raise ValueError(f"DEA arg[4]=reduction must be positive int, got {rd}")
  1499. return [ch, ks, pk, mk, rd]
  1500. def parse_model(d, ch, verbose=True, dataset_config=None):
  1501. """Parse the NeuroSeg-MF RT-DETR YAML into a PyTorch model.
  1502. Supported modules are intentionally limited to the current NeuroSeg-MF architecture:
  1503. ConvNormLayer, BasicBlock, Blocks, GateFusion, A2C2f, Conv, AIFI, Concat,
  1504. RepC3, RTDETRDecoder, nn.MaxPool2d and nn.Upsample.
  1505. """
  1506. import ast
  1507. max_channels = float("inf")
  1508. nc, act, scales = (d.get(x) for x in ("nc", "activation", "scales"))
  1509. depth, width = (d.get(x, 1.0) for x in ("depth_multiple", "width_multiple"))
  1510. if scales:
  1511. scale = d.get("scale") or tuple(scales.keys())[0]
  1512. if not d.get("scale"):
  1513. LOGGER.warning(f"no model scale passed. Assuming scale='{scale}'.")
  1514. depth, width, max_channels = scales[scale]
  1515. if act:
  1516. allowed_acts = {
  1517. "torch.nn.SiLU()": torch.nn.SiLU(),
  1518. "nn.SiLU()": torch.nn.SiLU(),
  1519. "torch.nn.ReLU()": torch.nn.ReLU(),
  1520. "nn.ReLU()": torch.nn.ReLU(),
  1521. }
  1522. if act not in allowed_acts:
  1523. raise ValueError(f"Unsupported activation expression in minimal parser: {act}")
  1524. Conv.default_act = allowed_acts[act]
  1525. if verbose:
  1526. LOGGER.info(f"{colorstr('activation:')} {act}")
  1527. if verbose:
  1528. LOGGER.info(f"\n{'':>3}{'from':>20}{'n':>3}{'params':>10} {'module':<45}{'arguments':<30}")
  1529. mm_router = None
  1530. if MULTIMODAL_AVAILABLE:
  1531. try:
  1532. config_dict = d.copy()
  1533. if dataset_config:
  1534. config_dict['dataset_config'] = dataset_config
  1535. mm_router = MultiModalRouter(config_dict, verbose=verbose)
  1536. except Exception as e:
  1537. if verbose:
  1538. LOGGER.warning(f"MultiModal router initialization failed: {e}")
  1539. ch = [ch]
  1540. layers, save = [], []
  1541. base_modules = {Conv, ConvNormLayer, A2C2f}
  1542. repeat_modules = {A2C2f, RepC3}
  1543. for i, layer_config in enumerate(d["backbone"] + d["head"]):
  1544. if mm_router:
  1545. c1, mm_input_source, mm_attributes = mm_router.parse_layer_config(layer_config, i, ch, verbose)
  1546. f, n, m, args = layer_config[:4]
  1547. else:
  1548. if len(layer_config) < 4:
  1549. raise ValueError(f"Invalid layer definition at index {i}: {layer_config}")
  1550. f, n, m, args = layer_config[:4]
  1551. c1, mm_input_source, mm_attributes = None, None, {}
  1552. m = getattr(torch.nn, m[3:]) if isinstance(m, str) and m.startswith("nn.") else globals().get(m)
  1553. if m is None:
  1554. raise ImportError(f"Module '{layer_config[2]}' used at layer {i} is not available in the NeuroSeg-MF minimal runtime.")
  1555. args = list(args)
  1556. for j, a in enumerate(args):
  1557. if isinstance(a, str):
  1558. if a == "nc":
  1559. args[j] = nc
  1560. elif a in globals():
  1561. args[j] = globals()[a]
  1562. else:
  1563. with contextlib.suppress(Exception):
  1564. args[j] = ast.literal_eval(a)
  1565. n_ = max(round(n * depth), 1) if n > 1 else n
  1566. n = n_
  1567. if m in base_modules:
  1568. if mm_input_source and c1 is not None:
  1569. c2 = args[0]
  1570. else:
  1571. c1, c2 = ch[f], args[0]
  1572. if c2 != nc:
  1573. c2 = make_divisible(min(c2, max_channels) * width, 8)
  1574. args = [c1, c2, *args[1:]]
  1575. if m in repeat_modules:
  1576. args.insert(2, n)
  1577. n = 1
  1578. if m is A2C2f:
  1579. args.extend((True, 1.2))
  1580. elif m is Blocks:
  1581. block_type = globals()[args[1]] if isinstance(args[1], str) else args[1]
  1582. c1, c2 = ch[f], args[0] * block_type.expansion
  1583. args = [c1, args[0], block_type, *args[2:]]
  1584. elif m is BasicBlock:
  1585. c1, c2 = ch[f], args[0] * BasicBlock.expansion
  1586. args = [c1, *args]
  1587. elif m is RepC3:
  1588. c1, c2 = ch[f], args[0]
  1589. if c2 != nc:
  1590. c2 = make_divisible(min(c2, max_channels) * width, 8)
  1591. args = [c1, c2, n, *args[1:]]
  1592. n = 1
  1593. elif m is AIFI:
  1594. c2 = ch[f]
  1595. args = [ch[f], *args]
  1596. elif m is Concat:
  1597. c2 = sum(ch[x] for x in f)
  1598. elif m is GateFusion:
  1599. c2 = ch[f[0]]
  1600. elif m is RTDETRDecoder:
  1601. args.insert(1, [ch[x] for x in f])
  1602. c2 = nc
  1603. elif m is torch.nn.BatchNorm2d:
  1604. args = [ch[f]]
  1605. c2 = ch[f]
  1606. elif m in {torch.nn.MaxPool2d, torch.nn.Upsample, torch.nn.Identity}:
  1607. c2 = ch[f]
  1608. elif m is Index:
  1609. c2 = ch[f]
  1610. else:
  1611. raise ImportError(f"Module '{layer_config[2]}' at layer {i} is not supported by the NeuroSeg-MF minimal parser.")
  1612. m_ = torch.nn.Sequential(*(m(*args) for _ in range(n))) if n > 1 else m(*args)
  1613. t = str(m)[8:-2].replace("__main__.", "")
  1614. m_.np = sum(x.numel() for x in m_.parameters())
  1615. m_.i, m_.f, m_.type = i, f, t
  1616. if mm_attributes and mm_router:
  1617. mm_router.set_module_attributes(m_, mm_attributes)
  1618. if verbose:
  1619. display_args = [a.__name__ if isinstance(a, type) else a for a in args]
  1620. LOGGER.info(f"{i:>3}{str(f):>20}{n_:>3}{m_.np:10.0f} {t:<45}{str(display_args):<30}")
  1621. save.extend(x % i for x in ([f] if isinstance(f, int) else f) if x != -1)
  1622. layers.append(m_)
  1623. if i == 0:
  1624. ch = []
  1625. ch.append(c2)
  1626. model = torch.nn.Sequential(*layers)
  1627. if mm_router:
  1628. model.multimodal_router = mm_router
  1629. return model, sorted(save)
  1630. def yaml_model_load(path):
  1631. """
  1632. Load a YOLOv8 model from a YAML file.
  1633. Args:
  1634. path (str | Path): Path to the YAML file.
  1635. Returns:
  1636. (dict): Model dictionary.
  1637. """
  1638. path = Path(path)
  1639. if path.stem in (f"yolov{d}{x}6" for x in "nsmlx" for d in (5, 8)):
  1640. new_stem = re.sub(r"(\d+)([nslmx])6(.+)?$", r"\1\2-p6\3", path.stem)
  1641. LOGGER.warning(f"Ultralytics YOLO P6 models now use -p6 suffix. Renaming {path.stem} to {new_stem}.")
  1642. path = path.with_name(new_stem + path.suffix)
  1643. unified_path = re.sub(r"(\d+)([nslmx])(.+)?$", r"\1\3", str(path)) # i.e. yolov8x.yaml -> yolov8.yaml
  1644. yaml_file = check_yaml(unified_path, hard=False) or check_yaml(path)
  1645. d = YAML.load(yaml_file) # model dict
  1646. d["scale"] = guess_model_scale(path)
  1647. d["yaml_file"] = str(path)
  1648. return d
  1649. def guess_model_scale(model_path):
  1650. """
  1651. Extract the size character n, s, m, l, or x of the model's scale from the model path.
  1652. Args:
  1653. model_path (str | Path): The path to the YOLO model's YAML file.
  1654. Returns:
  1655. (str): The size character of the model's scale (n, s, m, l, or x).
  1656. """
  1657. try:
  1658. return re.search(r"yolo(e-)?[v]?\d+([nslmx])", Path(model_path).stem).group(2) # noqa
  1659. except AttributeError:
  1660. return ""
  1661. def guess_model_task(model):
  1662. """
  1663. Guess the task of a PyTorch model from its architecture or configuration.
  1664. Args:
  1665. model (torch.nn.Module | dict): PyTorch model or model configuration in YAML format.
  1666. Returns:
  1667. (str): Task of the model ('detect', 'segment', 'classify', 'pose', 'obb').
  1668. """
  1669. def cfg2task(cfg):
  1670. """Guess from YAML dictionary."""
  1671. m = cfg["head"][-1][-2].lower() # output module name
  1672. if m in {"classify", "classifier", "cls", "fc"}:
  1673. return "classify"
  1674. if "detect" in m:
  1675. return "detect"
  1676. if "segment" in m:
  1677. return "segment"
  1678. if m == "pose":
  1679. return "pose"
  1680. if m == "obb":
  1681. return "obb"
  1682. # Guess from model cfg
  1683. if isinstance(model, dict):
  1684. with contextlib.suppress(Exception):
  1685. return cfg2task(model)
  1686. # Guess from PyTorch model
  1687. if isinstance(model, torch.nn.Module): # PyTorch model
  1688. def _resolve_attr(obj, dotted):
  1689. for name in dotted.split("."):
  1690. obj = getattr(obj, name)
  1691. return obj
  1692. for x in "model.args", "model.model.args", "model.model.model.args":
  1693. with contextlib.suppress(Exception):
  1694. return _resolve_attr(model, x)["task"]
  1695. for x in "model.yaml", "model.model.yaml", "model.model.model.yaml":
  1696. with contextlib.suppress(Exception):
  1697. return cfg2task(_resolve_attr(model, x))
  1698. for m in model.modules():
  1699. if isinstance(m, (Segment, YOLOESegment)):
  1700. return "segment"
  1701. elif isinstance(m, Classify):
  1702. return "classify"
  1703. elif isinstance(m, Pose):
  1704. return "pose"
  1705. elif isinstance(m, OBB):
  1706. return "obb"
  1707. elif isinstance(m, (Detect, WorldDetect, YOLOEDetect, v10Detect)):
  1708. return "detect"
  1709. # Guess from model filename
  1710. if isinstance(model, (str, Path)):
  1711. model = Path(model)
  1712. if "-seg" in model.stem or "segment" in model.parts:
  1713. return "segment"
  1714. elif "-cls" in model.stem or "classify" in model.parts:
  1715. return "classify"
  1716. elif "-pose" in model.stem or "pose" in model.parts:
  1717. return "pose"
  1718. elif "-obb" in model.stem or "obb" in model.parts:
  1719. return "obb"
  1720. elif "detect" in model.parts:
  1721. return "detect"
  1722. # Unable to determine task from model
  1723. LOGGER.warning(
  1724. "Unable to automatically guess model task, assuming 'task=detect'. "
  1725. "Explicitly define task for your model, i.e. 'task=detect', 'segment', 'classify','pose' or 'obb'."
  1726. )
  1727. return "detect" # assume detect

tasks.py at commit 4bfb1ef, no license · at the source

Overview

Authors: Zhehao Xu1, Weiyi Liu2, Shanshan Liang2, Hongbo Jia3,4, Xiaowei Chen2,5, Han Qin5, Xiang Liao1
ORCID iDs: Xiang Liao
  1. Center for Neurointelligence, School of Medicine, Chongqing University, Chongqing 400030, China
  2. Brain Research Center, State Key Laboratory of Trauma and Chemical Poisoning, Third Military Medical University, Chongqing 400038, China
  3. Jiangsu Key Laboratory for Advanced Theranostics and Medical Instrumentation, Suzhou Institute of Biomedical Engineering and Technology, Chinese Academy of Sciences, Suzhou 215163, Jiangsu, China
  4. Leibniz Institute for Neurobiology, Magdeburg 39118, Germany
  5. LFC Laboratory (Chongqing Key Laboratory of Brain and Aerospace Intelligence) and Chongqing Institute for Brain and Intelligence, Guangyang Bay Laboratory, Chongqing 400064, China
Journal: Biomedical optics express, volume 17, issue 7, pages 3727-3746
Dates: received 31 March 2026; accepted 9 June 2026; published online 17 June 2026
Type: Research article · Language: English
License: none stated
Identifiers: DOI 10.1364/boe.600665 · PMID 42460356 · PMCID PMC13372336 · OpenAlex W7164500274
Open access: gold, a free copy (OpenAlex)
Status: code verified
Methods: Statistics, Machine learning, Connectivity, fMRI & imaging
Topic: Advanced Fluorescence Microscopy Techniques (Biophysics, Biochemistry, Genetics and Molecular Biology), according to OpenAlex
Funding: National Natural Science Foundation of China (32171096, 32127801)
Citations: cited by 1 paper (Europe PMC); 57 references in the paper

Abstract

The abstract is not reproduced here: the paper's license (none stated) does not allow it. Read it in the paper, at the publisher or on Europe PMC.

Repository

Its files are read in the Code ↔ Paper reader above, with 11 matches between paragraphs and lines of code.

XZH-James/NeuroSeg-MF

License: none: the authors keep all their rights
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Commit: 4bfb1ef734acd2fa90b44960980d6233d71ed7cb, 16 June 2026
Languages: Python (100)
Size: 108 files, 100 scripts
Software Heritage: not archived
Found in: the references
Holds: README, environment (requirements.txt)
Not found: license file, CITATION.cff, tests, continuous integration, documentation
Tools: PyTorch (56 files), NumPy (39 files), OpenCV (20 files), Pillow (10 files), Matplotlib (5 files), pandas (4 files), SciPy (4 files), tifffile (2 files), seaborn (1 file), TensorFlow (1 file)
Availability: 1 check, the latest on 27 September 2026: the link answers
  • 27 September 2026: the link answers
101 files

The paper's code and data availability statement is in the Data section.

Tracing map

Proposed by the machine: these links were found in the paper and verified at the source, without human review. The map will receive a Zenodo DOI once one of the paper's authors has validated it with their ORCID.

What the map holds:

  • 1 repository of the authors' code, each at its verified commit, with its license and how the link was found in the paper;
  • 100 scripts, each with its path and the digest of its content;
  • 11 matches between paragraphs of the paper and lines of the code (method lexical-v1);
  • neither the text of the paper nor the code itself.

Its JSON (tracing-map.json) is deposited on Zenodo with its DOI once the map is validated.

Data

No dataset and no data link were found in the paper.

Code and data availability statement

The paper has a code and data availability statement. Its license (none stated) does not allow reproducing it here; in short, from what the harvester recognized in it:

  • it says that the data are available on request

Read it in the paper: doi.org/10.1364/boe.600665.

Versions

The history of this record: each version stored by the harvester or made by a correction of its authors or of the maintainers of its code, and what changed in its facts. The texts of the paper (its abstract, its availability statements) are not part of it; versions that changed only those are not listed.

Version 1, 27 September 2026: the first record

Recorded: type, language, journal, volume, issue, pages, dates, 7 authors, 1 funder, 45 references.

Cite

This paper

Xu, Z., Liu, W., Liang, S., Jia, H., Chen, X., Qin, H., & Liao, X. (2026). NeuroSeg-MF: robust neuron segmentation in two-photon Ca&lt;sup&gt;2+&lt;/sup&gt; imaging using multi-feature fusion and detection-guided SAM. Biomedical optics express, 17(7), 3727-3746. https://doi.org/10.1364/boe.600665

BibTeX

@article{xu2026neuroseg,
author = {Xu, Zhehao and Liu, Weiyi and Liang, Shanshan and Jia, Hongbo and Chen, Xiaowei and Qin, Han and Liao, Xiang},
title = {{NeuroSeg-MF: robust neuron segmentation in two-photon Ca\&lt;sup\&gt;2+\&lt;/sup\&gt; imaging using multi-feature fusion and detection-guided SAM}},
journal = {Biomedical optics express},
year = {2026},
month = jun,
volume = {17},
number = {7},
pages = {3727--3746},
publisher = {Optica Publishing Group},
issn = {2156-7085},
doi = {10.1364/boe.600665},
url = {https://doi.org/10.1364/boe.600665},
pmid = {42460356},
pmcid = {PMC13372336}
}

RIS

TY - JOUR
AU - Xu, Zhehao
AU - Liu, Weiyi
AU - Liang, Shanshan
AU - Jia, Hongbo
AU - Chen, Xiaowei
AU - Qin, Han
AU - Liao, Xiang
TI - NeuroSeg-MF: robust neuron segmentation in two-photon Ca&lt;sup&gt;2+&lt;/sup&gt; imaging using multi-feature fusion and detection-guided SAM
T2 - Biomedical optics express
J2 - Biomed Opt Express
PY - 2026
DA - 2026/06/17
VL - 17
IS - 7
SP - 3727
EP - 3746
SN - 2156-7085
PB - Optica Publishing Group
DO - 10.1364/boe.600665
UR - https://doi.org/10.1364/boe.600665
LA - en
ER -

CSL-JSON

{
"id": "10.1364/boe.600665",
"type": "article-journal",
"title": "NeuroSeg-MF: robust neuron segmentation in two-photon Ca&lt;sup&gt;2+&lt;/sup&gt; imaging using multi-feature fusion and detection-guided SAM",
"container-title": "Biomedical optics express",
"author": [
{
"family": "Xu",
"given": "Zhehao"
},
{
"family": "Liu",
"given": "Weiyi"
},
{
"family": "Liang",
"given": "Shanshan"
},
{
"family": "Jia",
"given": "Hongbo"
},
{
"family": "Chen",
"given": "Xiaowei"
},
{
"family": "Qin",
"given": "Han"
},
{
"family": "Liao",
"given": "Xiang"
}
],
"container-title-short": "Biomed Opt Express",
"volume": "17",
"issue": "7",
"page": "3727-3746",
"DOI": "10.1364/boe.600665",
"PMID": "42460356",
"PMCID": "PMC13372336",
"ISSN": "2156-7085",
"publisher": "Optica Publishing Group",
"URL": "https://doi.org/10.1364/boe.600665",
"language": "en",
"issued": {
"date-parts": [
[
2026,
6,
17
]
]
}
}

The tracing map gets a citation of its own once an author has validated it and it has a DOI.

Similar papers

The papers with a page that share the most with this one: the tools found in their code, their categories, datasets, cited references and authors, the rarest counting most.

[1] doi:10.21037/qims-2026-0792 [code]
An nnU-Net-based framework with adaptive feature representation for 3D brain tumor segmentation.
Journal: Quantitative imaging in medicine and surgery
In common: tifffile, OpenCV, PyTorch, 5 other tools, methods / tools, 5 references
[2] doi:10.1371/journal.pcbi.1014571 [code]
SynAPSeg: A novel dataset and image analysis framework for deep learning-based synapse detection and quantification.
Journal: PLoS computational biology
In common: tifffile, TensorFlow, OpenCV, 7 other tools, 3 references
[3] doi:10.1364/boe.605322 [code]
Generalized plaque digitization framework for multi-dimensional mesoscopic images.
Journal: Biomedical optics express
In common: tifffile, TensorFlow, OpenCV, 7 other tools, 2 references
[4] doi:10.1016/j.isci.2026.116206 [code]
Gut distension evokes rapid neural dynamics in vagal and hindbrain populations of larval zebrafish.
Journal: iScience
In common: tifffile, OpenCV, Pillow, 6 other tools, optical imaging (calcium, voltage, 2-photon), 3 references
[5] doi:10.1371/journal.pcbi.1013441 [code]
Large vision model framework for automated C. elegans analysis: From static morphometry to dynamic neural activity.
Journal: PLoS computational biology
In common: tifffile, OpenCV, Pillow, 5 other tools, optical imaging (calcium, voltage, 2-photon), methods / tools, 2 references
[6] doi:10.1016/j.patter.2026.101538 [code]
A multi-modal foundation model for brain disease diagnosis and medical imaging.
Journal: Patterns (New York, N.Y.)
In common: tifffile, TensorFlow, OpenCV, 6 other tools, 2 references
[7] doi:10.1016/j.xpro.2026.104659 [code]
Protocol for simultaneous in vivo two-photon imaging and locomotion quantification during olfactory stimulation and pharmacology in walking Drosophila.
Journal: STAR protocols
In common: tifffile, OpenCV, Pillow, 5 other tools, optical imaging (calcium, voltage, 2-photon), 2 references
[8] doi:10.1016/j.isci.2026.117010 [code]
Deep learning-assisted mapping of dendritic spines using sequential 2D two-photon calcium imaging.
Journal: iScience
In common: tifffile, TensorFlow, OpenCV, 6 other tools, optical imaging (calcium, voltage, 2-photon), 1 reference
[9] doi:10.7554/elife.111876 [code]
Distinct sensorimotor encoding in tuft dendrites and somata associated with action, correction, and learning.
Journal: eLife
In common: tifffile, Pillow, PyTorch, 5 other tools, 3 references
[10] doi:10.1126/sciadv.adv3770 [code]
Evolution of a central dopamine circuit underlies adaptation of a light-evoked sensorimotor response in the blind cavefish.
Journal: Science advances
In common: tifffile, TensorFlow, OpenCV, 6 other tools, 1 reference

Contribute

The authors of this paper can claim it, correct its record and validate its tracing map, and the maintainers of its code (its owner, or a public member of its organization) correct what it says of their repository; anyone signed in can ask for its removal. Every request goes to OSCR's own machine, which answers it; your account page follows them.

Sign in with ORCID to claim this paper as one of its authors, correct its record or validate its tracing map: when the paper's metadata lists your ORCID iD, you are recognized at once. Maintainers of its code: sign in with GitHub, then claim the repository on your account page.

Request its removal

To ask OSCR to remove this record, the copies of its authors' scripts or its tracing map, use the removal request page: signed in, you say who you are, what to remove and why, then review and confirm the request. Published rules decide every request (how).

Discussion, reproductions, activity

Discussion: questions and error reports about this paper and its code, from signed-in readers and its authors. It opens with sign-in.

Reproductions: reports from readers who ran the authors' code: what they reproduced, with which environment, commit and data. It opens with sign-in.

Activity: what happens around this paper: new versions of its record, its map's validation, discussions and reproductions. It opens with sign-in.