BiGRU_T_version / src /bigru_t /tools /tool_mutation.py
PowerMachine's picture
V6: Kohonen 4D SOM + EWC per-neuron + MTP+entropy + Xeon V6
79e8e52 verified
Raw History Blame Contribute Delete
3.95 kB
"""tool_mutation.py — V6: tool mutation system.
Permite que o modelo "mute" suas ferramentas (sub-rotinas) ao longo do
treinamento. Cada ferramenta tem uma versão atual + variantes mutantes.
A ferramenta com melhor performance substitui a versão atual.
Inspirado em programação genética, aplicado a sub-rotinas do modelo.
"""
from __future__ import annotations
import copy
import random
from typing import Optional, Dict, List, Any, Callable
from dataclasses import dataclass, field
import torch
import torch.nn as nn
@dataclass
class ToolVariant:
"""Uma variante mutante de uma ferramenta."""
name: str
parent_id: str
generation: int
module: nn.Module
score: float = 0.0
n_evaluations: int = 0
class ToolMutator:
"""V6: gerencia mutação de ferramentas.
Usage:
mutator = ToolMutator(mutation_rate=0.1)
mutator.register("encoder", encoder_module)
mutator.mutate("encoder")
# Após avaliação, mutator.update_score("encoder", variant_id, score)
mutator.promote_best("encoder")
"""
def __init__(
self,
mutation_rate: float = 0.1,
mutation_scale: float = 0.02,
max_variants: int = 4,
):
self.mutation_rate = mutation_rate
self.mutation_scale = mutation_scale
self.max_variants = max_variants
self.tools: Dict[str, List[ToolVariant]] = {}
self.generation_counter: Dict[str, int] = {}
def register(self, name: str, module: nn.Module) -> str:
"""Registra uma nova ferramenta com sua versão inicial."""
variant = ToolVariant(
name=name,
parent_id="root",
generation=0,
module=module,
)
self.tools[name] = [variant]
self.generation_counter[name] = 0
return f"{name}_v0"
def mutate(self, name: str) -> Optional[str]:
"""Cria uma variante mutante da melhor ferramenta atual."""
if name not in self.tools:
return None
if len(self.tools[name]) >= self.max_variants:
# Remove pior variante
self.tools[name].sort(key=lambda v: v.score, reverse=True)
self.tools[name].pop()
# Pega a melhor variante como pai
parent = max(self.tools[name], key=lambda v: v.score)
self.generation_counter[name] += 1
gen = self.generation_counter[name]
# Cria mutante
mutant_module = copy.deepcopy(parent.module)
with torch.no_grad():
for p in mutant_module.parameters():
if random.random() < self.mutation_rate:
noise = torch.randn_like(p) * self.mutation_scale
p.add_(noise)
variant = ToolVariant(
name=name,
parent_id=f"{name}_v{parent.generation}",
generation=gen,
module=mutant_module,
)
self.tools[name].append(variant)
return f"{name}_v{gen}"
def update_score(self, name: str, variant_idx: int, score: float) -> None:
"""Atualiza o score de uma variante (média móvel)."""
if name not in self.tools or variant_idx >= len(self.tools[name]):
return
v = self.tools[name][variant_idx]
v.score = (v.score * v.n_evaluations + score) / (v.n_evaluations + 1)
v.n_evaluations += 1
def promote_best(self, name: str) -> Optional[nn.Module]:
"""Promove a melhor variante para versão atual."""
if name not in self.tools:
return None
best = max(self.tools[name], key=lambda v: v.score)
return best.module
def get_status(self) -> Dict[str, Any]:
return {
name: [
{"gen": v.generation, "score": v.score, "n_eval": v.n_evaluations}
for v in variants
]
for name, variants in self.tools.items()
}
__all__ = ["ToolMutator", "ToolVariant"]