Skip to content

configure_metrics default kwargs causes errors as default with certain metrics #45

Description

@X02b3ar
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

Metadata

Metadata

Assignees

Labels

bugSomething isn't working

Type

No type

Projects

No projects

Milestone

No milestone

Relationships

None yet

Development

No branches or pull requests

Issue actions