Custom Metrics
Each task type ships with sensible default metrics (listed on the individual Model Configuration pages). Passing your own metrics replaces the defaults rather than adding to them, so include any default metric you still want. You can also change which metric is monitored for early stopping and checkpointing.
Adding Metrics
Pass a list of MetricParams or CustomMetric objects to TrainingParams.metrics.
MetricParams
Use MetricParams when the metric is available by name in BaseModel's predefined set or in TorchMetrics:
from monad.config import TrainingParams, MetricParams
training_params = TrainingParams(
checkpoint_dir="./model",
epochs=1,
metrics=[
MetricParams(alias="auroc", metric_name="AUROC", kwargs={"task": "binary", "average": None}),
MetricParams(alias="recall", metric_name="Recall", kwargs={"task": "binary"}),
],
)
| Field | Description |
|---|---|
alias | Name you assign to identify the metric in logs |
metric_name | Name from BaseModel's predefined metrics or TorchMetrics |
kwargs | Arguments passed to the metric constructor |
CustomMetric
Use CustomMetric when you need a fully initialized TorchMetrics instance:
from monad.config import CustomMetric, MetricMonitoringMode, TrainingParams
from torchmetrics.classification import F1Score
training_params = TrainingParams(
checkpoint_dir="./model",
epochs=1,
metrics=[
CustomMetric(alias="f1", metric=F1Score(task="binary")),
],
metric_to_monitor="val_f1_0",
metric_monitoring_mode=MetricMonitoringMode.MAX,
)
Custom metrics replace the task's default metrics. Here the binary task's default monitored key, val_auroc_0, is no longer logged, so the example monitors val_f1_0 instead; without it, fit() would raise InvalidConfigurationError before training starts because checkpoint_dir is set. See Monitoring a Metric.
| Field | Description |
|---|---|
alias | Name you assign to identify the metric in logs |
metric | An initialized torchmetrics.Metric instance |
Monitoring a Metric
Control which metric determines the best checkpoint and early stopping:
from monad.config import TrainingParams, MetricParams, MetricMonitoringMode
training_params = TrainingParams(
checkpoint_dir="./model",
epochs=3,
metrics=[
MetricParams(alias="auroc", metric_name="AUROC", kwargs={"task": "binary", "average": None}),
MetricParams(alias="recall", metric_name="Recall", kwargs={"task": "binary"}),
],
metric_to_monitor="val_recall_0",
metric_monitoring_mode=MetricMonitoringMode.MAX,
)
| Parameter | Description |
|---|---|
metric_to_monitor | Logged validation key to track: val_<alias>_<target_index> (for example val_recall_0) or val_loss. When checkpoint_dir or early_stopping is set and the key is not logged during validation, fit() raises InvalidConfigurationError before training starts. |
metric_monitoring_mode | MetricMonitoringMode.MAX (higher is better) or MetricMonitoringMode.MIN (lower is better) |
Predefined Metrics
These metrics are available by name via MetricParams.metric_name:
| Metric name | Task | Description |
|---|---|---|
MultipleTargetsRecall | Multiclass | Fraction of correct class predictions |
MultipleTargetsRecallPerClass | Multiclass | Recall per class |
PrecisionAtK | Recommendation | Fraction of top-K items that are relevant |
MeanAveragePrecisionAtK | Recommendation | Average precision within top-K, averaged over entities |
HitRateAtK | Recommendation | Fraction of entities where top-K contains at least one hit |
MeanReciprocalRank | Recommendation | How early the first relevant item appears |
NDCGAtK | Recommendation | Ranking quality for top-K, normalized to [0, 1] |
Any metric from TorchMetrics can also be used via either MetricParams or CustomMetric.