File size: 3,950 Bytes
79e8e52
e383cb5
79e8e52
 
 
e383cb5
79e8e52
e383cb5
 
79e8e52
 
 
e383cb5
 
79e8e52
 
 
e383cb5
 
 
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
"""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"]