fugu-lite / src /fugu_lite /config.py
tahsinsoyak's picture
Upload private Fugu-Lite V1 snapshot
88e15cd verified
Raw
History Blame
4.8 kB
from __future__ import annotations
from pathlib import Path
from typing import Any, Literal
import yaml
from pydantic import BaseModel, Field, field_validator, model_validator
class RewardConfig(BaseModel):
quality_weight: float = 1.0
cost_weight: float = 0.0
latency_weight: float = 0.0
cost_scale_usd: float = 0.01
latency_scale_seconds: float = 30.0
judge_model: str | None = None
judge_system_prompt: str = (
"Score the candidate answer from 0 to 1. Return only one decimal number."
)
@field_validator("cost_scale_usd", "latency_scale_seconds")
@classmethod
def positive_scales(cls, value: float) -> float:
if value <= 0:
raise ValueError("reward scales must be positive")
return value
class WorkerConfig(BaseModel):
id: str
provider: Literal["openrouter", "openai_compatible", "mock"] = "openrouter"
model: str
description: str = ""
base_url: str | None = None
api_key_env: str | None = None
api_key: str | None = Field(default=None, repr=False)
system_prompt: str = "You are a helpful expert. Follow the user's requested output format."
temperature: float = 0.1
max_tokens: int = 1024
extra_body: dict[str, Any] = Field(default_factory=dict)
mock_domains: list[str] = Field(default_factory=list)
mock_specialist_accuracy: float = 1.0
mock_general_accuracy: float = 0.0
class WorkerPoolConfig(BaseModel):
project_name: str = "benchgen-fugu-lite"
base_url: str = "https://openrouter.ai/api/v1"
api_key_env: str = "OPENROUTER_API_KEY"
concurrency: int = 4
timeout_seconds: float = 120.0
max_retries: int = 3
workers: list[WorkerConfig]
reward: RewardConfig = Field(default_factory=RewardConfig)
@model_validator(mode="after")
def unique_workers(self) -> WorkerPoolConfig:
ids = [worker.id for worker in self.workers]
if not ids:
raise ValueError("at least one worker is required")
if len(ids) != len(set(ids)):
raise ValueError("worker IDs must be unique")
return self
@property
def worker_ids(self) -> list[str]:
return [worker.id for worker in self.workers]
@property
def worker_descriptions(self) -> dict[str, str]:
return {worker.id: worker.description for worker in self.workers}
class ModelConfig(BaseModel):
base_model: str = "Qwen/Qwen3-0.6B"
max_length: int = 512
pooling: Literal["last_token", "mean"] = "last_token"
dropout: float = 0.05
freeze_backbone: bool = True
dtype: Literal["auto", "float32", "float16", "bfloat16"] = "auto"
class SFTTrainingConfig(BaseModel):
seed: int = 42
epochs: int = 5
batch_size: int = 8
gradient_accumulation_steps: int = 1
learning_rate: float = 1e-3
weight_decay: float = 0.01
max_grad_norm: float = 1.0
soft_target_temperature: float = 0.15
log_every: int = 10
class RLTrainingConfig(BaseModel):
seed: int = 43
epochs: int = 3
batch_size: int = 8
gradient_accumulation_steps: int = 1
learning_rate: float = 2e-4
weight_decay: float = 0.0
max_grad_norm: float = 1.0
estimator: Literal["reinforce", "expected_reward"] = "reinforce"
rollouts_per_prompt: int = 8
entropy_coefficient: float = 0.01
normalize_advantage: bool = False
log_every: int = 10
class ESTrainingConfig(BaseModel):
seed: int = 44
generations: int = 30
sigma: float = 0.03
population_size: int | None = 16
objective: Literal["greedy_utility", "expected_utility"] = "greedy_utility"
class SFTConfig(BaseModel):
model: ModelConfig = Field(default_factory=ModelConfig)
training: SFTTrainingConfig = Field(default_factory=SFTTrainingConfig)
class RLConfig(BaseModel):
model: ModelConfig = Field(default_factory=ModelConfig)
training: RLTrainingConfig = Field(default_factory=RLTrainingConfig)
class ESConfig(BaseModel):
model: ModelConfig = Field(default_factory=ModelConfig)
training: ESTrainingConfig = Field(default_factory=ESTrainingConfig)
def load_yaml(path: str | Path) -> dict[str, Any]:
with Path(path).open("r", encoding="utf-8") as handle:
data = yaml.safe_load(handle)
if not isinstance(data, dict):
raise TypeError(f"Expected a YAML object in {path}")
return data
def load_worker_pool(path: str | Path) -> WorkerPoolConfig:
return WorkerPoolConfig.model_validate(load_yaml(path))
def load_sft_config(path: str | Path) -> SFTConfig:
return SFTConfig.model_validate(load_yaml(path))
def load_rl_config(path: str | Path) -> RLConfig:
return RLConfig.model_validate(load_yaml(path))
def load_es_config(path: str | Path) -> ESConfig:
return ESConfig.model_validate(load_yaml(path))