"""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"]