| from __future__ import annotations |
|
|
| from typing import Any, Literal |
|
|
| from pydantic import BaseModel, Field, model_validator |
|
|
|
|
| class GraderSpec(BaseModel): |
| type: Literal["exact", "contains", "numeric", "regex", "llm_judge"] = "exact" |
| tolerance: float = 0.0 |
| pattern: str | None = None |
| rubric: str | None = None |
| case_sensitive: bool = False |
|
|
|
|
| class TaskRecord(BaseModel): |
| task_id: str |
| prompt: str |
| domain: str = "general" |
| reference_answer: str | None = None |
| grader: GraderSpec = Field(default_factory=GraderSpec) |
| split: Literal["train", "validation", "test"] | None = None |
| tags: list[str] = Field(default_factory=list) |
| metadata: dict[str, Any] = Field(default_factory=dict) |
|
|
|
|
| class WorkerResult(BaseModel): |
| worker_id: str |
| requested_model: str |
| served_model: str | None = None |
| response: str |
| quality: float |
| cost_usd: float | None = None |
| latency_ms: float |
| prompt_tokens: int | None = None |
| completion_tokens: int | None = None |
| generation_id: str | None = None |
| error: str | None = None |
| utility: float = 0.0 |
|
|
|
|
| class RewardRecord(BaseModel): |
| task_id: str |
| prompt: str |
| domain: str |
| split: Literal["train", "validation", "test"] | None = None |
| tags: list[str] = Field(default_factory=list) |
| metadata: dict[str, Any] = Field(default_factory=dict) |
| repetitions: int = Field(default=1, ge=1) |
| worker_ids: list[str] |
| rewards: list[float] |
| results: list[WorkerResult] = Field(default_factory=list) |
|
|
| @model_validator(mode="after") |
| def validate_vector(self) -> RewardRecord: |
| if not self.worker_ids: |
| raise ValueError("worker_ids cannot be empty") |
| if len(self.worker_ids) != len(self.rewards): |
| raise ValueError("worker_ids and rewards must have equal length") |
| if len(set(self.worker_ids)) != len(self.worker_ids): |
| raise ValueError("worker_ids must be unique") |
| return self |
|
|
|
|
| def route_text( |
| prompt: str, |
| domain: str = "general", |
| tags: list[str] | None = None, |
| worker_descriptions: dict[str, str] | None = None, |
| ) -> str: |
| """Create the router input. Do not include answers or worker outcomes.""" |
| tag_text = ", ".join(tags or []) or "none" |
| experts = "" |
| if worker_descriptions: |
| experts = "\nAvailable workers:\n" + "\n".join( |
| f"- {worker_id}: {description}" |
| for worker_id, description in worker_descriptions.items() |
| ) |
| return ( |
| "Select the best worker for this task.\n" |
| f"Domain: {domain}\n" |
| f"Tags: {tag_text}{experts}\n" |
| f"Task:\n{prompt}" |
| ) |
|
|
|
|