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