diff --git a/src/fairseq2/metrics/recorders/descriptor.py b/src/fairseq2/metrics/recorders/descriptor.py index 9def7254c..6b716a118 100644 --- a/src/fairseq2/metrics/recorders/descriptor.py +++ b/src/fairseq2/metrics/recorders/descriptor.py @@ -17,6 +17,11 @@ def __call__(self, value: object) -> str: ... @dataclass class MetricDescriptor: + """ + Represents a description of a metric including high level name, + name to display, and formatting + """ + name: str display_name: str priority: int @@ -32,8 +37,15 @@ class MetricDescriptor: @final class MetricDescriptorRegistry: + """ + Represents a way to store descriptors for multiple metrics in a composite metric + """ + def __init__(self, descriptors: Iterable[MetricDescriptor]) -> None: self._descriptors = {d.name: d for d in descriptors} def maybe_get(self, name: str) -> MetricDescriptor | None: + """ + Returns a metric descriptor if it exists + """ return self._descriptors.get(name) diff --git a/src/fairseq2/metrics/recorders/jsonl.py b/src/fairseq2/metrics/recorders/jsonl.py index 58c6619f2..bc56022f6 100644 --- a/src/fairseq2/metrics/recorders/jsonl.py +++ b/src/fairseq2/metrics/recorders/jsonl.py @@ -27,7 +27,7 @@ @final class JsonlMetricRecorder(MetricRecorder): - """Records metric values to JSONL files.""" + """Records metric values to JSONL files""" _CATEGORY_PART_REGEX: Final = re.compile("^[-_a-zA-Z0-9]+$") diff --git a/src/fairseq2/metrics/recorders/log.py b/src/fairseq2/metrics/recorders/log.py index 0aef7be76..1342f5ecc 100644 --- a/src/fairseq2/metrics/recorders/log.py +++ b/src/fairseq2/metrics/recorders/log.py @@ -32,6 +32,12 @@ def __init__(self, metric_descriptors: MetricDescriptorRegistry) -> None: def record_metric_values( self, category: str, values: Mapping[str, object], step_nr: int | None = None ) -> None: + """ + Retrieves names, descriptors, and values for all metrics from a :class:`Mapping`. + Stores and sorts (values, descriptors) as ``tuples`` and formats a ``str`` + with | separated tabular values. Splits metrics into multiple category parts + where they exist. + """ if not log.is_enabled_for_info(): return @@ -81,4 +87,7 @@ def record_metric_values( @override def close(self) -> None: + """ + Close the :class:`LogMetricRecorder` + """ pass diff --git a/src/fairseq2/metrics/recorders/tensorboard.py b/src/fairseq2/metrics/recorders/tensorboard.py index 14e88d7c7..42f653962 100644 --- a/src/fairseq2/metrics/recorders/tensorboard.py +++ b/src/fairseq2/metrics/recorders/tensorboard.py @@ -51,6 +51,13 @@ def __init__( def record_metric_values( self, category: str, values: Mapping[str, object], step_nr: int | None = None ) -> None: + """ + Retrieves TensorBoard :class:`SummaryWriter` ``objects`` for the metric scalars. + Maps descriptors to metric, step number, name, and scalar value and + flushes output + :raises OSError: If an operational system error occurs (FileNotFound, PermissionError, ConnectionError) + :raises RuntimeError: If the values cannot be written + """ writer = self._get_writer(category) try: @@ -74,6 +81,11 @@ def record_metric_values( def _add_value( self, writer: SummaryWriter, step_nr: int | None, name: str, value: object ) -> None: + """ + Adds a `value` ``object`` to the `SummaryWriter`. + Can be text, a scalar, or torch Tensor + :raises ValueError: If `value` is not of type `int`, `float`, `Tensor` or `str` + """ if isinstance(value, str): writer.add_text(name, value, step_nr) @@ -95,6 +107,10 @@ def _add_value( ) def _get_writer(self, category: str) -> SummaryWriter: + """ + Instantiates a `SummaryWriter` and adds it to the ``writers`` `dict` + Stores `writer` under its category as key + """ writer = self._writers.get(category) if writer is None: path = self._output_dir.joinpath(category) @@ -107,6 +123,9 @@ def _get_writer(self, category: str) -> SummaryWriter: @override def close(self) -> None: + """ + Close out and destroy every ``writer`` + """ for writer in self._writers.values(): writer.close() diff --git a/src/fairseq2/metrics/recorders/wandb.py b/src/fairseq2/metrics/recorders/wandb.py index c852dfd8d..c2a114120 100644 --- a/src/fairseq2/metrics/recorders/wandb.py +++ b/src/fairseq2/metrics/recorders/wandb.py @@ -40,6 +40,11 @@ def __init__( def record_metric_values( self, category: str, values: Mapping[str, object], step_nr: int | None = None ) -> None: + """ + Retrieves and stores a `descriptor` and `value` as a ``Mapping`` for each metric + :raises OSError: If an operational system error occurs (file not found, permission issue, connection problem) + :raises RuntimeError: If unable to write the metrics to wandb + """ output: dict[str, object] = {} for name, value in values.items(): @@ -61,6 +66,10 @@ def record_metric_values( ) from ex def _add_value(self, name: str, value: object, output: dict[str, object]) -> None: + """ + Adds a value to the output dictionary + :raises ValueError: If `values` are not of type int, float, Tensor or str + """ if isinstance(value, (int, float, Tensor, str)): output[name] = value @@ -78,4 +87,7 @@ def _add_value(self, name: str, value: object, output: dict[str, object]) -> Non @override def close(self) -> None: + """ + Close wandb logger and end run + """ self._run.finish()