From 3758fde1016c28ff89c9aaf1b3d20bb640dc408c Mon Sep 17 00:00:00 2001 From: luandalmazo Date: Wed, 29 Jul 2026 10:06:16 -0300 Subject: [PATCH] add _repr_utils and refactored base_entity --- datamint/_repr_utils.py | 80 ++++++++++++++++++ datamint/dataset/base.py | 36 +++++--- datamint/dataset/image_dataset.py | 5 -- datamint/dataset/sliced_dataset.py | 5 +- datamint/dataset/sliced_video_dataset.py | 4 - datamint/dataset/video_dataset.py | 4 - datamint/dataset/volume_dataset.py | 4 - datamint/entities/base_entity.py | 82 +------------------ datamint/lightning/trainers/base_trainer.py | 33 ++++++++ .../trainers/classification_trainer.py | 10 +++ datamint/lightning/trainers/seg2d_trainer.py | 5 ++ datamint/lightning/trainers/seg3d_trainer.py | 13 +++ .../trainers/specialized/deeplabv3plus.py | 10 +++ .../trainers/specialized/nnunet/trainer.py | 14 ++++ .../trainers/specialized/transunet.py | 7 ++ .../lightning/trainers/specialized/unetpp.py | 6 ++ .../lightning/trainers/specialized/unetrpp.py | 10 +++ .../lightning/trainers/specialized/yolox.py | 12 +++ .../lightning/trainers/vol_seg_trainer.py | 4 + 19 files changed, 232 insertions(+), 112 deletions(-) create mode 100644 datamint/_repr_utils.py diff --git a/datamint/_repr_utils.py b/datamint/_repr_utils.py new file mode 100644 index 00000000..0212a1b7 --- /dev/null +++ b/datamint/_repr_utils.py @@ -0,0 +1,80 @@ +"""Shared plain-text / Jupyter HTML repr rendering for entities, trainers, and datasets. + +Any class that can produce a ``(label, value)`` field list gets a consistent +`print()` block and a consistent HTML card in Jupyter for free. +""" + +# --------------------------------------------------------------------------- +# Jinja2 HTML template for the Jupyter card repr +# --------------------------------------------------------------------------- +_CARD_HTML_TEMPLATE = """\ +
+ + {# ---- Header ---- #} +
+
{{ kind }}
+
+

{{ name }}

+
+
+ + {# ---- Fields table ---- #} + {%- if fields %} +
+ + {%- for label, value in fields %} + + + + + {%- endfor %} +
{{ label }} + {{ value }} +
+
+ {%- else %} +
No non-empty fields to display.
+ {%- endif %} + +
+""" + +_card_template = None + + +def _get_card_template(): + """Lazily compile and cache the Jinja2 card template.""" + global _card_template + if _card_template is None: + from jinja2 import Environment + _card_template = Environment(autoescape=True).from_string(_CARD_HTML_TEMPLATE) + return _card_template + + +def render_text_block(header: str, fields: list[tuple[str, str]], empty_message: str = "(no non-empty fields)") -> str: + """Plain-text ``Header\\n Label: value`` block, used by ``__str__``/``__repr__``.""" + if not fields: + return f"{header}\n {empty_message}" + lines = [header] + [f" {label}: {value}" for label, value in fields] + return "\n".join(lines) + + +def render_html_card(kind: str, name: str, fields: list[tuple[str, str]]) -> str: + """Styled HTML card for Jupyter's ``_repr_html_`` display hook.""" + return _get_card_template().render(kind=kind, name=name, fields=fields) diff --git a/datamint/dataset/base.py b/datamint/dataset/base.py index 043e4eef..3a9fa1e6 100644 --- a/datamint/dataset/base.py +++ b/datamint/dataset/base.py @@ -16,6 +16,7 @@ import numpy as np from datamint.entities.annotation_worklist import AnnotationWorklist from datamint.exceptions import DatamintException, ItemNotFoundError +from datamint._repr_utils import render_text_block, render_html_card from .annotation_processor import AnnotationProcessor, MergeStrategy from datamint.entities.annotations.annotation_spec import AnnotationSpec, CategoryAnnotationSpec from datamint.entities.annotations import AnnotationType @@ -1225,19 +1226,22 @@ def subset(self, indices: list[int]) -> 'DatamintBaseDataset': raise IndexError(f"Subset indices out of bounds for dataset of length {len(self)}.") from e return new_ds - def __repr__(self) -> str: + def _extra_repr_fields(self) -> list[tuple[str, str]]: + """Hook for subclasses to insert extra ``(label, value)`` lines into the repr.""" + return [] + + def _repr_fields(self) -> list[tuple[str, str]]: name = self.project.name if self.project else "" - head = f"Dataset {name}" - body = [f"Number of datapoints: {len(self)}"] + fields = [ + ("Project", name), + ("Number of datapoints", str(len(self))), + ] if self.split_name is not None: - body.append(f"Split: {self.split_name}") + fields.append(("Split", str(self.split_name))) if self.split_source is not None: - body.append(f"Split source: {self.split_source}") + fields.append(("Split source", str(self.split_source))) if self.split_as_of_timestamp is not None: - body.append(f"Split as of: {self.split_as_of_timestamp}") - - # if self.manager.root is not None: - # body.append(f"Location: {self.manager.dataset_dir}") + fields.append(("Split as of", str(self.split_as_of_timestamp))) filters = [ (self.include_annotators, "Including annotators"), @@ -1249,13 +1253,19 @@ def __repr__(self) -> str: (self.include_frame_label_names, "Including frame labels"), (self.exclude_frame_label_names, "Excluding frame labels"), ] - for value, desc in filters: if value is not None: - body.append(f"{desc}: {value}") + fields.append((desc, str(value))) + + fields.extend(self._extra_repr_fields()) + return fields + + def __repr__(self) -> str: + return render_text_block(self.__class__.__name__, self._repr_fields()) - lines = [head] + [" " + line for line in body] - return "\n".join(lines) + def _repr_html_(self) -> str: + """HTML representation for Jupyter Notebooks.""" + return render_html_card(kind="Dataset", name=self.__class__.__name__, fields=self._repr_fields()) def split( self, diff --git a/datamint/dataset/image_dataset.py b/datamint/dataset/image_dataset.py index f53ad1c9..75c47d58 100644 --- a/datamint/dataset/image_dataset.py +++ b/datamint/dataset/image_dataset.py @@ -125,11 +125,6 @@ def apply_alb_transform( result['image'] = aug_img return result - @override - def __repr__(self) -> str: - base = super(VolumeDataset, self).__repr__() - return f"ImageDataset\n{base}" - def detection_collate_fn(batch: list[dict]) -> dict: """Collate a list of detection items into a batch. diff --git a/datamint/dataset/sliced_dataset.py b/datamint/dataset/sliced_dataset.py index c80ee2d8..38eb0fbc 100644 --- a/datamint/dataset/sliced_dataset.py +++ b/datamint/dataset/sliced_dataset.py @@ -584,6 +584,5 @@ def apply_alb_transform( result['masks'] = aug_segmentations return result - def __repr__(self) -> str: - base = super().__repr__() - return f"SlicedVolumeDataset (axis={self._slice_axis})\n{base}" + def _extra_repr_fields(self) -> list[tuple[str, str]]: + return [*super()._extra_repr_fields(), ("Slice axis", str(self._slice_axis))] diff --git a/datamint/dataset/sliced_video_dataset.py b/datamint/dataset/sliced_video_dataset.py index b76acb91..7621b8bb 100644 --- a/datamint/dataset/sliced_video_dataset.py +++ b/datamint/dataset/sliced_video_dataset.py @@ -454,7 +454,3 @@ def apply_alb_transform( 'image': aug_img, 'segmentations': aug_segmentations, } - - def __repr__(self) -> str: - base = super().__repr__() - return f"SlicedVideoDataset\n{base}" diff --git a/datamint/dataset/video_dataset.py b/datamint/dataset/video_dataset.py index 48ba7e7b..c1b75c93 100644 --- a/datamint/dataset/video_dataset.py +++ b/datamint/dataset/video_dataset.py @@ -36,10 +36,6 @@ class VideoDataset(MultiFrameDataset): print(frame_ds[0]['image'].shape) # (C, H, W) """ - def __repr__(self) -> str: - base = super().__repr__() - return f"VideoDataset\n{base}" - def frame_by_frame(self) -> 'SlicedVideoDataset': """Create a 2D dataset iterating over individual video frames. diff --git a/datamint/dataset/volume_dataset.py b/datamint/dataset/volume_dataset.py index 6f8dacd0..9492cc26 100644 --- a/datamint/dataset/volume_dataset.py +++ b/datamint/dataset/volume_dataset.py @@ -30,10 +30,6 @@ class VolumeDataset(MultiFrameDataset): Inherits multi-frame loading and augmentation from :class:`MultiFrameDataset`. """ - def __repr__(self) -> str: - base = super().__repr__() - return f"VolumeDataset\n{base}" - def slice(self, axis: str | int = 'axial') -> 'SlicedVolumeDataset': """Create a 2D dataset by slicing this volume along an axis. diff --git a/datamint/entities/base_entity.py b/datamint/entities/base_entity.py index 16b5fdf9..da6473f6 100644 --- a/datamint/entities/base_entity.py +++ b/datamint/entities/base_entity.py @@ -6,6 +6,7 @@ from pydantic import BaseModel, ConfigDict, PrivateAttr from datamint.types import CacheMode +from datamint._repr_utils import render_text_block, render_html_card if TYPE_CHECKING: from datamint.api.entity_base_api import EntityBaseApi @@ -21,68 +22,6 @@ # Track logged warnings to avoid duplicates _LOGGED_WARNINGS: set[tuple[str, str]] = set() -# --------------------------------------------------------------------------- -# Jinja2 HTML template for BaseEntity Jupyter repr -# --------------------------------------------------------------------------- -_ENTITY_HTML_TEMPLATE = """\ -
- - {# ---- Header ---- #} -
-
Entity
-
-

{{ entity_name }}

-
-
- - {# ---- Fields table ---- #} - {%- if fields %} -
- - {%- for name, value in fields %} - - - - - {%- endfor %} -
{{ name }} - {{ value }} -
-
- {%- else %} -
No non-empty fields to display.
- {%- endif %} - -
-""" - -_entity_template = None - - -def _get_entity_template(): - """Lazily compile and cache the Jinja2 entity template.""" - global _entity_template - if _entity_template is None: - from jinja2 import Environment - _entity_template = Environment(autoescape=True).from_string(_ENTITY_HTML_TEMPLATE) - return _entity_template - class BaseEntityModel(BaseModel): """Shared lightweight Pydantic base for Datamint entities and DTOs.""" @@ -128,25 +67,10 @@ def _get_display_fields(self, max_value_len: int = 120) -> list[tuple[str, str]] def _repr_html_(self) -> str: """HTML representation for Jupyter Notebooks.""" - entity_id = getattr(self, 'id', None) - fields = self._get_display_fields() - - return _get_entity_template().render( - entity_name=self.__class__.__name__, - entity_id=str(entity_id) if entity_id else None, - fields=fields, - ) + return render_html_card(kind='Entity', name=self.__class__.__name__, fields=self._get_display_fields()) def __str__(self) -> str: - fields = self._get_display_fields() - - header = self.__class__.__name__ - - if not fields: - return f"{header}\n (no non-empty fields)" - - lines = [header] + [f" {name}: {value}" for name, value in fields] - return "\n".join(lines) + return render_text_block(self.__class__.__name__, self._get_display_fields()) def __init__(self, **data): super().__init__(**data) diff --git a/datamint/lightning/trainers/base_trainer.py b/datamint/lightning/trainers/base_trainer.py index fb468fcf..c0d7e6ea 100644 --- a/datamint/lightning/trainers/base_trainer.py +++ b/datamint/lightning/trainers/base_trainer.py @@ -20,6 +20,7 @@ from datamint.mlflow import set_project from datamint.mlflow.flavors.model import BaseDatamintModel from datamint.lightning.trainers.lightning_modules.base import DatamintLightningModule +from datamint._repr_utils import render_text_block, render_html_card if TYPE_CHECKING: from albumentations import BaseCompose @@ -137,6 +138,38 @@ def _project_name(self) -> str: def experiment_name(self) -> str: return self.mlflow_experiment_name or f"{self._project_name}_training" + def _model_description(self) -> str: + """Short human-readable model description for :meth:`__repr__`. Override per architecture.""" + if self._user_model is not None: + model = self._user_model + cls = model if isinstance(model, type) else model.__class__ + return f"Custom ({cls.__name__})" + return self.__class__.__name__.removesuffix("Trainer") + + def _extra_repr_fields(self) -> list[tuple[str, str]]: + """Architecture-specific ``(label, value)`` lines inserted between Batch size and Early stopping patience.""" + return [] + + def _repr_fields(self) -> list[tuple[str, str]]: + """Fields shown by :meth:`__repr__`/:meth:`_repr_html_`. Cheap and side-effect free: never resolves the dataset or builds the model.""" + return [ + ("Project", self._project_name), + ("Model", self._model_description()), + ("Max epochs", str(self.max_epochs)), + ("Batch size", str(self.batch_size)), + *self._extra_repr_fields(), + ("Early stopping patience", str(self.early_stopping_patience) if self.early_stopping_patience else "disabled"), + ("MLflow experiment", self.experiment_name), + ("Auto-deploy adapter", "enabled" if self.auto_deploy_adapter else "disabled"), + ] + + def __repr__(self) -> str: + return render_text_block(self.__class__.__name__, self._repr_fields()) + + def _repr_html_(self) -> str: + """HTML representation for Jupyter Notebooks.""" + return render_html_card(kind="Trainer", name=self.__class__.__name__, fields=self._repr_fields()) + def _with_project(self): set_project(self._project_name) diff --git a/datamint/lightning/trainers/classification_trainer.py b/datamint/lightning/trainers/classification_trainer.py index 943b0501..17c077bb 100644 --- a/datamint/lightning/trainers/classification_trainer.py +++ b/datamint/lightning/trainers/classification_trainer.py @@ -95,6 +95,16 @@ def __init__( else: self.image_size = image_size + def _model_description(self) -> str: + if self._user_model is not None: + return super()._model_description() + pretrained_str = "pretrained" if self.pretrained else "random init" + return f"{self.architecture} ({pretrained_str})" + + def _extra_repr_fields(self) -> list[tuple[str, str]]: + image_size = f"{self.image_size[0]}×{self.image_size[1]}" if self.image_size else "auto (no resize)" + return [*super()._extra_repr_fields(), ("Image size", image_size)] + # ── Template hooks ────────────────────────────────────────── def _build_dataset(self, project: 'str | Project', **kwargs: Any) -> ImageDataset: diff --git a/datamint/lightning/trainers/seg2d_trainer.py b/datamint/lightning/trainers/seg2d_trainer.py index 13acfc97..8232bb85 100644 --- a/datamint/lightning/trainers/seg2d_trainer.py +++ b/datamint/lightning/trainers/seg2d_trainer.py @@ -77,6 +77,11 @@ def __init__( else: self.image_size = image_size + @override + def _extra_repr_fields(self) -> list[tuple[str, str]]: + image_size = f"{self.image_size[0]}×{self.image_size[1]}" if self.image_size else "auto (no resize)" + return [*super()._extra_repr_fields(), ("Image size", image_size)] + def _build_dataset(self, project: 'str | Project', **kwargs: Any) -> ImageDataset | SlicedVolumeDataset: default_params = dict( return_as_semantic_segmentation=True, diff --git a/datamint/lightning/trainers/seg3d_trainer.py b/datamint/lightning/trainers/seg3d_trainer.py index 722dd476..52530ddb 100644 --- a/datamint/lightning/trainers/seg3d_trainer.py +++ b/datamint/lightning/trainers/seg3d_trainer.py @@ -57,6 +57,19 @@ def __init__( else: self.image_size = image_size + def _model_description(self) -> str: + if self._user_model is not None: + return super()._model_description() + return f"{self.encoder_name} encoder" + + def _extra_repr_fields(self) -> list[tuple[str, str]]: + image_size = f"{self.image_size[0]}×{self.image_size[1]}" if self.image_size else "auto (original slice size)" + return [ + *super()._extra_repr_fields(), + ("Slice axis", str(self.slice_axis)), + ("Image size", image_size), + ] + # ── Template hooks ────────────────────────────────────────── def _build_dataset(self, project: 'str | Project', **kwargs: Any): diff --git a/datamint/lightning/trainers/specialized/deeplabv3plus.py b/datamint/lightning/trainers/specialized/deeplabv3plus.py index 4e9e2b56..00746c33 100644 --- a/datamint/lightning/trainers/specialized/deeplabv3plus.py +++ b/datamint/lightning/trainers/specialized/deeplabv3plus.py @@ -127,6 +127,16 @@ def __init__( self.encoder_name = encoder_name self.decoder_atrous_rates = decoder_atrous_rates + @override + def _model_description(self) -> str: + if self._user_model is not None: + return super()._model_description() + return f"DeepLabV3+ ({self.encoder_name} encoder)" + + @override + def _extra_repr_fields(self) -> list[tuple[str, str]]: + return [*super()._extra_repr_fields(), ("Decoder atrous rates", str(self.decoder_atrous_rates))] + @override def _build_model( self, diff --git a/datamint/lightning/trainers/specialized/nnunet/trainer.py b/datamint/lightning/trainers/specialized/nnunet/trainer.py index c2b62ccf..06e3613d 100644 --- a/datamint/lightning/trainers/specialized/nnunet/trainer.py +++ b/datamint/lightning/trainers/specialized/nnunet/trainer.py @@ -84,6 +84,20 @@ def __init__( self.channel_names = channel_names or {'0': 'CT'} self.num_processes_preprocessing = num_processes_preprocessing + def _repr_fields(self) -> list[tuple[str, str]]: + # nnUNet bypasses the Lightning pipeline entirely, so batch size, + # early stopping, and auto-deploy-adapter from BaseTrainer don't + # apply here (nnUNet manages batch size internally via its plans, + # has no early stopping, and always builds a deploy adapter). + return [ + ("Project", self._project_name), + ("Model", f"nnU-Net ({self.configuration})"), + ("Fold", str(self.fold)), + ("Max epochs", str(self.max_epochs)), + ("Continue training", "yes" if self.continue_training else "no"), + ("MLflow experiment", self.experiment_name), + ] + # ── BaseTrainer abstract methods bypassed by nnUNet ─────────────────────── def _build_dataset(self, project: 'str | Project', **kwargs) -> VolumeDataset: diff --git a/datamint/lightning/trainers/specialized/transunet.py b/datamint/lightning/trainers/specialized/transunet.py index 55f3a418..a7d956b4 100644 --- a/datamint/lightning/trainers/specialized/transunet.py +++ b/datamint/lightning/trainers/specialized/transunet.py @@ -137,6 +137,13 @@ def __init__( self.variant = variant self.pretrained = pretrained + @override + def _model_description(self) -> str: + if self._user_model is not None: + return super()._model_description() + pretrained_str = "pretrained" if self.pretrained else "random init" + return f"TransUNet ({self.variant}, {pretrained_str})" + @override def _train_transform(self) -> 'BaseCompose': import albumentations as A diff --git a/datamint/lightning/trainers/specialized/unetpp.py b/datamint/lightning/trainers/specialized/unetpp.py index 1f9bf05b..72df57d0 100644 --- a/datamint/lightning/trainers/specialized/unetpp.py +++ b/datamint/lightning/trainers/specialized/unetpp.py @@ -121,6 +121,12 @@ def __init__( ) self.encoder_name = encoder_name + @override + def _model_description(self) -> str: + if self._user_model is not None: + return super()._model_description() + return f"UNet++ ({self.encoder_name} encoder)" + @override def _train_transform(self) -> 'BaseCompose': import albumentations as A diff --git a/datamint/lightning/trainers/specialized/unetrpp.py b/datamint/lightning/trainers/specialized/unetrpp.py index d9099cbc..8f41ea84 100644 --- a/datamint/lightning/trainers/specialized/unetrpp.py +++ b/datamint/lightning/trainers/specialized/unetrpp.py @@ -109,6 +109,16 @@ def __init__( self.sw_overlap = sw_overlap self.in_channels = in_channels + @override + def _model_description(self) -> str: + if self._user_model is not None: + return super()._model_description() + return f"UNETR++ (feature size {self.feature_size}, {self.num_heads} heads)" + + @override + def _extra_repr_fields(self) -> list[tuple[str, str]]: + return [*super()._extra_repr_fields(), ("Depths", str(self.depths)), ("SW overlap", str(self.sw_overlap))] + @override def _build_model( self, diff --git a/datamint/lightning/trainers/specialized/yolox.py b/datamint/lightning/trainers/specialized/yolox.py index dbd2475c..39ed515a 100644 --- a/datamint/lightning/trainers/specialized/yolox.py +++ b/datamint/lightning/trainers/specialized/yolox.py @@ -108,6 +108,18 @@ def __init__( else: self.image_size = image_size + def _model_description(self) -> str: + if self._user_model is not None: + return super()._model_description() + return f"YOLOX-{self.model_size}" + + def _extra_repr_fields(self) -> list[tuple[str, str]]: + return [ + *super()._extra_repr_fields(), + ("Image size", f"{self.image_size[0]}×{self.image_size[1]}"), + ("Conf / NMS threshold", f"{self.conf_thre} / {self.nms_thre}"), + ] + # ------------------------------------------------------------------ # Transforms # ------------------------------------------------------------------ diff --git a/datamint/lightning/trainers/vol_seg_trainer.py b/datamint/lightning/trainers/vol_seg_trainer.py index 5ae96d63..742edfec 100644 --- a/datamint/lightning/trainers/vol_seg_trainer.py +++ b/datamint/lightning/trainers/vol_seg_trainer.py @@ -90,6 +90,10 @@ def __init__( super().__init__(batch_size=batch_size, **kwargs) self.patch_crop_size = patch_crop_size + def _extra_repr_fields(self) -> list[tuple[str, str]]: + d, h, w = self.patch_crop_size + return [*super()._extra_repr_fields(), ("Patch crop size", f"{d}×{h}×{w}")] + # ── Template hooks ─────────────────────────────────────────── def _build_dataset(self, project: 'str | Project', **kwargs: Any) -> VolumeDataset: