Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
142 changes: 105 additions & 37 deletions plugins/nemo-evaluator/src/nemo_evaluator/jobs/evaluate.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,12 +13,16 @@
from nemo_evaluator.jobs.utils import resolve_run_dataset
from nemo_evaluator.resolvers import PlatformModelResolver
from nemo_evaluator.sdk.values.filesets import FilesetRef
from nemo_evaluator.shared.metric_bundles.bundles import MetricBundle, unbundle_metric
from nemo_evaluator.shared.metric_bundles.bundles import (
MetricBundle,
bundle_metric,
metric_bundle_packager_for_payload,
unbundle_metric,
)
from nemo_evaluator.shared.metric_bundles.cloudpickle import CloudpickleMetricPayload # noqa: F401
from nemo_evaluator_sdk import Evaluator
from nemo_evaluator_sdk.execution._protocols import JobParamsConfigurableMetric
from nemo_evaluator_sdk.execution.config import normalize_params
from nemo_evaluator_sdk.execution.metric_execution import run_sync
from nemo_evaluator_sdk.metrics.protocol import Metric, MetricWithModels
from nemo_evaluator_sdk.values import (
Agent,
Expand Down Expand Up @@ -61,8 +65,51 @@ class EvaluationResultFiles:
artifacts_dir: Path


class EvaluateSpec(BaseModel):
"""Inline SDK evaluation input for the first evaluator plugin job."""
def _hydrate_metrics(metrics: list[MetricBundle]) -> list[Metric]:
return [unbundle_metric(bundle) for bundle in metrics]


def _unresolved_model_refs(metrics: list[Metric]) -> list[str]:
refs = [
model_ref.root
for item in metrics
if isinstance(item, MetricWithModels)
for model_ref in item.model_refs().values()
]
return sorted(refs)


async def _resolve_metric_models(
metrics: list[Metric],
resolver: PlatformModelResolver,
) -> None:
"""Resolve ModelRef fields on metric configs before SDK execution."""
for item in metrics:
if isinstance(item, MetricWithModels):
await item.resolve_models(resolver)


def _apply_metric_job_params(
metrics: list[Metric],
params: RunConfig | RunConfigOnline | RunConfigOnlineModel,
) -> bool:
"""Apply evaluation job params to metrics that support runtime configuration."""
applied = False
for item in metrics:
if isinstance(item, JobParamsConfigurableMetric):
item.apply_evaluation_job_params(params)
applied = True
return applied


def _bundle_resolved_metric(metric: Metric, source_bundle: MetricBundle) -> MetricBundle:
packager = metric_bundle_packager_for_payload(source_bundle.payload)
resolved_bundle = bundle_metric(metric, packager)
return resolved_bundle.model_copy(update={"metadata": source_bundle.metadata})


class EvaluateInputSpec(BaseModel):
"""Submitter-facing SDK evaluation input for the evaluator plugin job."""

model_config = ConfigDict(extra="forbid")

Expand All @@ -84,38 +131,29 @@ def normalize_params_for_target(self) -> Self:
return self


class EvaluateSpec(EvaluateInputSpec):
"""Canonical SDK evaluation spec with platform model references resolved."""

@model_validator(mode="after")
def reject_unresolved_metric_model_refs(self) -> Self:
unresolved_refs = _unresolved_model_refs(_hydrate_metrics(self.metrics))
if unresolved_refs:
raise ValueError(
"EvaluateSpec metric models must be resolved before compile/run: " + ", ".join(unresolved_refs)
)
return self


class EvaluateJob(NemoJob):
"""Run one evaluator SDK metric against inline rows."""

name: ClassVar[str] = "evaluate"
description: ClassVar[str] = "Run an inline evaluator SDK metric against inline dataset rows."
container: ClassVar[str] = "cpu-tasks"
input_spec_schema: ClassVar[type[BaseModel] | None] = EvaluateInputSpec
spec_schema: ClassVar[type[BaseModel] | None] = EvaluateSpec
job_collection_path: ClassVar[str | None] = "/evaluate/jobs"

@staticmethod
async def _resolve_metric_models(
metrics: list[Metric],
resolver: PlatformModelResolver,
params: RunConfig | RunConfigOnline | RunConfigOnlineModel,
) -> None:
"""Resolve ModelRef fields on metric configs before local SDK execution."""
for item in metrics:
if isinstance(item, JobParamsConfigurableMetric):
item.apply_evaluation_job_params(params)
if isinstance(item, MetricWithModels):
await item.resolve_models(resolver)

@staticmethod
def _unresolved_model_refs(metrics: list[Metric]) -> list[str]:
refs = [
model_ref.root
for item in metrics
if isinstance(item, MetricWithModels)
for model_ref in item.model_refs().values()
]
return sorted(refs)

@classmethod
async def compile(
cls,
Expand All @@ -137,7 +175,7 @@ async def compile(

@staticmethod
def _hydrate_metrics(metrics: MetricSpec) -> list[Metric]:
return [unbundle_metric(bundle) for bundle in metrics]
return _hydrate_metrics(metrics)

@staticmethod
def _write_result_files(result: EvaluationArtifactResult, persistent_dir: Path) -> EvaluationResultFiles:
Expand All @@ -162,21 +200,51 @@ def _write_result_files(result: EvaluationArtifactResult, persistent_dir: Path)
artifacts_dir=artifacts_dir,
)

@classmethod
async def to_spec(
cls,
input_spec: BaseModel,
*,
workspace: str,
entity_client: object,
async_sdk: AsyncNeMoPlatform | None,
is_local: bool,
) -> BaseModel:
"""Resolve submitter-facing model references into the canonical evaluation spec."""
del workspace, entity_client, is_local
submit_spec = (
input_spec.model_copy(deep=True)
if isinstance(input_spec, EvaluateInputSpec)
else EvaluateInputSpec.model_validate(input_spec.model_dump())
)
metrics = _hydrate_metrics(submit_spec.metrics)
applied_params = _apply_metric_job_params(
metrics,
normalize_params(submit_spec.params, submit_spec.target),
)
unresolved_refs = _unresolved_model_refs(metrics)
if unresolved_refs:
if async_sdk is None:
raise ValueError(
"ModelRef metrics require `async_sdk` for spec resolution: " + ", ".join(unresolved_refs)
)
await _resolve_metric_models(
metrics,
PlatformModelResolver(async_sdk),
)
if applied_params or unresolved_refs:
submit_spec.metrics = [
_bundle_resolved_metric(metric, bundle)
for metric, bundle in zip(metrics, submit_spec.metrics, strict=True)
]
return EvaluateSpec.model_validate(submit_spec.model_dump(mode="python"))

def run(self, config: dict, *, ctx: JobContext, sdk: object | None = None, async_sdk: object | None = None) -> dict:
"""Run the evaluator job locally and persist its result artifact."""
spec = EvaluateSpec.model_validate(config)
evaluator = Evaluator()
platform_sdk = async_sdk or sdk
params = normalize_params(spec.params, spec.target)
metrics = self._hydrate_metrics(spec.metrics)
if platform_sdk is None:
unresolved_refs = self._unresolved_model_refs(metrics)
if unresolved_refs:
raise ValueError(
"ModelRef metrics require `sdk` or `async_sdk` for local execution: " + ", ".join(unresolved_refs)
)
else:
run_sync(lambda: self._resolve_metric_models(metrics, PlatformModelResolver(platform_sdk), params))
dataset = resolve_run_dataset(
spec.dataset,
ctx=ctx,
Expand Down
Loading
Loading