BiGRU_T_version / src /bigru_t /tools /tool_agent.py
PowerMachine's picture
V6: Kohonen 4D SOM + EWC per-neuron + MTP+entropy + Xeon V6
79e8e52 verified
Raw History Blame Contribute Delete
4.6 kB
"""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"]