Skip to content

ProgramOperationalMetric metric

Bases: Metric

Base class for program-wide runtime-counter metrics.

Subclasses set _phase to one of "inference", "reward", or "optimizer" to read the corresponding counter set. Counters are populated based on the active op_scope (contextvar) the trainer sets for each phase.

Example:

program.compile(
    metrics=[
        synalinks.metrics.ProgramOperationalMetric(),
    ],
)
Source code in synalinks/src/metrics/program_metrics.py
@synalinks_export(
    [
        "synalinks.metrics.ProgramOperationalMetric",
        "synalinks.ProgramOperationalMetric",
    ]
)
class ProgramOperationalMetric(Metric):
    """Base class for program-wide runtime-counter metrics.

    Subclasses set `_phase` to one of ``"inference"``, ``"reward"``, or
    ``"optimizer"`` to read the corresponding counter set. Counters are
    populated based on the active ``op_scope`` (contextvar) the trainer
    sets for each phase.

    Example:

    ```python
    program.compile(
        metrics=[
            synalinks.metrics.ProgramOperationalMetric(),
        ],
    )
    ```
    """

    _phase = "inference"
    direction = "down"

    def __init__(self, name=None):
        super().__init__(name=name)
        self._program = None
        self._language_models = []
        self._embedding_models = []
        self._program_baselines = {suffix: 0 for suffix in _PROGRAM_SUFFIXES}
        self._model_cost_baseline = 0.0

    @property
    def program(self):
        return self._program

    @property
    def language_models(self):
        return list(self._language_models)

    @property
    def embedding_models(self):
        return list(self._embedding_models)

    def bind_program(self, program):
        self._program = program
        self._language_models = _collect_language_models(program)
        self._embedding_models = _collect_embedding_models(program)
        self._snapshot()

    def _attr(self, suffix):
        return f"{self._phase}_cumulated_{suffix}"

    def _read_program(self, suffix):
        if self._program is None:
            return 0
        return getattr(self._program, self._attr(suffix), 0)

    def _read_model_cost(self):
        attr = self._attr("cost")
        lm_cost = sum(getattr(lm, attr, 0) for lm in self._language_models)
        em_cost = sum(getattr(em, attr, 0) for em in self._embedding_models)
        return lm_cost + em_cost

    def _snapshot(self):
        for suffix in _PROGRAM_SUFFIXES:
            self._program_baselines[suffix] = self._read_program(suffix)
        self._model_cost_baseline = self._read_model_cost()

    def _delta_program(self, suffix):
        return self._read_program(suffix) - self._program_baselines.get(suffix, 0)

    def _delta_model_cost(self):
        return self._read_model_cost() - self._model_cost_baseline

    def reset_state(self):
        self._snapshot()

    async def update_state(self, *args, **kwargs):
        return

    def result(self):
        raise NotImplementedError

    def get_config(self):
        return {"name": self.name}