File size: 8,776 Bytes
872cf4d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
201
202
203
204
205
206
207
208
209
"""Latent-to-pixel decoder — a diagnostic, never part of the control loop.



Reimplements the decoder described in LeWM appendix D ("Decoder (Visualization

Only)", arXiv:2603.19312). **It is not shipped with the model**: the released

``quentinll/lewm-pusht`` checkpoint holds 303 tensors under

``encoder/predictor/projector/pred_proj/action_encoder`` and nothing decoder

shaped, and neither ``stable-worldmodel`` nor ``lucas-maes/le-wm`` contains the

code. So it has to be trained here — see ``scripts/train_decoder.py``.



The paper's description, restated as shapes::



    z          (B, 192)                 one latent vector, the whole input

      -> to_hidden                      (B, 1, W)      memory: keys/values

    queries    (B, P, W)                P learned tokens, one per output patch

      -> depth x [cross-attn(queries <- memory); residual MLP]

      -> to_pixels                      (B, P, p*p*3)

      -> unpatchify                     (B, 3, 224, 224)



``P = (image_size / patch_size)^2``. The paper's worked example is

``(224/16)^2 = 196``. That 16 is the **decoder's own output tiling** and is

unrelated to the encoder's patch size (which is 14 for this checkpoint): the

decoder never sees patch tokens, only the single pooled vector, so it is free

to carve its output image however it likes.



### One thing worth knowing before you read the attention maps



With ``memory_tokens=1`` — exactly what the paper describes — the cross

attention has a single key. ``softmax`` over one element is identically 1, so

every query gets the same value vector and the layer reduces to



    x <- x + W_o V(z)          (the same increment added to all P queries)



i.e. FiLM-style global conditioning, not a spatial lookup. That is fine for the

job (all spatial structure lives in the learned queries), but attention weights

carry no information and plotting them would be meaningless. Set

``memory_tokens > 1`` to expand ``z`` into several distinct memory slots if you

want attention that actually selects.

"""

import torch
from torch import nn


class CrossAttentionBlock(nn.Module):
    """Pre-norm cross-attention to a fixed memory, then a residual MLP."""

    def __init__(self, dim, heads=6, mlp_ratio=4.0, dropout=0.0):
        super().__init__()
        self.norm_q = nn.LayerNorm(dim)
        self.norm_kv = nn.LayerNorm(dim)
        self.attn = nn.MultiheadAttention(
            dim, heads, dropout=dropout, batch_first=True
        )
        self.norm_mlp = nn.LayerNorm(dim)
        hidden = int(dim * mlp_ratio)
        self.mlp = nn.Sequential(
            nn.Linear(dim, hidden),
            nn.GELU(),
            nn.Dropout(dropout),
            nn.Linear(hidden, dim),
            nn.Dropout(dropout),
        )

    def forward(self, x, memory):
        q = self.norm_q(x)
        kv = self.norm_kv(memory)
        x = x + self.attn(q, kv, kv, need_weights=False)[0]
        return x + self.mlp(self.norm_mlp(x))


class LatentDecoder(nn.Module):
    """Decode one latent vector into a 224x224 RGB image.



    Args:

        latent_dim: Width of the latent being decoded (192 here).

        hidden_dim: Decoder working width.

        image_size / patch_size: Output tiling; ``P = (image_size/patch_size)^2``

            learned query tokens, each emitting one ``patch_size^2 * 3`` patch.

        depth: Number of cross-attention + MLP blocks.

        heads: Attention heads.

        memory_tokens: Slots the latent is expanded into. 1 reproduces the

            paper (and makes attention degenerate — see the module docstring).

        out_range: ``'unit'`` clamps nothing and expects targets in ``[0, 1]``

            (the paper's "linearly projected to pixels"); ``'sigmoid'`` bounds

            the output instead. Linear is the faithful default.

    """

    def __init__(

        self,

        latent_dim=192,

        hidden_dim=384,

        image_size=224,

        patch_size=16,

        depth=4,

        heads=6,

        mlp_ratio=4.0,

        dropout=0.0,

        memory_tokens=1,

        out_range='unit',

    ):
        super().__init__()
        assert image_size % patch_size == 0, (
            f'image_size {image_size} not divisible by patch_size {patch_size}'
        )
        self.image_size = image_size
        self.patch_size = patch_size
        self.grid = image_size // patch_size
        self.num_patches = self.grid ** 2
        self.memory_tokens = memory_tokens
        self.out_range = out_range

        # The latent is the decoder's entire input: it becomes the key/value
        # memory that every query reads from.
        self.to_hidden = nn.Linear(latent_dim, hidden_dim * memory_tokens)

        # One learned query per output patch. These carry all the spatial
        # structure — the latent itself has no spatial layout at all.
        self.queries = nn.Parameter(
            torch.randn(1, self.num_patches, hidden_dim) * 0.02
        )

        self.blocks = nn.ModuleList(
            CrossAttentionBlock(hidden_dim, heads, mlp_ratio, dropout)
            for _ in range(depth)
        )
        self.norm = nn.LayerNorm(hidden_dim)
        self.to_pixels = nn.Linear(hidden_dim, patch_size * patch_size * 3)

    @torch.no_grad()
    def init_output_at(self, mean_pixel):
        """Start the decoder at "always predict the mean colour".



        Default init emits roughly zero-mean patches against targets in

        ``[0, 1]``, so step 0 sits ~300x worse than simply predicting the mean

        image and the first few hundred steps are spent rediscovering the

        background. Zeroing the output weight and parking the bias on the mean

        colour starts training at that baseline instead, so every subsequent

        step is spent on structure. Gradients are unaffected — the weight is

        zero, not frozen.



        Args:

            mean_pixel: Per-channel mean in ``[0, 1]``, shape ``(3,)``.

        """
        mean_pixel = torch.as_tensor(mean_pixel, dtype=torch.float32).view(3)
        patch = mean_pixel.repeat(self.patch_size * self.patch_size)
        self.to_pixels.weight.zero_()
        self.to_pixels.bias.copy_(patch)

    def unpatchify(self, patches):
        """``(B, P, p*p*3)`` -> ``(B, 3, H, W)``, row-major over the grid."""
        B = patches.size(0)
        g, p = self.grid, self.patch_size
        x = patches.view(B, g, g, p, p, 3)
        x = x.permute(0, 5, 1, 3, 2, 4)          # (B, 3, gy, py, gx, px)
        return x.reshape(B, 3, g * p, g * p)

    def patchify(self, images):
        """``(B, 3, H, W)`` -> ``(B, P, p*p*3)``; inverse of :meth:`unpatchify`."""
        B = images.size(0)
        g, p = self.grid, self.patch_size
        x = images.view(B, 3, g, p, g, p)
        x = x.permute(0, 2, 4, 3, 5, 1)          # (B, gy, gx, py, px, c)
        return x.reshape(B, self.num_patches, p * p * 3)

    def forward(self, z):
        """Decode ``(B, D)`` or ``(B, T, D)`` latents to images.



        A ``(B, T, D)`` input returns ``(B, T, 3, H, W)``, which is what you

        want for a plan rollout: one frame per horizon step.

        """
        squeeze_time = z.dim() == 2
        if squeeze_time:
            z = z.unsqueeze(1)
        B, T, _ = z.shape
        z = z.reshape(B * T, -1)

        memory = self.to_hidden(z).view(B * T, self.memory_tokens, -1)
        x = self.queries.expand(B * T, -1, -1)
        for block in self.blocks:
            x = block(x, memory)

        images = self.unpatchify(self.to_pixels(self.norm(x)))
        if self.out_range == 'sigmoid':
            images = torch.sigmoid(images)

        images = images.view(B, T, 3, self.image_size, self.image_size)
        return images.squeeze(1) if squeeze_time else images


def reconstruction_loss(pred, target):
    """Plain pixel MSE.



    The paper never writes the reconstruction objective down — it only says a

    decoder was "trained to reconstruct pixel observations" and, in the

    appendix G ablation, that adding "a reconstruction loss" to LeWM training

    *hurt* control (PushT 96.0 -> 86.0 SR). MSE is the default reading and the

    one used here; nothing downstream depends on the choice, since this decoder

    never touches the world model's gradients.

    """
    return torch.nn.functional.mse_loss(pred, target)


def load_decoder(path, device='cpu'):
    """Rebuild a decoder from a training checkpoint."""
    ckpt = torch.load(path, map_location=device, weights_only=False)
    decoder = LatentDecoder(**ckpt['config'])
    decoder.load_state_dict(ckpt['state_dict'])
    return decoder.to(device).eval(), ckpt