def configure_metrics(self, metric, metric_kwargs):
"""Configure metrics
Keeping this a function allows for the metric to be reconfigured
in inherited classes TODO: add support for multiple metrics
Returns:
----------
torchmetrics.Metric:
metric
"""
metric_name = DEFAULT_TASK_METRICS[self.task] if metric is None else metric
metric_kwargs = (
metric_kwargs
if metric_kwargs is not None
else DEFAULT_METRIC_KWARGS[self.task]
)
metric = METRIC_REGISTRY[metric_name](
num_outputs_=self.output_dim, **metric_kwargs
)
return metric, metric_kwargs, metric_name
num_outputs may have some fault