File size: 3,273 Bytes
60efcc3
 
 
 
 
 
 
 
 
 
 
 
 
 
 
42801ac
60efcc3
 
 
42801ac
 
60efcc3
 
 
 
 
 
 
 
 
 
 
42801ac
 
 
 
60efcc3
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
from __future__ import annotations

import argparse
import json
from pathlib import Path
from typing import Dict, List

import numpy as np
from sentence_transformers import SentenceTransformer


ROOT = Path(__file__).resolve().parent


class AgentWeaveSemanticRouter:
    """Prototype-based semantic capability router built on frozen MiniLM embeddings.

    This is an experimental semantic companion to AgentWeave's default
    deterministic routing path. It does not replace AgentWeave policy,
    authorization, or execution controls. The upstream encoder is loaded as a
    runtime dependency; this repository is not a fine-tuned MiniLM model.
    """

    def __init__(
        self,
        config_path: str | Path = ROOT / "config.json",
        prototypes_path: str | Path = ROOT / "route_prototypes.json",
    ) -> None:
        self.config = json.loads(Path(config_path).read_text(encoding="utf-8"))
        self.prototypes: Dict[str, List[str]] = json.loads(
            Path(prototypes_path).read_text(encoding="utf-8")
        )
        encoder_model = self.config.get("encoder_model") or self.config.get("base_model")
        if not encoder_model:
            raise ValueError("config.json must define 'encoder_model'")
        self.model = SentenceTransformer(str(encoder_model), device="cpu")

        texts: List[str] = []
        labels: List[str] = []
        for label, examples in self.prototypes.items():
            for example in examples:
                labels.append(label)
                texts.append(example)

        self._prototype_labels = labels
        self._prototype_embeddings = self.model.encode(
            texts,
            normalize_embeddings=bool(self.config.get("normalize_embeddings", True)),
            convert_to_numpy=True,
            show_progress_bar=False,
        )

    def route(self, query: str, top_k: int | None = None) -> List[dict]:
        if not query or not query.strip():
            raise ValueError("query must be a non-empty string")

        top_k = int(top_k or self.config.get("default_top_k", 3))
        query_embedding = self.model.encode(
            [query],
            normalize_embeddings=bool(self.config.get("normalize_embeddings", True)),
            convert_to_numpy=True,
            show_progress_bar=False,
        )[0]

        similarities = self._prototype_embeddings @ query_embedding
        best_by_label: Dict[str, float] = {}
        for label, score in zip(self._prototype_labels, similarities):
            best_by_label[label] = max(best_by_label.get(label, -1.0), float(score))

        ranked = sorted(best_by_label.items(), key=lambda item: item[1], reverse=True)
        return [
            {"route": label, "score": round(score, 6)}
            for label, score in ranked[: max(1, min(top_k, len(ranked)))]
        ]


def main() -> None:
    parser = argparse.ArgumentParser(description="AgentWeave MiniLM semantic router")
    parser.add_argument("query", help="Task or request to route")
    parser.add_argument("--top-k", type=int, default=None, help="Number of routes to return")
    args = parser.parse_args()

    router = AgentWeaveSemanticRouter()
    print(json.dumps(router.route(args.query, args.top_k), indent=2))


if __name__ == "__main__":
    main()