File size: 3,845 Bytes
e4f7326
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""MLX rolling canvas, detached history, row tape, and commit-only slot state."""

from __future__ import annotations

from dataclasses import dataclass

import mlx.core as mx

from .mlx_gdn2_trajectory import GDN2TrajectoryState


def _overwritten(lengths: mx.array, head: mx.array, canvas: int) -> mx.array:
    return ((mx.arange(canvas)[None, :] - head[:, None]) % canvas) < lengths[:, None]


@dataclass(frozen=True)
class MLXLatentState:
    memory_slots: mx.array
    confidence: mx.array
    entropy: mx.array
    age: mx.array
    token_changed: mx.array
    confidence_delta: mx.array
    entropy_delta: mx.array
    ponder_steps: mx.array
    stagnation_steps: mx.array
    gdn2: GDN2TrajectoryState | None = None

    @classmethod
    def empty(cls, batch: int, canvas: int, slots: int, latent: int,
              *, dtype: mx.Dtype = mx.bfloat16,
              enable_gdn2: bool = False) -> "MLXLatentState":
        zeros = lambda: mx.zeros((batch, canvas), mx.float32)
        persistent = (mx.zeros((batch, 16, 128, 128), mx.float32)
                      if enable_gdn2 else None)
        return cls(persistent if persistent is not None
                   else mx.zeros((batch, slots, latent), dtype),
                   zeros(), zeros(), mx.zeros((batch, canvas), mx.int32),
                   zeros(), zeros(), zeros(),
                   mx.zeros((batch,), mx.int32), mx.zeros((batch,), mx.int32),
                   GDN2TrajectoryState(
                       mx.zeros((batch, canvas, 16, 64, 64), mx.float32),
                       mx.zeros((batch, 16, 64, 64), mx.float32),
                       persistent,
                       mx.zeros((batch, canvas), mx.bool_),
                   ) if enable_gdn2 else None)

    def advance_ring(self, lengths: mx.array, head: mx.array,
                     *, entropy_fill_value: float,
                     reset_clocks_mask: mx.array | None = None) -> "MLXLatentState":
        mask = _overwritten(lengths, head, self.confidence.shape[-1])
        committed = lengths > 0 if reset_clocks_mask is None else reset_clocks_mask
        return MLXLatentState(
            self.memory_slots,
            mx.where(mask, 0.0, self.confidence),
            mx.where(mask, entropy_fill_value, self.entropy),
            mx.where(mask, 0, self.age),
            mx.where(mask, 0.0, self.token_changed),
            mx.where(mask, 0.0, self.confidence_delta),
            mx.where(mask, 0.0, self.entropy_delta),
            mx.where(committed, 0, self.ponder_steps).astype(mx.int32),
            mx.where(committed, 0, self.stagnation_steps).astype(mx.int32),
            None if self.gdn2 is None else self.gdn2.clear_refill(lengths, head),
        )


@dataclass(frozen=True)
class MLXRollingState:
    canvas: mx.array
    latent: MLXLatentState
    head: mx.array

    def advance_ring(self, commit_lengths: mx.array, tail_tokens: mx.array,
                     *, entropy_fill_value: float,
                     reset_clocks_mask: mx.array | None = None) -> "MLXRollingState":
        batch, canvas = self.canvas.shape
        if commit_lengths.shape != (batch,) or tail_tokens.shape != self.canvas.shape:
            raise ValueError("Commit lengths or tail tokens do not match the canvas.")
        overwritten = _overwritten(commit_lengths, self.head, canvas)
        old_offsets = (mx.arange(canvas)[None, :] - self.head[:, None]) % canvas
        tail_source = mx.take_along_axis(tail_tokens, old_offsets, axis=1)
        return MLXRollingState(
            canvas=mx.where(overwritten, tail_source, self.canvas),
            latent=self.latent.advance_ring(
                commit_lengths, self.head, entropy_fill_value=entropy_fill_value,
                reset_clocks_mask=reset_clocks_mask,
            ),


            head=(self.head + commit_lengths) % canvas,
        )