From d2e91089547bf0da358847887be0e1441f401da5 Mon Sep 17 00:00:00 2001 From: Tanisha-Katara <118114217+Tanisha-Katara@users.noreply.github.com> Date: Thu, 6 Aug 2026 20:15:19 +0400 Subject: [PATCH] Preserve provider usage metadata --- .../endpoints/inference_providers_model.py | 25 +++++++++ src/lighteval/models/model_output.py | 3 + .../test_inference_providers_model.py | 56 +++++++++++++++++++ tests/unit/models/test_model_output.py | 34 +++++++++++ 4 files changed, 118 insertions(+) create mode 100644 tests/unit/models/endpoints/test_inference_providers_model.py create mode 100644 tests/unit/models/test_model_output.py diff --git a/src/lighteval/models/endpoints/inference_providers_model.py b/src/lighteval/models/endpoints/inference_providers_model.py index 54790e45b..97cfc793d 100644 --- a/src/lighteval/models/endpoints/inference_providers_model.py +++ b/src/lighteval/models/endpoints/inference_providers_model.py @@ -22,6 +22,8 @@ import asyncio import logging +from collections.abc import Mapping +from dataclasses import asdict, is_dataclass from typing import Any, List, Optional from huggingface_hub import AsyncInferenceClient, ChatCompletionOutput @@ -42,6 +44,28 @@ logger = logging.getLogger(__name__) +def _jsonable_usage_value(value: Any) -> Any: + if value is None: + return None + if hasattr(value, "model_dump"): + return _jsonable_usage_value(value.model_dump(exclude_none=True)) + if hasattr(value, "dict"): + return _jsonable_usage_value(value.dict(exclude_none=True)) + if is_dataclass(value): + return _jsonable_usage_value(asdict(value)) + if isinstance(value, Mapping): + return {str(key): _jsonable_usage_value(item) for key, item in value.items() if item is not None} + if isinstance(value, (list, tuple)): + return [_jsonable_usage_value(item) for item in value] + return value + + +def _usage_metadata_from_response(response: Any) -> dict[str, Any]: + usage = response.get("usage") if isinstance(response, Mapping) else getattr(response, "usage", None) + metadata = _jsonable_usage_value(usage) + return metadata if isinstance(metadata, dict) else {} + + class InferenceProvidersModelConfig(ModelConfig): """Configuration class for HuggingFace's inference providers (like Together AI, Anyscale, etc.). @@ -234,6 +258,7 @@ def greedy_until( # In empty responses, the model should return an empty string instead of None text=result if result[0] else [""], input=context, + usage_metadata=_usage_metadata_from_response(response), ) results.append(cur_response) diff --git a/src/lighteval/models/model_output.py b/src/lighteval/models/model_output.py index b10ce7f56..109a25971 100644 --- a/src/lighteval/models/model_output.py +++ b/src/lighteval/models/model_output.py @@ -21,6 +21,7 @@ # SOFTWARE. from dataclasses import dataclass, field +from typing import Any import torch @@ -137,6 +138,7 @@ class ModelResponse: # Other metadata truncated_tokens_count: int = 0 # How many tokens truncated padded_tokens_count: int = 0 # How many tokens of padding + usage_metadata: dict[str, Any] = field(default_factory=dict) # Provider-reported billing/token usage fields @property def final_text(self) -> list[str]: @@ -156,6 +158,7 @@ def __getitem__(self, index: int) -> "ModelResponse": unconditioned_logprobs=[self.unconditioned_logprobs[index]] if self.unconditioned_logprobs else None, truncated_tokens_count=self.truncated_tokens_count, padded_tokens_count=self.padded_tokens_count, + usage_metadata=self.usage_metadata, ) diff --git a/tests/unit/models/endpoints/test_inference_providers_model.py b/tests/unit/models/endpoints/test_inference_providers_model.py new file mode 100644 index 000000000..46b48132b --- /dev/null +++ b/tests/unit/models/endpoints/test_inference_providers_model.py @@ -0,0 +1,56 @@ +from dataclasses import dataclass + +from lighteval.models.endpoints.inference_providers_model import _usage_metadata_from_response + + +@dataclass +class UsageDetails: + cached_tokens: int + cache_write_tokens: int | None = None + + +@dataclass +class Usage: + prompt_tokens: int + completion_tokens: int + total_tokens: int + prompt_tokens_details: UsageDetails + + +@dataclass +class Response: + usage: Usage + + +def test_usage_metadata_from_response_preserves_nested_cache_counters(): + response = Response( + usage=Usage( + prompt_tokens=100, + completion_tokens=20, + total_tokens=120, + prompt_tokens_details=UsageDetails(cached_tokens=75), + ) + ) + + assert _usage_metadata_from_response(response) == { + "prompt_tokens": 100, + "completion_tokens": 20, + "total_tokens": 120, + "prompt_tokens_details": {"cached_tokens": 75}, + } + + +def test_usage_metadata_from_response_accepts_mapping_responses(): + response = { + "usage": { + "input_tokens": 100, + "output_tokens": 20, + "input_tokens_details": {"cached_tokens": 40, "cache_write_tokens": None}, + } + } + + assert _usage_metadata_from_response(response) == { + "input_tokens": 100, + "output_tokens": 20, + "input_tokens_details": {"cached_tokens": 40}, + } diff --git a/tests/unit/models/test_model_output.py b/tests/unit/models/test_model_output.py new file mode 100644 index 000000000..4e0e032cc --- /dev/null +++ b/tests/unit/models/test_model_output.py @@ -0,0 +1,34 @@ +from dataclasses import asdict + +from lighteval.models.model_output import ModelResponse + + +def test_model_response_preserves_usage_metadata_when_sliced(): + response = ModelResponse( + text=["a", "b"], + output_tokens=[[1], [2]], + usage_metadata={ + "prompt_tokens": 100, + "completion_tokens": 10, + "prompt_tokens_details": {"cached_tokens": 80}, + }, + ) + + sliced = response[1] + + assert sliced.text == ["b"] + assert sliced.usage_metadata == response.usage_metadata + + +def test_model_response_usage_metadata_round_trips_through_dict(): + response = ModelResponse( + text=["answer"], + usage_metadata={ + "prompt_tokens": 42, + "prompt_tokens_details": {"cached_tokens": 12}, + }, + ) + + restored = ModelResponse(**asdict(response)) + + assert restored.usage_metadata == response.usage_metadata