From b6863b3ba0ef275850a2960e3b36fe6a4fe17f21 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=EC=A0=95=EC=8B=9C=EC=9B=90?= Date: Sun, 10 May 2026 14:42:32 +0900 Subject: [PATCH] Fix ColorSharedNetwork to use real per-color encoding layout ColorSharedNetwork previously sliced the input vector into n_colors equal chunks (input_dim // n_colors). The slice boundaries do not align with the actual encoding layout: adjacent slices contain phase flags, hand slots, expedition state, scores, etc. mixed together. The "color-shared" encoder was therefore sharing weights across semantically unrelated chunks, not across per-color blocks. The single archived run that exercised this path (2026-05-07_092137_color_shared_attention_1000iter) was killed at iter 41 and produced no eval data, so we have no measurement of whether a real per-color architecture would help. Adds compute_lost_cities_color_layout(input_dim) which returns explicit per-color and common index lists for the standard Lost Cities encoding (n_colors=5, hand_size=8, n_ranks=9). It recognises input_dim values 171, 219, 249, 297 across derived_playability and slot_aware_playability flag combinations. Per-color block (39 dims when derived_playability is on): both players' expedition state for that color, discard top, public-histogram row, pending-discard one-hot bit, legal-action draw-pile bit, and the derived_playability per-color block. Slot-aware features are slot-major and stay in common. ColorSharedNetwork.forward now indexes per-color blocks via the layout when input_dim matches a known schema. For other dims (unit tests, non-Lost Cities use), it falls back to chunked slicing with a UserWarning - preserves backward compatibility for tests but makes the legacy behaviour visible. No fair test of the new architecture was run as part of this commit. Documented in docs/plans/deep-cfr-selectivity.md section 7. Co-Authored-By: Claude Opus 4.7 (1M context) --- docs/plans/deep-cfr-selectivity.md | 44 +++- .../games/classic/deep_cfr/networks.py | 205 +++++++++++++++--- tests/games/classic/test_networks.py | 67 ++++++ 3 files changed, 289 insertions(+), 27 deletions(-) diff --git a/docs/plans/deep-cfr-selectivity.md b/docs/plans/deep-cfr-selectivity.md index 077f1e5..eb1f61f 100644 --- a/docs/plans/deep-cfr-selectivity.md +++ b/docs/plans/deep-cfr-selectivity.md @@ -566,7 +566,49 @@ Deep CFR can do in this game from raw input. The model may stay in the phase 1 trap longer or fail more visibly, both of which are useful information. -### 7. Short open-selectivity ablation +### 7. ColorSharedNetwork chunked-layout bug (2026-05-10) + +While re-examining the archived `2026-05-07_092137_color_shared_attention_1000iter` +run (killed at iter 41), discovered the `ColorSharedNetwork` implementation +in `networks.py` was not actually color-aware. The forward pass split the +input vector into `input_dim // n_colors` contiguous slices and ran them +through a shared encoder. The slice boundaries do not align with the actual +encoding layout — adjacent slices contain phase flags, hand slots, +expedition state, scores, etc. mixed together. The "shared color encoder" +was therefore sharing weights across semantically unrelated chunks, not +across per-color blocks. + +This means the prior conclusion that "color_shared / attention archive run +was inconclusive" was charitable. The architecture being measured was a +chunked-input network mislabeled `color_shared`, not a real per-color +shared architecture. We have *no* signal on whether a properly per-color +architecture would help. + +Fix landed in this same session: + +- Added `compute_lost_cities_color_layout(input_dim)` in `networks.py`. For + the standard Lost Cities schema (n_colors=5, hand_size=8, n_ranks=9), it + recognises the four valid `input_dim` values (171, 219, 249, 297 across + derived/slot-aware flag combinations) and returns per-color and common + index lists derived from the actual encoding layout. +- Per-color block (39 dims with `derived_playability` on): both players' + expedition state for that color, discard top metadata, public-histogram + row, pending-discard one-hot bit, legal-action draw-pile bit, and the + `derived_playability` per-color block. Slot-aware features stay in + common because they are slot-major, not color-major. +- `ColorSharedNetwork.forward` now indexes per-color blocks via the layout + when `input_dim` matches a known schema. For other input dims (unit + tests, non-Lost Cities use), it falls back to the old chunked slicing + with a `UserWarning`, preserving backward compatibility for tests but + making the legacy behaviour visible. +- New unit tests cover both branches and the layout helper. + +This is purely an implementation correctness fix; no fair test of the +architecture has been run yet. Fair test deferred — the diagnosis from +sections 4–5 (selection bias, post-open behaviour) suggests that even a +correct color-aware encoder would not break the closed loop on its own. + +### 8. Short open-selectivity ablation Run a 200-300 iteration ablation only after the target audit identifies a specific change. Candidate changes include: diff --git a/src/coolrl_lost_cities/games/classic/deep_cfr/networks.py b/src/coolrl_lost_cities/games/classic/deep_cfr/networks.py index f173b1e..e0b1683 100644 --- a/src/coolrl_lost_cities/games/classic/deep_cfr/networks.py +++ b/src/coolrl_lost_cities/games/classic/deep_cfr/networks.py @@ -1,5 +1,8 @@ from __future__ import annotations +import warnings +from dataclasses import dataclass + import torch from torch import nn @@ -32,6 +35,118 @@ def _build_mlp( return nn.Sequential(*layers) +@dataclass(frozen=True) +class ColorLayout: + """Per-color and common index lists for a Lost Cities encoding vector.""" + + per_color_indices: tuple[tuple[int, ...], ...] + common_indices: tuple[int, ...] + + @property + def n_colors(self) -> int: + return len(self.per_color_indices) + + @property + def color_block_size(self) -> int: + return len(self.per_color_indices[0]) + + @property + def common_size(self) -> int: + return len(self.common_indices) + + +def compute_lost_cities_color_layout( + input_dim: int, + *, + n_colors: int = 5, + hand_size: int = 8, + n_ranks: int = 9, +) -> ColorLayout | None: + """Map encoding offsets to per-color blocks for the standard Lost Cities schema. + + Returns ``None`` when ``input_dim`` does not match a recognised combination of + encoding flags (``derived_playability``, ``slot_aware_playability``) under the + given schema. The standard Lost Cities tier-3 game uses the defaults + n_colors=5, hand_size=8, n_ranks=9, which produces 171 / 219 / 249 / 297 dims. + + Per-color indices, in semantic order, gather: + - both players' expedition state for the color (4 dims each) + - discard top metadata for the color (4 dims) + - public card-type histogram for the color (n_ranks + 1 dims) + - pending-discard one-hot bit for the color + - legal-action draw-pile bit for the color + - derived_playability per-color block (15 dims) when enabled + Slot-aware features are slot-major (not color-major) and stay in common. + """ + base_dim = ( + 5 + + hand_size * 3 + + 2 * n_colors * 4 + + n_colors * 4 + + n_colors * (n_ranks + 1) + + 3 + + 1 + + (n_colors + 1) + + (2 * hand_size + 1 + n_colors) + ) + derived_size = n_colors * 15 + 3 + slot_size = hand_size * 6 + + has_derived = False + if input_dim == base_dim: + pass + elif input_dim == base_dim + derived_size: + has_derived = True + elif input_dim == base_dim + slot_size: + pass + elif input_dim == base_dim + derived_size + slot_size: + has_derived = True + else: + return None + + per_color: list[list[int]] = [[] for _ in range(n_colors)] + + expedition_start = 5 + hand_size * 3 + for player in range(2): + for color in range(n_colors): + base_idx = expedition_start + player * n_colors * 4 + color * 4 + per_color[color].extend(range(base_idx, base_idx + 4)) + + discard_start = expedition_start + 2 * n_colors * 4 + for color in range(n_colors): + per_color[color].extend(range(discard_start + color * 4, discard_start + color * 4 + 4)) + + histogram_start = discard_start + n_colors * 4 + for color in range(n_colors): + block_start = histogram_start + color * (n_ranks + 1) + per_color[color].extend(range(block_start, block_start + n_ranks + 1)) + + pending_start = histogram_start + n_colors * (n_ranks + 1) + 3 + 1 + for color in range(n_colors): + per_color[color].append(pending_start + color) + + legal_start = pending_start + n_colors + 1 + draw_pile_start = legal_start + 2 * hand_size + 1 + for color in range(n_colors): + per_color[color].append(draw_pile_start + color) + + if has_derived: + derived_start = base_dim + for color in range(n_colors): + block_start = derived_start + color * 15 + per_color[color].extend(range(block_start, block_start + 15)) + + color_set: set[int] = set() + for indices in per_color: + color_set.update(indices) + common = [i for i in range(input_dim) if i not in color_set] + + return ColorLayout( + per_color_indices=tuple(tuple(indices) for indices in per_color), + common_indices=tuple(common), + ) + + class DeepCFRMLP(nn.Module): def __init__( self, @@ -70,13 +185,27 @@ class DeepCFRMLP(nn.Module): class ColorSharedNetwork(nn.Module): - """Color-shared architecture that splits input into per-color blocks. + """Color-shared architecture for Lost Cities encodings. - Splits the input into n_colors equal parts, encodes each with shared weights, - pools the color embeddings, and concatenates with the original input. + For a recognised Lost Cities encoding (``input_dim`` matching a known + combination of base, derived_playability, and slot_aware_playability), the + forward pass gathers the per-color feature indices computed by + :func:`compute_lost_cities_color_layout`, runs a shared encoder over each + color block, mean+max pools, and concatenates the result with the + color-independent ("common") features before the final head. + + For any other ``input_dim`` (e.g. unit tests using ``input_dim=100``), the + network falls back to a *chunked* layout that splits the input into + ``input_dim // n_colors`` equal slices. The chunked layout was the only + behaviour shipped before 2026-05-10 and does **not** correspond to actual + per-color blocks in the encoding — adjacent slices contain unrelated + features (phase flags, hand slots, scores, etc.). It is preserved purely + for backward compatibility with older checkpoints and tests; new training + runs should always use the standard Lost Cities encoding so the proper + layout is selected automatically. """ - N_COLORS = 5 + DEFAULT_N_COLORS = 5 def __init__( self, @@ -93,15 +222,36 @@ class ColorSharedNetwork(nn.Module): self.input_dim = input_dim self.output_dim = output_dim self.hidden_size = hidden_size - self.n_colors = self.N_COLORS self.color_attention_layers = color_attention_layers self.color_attention_heads = color_attention_heads - color_block_size = input_dim // self.n_colors - self.color_block_size = color_block_size + layout = compute_lost_cities_color_layout(input_dim) + + if layout is not None: + self.use_chunked_fallback = False + self.n_colors = layout.n_colors + self.color_block_size = layout.color_block_size + self.common_size = layout.common_size + per_color_idx = torch.tensor( + [list(indices) for indices in layout.per_color_indices], dtype=torch.long + ) + common_idx = torch.tensor(list(layout.common_indices), dtype=torch.long) + self.register_buffer("per_color_indices", per_color_idx, persistent=False) + self.register_buffer("common_indices", common_idx, persistent=False) + else: + warnings.warn( + f"ColorSharedNetwork: input_dim={input_dim} does not match the " + "standard Lost Cities encoding schema; falling back to chunked " + "input slicing (legacy behaviour, semantically not per-color).", + stacklevel=2, + ) + self.use_chunked_fallback = True + self.n_colors = self.DEFAULT_N_COLORS + self.color_block_size = input_dim // self.n_colors + self.common_size = input_dim - self.n_colors * self.color_block_size self.color_encoder = _build_mlp( - color_block_size, + self.color_block_size, hidden_size, hidden_size, num_layers, @@ -117,8 +267,7 @@ class ColorSharedNetwork(nn.Module): activation=activation, ) - final_input_dim = hidden_size * 2 + input_dim % self.n_colors - + final_input_dim = hidden_size * 2 + self.common_size self.final_net = _build_mlp( final_input_dim, output_dim, @@ -128,27 +277,31 @@ class ColorSharedNetwork(nn.Module): ) def forward(self, x: torch.Tensor) -> torch.Tensor: - color_embeddings = [] - for i in range(self.n_colors): - start = i * self.color_block_size - end = start + self.color_block_size - block = x[:, start:end] - embedding = self.color_encoder(block) - color_embeddings.append(embedding) + if self.use_chunked_fallback: + color_blocks = [] + for i in range(self.n_colors): + start = i * self.color_block_size + end = start + self.color_block_size + color_blocks.append(x[:, start:end]) + stacked = torch.stack(color_blocks, dim=1) + common = x[:, self.n_colors * self.color_block_size :] + else: + stacked = x[:, self.per_color_indices] + common = x[:, self.common_indices] - color_embeddings = torch.stack(color_embeddings, dim=1) + batch_size = stacked.shape[0] + flat = stacked.reshape(batch_size * self.n_colors, self.color_block_size) + encoded_flat = self.color_encoder(flat) + encoded = encoded_flat.reshape(batch_size, self.n_colors, self.hidden_size) if self.color_attention is not None: - color_embeddings = self.color_attention(color_embeddings) + encoded = self.color_attention(encoded) - mean_pooled = color_embeddings.mean(dim=1) - max_pooled = color_embeddings.max(dim=1)[0] + mean_pooled = encoded.mean(dim=1) + max_pooled = encoded.max(dim=1)[0] - remainder = x[:, self.n_colors * self.color_block_size :] - final_features = torch.cat([mean_pooled, max_pooled, remainder], dim=1) - - logits = self.final_net(final_features) - return logits + final_features = torch.cat([mean_pooled, max_pooled, common], dim=1) + return self.final_net(final_features) class ColorAttention(nn.Module): diff --git a/tests/games/classic/test_networks.py b/tests/games/classic/test_networks.py index a1aac42..1b88067 100644 --- a/tests/games/classic/test_networks.py +++ b/tests/games/classic/test_networks.py @@ -308,3 +308,70 @@ class TestNetworkIntegration: x = torch.randn(8, dim) output = network(x) assert output.shape == (8, action_size) + + +class TestComputeLostCitiesColorLayout: + def test_layout_for_full_encoding_input_dim(self) -> None: + from coolrl_lost_cities.games.classic.deep_cfr.networks import ( + compute_lost_cities_color_layout, + ) + + layout = compute_lost_cities_color_layout(297) + assert layout is not None + assert layout.n_colors == 5 + assert layout.color_block_size == 39 + assert layout.common_size == 297 - 5 * 39 + + all_color_idx: set[int] = set() + for indices in layout.per_color_indices: + assert len(indices) == 39 + all_color_idx.update(indices) + assert len(all_color_idx) == 5 * 39 + assert set(layout.common_indices).isdisjoint(all_color_idx) + assert all_color_idx | set(layout.common_indices) == set(range(297)) + + def test_layout_returns_none_for_unknown_input_dim(self) -> None: + from coolrl_lost_cities.games.classic.deep_cfr.networks import ( + compute_lost_cities_color_layout, + ) + + assert compute_lost_cities_color_layout(100) is None + assert compute_lost_cities_color_layout(150) is None + assert compute_lost_cities_color_layout(296) is None + + def test_layout_recognises_all_four_flag_combinations(self) -> None: + from coolrl_lost_cities.games.classic.deep_cfr.networks import ( + compute_lost_cities_color_layout, + ) + + for dim in (171, 219, 249, 297): + layout = compute_lost_cities_color_layout(dim) + assert layout is not None, f"layout missing for input_dim={dim}" + assert layout.n_colors == 5 + + def test_color_shared_uses_proper_layout_for_real_encoding(self) -> None: + import warnings + + from coolrl_lost_cities.games.classic.deep_cfr.networks import ColorSharedNetwork + + network = ColorSharedNetwork(input_dim=297, output_dim=22, hidden_size=64) + assert network.use_chunked_fallback is False + assert network.color_block_size == 39 + + with warnings.catch_warnings(record=True) as caught: + warnings.simplefilter("always") + x = torch.randn(4, 297) + out = network(x) + assert out.shape == (4, 22) + assert not any("chunked" in str(w.message).lower() for w in caught) + + def test_color_shared_warns_on_non_lost_cities_input_dim(self) -> None: + import warnings + + from coolrl_lost_cities.games.classic.deep_cfr.networks import ColorSharedNetwork + + with warnings.catch_warnings(record=True) as caught: + warnings.simplefilter("always") + network = ColorSharedNetwork(input_dim=100, output_dim=20, hidden_size=64) + assert network.use_chunked_fallback is True + assert any("chunked" in str(w.message).lower() for w in caught)