| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| """Base classes for Metrics.""" |
| import dataclasses |
| from typing import Callable |
| from dinosaur import typing |
| import jax |
| import jax.numpy as jnp |
| import model.reference_code.metrics_util as metrics_util |
|
|
|
|
| Pytree = typing.Pytree |
| TrajectoryRepresentations = typing.TrajectoryRepresentations |
|
|
| tree_leaves = jax.tree_util.tree_leaves |
| tree_map = jax.tree_util.tree_map |
|
|
|
|
| @dataclasses.dataclass |
| class Evaluator: |
| """Class that evaluates on (prediction, trajectory) returning Pytree.""" |
|
|
| def evaluate( |
| self, |
| prediction: TrajectoryRepresentations, |
| target: TrajectoryRepresentations, |
| ) -> Pytree: |
| """Evaluates giving values of interest.""" |
| raise NotImplementedError() |
|
|
|
|
| @dataclasses.dataclass |
| class EvaluateFunctionWrapper(Evaluator): |
| """Wraps `evaluate_fn` function to be used as an Evaluator.""" |
|
|
| def __init__( |
| self, |
| evaluate_fn: Callable[ |
| [TrajectoryRepresentations, TrajectoryRepresentations], Pytree |
| ], |
| ): |
| self._evaluate_fn = evaluate_fn |
|
|
| def evaluate( |
| self, |
| prediction: TrajectoryRepresentations, |
| target: TrajectoryRepresentations, |
| ) -> Pytree: |
| return self._evaluate_fn(prediction, target) |
|
|
|
|
| class MetricRuntimeError(Exception): |
| """Generic error for Metrics to raise in place of generic RuntimeError.""" |
|
|
|
|
| @dataclasses.dataclass |
| class Metric(Evaluator): |
| """An Evaluator that derives information from a TrajectorySpec.""" |
|
|
| trajectory_spec: metrics_util.TrajectorySpec |
| is_nodal: bool = dataclasses.field(default=True, kw_only=True) |
| is_encoded: bool = dataclasses.field(default=False, kw_only=True) |
|
|
| def get_representation(self, x: TrajectoryRepresentations) -> Pytree: |
| x_rep = x.get_representation( |
| is_nodal=self.is_nodal, is_encoded=self.is_encoded |
| ) |
| if x_rep is None: |
| raise MetricRuntimeError( |
| 'Desired representation of `x` was None. ' |
| f'{self.is_nodal=}, {self.is_encoded=}' |
| ) |
| return x_rep |
|
|
| def surface_mean(self, trajectory: Pytree) -> Pytree: |
| if self.is_encoded: |
| coords = self.trajectory_spec.coords |
| else: |
| coords = self.trajectory_spec.data_coords |
| if self.is_nodal: |
| |
| |
| fn = lambda x: metrics_util.nodal_surface_mean(x, coords) |
| else: |
| fn = lambda x: metrics_util.modal_surface_mean(x, coords) |
| return tree_map(fn, trajectory) |
|
|
| def mean_per_variable(self, trajectory: Pytree) -> Pytree: |
| |
| return tree_map(jnp.mean, self.surface_mean(trajectory)) |
|
|
|
|
| class ScalarMetric(Metric): |
| """Metric that compute scalar quantities.""" |
|
|
|
|
| @dataclasses.dataclass |
| class Loss(ScalarMetric): |
| """Metric that can be used as a loss.""" |
|
|
| trajectory_spec: metrics_util.TrajectorySpec |
| is_nodal: bool = dataclasses.field(default=True, kw_only=True) |
| is_encoded: bool = dataclasses.field(default=False, kw_only=True) |
| time_step: int | slice | None = dataclasses.field(default=None, kw_only=True) |
|
|
| def evaluate_per_variable( |
| self, |
| prediction: TrajectoryRepresentations, |
| target: TrajectoryRepresentations, |
| ) -> Pytree: |
| raise NotImplementedError() |
|
|
| def evaluate( |
| self, |
| prediction: TrajectoryRepresentations, |
| target: TrajectoryRepresentations, |
| ) -> jnp.ndarray: |
| error_per_variable = self.evaluate_per_variable(prediction, target) |
| return sum(tree_leaves(error_per_variable)) |
|
|
| def debug_loss_terms_instance(self) -> EvaluateFunctionWrapper: |
| """Returns class that evaluates relative loss per variable.""" |
|
|
| def evaluate_fn( |
| prediction: TrajectoryRepresentations, |
| target: TrajectoryRepresentations, |
| ) -> Pytree: |
| |
| |
| loss_per_variable = self.evaluate_per_variable(prediction, target) |
| |
| |
| sum_of_all_terms = sum(tree_leaves(loss_per_variable)) |
| relative_loss = tree_map( |
| lambda x: x / sum_of_all_terms, loss_per_variable |
| ) |
| return {'relative_loss': relative_loss} |
|
|
| return EvaluateFunctionWrapper(evaluate_fn) |
|
|