fugu-lite / src /fugu_lite /orchestrator.py
tahsinsoyak's picture
Upload private Fugu-Lite V1 snapshot
88e15cd verified
Raw
History Blame
2.82 kB
from __future__ import annotations
from pathlib import Path
import torch
from .config import WorkerPoolConfig, load_worker_pool
from .model import RouterModel
from .providers import WorkerCallError, WorkerRunner
from .schemas import TaskRecord, route_text
from .training_common import resolve_device
class FuguLiteOrchestrator:
def __init__(self, checkpoint: str | Path, worker_config: str | Path):
self.pool: WorkerPoolConfig = load_worker_pool(worker_config)
self.model = RouterModel.from_checkpoint(checkpoint, device=resolve_device())
if self.model.worker_ids != self.pool.worker_ids:
raise ValueError(
f"Checkpoint workers {self.model.worker_ids} do not match "
f"configured workers {self.pool.worker_ids}"
)
self.model.eval()
self.runner = WorkerRunner(self.pool)
@torch.no_grad()
def route(self, prompt: str, domain: str = "general", tags: list[str] | None = None) -> dict:
text = route_text(prompt, domain, tags or [])
encoded = self.model.tokenizer(
[text],
padding=True,
truncation=True,
max_length=self.model.router_config.max_length,
return_tensors="pt",
)
inputs = {
"input_ids": encoded["input_ids"].to(self.model.device_ref),
"attention_mask": encoded["attention_mask"].to(self.model.device_ref),
}
logits = self.model(**inputs)[0]
probabilities = torch.softmax(logits, dim=-1)
order = torch.argsort(probabilities, descending=True).tolist()
ranked = [
{"worker_id": self.model.worker_ids[index], "probability": float(probabilities[index])}
for index in order
]
return {"selected_worker": ranked[0]["worker_id"], "ranking": ranked}
async def answer(
self, prompt: str, domain: str = "general", tags: list[str] | None = None
) -> dict:
routing = self.route(prompt, domain, tags)
worker_map = {worker.id: worker for worker in self.pool.workers}
task = TaskRecord(
task_id="inference",
prompt=prompt,
domain=domain,
tags=tags or [],
)
errors = []
for candidate in routing["ranking"]:
worker = worker_map[candidate["worker_id"]]
try:
result = await self.runner.complete(worker, task)
return {"routing": routing, "result": result.model_dump(), "fallback_errors": errors}
except WorkerCallError as exc:
errors.append({"worker_id": worker.id, "error": str(exc)})
raise WorkerCallError(f"Every routed worker failed: {errors}")
async def close(self) -> None:
await self.runner.close()