Download src/bigru_t/tools/tool_agent.py from PowerMachine/BiGRU_T_version: direct link, hf CLI and curl.
- Browser
- Download file 4.6 kB
-
https://huggingface.co/PowerMachine/BiGRU_T_version/resolve/main/src/bigru_t/tools/tool_agent.py
- Command line
-
hf download hf://PowerMachine/BiGRU_T_version/src/bigru_t/tools/tool_agent.py
-
curl -L -o tool_agent.py https://huggingface.co/PowerMachine/BiGRU_T_version/resolve/main/src/bigru_t/tools/tool_agent.py
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 | |
| 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"] | |