File size: 2,819 Bytes
88e15cd
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
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()