File size: 7,502 Bytes
f6ef360
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
0586ff6
 
 
 
 
 
 
 
 
f6ef360
 
 
 
0586ff6
f6ef360
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
0586ff6
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
f6ef360
0586ff6
 
 
 
 
 
 
 
 
 
 
f6ef360
0586ff6
 
 
 
 
 
f6ef360
0586ff6
 
 
f6ef360
 
0586ff6
f6ef360
 
 
 
0586ff6
 
 
 
f6ef360
 
 
 
 
0586ff6
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
f6ef360
 
 
 
 
 
0586ff6
f6ef360
 
 
 
 
0586ff6
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
f6ef360
 
 
0586ff6
 
 
f6ef360
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
"""Spatial-only autoencoder (v3): compress only what is actually big.

Lesson from v2: squeezing tiny exact state (pairwise alliances, player
scalars) through the latent and reconstructing it lossily is manufactured
difficulty — three training runs got alliance precision to 0.55 when the
policy could just read the bit. The bypass split:

  AE compresses (high-dimensional, spatial):
    - tile ownership grid + terrain + fallout
    - static structure planes (city/port/defense post/silo/SAM/factory)

  Bypass straight to the policy (small, exact):
    - pairwise diplomacy (alliances, embargoes, expiry, pending requests)
    - per-player scalars (troops, gold, tiles, alive, ...)
    - transient units (nukes/transports/warships) as (owner, pos, target)
    - attack aggregates, globals, legality masks

Latent: z_grid (B, latent_c, H/16, W/16). No global vector: the policy can
pool spatially itself, and per-player/global state arrives via bypass.

v3.1 additions (all off by default so old v3 checkpoints load unchanged):
  - terrain_cond: the decoder consumes the 2 STATIC terrain planes (land,
    magnitude) as free side-information at every scale, so the latent only
    has to encode ownership relative to terrain. Fallout is dynamic state
    and is never fed to the decoder.
  - upsample_decoder: nearest-upsample + 3x3 conv stages (no checkerboard)
    plus a full-resolution 3x3 refinement block before the classifier.
  - latent_down: 8 or 16; latent grid at 1/8 or 1/16 resolution.
"""

import torch
import torch.nn as nn
import torch.nn.functional as F

from ae.units import STATIC_CLASSES

MAX_SLOTS = 128
OWNER_EMB_DIM = 8
TERRAIN_CHANNELS = 3  # land, magnitude, fallout
NUM_STATIC = len(STATIC_CLASSES)

# BCE positive-weight multiplier per static class (× --unit-pos-weight).
STATIC_CLASS_WEIGHTS = [1.0, 1.0, 1.0, 4.0, 4.0, 2.0]


def conv_block(c_in: int, c_out: int, stride: int) -> nn.Sequential:
    return nn.Sequential(
        nn.Conv2d(c_in, c_out, kernel_size=3, stride=stride, padding=1),
        nn.GroupNorm(8, c_out),
        nn.SiLU(),
    )


def deconv_block(c_in: int, c_out: int) -> nn.Sequential:
    return nn.Sequential(
        nn.ConvTranspose2d(c_in, c_out, kernel_size=4, stride=2, padding=1),
        nn.GroupNorm(8, c_out),
        nn.SiLU(),
    )


class UpsampleBlock(nn.Module):
    """Nearest 2x upsample + 3x3 conv (no ConvTranspose checkerboard).

    Optionally concatenates static terrain planes (at the post-upsample
    resolution) before the conv.
    """

    def __init__(self, c_in: int, c_out: int, extra_c: int = 0):
        super().__init__()
        self.conv = conv_block(c_in + extra_c, c_out, stride=1)

    def forward(self, x: torch.Tensor, extra: torch.Tensor | None = None):
        x = F.interpolate(x, scale_factor=2, mode="nearest")
        if extra is not None:
            x = torch.cat([x, extra], dim=1)
        return self.conv(x)


STATIC_TERRAIN_C = 2  # land, magnitude (fallout is dynamic: never decoded from)


class SpatialAE(nn.Module):
    """Defaults preserve the original v3 architecture (old checkpoints load
    with strict state_dict matching). v3.1 runs set terrain_cond=True and
    upsample_decoder=True (and optionally latent_down=8)."""

    def __init__(
        self,
        latent_c: int = 64,
        terrain_cond: bool = False,
        upsample_decoder: bool = False,
        latent_down: int = 16,
    ):
        super().__init__()
        if latent_down not in (8, 16):
            raise ValueError(f"latent_down must be 8 or 16, got {latent_down}")
        if latent_down == 8 and not upsample_decoder:
            raise ValueError("latent_down=8 requires the v3.1 upsample decoder")
        if terrain_cond and not upsample_decoder:
            raise ValueError("terrain_cond requires the v3.1 upsample decoder")
        self.latent_c = latent_c
        self.terrain_cond = terrain_cond
        self.upsample_decoder = upsample_decoder
        self.latent_down = latent_down
        self.owner_emb = nn.Embedding(MAX_SLOTS, OWNER_EMB_DIM)

        stem = [
            conv_block(OWNER_EMB_DIM + TERRAIN_CHANNELS, 32, stride=1),
            conv_block(32, 64, stride=2),
            conv_block(64, 96, stride=2),
            conv_block(96, 128, stride=2),
        ]
        if latent_down == 16:
            stem.append(conv_block(128, 128, stride=2))  # -> 1/16
        self.enc_stem = nn.Sequential(*stem)
        self.enc_fuse = nn.Sequential(
            conv_block(128 + NUM_STATIC, 128, stride=1),
            nn.Conv2d(128, latent_c, kernel_size=1),
        )

        cond_c = STATIC_TERRAIN_C if terrain_cond else 0
        self.dec_in = conv_block(latent_c + cond_c, 128, stride=1)
        if upsample_decoder:
            chans = [128, 128, 96, 64, 32] if latent_down == 16 else [128, 96, 64, 32]
            self.dec_up = nn.ModuleList(
                UpsampleBlock(chans[i], chans[i + 1], extra_c=cond_c)
                for i in range(len(chans) - 1)
            )
            self.dec_refine = conv_block(32 + cond_c, 32, stride=1)
            self.dec_out = nn.Conv2d(32, MAX_SLOTS, kernel_size=1)
        else:
            self.dec_tiles = nn.Sequential(
                deconv_block(128, 128),
                deconv_block(128, 96),
                deconv_block(96, 64),
                deconv_block(64, 32),
                nn.Conv2d(32, MAX_SLOTS, kernel_size=1),
            )
        # Static structure occupancy logits at latent resolution.
        self.dec_units = nn.Conv2d(128, NUM_STATIC, kernel_size=1)

    def encode(
        self,
        owners: torch.Tensor,  # (B, H, W) int64
        terrain: torch.Tensor,  # (B, 3, H, W)
        static_planes: torch.Tensor,  # (B, NUM_STATIC, H/down, W/down)
    ) -> torch.Tensor:
        emb = self.owner_emb(owners).permute(0, 3, 1, 2)
        g = self.enc_stem(torch.cat([emb, terrain], dim=1))
        return self.enc_fuse(torch.cat([g, static_planes], dim=1))

    def decode(
        self,
        z_grid: torch.Tensor,
        terrain: torch.Tensor | None = None,  # (B, >=2, H, W); only [:, :2] used
    ) -> tuple[torch.Tensor, torch.Tensor]:
        if not self.terrain_cond:
            h = self.dec_in(z_grid)
            if self.upsample_decoder:
                x = h
                for up in self.dec_up:
                    x = up(x)
                return self.dec_out(self.dec_refine(x)), self.dec_units(h)
            return self.dec_tiles(h), self.dec_units(h)

        if terrain is None:
            raise ValueError("terrain_cond model needs terrain in decode()")
        # Static side-information pyramid: full res, 1/2, 1/4, ... latent res.
        static_t = terrain[:, :STATIC_TERRAIN_C]
        pyramid = {1: static_t}
        down = 2
        while down <= self.latent_down:
            pyramid[down] = F.avg_pool2d(static_t, kernel_size=down)
            down *= 2

        h = self.dec_in(torch.cat([z_grid, pyramid[self.latent_down]], dim=1))
        x = h
        scale = self.latent_down
        for up in self.dec_up:
            scale //= 2
            x = up(x, pyramid[scale])
        x = self.dec_refine(torch.cat([x, pyramid[1]], dim=1))
        return self.dec_out(x), self.dec_units(h)

    def forward(self, owners, terrain, static_planes):
        z_grid = self.encode(owners, terrain, static_planes)
        tile_logits, unit_logits = self.decode(
            z_grid, terrain if self.terrain_cond else None
        )
        return tile_logits, unit_logits, z_grid