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 %}
+
+ | {{ label }} |
+
+ {{ value }}
+ |
+
+ {%- endfor %}
+
+
+ {%- 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 %}
-
- | {{ name }} |
-
- {{ value }}
- |
-
- {%- endfor %}
-
-
- {%- 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: