File size: 9,579 Bytes
9b51f42
 
 
 
 
 
 
 
 
4b65c52
 
 
 
9b51f42
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
4b65c52
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
9b51f42
 
 
 
 
4b65c52
 
 
 
 
 
 
 
 
9b51f42
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
"""FlyGPT: a character-level language model whose recurrent core is a real subgraph of the
fruit-fly connectome (MaleCNS v1.0). Hugging Face `transformers` implementation; self-contained.

Dynamics (one scalar state per neuron, plan.md §7 of the FlyGPT spec):

    proposal_i = tanh( sum_j W_ij h_j / sqrt(in_degree_i) + external_input_i + bias_i )
    h_i_new    = (1 - leak_i) * h_i + leak_i * proposal_i

The connectome is stored in `model.safetensors` as integer tensors (`graph.*`); only the learned
per-edge values and the adapters are floating point (bf16 on disk). The sparse recurrent matmul runs
in fp32: through the fused kernels of the `connectome-kernels` package when it is installed and a
CUDA device is used (training speed), else through torch.sparse COO (rows = destination). Both give
the same logits and gradients.
"""
from __future__ import annotations

import math
from dataclasses import dataclass
from typing import Optional

import torch
import torch.nn as nn
import torch.nn.functional as F
from transformers import PreTrainedModel
from transformers.generation import GenerationMixin
from transformers.utils import ModelOutput

from .configuration_flygpt import FlyGPTConfig


@dataclass
class FlyGPTOutput(ModelOutput):
    loss: Optional[torch.FloatTensor] = None
    logits: Optional[torch.FloatTensor] = None
    state: Optional[torch.FloatTensor] = None   # [B, N] neuron states after the last character


class FlyGraph(nn.Module):
    """The anatomy. Integer buffers only; never trained."""

    def __init__(self, num_neurons: int, num_edges: int, num_input: int, num_output: int):
        super().__init__()
        self.register_buffer("edge_index", torch.zeros(2, num_edges, dtype=torch.int32))   # [source, destination]
        self.register_buffer("synapse_count", torch.zeros(num_edges, dtype=torch.int32))   # MaleCNS synaptic contacts
        self.register_buffer("node_id", torch.zeros(num_neurons, dtype=torch.int64))       # MaleCNS body ids
        self.register_buffer("input_nodes", torch.zeros(num_input, dtype=torch.int64))
        self.register_buffer("output_nodes", torch.zeros(num_output, dtype=torch.int64))


class FlyRecurrentCore(nn.Module):
    """The learned state: one value per real edge, plus per-neuron bias and leak."""

    def __init__(self, num_neurons: int, num_edges: int, leak_init: float, learned_leak: bool):
        super().__init__()
        self.edge_values = nn.Parameter(torch.zeros(num_edges))
        self.bias = nn.Parameter(torch.zeros(num_neurons))
        self.raw_leak = nn.Parameter(torch.full((num_neurons,), math.log(leak_init / (1 - leak_init))),
                                     requires_grad=learned_leak)


class FlyGPTPreTrainedModel(PreTrainedModel):
    config_class = FlyGPTConfig
    base_model_prefix = "flygpt"
    _is_stateful = True
    _supports_cache_class = False
    supports_gradient_checkpointing = False

    def _init_weights(self, module):
        if isinstance(module, FlyRecurrentCore):
            nn.init.normal_(module.edge_values, std=self.config.init_scale)
            nn.init.zeros_(module.bias)
        elif isinstance(module, nn.Linear):
            nn.init.normal_(module.weight, std=0.02)
            if module.bias is not None:
                nn.init.zeros_(module.bias)
        elif isinstance(module, nn.Embedding):
            nn.init.normal_(module.weight, std=1.0)


class FlyGPTForCausalLM(FlyGPTPreTrainedModel, GenerationMixin):
    def __init__(self, config: FlyGPTConfig):
        super().__init__(config)
        c = config
        self.graph = FlyGraph(c.num_neurons, c.num_edges, c.num_input_neurons, c.num_output_neurons)
        self.recurrent = FlyRecurrentCore(c.num_neurons, c.num_edges, c.leak_init, c.learned_leak)
        self.embed = nn.Embedding(c.vocab_size, c.embedding_dim)
        self.input_proj = nn.Linear(c.embedding_dim, c.num_input_neurons)
        self.lm_head = nn.Linear(c.num_output_neurons, c.vocab_size)
        self.post_init()

    # ---- sparse recurrent matrix -------------------------------------------------------------
    @property
    def num_neurons(self) -> int:
        return self.config.num_neurons

    def edge_scale(self) -> torch.Tensor:
        """1/sqrt(in_degree) per edge (degree normalization), or ones."""
        dst = self.graph.edge_index[1].long()
        if not self.config.degree_normalization:
            return torch.ones_like(dst, dtype=torch.float32)
        in_deg = torch.bincount(dst, minlength=self.num_neurons).clamp(min=1).float()
        return 1.0 / in_deg[dst].sqrt()

    def sparse_weight(self) -> torch.Tensor:
        src, dst = self.graph.edge_index[0].long(), self.graph.edge_index[1].long()
        values = self.recurrent.edge_values.float() * self.edge_scale()
        return torch.sparse_coo_tensor(torch.stack([dst, src]), values, (self.num_neurons, self.num_neurons))

    def dense_weight(self) -> torch.Tensor:
        """Convenience for analysis; [N, N] with rows = destination. Never used in the forward pass."""
        return self.sparse_weight().to_dense()

    @property
    def leak(self) -> torch.Tensor:
        return torch.sigmoid(self.recurrent.raw_leak.float())

    # ---- dynamics ------------------------------------------------------------------------------
    def init_state(self, batch: int, device=None) -> torch.Tensor:
        return torch.zeros(batch, self.num_neurons, device=device or self.recurrent.edge_values.device)

    def drive(self, x: torch.Tensor) -> torch.Tensor:
        d = torch.zeros(x.shape[0], self.num_neurons, device=x.device, dtype=torch.float32)
        d[:, self.graph.input_nodes] = self.input_proj(self.embed(x)).float()
        return d

    def step(self, state: torch.Tensor, x: torch.Tensor, W: torch.Tensor | None = None) -> torch.Tensor:
        W = self.sparse_weight() if W is None else W
        drive, leak, bias = self.drive(x), self.leak, self.recurrent.bias.float()
        for _ in range(self.config.microsteps):
            incoming = torch.sparse.mm(W, state.float().T).T
            proposal = torch.tanh(incoming + drive + bias)
            state = (1 - leak) * state + leak * proposal
        return state

    def logits_from_state(self, state: torch.Tensor) -> torch.Tensor:
        return self.lm_head(state[:, self.graph.output_nodes].to(self.lm_head.weight.dtype)).float()

    # ---- fused CUDA path via the connectome-kernels package (optional, used for training) ------------
    def _fused_graph(self):
        from connectome_kernels import SparseGraph
        dev = self.graph.edge_index.device
        if getattr(self, "_fg", None) is None or self._fg_device != dev:
            self._fg = SparseGraph(self.graph.edge_index[0].long(), self.graph.edge_index[1].long(), self.num_neurons,
                                   self.graph.input_nodes)
            self._fg_device = dev
        return self._fg

    def _fused_available(self, device) -> bool:
        if device.type != "cuda":
            return False
        if not hasattr(self, "_fused_ok"):
            try:
                from connectome_kernels import available
                self._fused_ok = available()
            except Exception:
                self._fused_ok = False
        return self._fused_ok

    def _forward_fused(self, input_ids, state):
        from connectome_kernels import sparse_recurrence
        drives = self.input_proj(self.embed(input_ids)).float().permute(1, 2, 0).contiguous()       # [T, n_in, B]
        vals = self.recurrent.edge_values.float() * self.edge_scale()
        out = sparse_recurrence(vals, self.leak, self.recurrent.bias.float(), drives, state,
                                self._fused_graph(), self.config.microsteps)                         # [T, B, N]
        logits = self.lm_head(out[:, :, self.graph.output_nodes].to(self.lm_head.weight.dtype)).float().permute(1, 0, 2)
        return logits, out[-1]

    def forward(self, input_ids: torch.LongTensor, state: Optional[torch.Tensor] = None,
                labels: Optional[torch.LongTensor] = None, use_cache: Optional[bool] = None,
                return_dict: Optional[bool] = None, **kwargs) -> FlyGPTOutput:
        B, T = input_ids.shape
        state = self.init_state(B, input_ids.device) if state is None else state
        if self._fused_available(input_ids.device):
            logits, state = self._forward_fused(input_ids, state)
        else:
            W = self.sparse_weight()
            outs = []
            for t in range(T):
                state = self.step(state, input_ids[:, t], W)
                outs.append(self.logits_from_state(state))
            logits = torch.stack(outs, 1)
        loss = None
        if labels is not None:
            loss = F.cross_entropy(logits[:, :-1].reshape(-1, logits.shape[-1]), labels[:, 1:].reshape(-1))
        return FlyGPTOutput(loss=loss, logits=logits, state=state)

    # ---- generation: carry the neuron state instead of a KV cache ------------------------------
    @classmethod
    def _supports_default_dynamic_cache(cls) -> bool:
        return False  # stateful recurrent model: no KV cache, the neuron state is carried in `state`

    def prepare_inputs_for_generation(self, input_ids, state=None, **kwargs):
        if state is not None:
            input_ids = input_ids[:, -1:]
        return {"input_ids": input_ids, "state": state}

    def _update_model_kwargs_for_generation(self, outputs, model_kwargs, is_encoder_decoder=False, **kwargs):
        model_kwargs["state"] = outputs.state
        return model_kwargs