File size: 3,990 Bytes
b88f761
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Dual-timescale matrix state with denoise and commit lifetimes.

This module is independent of the decoder and commit policy.  It provides the
reference state transition used when replacing the old history and slot paths.
"""

from __future__ import annotations

from dataclasses import dataclass

import torch
from torch import nn

from .gdn2_memory import GDN2Memory


@dataclass(frozen=True)
class GDN2TrajectoryState:
    cells: torch.Tensor
    row: torch.Tensor
    persistent: torch.Tensor
    seen: torch.Tensor


    def shift(self, lengths: torch.Tensor) -> "GDN2TrajectoryState":
        """Shift a non-ring canvas and zero its newly filled tail."""
        batch, canvas = self.seen.shape
        if lengths.shape != (batch,):
            raise ValueError("Commit lengths must be per row.")
        physical = torch.arange(canvas, device=self.seen.device)[None, :] + lengths[:, None]
        kept = physical < canvas
        selected = physical.clamp_max(canvas - 1)
        cells = self.cells.gather(
            1, selected[..., None, None, None].expand_as(self.cells)
        )
        seen = self.seen.gather(1, selected)
        return GDN2TrajectoryState(
            cells.masked_fill(~kept[..., None, None, None], 0.0),
            self.row, self.persistent, seen & kept,
        )


class GDN2TrajectoryMemory(nn.Module):
    def __init__(self, width: int, *, probes: int = 4,
                 working_heads: int = 16, working_key: int = 64,
                 working_value: int = 64, persistent_heads: int = 16,
                 persistent_key: int = 128, persistent_value: int = 128,
                 persistent_observation_dim: int | None = None) -> None:
        super().__init__()
        if probes <= 0:
            raise ValueError("Probe count must be positive.")
        self.probes = probes
        self.cell = GDN2Memory(width, working_heads, working_key, working_value)
        self.row = GDN2Memory(width, working_heads, working_key, working_value)
        self.persistent = GDN2Memory(width, persistent_heads, persistent_key, persistent_value,
                                    observation_dim=persistent_observation_dim)
        self.probe_embed = nn.Embedding(probes, width)
        nn.init.normal_(self.probe_embed.weight, std=0.02)

    def read(self, state: GDN2TrajectoryState, query: torch.Tensor) -> torch.Tensor:
        result = (self.cell.read(state.cells, query)
                  + self.row.read_shared(state.row, query)
                  + self.persistent.read_shared(state.persistent, query))
        return torch.where(state.seen[..., None], result, torch.zeros_like(result))

    def observe(self, state: GDN2TrajectoryState, observation: torch.Tensor,
                live: torch.Tensor, head: torch.Tensor) -> GDN2TrajectoryState:
        batch, canvas, width = observation.shape
        if live.shape != (batch, canvas) or head.shape != (batch,):
            raise ValueError("Observation mask and head have incorrect shapes.")
        cells = self.cell.transition(state.cells, observation, live)
        seen = state.seen | live
        logical_idx = (head[:, None] + torch.arange(canvas, device=head.device)[None, :]) % canvas
        logical = observation.gather(1, logical_idx[..., None].expand(-1, -1, width))
        logical_live = live.gather(1, logical_idx)
        row_state = state.row
        for probe in range(self.probes):
            lo = canvas * probe // self.probes
            hi = canvas * (probe + 1) // self.probes
            selected = logical_live[:, lo:hi]
            count = selected.sum(dim=1, keepdim=True)
            pooled = (logical[:, lo:hi].float() * selected[..., None]).sum(dim=1)
            pooled = (pooled / count.clamp_min(1)).to(observation.dtype)
            pooled = pooled + self.probe_embed.weight[probe].to(observation.dtype)
            row_state = self.row.transition(row_state, pooled, count[:, 0] > 0)
        return GDN2TrajectoryState(cells, row_state, state.persistent, seen)