Skip to content
Open
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
25 changes: 25 additions & 0 deletions src/lighteval/models/endpoints/inference_providers_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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.).

Expand Down Expand Up @@ -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)

Expand Down
3 changes: 3 additions & 0 deletions src/lighteval/models/model_output.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,7 @@
# SOFTWARE.

from dataclasses import dataclass, field
from typing import Any

import torch

Expand Down Expand Up @@ -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]:
Expand All @@ -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,
)


Expand Down
56 changes: 56 additions & 0 deletions tests/unit/models/endpoints/test_inference_providers_model.py
Original file line number Diff line number Diff line change
@@ -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},
}
34 changes: 34 additions & 0 deletions tests/unit/models/test_model_output.py
Original file line number Diff line number Diff line change
@@ -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