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()