Skip to content
Draft
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
53 changes: 53 additions & 0 deletions docs/inference.md
Original file line number Diff line number Diff line change
Expand Up @@ -46,6 +46,53 @@ pipeline = mdr.read_jsonl("input.jsonl").map_async(

Set `OPENAI_API_KEY` in the worker environment before execution. For cloud jobs, pass it through `secrets={"OPENAI_API_KEY": None}`.

### Request limits

Refiner uses adaptive request limiting by default. It starts below the configured
maximum, grows after clean success windows, and backs off when an endpoint
returns HTTP `429`.

```python
import refiner as mdr

endpoint = mdr.inference.OpenAIEndpointProvider(
base_url="https://api.openai.com",
model="gpt-5-mini",
)

pipeline = mdr.read_jsonl("input.jsonl").map_async(
mdr.inference.generate(
fn=summarize,

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

medium

The summarize function is used here but is not defined in this code example. To make the example self-contained and easier for users to understand and copy, please define the summarize function within the python code block, similar to how it's done in the preceding example.

provider=endpoint,
rate_limit=mdr.inference.AdaptiveRateLimit(
max_concurrency=256,
initial_concurrency=16,
),
),
max_in_flight=256,
)
```

When `rate_limit` is omitted or set to `None`, Refiner uses
`AdaptiveRateLimit()` with its defaults. Growth uses a shrinking multiplier: the
default starts at `2.0`, subtracts `0.1` after each successful growth window,
and never drops below `1.05`. A rate-limit response halves the current request
concurrency by default and subtracts `0.2` from the growth multiplier. If the
endpoint sends `Retry-After`, Refiner waits before starting more requests.

Use `StaticRateLimit` when you want the old fixed-concurrency behavior:

```python
pipeline = mdr.read_jsonl("input.jsonl").map_async(
mdr.inference.generate(
fn=summarize,
provider=endpoint,
rate_limit=mdr.inference.StaticRateLimit(max_concurrency=64),
),
max_in_flight=64,
)
```

### Refiner managed VLLM runtime

Use `VLLMProvider` when launching on Refiner Cloud and you want the platform to start and manage a VLLM server for you.
Expand Down Expand Up @@ -96,3 +143,9 @@ Only the following models are currently supported. If you are missing one, pleas

- Direct endpoint example: [`examples/inference_endpoint.py`](../examples/inference_endpoint.py)
- VLLM-backed example: [`examples/inference_vllm.py`](../examples/inference_vllm.py)

## Internal Notes

Adaptive request limiting is process-local. In multi-worker jobs, each worker
adapts independently to the same provider. This reduces per-worker bursts but is
not a global quota coordinator across the full job.
4 changes: 3 additions & 1 deletion examples/lerobot/sarm_annotation.py
Original file line number Diff line number Diff line change
Expand Up @@ -187,7 +187,9 @@ async def annotate_dense_subtasks(row, generate):
fn=annotate_dense_subtasks,
provider=PROVIDER,
default_generation_params={"temperature": 0.1},
max_concurrent_requests=MAX_IN_FLIGHT,
rate_limit=mdr.inference.AdaptiveRateLimit(
max_concurrency=MAX_IN_FLIGHT
),
),
max_in_flight=MAX_IN_FLIGHT,
)
Expand Down
6 changes: 5 additions & 1 deletion src/refiner/inference/__init__.py
Original file line number Diff line number Diff line change
@@ -1,10 +1,14 @@
from refiner.inference.generate import generate
from refiner.inference.client import InferenceResponse
from refiner.inference.client import GenerationRateLimitError, InferenceResponse
from refiner.inference.providers import OpenAIEndpointProvider, VLLMProvider
from refiner.inference.rate_limit import AdaptiveRateLimit, StaticRateLimit

__all__ = [
"generate",
"AdaptiveRateLimit",
"GenerationRateLimitError",
"InferenceResponse",
"OpenAIEndpointProvider",
"StaticRateLimit",
"VLLMProvider",
]
73 changes: 66 additions & 7 deletions src/refiner/inference/_runtime.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,11 +2,19 @@

import asyncio
import inspect
import os
from collections.abc import Awaitable, Callable, Mapping
from dataclasses import asdict
from typing import Any, TypeAlias, cast

from refiner.inference.client import _OpenAIEndpointClient
from refiner.inference.client import GenerationRateLimitError, _OpenAIEndpointClient
from refiner.inference.providers import OpenAIEndpointProvider, VLLMProvider
from refiner.inference.rate_limit import (
AdaptiveRateLimit,
AdaptiveRateLimiter,
RateLimit,
StaticRateLimit,
)
from refiner.pipeline.data.row import Row
from refiner.pipeline.steps import MapResult
from refiner.services import VLLMRuntimeServiceBinding
Expand All @@ -31,14 +39,28 @@ def inference_map(
defaults: Mapping[str, Any] | None,
defaults_key: str | None = None,
max_concurrent_requests: int = 256,
rate_limit: RateLimit | None = None,
rate_limit_key: str | None = None,
call: ClientCall,
record: Callable[[Row, Any], None] | None = None,
) -> Callable[[Row], Awaitable[MapResult]]:
if max_concurrent_requests <= 0:
raise ValueError("max_concurrent_requests must be > 0")
resolved_rate_limit = rate_limit or StaticRateLimit(
max_concurrency=max_concurrent_requests
)
client: _OpenAIEndpointClient | None = None
client_lock = asyncio.Lock()
semaphore = asyncio.Semaphore(max_concurrent_requests)
adaptive_limiter: AdaptiveRateLimiter | None = (
AdaptiveRateLimiter(resolved_rate_limit)
if isinstance(resolved_rate_limit, AdaptiveRateLimit)
else None
)
semaphore = (
None
if adaptive_limiter is not None
else asyncio.Semaphore(resolved_rate_limit.max_concurrency)
)
gauges_registered = False
waiting_requests = 0
running_requests = 0
Expand All @@ -49,6 +71,12 @@ def _register_metrics() -> None:
return
register_gauge("waiting_requests", lambda: waiting_requests, unit="requests")
register_gauge("running_requests", lambda: running_requests, unit="requests")
if adaptive_limiter is not None:
register_gauge(
"adaptive_concurrency",
lambda: adaptive_limiter.limit,
unit="requests",
)
gauges_registered = True

async def _client() -> _OpenAIEndpointClient:
Expand All @@ -59,7 +87,10 @@ async def _client() -> _OpenAIEndpointClient:
if client is not None:
return client
if isinstance(provider, OpenAIEndpointProvider):
client = _OpenAIEndpointClient(base_url=provider.base_url)
client = _OpenAIEndpointClient(
base_url=provider.base_url,
api_key=os.environ.get(provider.api_key_env_var),
)
else:
service_name = provider.service_definition().name
service_manager = get_active_service_manager()
Expand Down Expand Up @@ -88,17 +119,37 @@ async def _request(row: Row, payload: Mapping[str, Any]) -> Any:
}
resolved_client = await _client()
waiting_requests += 1
await semaphore.acquire()
waiting_requests -= 1
acquired_adaptive = False
try:
if adaptive_limiter is not None:
await adaptive_limiter.acquire()
acquired_adaptive = True
else:
assert semaphore is not None
await semaphore.acquire()
finally:
waiting_requests -= 1
running_requests += 1
try:
response = await call(resolved_client, request_payload)
except GenerationRateLimitError as err:
if adaptive_limiter is not None:
await adaptive_limiter.record_rate_limit(err.retry_after_seconds)
row.log_throughput("rate_limited_requests", 1, unit="requests")
row.log_throughput("failed_requests", 1, unit="requests")
raise
except Exception:
row.log_throughput("failed_requests", 1, unit="requests")
raise
finally:
running_requests -= 1
semaphore.release()
if acquired_adaptive:
assert adaptive_limiter is not None
await adaptive_limiter.release()
if semaphore is not None:
semaphore.release()
if adaptive_limiter is not None:
await adaptive_limiter.record_success()
row.log_throughput("successful_requests", 1, unit="requests")
if record is not None:
record(row, response)
Expand All @@ -114,8 +165,16 @@ async def _wrapped(row: Row) -> MapResult:
args: dict[str, Any] = {
"fn": fn,
"provider": provider.to_builtin_args(),
"max_concurrent_requests": max_concurrent_requests,
}
if rate_limit_key is None:
args["max_concurrent_requests"] = resolved_rate_limit.max_concurrency
else:
args[rate_limit_key] = {
"type": "adaptive"
if isinstance(resolved_rate_limit, AdaptiveRateLimit)
else "static",
**asdict(resolved_rate_limit),
}
if defaults_key is not None:
args[defaults_key] = dict(defaults or {})
setattr(
Expand Down
34 changes: 33 additions & 1 deletion src/refiner/inference/client.py
Original file line number Diff line number Diff line change
@@ -1,9 +1,11 @@
from __future__ import annotations

import asyncio
import email.utils
import logging
import os
import random
import time
from collections.abc import Mapping, Sequence
from dataclasses import dataclass, field
from typing import Any
Expand All @@ -17,6 +19,12 @@
logger = logging.getLogger(__name__)


class GenerationRateLimitError(RuntimeError):
def __init__(self, message: str, *, retry_after_seconds: float | None = None):
super().__init__(message)
self.retry_after_seconds = retry_after_seconds


@dataclass(frozen=True, slots=True)
class InferenceResponse:
text: str
Expand Down Expand Up @@ -124,6 +132,11 @@ async def _post_json(
message = f"{operation} request failed with HTTP {err.response.status_code}"
if detail:
message = f"{message}: {detail}"
if err.response.status_code in {429, 503}:
raise GenerationRateLimitError(
message,
retry_after_seconds=_retry_after_seconds(err.response),
) from err
raise RuntimeError(message) from err
return response.json()

Expand Down Expand Up @@ -198,4 +211,23 @@ def _retry_delay_seconds(attempt: int) -> float:
return base_delay * (1.0 + jitter)


__all__ = ["InferenceResponse"]
def _retry_after_seconds(response: httpx.Response) -> float | None:
raw = response.headers.get("Retry-After")
if raw is None:
return None
value = raw.strip()
if not value:
return None
try:
return max(0.0, float(value))
except ValueError:
try:
parsed = email.utils.parsedate_to_datetime(value)
except (TypeError, ValueError):
return None
if parsed is None:
return None
return max(0.0, parsed.timestamp() - time.time())


__all__ = ["GenerationRateLimitError", "InferenceResponse"]
6 changes: 4 additions & 2 deletions src/refiner/inference/generate.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
from refiner.inference._runtime import inference_map
from refiner.inference.client import InferenceResponse, _OpenAIEndpointClient
from refiner.inference.providers import OpenAIEndpointProvider, VLLMProvider
from refiner.inference.rate_limit import AdaptiveRateLimit, RateLimit
from refiner.pipeline.data.row import Row
from refiner.pipeline.steps import MapResult

Expand All @@ -19,15 +20,16 @@ def generate(
fn: InferenceFn,
provider: OpenAIEndpointProvider | VLLMProvider,
default_generation_params: Mapping[str, Any] | None = None,
max_concurrent_requests: int = 256,
rate_limit: RateLimit | None = None,
) -> Callable[[Row], Awaitable[MapResult]]:
return inference_map(
name="inference.generate",
fn=fn,
provider=provider,
defaults=default_generation_params,
defaults_key="default_generation_params",
max_concurrent_requests=max_concurrent_requests,
rate_limit=rate_limit or AdaptiveRateLimit(),
rate_limit_key="rate_limit",
call=_generate,
record=_record_usage,
)
Expand Down
4 changes: 4 additions & 0 deletions src/refiner/inference/providers.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,12 +8,15 @@
class OpenAIEndpointProvider:
base_url: str
model: str
api_key_env_var: str = "OPENAI_API_KEY"

def __post_init__(self) -> None:
if not self.base_url.strip():
raise ValueError("base_url must be non-empty")
if not self.model.strip():
raise ValueError("model must be non-empty")
if not self.api_key_env_var.strip():
raise ValueError("api_key_env_var must be non-empty")

def service_definition(self) -> None:
return None
Expand All @@ -23,6 +26,7 @@ def to_builtin_args(self) -> dict[str, object]:
"type": "openai_endpoint",
"base_url": self.base_url,
"model": self.model,
"api_key_env_var": self.api_key_env_var,
}
return payload

Expand Down
Loading
Loading