File size: 4,599 Bytes
79e8e52
e383cb5
79e8e52
 
e383cb5
 
79e8e52
e383cb5
 
79e8e52
 
 
e383cb5
79e8e52
e383cb5
 
79e8e52
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
e383cb5
 
 
79e8e52
 
 
 
e383cb5
79e8e52
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
e383cb5
79e8e52
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
e383cb5
79e8e52
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
e383cb5
79e8e52
 
 
e383cb5
 
 
79e8e52
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
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
"""tool_agent.py — V6: tool agent for orchestrating multiple tools.

Agent que recebe uma query, decide quais ferramentas invocar, em que ordem,
e como combinar os resultados. Usa o ToolCache para evitar recomputação.
"""
from __future__ import annotations
from typing import Optional, Dict, Any, List, Callable
from dataclasses import dataclass, field

import torch
import torch.nn as nn
import torch.nn.functional as F

from .tool_cache import ToolCache


@dataclass
class ToolSpec:
    """Spec de uma ferramenta."""
    name: str
    description: str
    compute_fn: Callable
    input_dim: int
    output_dim: int


class ToolAgent(nn.Module):
    """V6: agent que orquestra ferramentas.

    Arquitetura:
    1. Encoder da query → representação
    2. Router: decide quais ferramentas invocar (top-k)
    3. Executor: invoca ferramentas (com cache)
    4. Combiner: combina resultados das ferramentas

    Forward:
        query: (B, input_dim) → result: (B, output_dim), tools_used: list
    """
    def __init__(
        self,
        tools: List[ToolSpec],
        d_model: int = 128,
        top_k: int = 3,
        cache_size: int = 1024,
    ):
        super().__init__()
        self.tools = {t.name: t for t in tools}
        self.tool_names = [t.name for t in tools]
        self.top_k = min(top_k, len(tools))
        self.d_model = d_model
        # Router: input → distribuição sobre ferramentas
        self.router = nn.Linear(d_model, len(tools))
        # Combiner: combina outputs das top-k ferramentas
        self.combiner = nn.Linear(self.top_k * d_model, d_model)
        # Cache
        self.cache = ToolCache(max_size=cache_size)
        # Output proj (se output_dim != d_model)
        if tools:
            self.output_proj = nn.Linear(d_model, tools[0].output_dim)
        else:
            self.output_proj = nn.Identity()

    def forward(
        self,
        query: torch.Tensor,
    ) -> tuple:
        """query: (B, d_model) → (combined, tools_used_per_sample)."""
        B, D = query.shape
        # 1. Router: logits sobre ferramentas
        router_logits = self.router(query)  # (B, n_tools)
        router_probs = F.softmax(router_logits, dim=-1)
        # 2. Top-k tools por sample
        topk_probs, topk_idx = router_probs.topk(self.top_k, dim=-1)  # (B, k)
        # 3. Executa tools (com cache)
        # Para simplicidade, executamos todas as top-k tools do batch
        # (em produção, faria-se por-sample com máscara)
        tool_outputs = []
        for k in range(self.top_k):
            # Para cada tool_id no batch, executa a tool correspondente
            batch_outputs = []
            for b in range(B):
                tool_id = int(topk_idx[b, k].item())
                tool_name = self.tool_names[tool_id]
                tool = self.tools[tool_name]
                # Cache key: (tool_name, query[b])
                result = self.cache.get_or_compute(
                    tool_name,
                    query[b].detach().cpu().numpy().tobytes(),
                    lambda: tool.compute_fn(query[b:b+1]),
                )
                batch_outputs.append(result)
            # Stack: (B, output_dim)
            try:
                stacked = torch.cat(batch_outputs, dim=0)
            except (RuntimeError, TypeError):
                # Fallback: zeros
                stacked = torch.zeros(B, D, device=query.device)
            # Projeta para d_model se necessário
            if stacked.size(-1) != D:
                # Mean pooling se dimensão diferente
                if stacked.dim() > 2:
                    stacked = stacked.mean(dim=tuple(range(1, stacked.dim() - 1)))
                # Pad/truncate
                if stacked.size(-1) < D:
                    pad = torch.zeros(B, D - stacked.size(-1), device=query.device)
                    stacked = torch.cat([stacked, pad], dim=-1)
                else:
                    stacked = stacked[..., :D]
            tool_outputs.append(stacked)
        # 4. Combina top-k outputs
        combined_input = torch.cat(tool_outputs, dim=-1)  # (B, k*D)
        combined = self.combiner(combined_input)  # (B, D)
        # Weighted by router probs
        weighted = combined * topk_probs.mean(dim=-1, keepdim=True)
        output = self.output_proj(weighted)
        return output, topk_idx.detach().cpu().tolist()

    def get_stats(self) -> Dict[str, Any]:
        return {
            "cache": self.cache.get_stats(),
            "n_tools": len(self.tools),
            "top_k": self.top_k,
        }


__all__ = ["ToolAgent", "ToolSpec"]