From 763f6258d253f857d353fbe171e69bda73768865 Mon Sep 17 00:00:00 2001 From: unanmed <1319491857@qq.com> Date: Tue, 4 Aug 2026 19:46:50 +0800 Subject: [PATCH] =?UTF-8?q?refactor:=20=E9=87=8D=E6=9E=84=E4=BB=A3?= =?UTF-8?q?=E7=A0=81=E7=BB=93=E6=9E=84?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- prompt.md => AGENTS.md | 4 + docs/重训问题跟踪.md | 126 ++++ ginka/dataset.py | 8 +- ginka/diagnose_codebook.py | 55 ++ ginka/maskGIT/model.py | 23 +- ginka/model.py | 193 +++++ ginka/sample.py | 144 ++++ ginka/train_seperated.py | 1396 +++++------------------------------- ginka/utils.py | 183 +++-- ginka/vqvae/model.py | 50 -- ginka/vqvae/quantize.py | 26 +- shared/image.py | 28 +- test_refactor.py | 336 +++++++++ 需要重训的问题.md | 445 ++++++++++++ 14 files changed, 1648 insertions(+), 1369 deletions(-) rename prompt.md => AGENTS.md (96%) create mode 100644 docs/重训问题跟踪.md create mode 100644 ginka/diagnose_codebook.py create mode 100644 ginka/model.py create mode 100644 ginka/sample.py create mode 100644 test_refactor.py create mode 100644 需要重训的问题.md diff --git a/prompt.md b/AGENTS.md similarity index 96% rename from prompt.md rename to AGENTS.md index 9ed041d..62d2fdf 100644 --- a/prompt.md +++ b/AGENTS.md @@ -78,3 +78,7 @@ - 编写验证代码时,优先输出可视化结果(图片文件),使用 `shared/image.py` 中的工具 - 验证阶段应对不同条件(不同 z 采样)分别生成图片,便于直观对比模型效果 + +## 其他 + +`app` 目录下的内容不用管,目前尚在训练阶段,还未到达推理发布阶段。 diff --git a/docs/重训问题跟踪.md b/docs/重训问题跟踪.md new file mode 100644 index 0000000..e8f56b6 --- /dev/null +++ b/docs/重训问题跟踪.md @@ -0,0 +1,126 @@ +# 重训问题跟踪 + +> 基于 `需要重训的问题.md` 整理,标注了当前重构状态(2026-08-04)。 +> 代码基准:当前 HEAD(已删除距离场通路、掩码位 CE、修正子集 3)。 + +--- + +## 已完成(本次重构已解决) + +| 问题 | 内容 | 措施 | +|------|------|------| +| **A 方案 1** | 距离场整条通路空转 1069 万参数 | 已删除 `DistFieldEncoder` + `dist_quantizer` + `z_dist` 全部通路 | +| **D** | CE 在全部 169 格上算,一半是"抄写" | 已在 `cross_entropy_loss` 加入 mask 参数,仅对被掩码位置计算 | +| **H-1** | 子集 3 空操作,20% 数据重复 | 已删除无效的 `out[0][out[0] == ENTRANCE] = MASK_ID` 行 | + +--- + +## 待解决(按优先级排序) + +### 1. 码本防塌缩(问题 B) + +**现象**:`quantizer1` 利用率 37.2%(16/32 存活),死去码字范数为 0。`quantizer3` 利用率 37.3%。`quantizer2` 反而 96.1% 健康。 + +**根因**:`VectorQuantizer` 只有 EMA 更新,无任何防塌缩机制。`VQ_GAMMA = 0.0`(entropy loss 未启用),`VQDecodeHead` 从未被训练脚本 import。 + +**与症状的对应**: + +| 阶段 | 条件通道状态 | 表现 | +|------|-------------|------| +| stage1 墙 | z1 用 37% + 距离场已删 | 差 | +| stage2 门/怪/入口 | z2 96.1% 健康 | 好 | +| stage3 资源 | z3 37.3% | 尚可 | + +**建议措施**: + +- ① 死码重启(dead-code restart):在 `VectorQuantizer.ema_update` 中检测 `ema_cluster_size` 长期低于阈值的码字,用随机真实 z_e 重新初始化 +- ② 启用 `VQ_GAMMA = 0.1`,打开 entropy loss +- ③ 或恢复两段式课程:先用 `VQDecodeHead` 做重建预训练,再冻结 VQ 训 MaskGIT + +**需重训**:是。已死的码字(范数为 0,离任何 z_e 都是 ~7.7)在当前 checkpoint 上不可恢复。 + +**把握度**:中 — 码本利用率必涨,但对生成质量的改善幅度未知。 + +--- + +### 2. 码本监控(问题 C) + +**现象**:`dist_quantizer` 塌缩了 320 epoch,训练日志从未提及。当前日志合并打印三个码本的 PPL(`PPL: 31.1 / 64`),掩盖了 `quantizer1` 只有 37% 的真相。 + +**措施**: + +- ① 分开打印每个码本的 PPL/利用率/存活码字数 +- ② 加离线诊断脚本,直接从 `.pth` 读 `ema_cluster_size` 和 `codebook.weight`,不用跑前向 + +**需重训**:否。**现在就能做。** + +**把握度**:确定 — 纯粹是监控完善,无副作用。 + +--- + +### 3. 用上数据集中的额外字段(问题 G) + +**现象**:`ginka-dataset.json` 每条记录有 `roomCount`、`highDegBranchCount`、`val`(16 维浮点,6098 种取值),但 `dataset.py` 只读了 `map` 和 `outerWall`。这些是数据侧已算好的拓扑信息,恰好是"墙该怎么排"的直接监督信号。 + +| 字段 | 内容 | 取值种类 | +|------|------|----------| +| `val` | 16 维浮点向量 | 6098 种(6731 张中) | +| `roomCount` | 房间数 | 17 种 | +| `highDegBranchCount` | 高度分支数 | 35 种 | + +**措施**:将 `roomCount` / `highDegBranchCount` / `val` 加进 `struct_inject`(当前只有对称性 3 bit + outerWall 1 bit),或作为额外条件 token。 + +**需重训**:是。改条件输入。 + +**把握度**:中 — 比已删除的距离场有信息量得多(6098 种取值 vs 距离场实测只用到 3 种),是所有方案里信息量最明确的一条。 + +--- + +### 4. 条件注入方式改为 cross-attention(问题 F) + +**现象**:所有 z 条件 token 被 flatten 成一个 d_model 维全局向量,通过 AdaLN 做全局 scale/shift 注入。z 中的 86 bit 只能以"全局风格"形式影响输出,没有任何机制传达空间位置信息。 + +**当前代码**(`maskGIT/model.py`): +``` +cond_seq = torch.cat([z_proj, e_struct, e_remain], dim=1) +c = self.cond_proj(cond_seq.reshape(B, -1)) # [B, d_model] +``` + +**措施**:改为 cross-attention — z 作为 memory,地图 token 作为 query。 + +**需重训**:是。改模型结构,权重完全不兼容。 + +**把握度**:低 — 此为推测,无实验证据证明它是主要瓶颈。工程量最大,建议放到最后。 + +--- + +### 5. 数据集去重(问题 H-2) + +**现象**:6731 张训练数据中唯一的有 6379 张,重复 352 张(5.23%)。 + +**措施**:去重。不算严重,但免费。 + +**需重训**:否。改数据,可从头训。 + +--- + +### 6. 连通性损失(问题 E) + +**原始建议**:不做。推理侧拒绝采样 + repair_connectivity 已做到可用率 100%。可微的连通性损失很难做,投入产出比不高。 + +**状态**:搁置。 + +--- + +## 建议实施顺序 + +| 顺序 | 事项 | 需重训 | 把握度 | 说明 | +|------|------|--------|--------|------| +| **1** | 码本监控(问题 C) | 否 | 确定 | 现在就能加,一行代码 | +| **2** | 码本防塌缩(问题 B ①+②) | 是 | 中 | 利用率必涨 | +| **3** | 用上额外字段(问题 G) | 是 | 中 | 信息量最明确 | +| **4** | 数据集去重(问题 H-2) | 否 | 确定 | 顺便做 | +| **5** | cross-attention(问题 F) | 是 | 未验证 | 工程量最大,最后做 | +| — | 连通性损失(问题 E) | — | — | 搁置 | + +**建议 2+3 一起改、跑一次训练;1+4 不用重训,随时能做。** diff --git a/ginka/dataset.py b/ginka/dataset.py index 6f39afb..205f9d8 100644 --- a/ginka/dataset.py +++ b/ginka/dataset.py @@ -3,7 +3,6 @@ import random import torch import numpy as np from torch.utils.data import Dataset -from shared.distance import compute_distance_field def rect_mask(ratio: float, map_size: int = 169) -> np.ndarray: # 连续矩形分块掩码,反复放置随机矩形直到掩码格数达标 @@ -144,8 +143,6 @@ class GinkaSeperatedDataset(Dataset): return enc1, enc2, enc3 def pack_sample(self, item: dict, map_np: np.ndarray, out: tuple) -> dict: - # out[2] = encoder_stage1,含完整墙壁,据此计算距离场 - dist_field = compute_distance_field(out[2]) return { "input_stage1": torch.LongTensor(out[0]), "target_stage1": torch.LongTensor(out[1]), @@ -158,7 +155,6 @@ class GinkaSeperatedDataset(Dataset): "encoder_stage3": torch.LongTensor(out[8]), "struct_inject": self.build_struct_inject(map_np, item['outerWall']), "target_density": self.build_target_density(item['map']), - "distance_field": torch.LongTensor(dist_field) } def random_sample_map(self, idx: int | None = None) -> dict: @@ -176,7 +172,6 @@ class GinkaSeperatedDataset(Dataset): "struct_inject": self.build_struct_inject(map_np, item['outerWall']), "target_density": self.build_target_density(item['map']), "raw_map": torch.LongTensor(map_np), - "distance_field": torch.LongTensor(compute_distance_field(enc1)) } sample['sample_idx'] = idx sample['map_name'] = self.map_names[idx] @@ -272,9 +267,8 @@ class GinkaSeperatedDataset(Dataset): return inp1, target1, enc1, inp2, target2, enc2, inp3, target3, enc3 def apply_subset3(self, raw: np.ndarray): - # 子集 3:在 2 的基础上掩码入口 + # 子集 3:与子集 2 相同(entry 已在 stage2 全掩码中覆盖) out = self.apply_subset2(raw) - out[0][out[0] == self.ENTRANCE] = self.MASK_ID return out def __getitem__(self, idx): diff --git a/ginka/diagnose_codebook.py b/ginka/diagnose_codebook.py new file mode 100644 index 0000000..aa6a8a7 --- /dev/null +++ b/ginka/diagnose_codebook.py @@ -0,0 +1,55 @@ +import sys +import torch + +if len(sys.argv) < 2: + print("用法: python diagnose_codebook.py result/seperated/sep-XXX.pth") + sys.exit(1) + +ckpt_path = sys.argv[1] +ckpt = torch.load(ckpt_path, map_location="cpu") + +names = ["quantizer1", "quantizer2", "quantizer3"] +labels = ["q1(stage1墙)", "q2(stage2门/怪/入口)", "q3(stage3资源)"] + +print(f"Checkpoint: {ckpt_path}") +print(f"Epoch: {ckpt.get('epoch', '?')}") +print() +print(f"{'名称':<20s} {'存活/总量':>10s} {'利用率':>8s} {'perplexity':>12s} {'最近大小中位':>14s}") +print("-" * 68) + +for name, label in zip(names, labels): + sd = ckpt.get(name) + if sd is None: + print(f"{label:<20s} {'未找到':>10s}") + continue + cs = sd["ema_cluster_size"] + K = cs.numel() + alive = int((cs > 1.0).sum()) + p = cs / cs.sum() + ppl = float(torch.exp(-(p * torch.log(p.clamp_min(1e-10))).sum())) + usage = alive / K * 100 + + # 最近(cluster_size 中位数,反映码字被使用频率) + median_size = float(cs.median()) + + print(f"{label:<20s} {alive:>3d}/{K:<3d} {usage:>5.1f}% {ppl:>8.2f}/{K:<8d} {median_size:>10.4f}") + +# 检查是否有码字范数为零的死码 +print() +print("--- 死码检查(范数=0 的码字) ---") +for name, label in zip(names, labels): + sd = ckpt.get(name) + if sd is None: + continue + w = sd["codebook.weight"] + if w is None: + w = sd.get("weight") # 可能的备用 key + if w is None: + print(f"{label}: 无法读取 codebook.weight") + continue + norms = w.norm(dim=1) + dead = int((norms < 1e-6).sum().item()) + if dead > 0: + print(f"{label}: {dead} / {w.size(0)} 死码(范数=0)") + else: + print(f"{label}: 全部码字有非零范数") diff --git a/ginka/maskGIT/model.py b/ginka/maskGIT/model.py index 991c67e..b50444a 100644 --- a/ginka/maskGIT/model.py +++ b/ginka/maskGIT/model.py @@ -7,13 +7,12 @@ from .maskGIT import Transformer # 结构标签词表大小 SYM_VOCAB = 8 # symmetryH/V/C 三位组合 0-7 OUTER_VOCAB = 2 # outerWall 0-1 -L_DIST = 4 # 距离场码字序列长度 class GinkaMaskGIT(nn.Module): def __init__( self, num_classes: int = 16, d_model: int = 192, dim_ff: int = 512, nhead: int = 8, num_layers: int = 4, map_h: int = 13, map_w: int = 13, - d_z: int = 64, z_seq_len: int = 6, z_dist_len: int = L_DIST + d_z: int = 64, z_seq_len: int = 6 ): super().__init__() self.map_h = map_h @@ -34,11 +33,8 @@ class GinkaMaskGIT(nn.Module): # z 投影:逐 token 线性变换,保持序列结构 self.z_proj = nn.Linear(d_z, d_z) - # 距离场 z 投影 - self.z_dist_proj = nn.Linear(d_z, d_z) - - # 条件融合投影:z_seq_len 个 z token + z_dist_len 个距离场 token + 2 个结构 token + 5 个剩余密度 token - self.cond_proj = nn.Linear((z_seq_len + z_dist_len + 2 + 5) * d_z, d_model) + # 条件融合投影:z_seq_len 个 z token + 2 个结构 token + 5 个剩余密度 token + self.cond_proj = nn.Linear((z_seq_len + 2 + 5) * d_z, d_model) # 纯 encoder Transformer,条件向量 c 通过 AdaLN 注入每一层 self.transformer = Transformer( @@ -51,13 +47,11 @@ class GinkaMaskGIT(nn.Module): self, map: torch.Tensor, z: torch.Tensor, - z_dist: torch.Tensor, struct: torch.Tensor, remain: torch.Tensor ) -> torch.Tensor: # map: [B, H * W] # z: [B, z_seq_len, d_z] - # z_dist: [B, z_dist_len, d_z] # struct: [B, 2] — [cond_sym(0-7), cond_outer(0-1)] # remain: [B, 5] float — [wall, door, monster, entrance, resource] 剩余密度 @@ -73,11 +67,8 @@ class GinkaMaskGIT(nn.Module): # z:逐 token 投影,保留序列结构 [B, z_seq_len, d_z] z_proj = self.z_proj(z) - # 距离场 z 投影 [B, z_dist_len, d_z] - zd_proj = self.z_dist_proj(z_dist) - # 拼接所有条件 token → 展平后投影到 d_model - cond_seq = torch.cat([z_proj, zd_proj, e_struct, e_remain], dim=1) + cond_seq = torch.cat([z_proj, e_struct, e_remain], dim=1) c = self.cond_proj(cond_seq.reshape(cond_seq.size(0), -1)) # [B, d_model] # tile embedding + 位置编码 @@ -96,7 +87,7 @@ if __name__ == "__main__": device = torch.device("cuda") if torch.cuda.is_available() else torch.device("cpu") map_input = torch.randint(0, 7, (4, 13 * 13)).to(device) # [4, 169] - z_input = torch.randn(4, 6, 64).to(device) # [4, L*3, 64] + z_input = torch.randn(4, 6, 64).to(device) # [4, 6, 64] struct_input = torch.tensor([ [3, 1], [0, 0], @@ -122,12 +113,10 @@ if __name__ == "__main__": z_seq_len=6 ).to(device) - z_dist_input = torch.randn(4, L_DIST, 64).to(device) # [4, L_DIST, 64] - print_memory(device, "初始化后") start = time.perf_counter() - logits = model(map_input, z_input, z_dist_input, struct_input, remain_input) + logits = model(map_input, z_input, struct_input, remain_input) end = time.perf_counter() print_memory(device, "前向传播后") diff --git a/ginka/model.py b/ginka/model.py new file mode 100644 index 0000000..a50175f --- /dev/null +++ b/ginka/model.py @@ -0,0 +1,193 @@ +import torch +import torch.nn as nn +import torch.optim as optim + +from .vqvae.quantize import VectorQuantizer +from .vqvae.model import GinkaVQVAE +from .maskGIT.model import GinkaMaskGIT + +# 共用 VQ-VAE 超参(共享的编码维度) +VQ_D_Z = 64 # 码字维度 +VQ_GAMMA = 0.1 # entropy loss 权重,鼓励码本使用均匀 + +# 三通道 VQ 各自独立超参(L、K、层数、维度等均独立配置) +# Stage1 墙壁骨架 — 结构最复杂,模型容量最大 +VQ_L1 = 24 +VQ_K1 = 32 +VQ_D_MODEL1 = 384 +VQ_NHEAD1 = 8 +VQ_LAYERS1 = 6 +VQ_DIM_FF1 = 1536 + +# Stage2 功能元素 — 中等复杂度 +VQ_L2 = 12 +VQ_K2 = 16 +VQ_D_MODEL2 = 256 +VQ_NHEAD2 = 4 +VQ_LAYERS2 = 6 +VQ_DIM_FF2 = 1024 + +# Stage3 资源分布 — 最简单,模型容量最小 +VQ_L3 = 8 +VQ_K3 = 16 +VQ_D_MODEL3 = 192 +VQ_NHEAD3 = 4 +VQ_LAYERS3 = 4 +VQ_DIM_FF3 = 768 + +# 第一阶段 MaskGIT 超参 +STAGE1_MG_DMODEL = 512 +STAGE1_MG_NHEAD = 4 +STAGE1_MG_NUM_LAYERS = 8 +STAGE1_MG_DIM_FF = 2048 + +# 第二阶段 MaskGIT 超参 +STAGE2_MG_DMODEL = 256 +STAGE2_MG_NHEAD = 4 +STAGE2_MG_NUM_LAYERS = 6 +STAGE2_MG_DIM_FF = 1024 + +# 第三阶段 MaskGIT 超参 +STAGE3_MG_DMODEL = 256 +STAGE3_MG_NHEAD = 4 +STAGE3_MG_NUM_LAYERS = 6 +STAGE3_MG_DIM_FF = 1024 + +# 各阶段 VQ commit loss 权重(当前未单独使用,统一由 VQ_BETA 控制) +STAGE1_VQ_WEIGHT = 0.5 +STAGE2_VQ_WEIGHT = 0.5 +STAGE3_VQ_WEIGHT = 0.5 + +# 全局参数 +NUM_CLASSES = 8 # 图块类型数 +MASK_TOKEN = 7 # 掩码图块 +TOTAL_K = VQ_K1 + VQ_K2 + VQ_K3 # 预计算,供日志输出使用 +MAP_W = 13 # 地图宽度 +MAP_H = 13 # 地图高度 + +LR = 1e-4 # AdamW 初始学习率 +MIN_LR = 1e-6 # 余弦退火最低学习率 +WEIGHT_DECAY = 1e-4 # L2 正则化系数 +EPOCHS = 400 # 总训练轮数 + + +class SeperatedModels: + # 三阶段级联模型集合,封装所有子模块、优化器和调度器 + vq1: GinkaVQVAE + vq2: GinkaVQVAE + vq3: GinkaVQVAE + mg1: GinkaMaskGIT + mg2: GinkaMaskGIT + mg3: GinkaMaskGIT + quantizers: tuple[VectorQuantizer, VectorQuantizer, VectorQuantizer] + quantizer1: VectorQuantizer + quantizer2: VectorQuantizer + quantizer3: VectorQuantizer + optimizer: optim.AdamW + scheduler: optim.lr_scheduler.CosineAnnealingLR + latent_mask_embedding: nn.Parameter + + def __init__(self, device: torch.device): + # 三组 VQ-VAE 编码器:各自独立编码一个阶段的地图上下文 + self.vq1 = GinkaVQVAE( + num_classes=NUM_CLASSES, L=VQ_L1, K=VQ_K1, d_model=VQ_D_MODEL1, nhead=VQ_NHEAD1, + num_layers=VQ_LAYERS1, dim_ff=VQ_DIM_FF1, map_h=MAP_H, map_w=MAP_W + ).to(device) + self.vq2 = GinkaVQVAE( + num_classes=NUM_CLASSES, L=VQ_L2, K=VQ_K2, d_model=VQ_D_MODEL2, nhead=VQ_NHEAD2, + num_layers=VQ_LAYERS2, dim_ff=VQ_DIM_FF2, map_h=MAP_H, map_w=MAP_W + ).to(device) + self.vq3 = GinkaVQVAE( + num_classes=NUM_CLASSES, L=VQ_L3, K=VQ_K3, d_model=VQ_D_MODEL3, nhead=VQ_NHEAD3, + num_layers=VQ_LAYERS3, dim_ff=VQ_DIM_FF3, map_h=MAP_H, map_w=MAP_W + ).to(device) + + # 三个独立 MaskGIT 解码器,分别接收各自阶段的 z_q 作为条件 + self.mg1 = GinkaMaskGIT( + num_classes=NUM_CLASSES, d_model=STAGE1_MG_DMODEL, d_z=VQ_D_Z, dim_ff=STAGE1_MG_DIM_FF, + nhead=STAGE1_MG_NHEAD, num_layers=STAGE1_MG_NUM_LAYERS, map_h=MAP_H, map_w=MAP_W, + z_seq_len=VQ_L1 + ).to(device) + self.mg2 = GinkaMaskGIT( + num_classes=NUM_CLASSES, d_model=STAGE2_MG_DMODEL, d_z=VQ_D_Z, dim_ff=STAGE2_MG_DIM_FF, + nhead=STAGE2_MG_NHEAD, num_layers=STAGE2_MG_NUM_LAYERS, map_h=MAP_H, map_w=MAP_W, + z_seq_len=VQ_L2 + ).to(device) + self.mg3 = GinkaMaskGIT( + num_classes=NUM_CLASSES, d_model=STAGE3_MG_DMODEL, d_z=VQ_D_Z, dim_ff=STAGE3_MG_DIM_FF, + nhead=STAGE3_MG_NHEAD, num_layers=STAGE3_MG_NUM_LAYERS, map_h=MAP_H, map_w=MAP_W, + z_seq_len=VQ_L3 + ).to(device) + + # 三个独立 VectorQuantizer:各阶段使用自己的码本大小 + self.quantizer1 = VectorQuantizer(K=VQ_K1, d_z=VQ_D_Z).to(device) + self.quantizer2 = VectorQuantizer(K=VQ_K2, d_z=VQ_D_Z).to(device) + self.quantizer3 = VectorQuantizer(K=VQ_K3, d_z=VQ_D_Z).to(device) + self.quantizers = (self.quantizer1, self.quantizer2, self.quantizer3) + + # latent dropout 用可学习 mask token,各阶段共享 + self.latent_mask_embedding = nn.Parameter( + torch.randn(1, 1, VQ_D_Z, device=device) * 0.02 + ) + + # 所有模块参数合并到同一优化器,余弦退火衰减至 MIN_LR + all_params = ( + list(self.vq1.parameters()) + list(self.vq2.parameters()) + list(self.vq3.parameters()) + + list(self.mg1.parameters()) + list(self.mg2.parameters()) + list(self.mg3.parameters()) + + list(self.quantizer1.parameters()) + list(self.quantizer2.parameters()) + list(self.quantizer3.parameters()) + + [self.latent_mask_embedding] + ) + self.optimizer = optim.AdamW(all_params, lr=LR, weight_decay=WEIGHT_DECAY) + self.scheduler = optim.lr_scheduler.CosineAnnealingLR( + self.optimizer, T_max=EPOCHS, eta_min=MIN_LR + ) + + def __iter__(self): + # 向后兼容:支持元组解包 + return iter(( + self.vq1, self.vq2, self.vq3, + self.mg1, self.mg2, self.mg3, + self.quantizers, self.optimizer, self.scheduler, + self.latent_mask_embedding + )) + + def __getitem__(self, idx): + return list(self)[idx] + + def load(self, ckpt_path: str, load_optim: bool = True, map_location: str = "cpu") -> int: + # 从检查点加载模型权重和训练状态,返回恢复的 epoch 编号 + ckpt = torch.load(ckpt_path, map_location=map_location) + self.vq1.load_state_dict(ckpt["vq1"]) + self.vq2.load_state_dict(ckpt["vq2"]) + self.vq3.load_state_dict(ckpt["vq3"]) + self.mg1.load_state_dict(ckpt["mg1"]) + self.mg2.load_state_dict(ckpt["mg2"]) + self.mg3.load_state_dict(ckpt["mg3"]) + self.quantizer1.load_state_dict(ckpt["quantizer1"]) + self.quantizer2.load_state_dict(ckpt["quantizer2"]) + self.quantizer3.load_state_dict(ckpt["quantizer3"]) + if "latent_mask_embedding" in ckpt: + self.latent_mask_embedding.data.copy_(ckpt["latent_mask_embedding"]) + if load_optim and "optimizer" in ckpt: + self.optimizer.load_state_dict(ckpt["optimizer"]) + if load_optim and "scheduler" in ckpt: + self.scheduler.load_state_dict(ckpt["scheduler"]) + return ckpt.get("epoch", 0) + + def save(self, path: str, epoch: int): + # 保存完整检查点(模型权重 + 优化器/调度器状态) + torch.save({ + "epoch": epoch, + "vq1": self.vq1.state_dict(), + "vq2": self.vq2.state_dict(), + "vq3": self.vq3.state_dict(), + "mg1": self.mg1.state_dict(), + "mg2": self.mg2.state_dict(), + "mg3": self.mg3.state_dict(), + "quantizer1": self.quantizer1.state_dict(), + "quantizer2": self.quantizer2.state_dict(), + "quantizer3": self.quantizer3.state_dict(), + "latent_mask_embedding": self.latent_mask_embedding.data, + "optimizer": self.optimizer.state_dict(), + "scheduler": self.scheduler.state_dict(), + }, path) diff --git a/ginka/sample.py b/ginka/sample.py new file mode 100644 index 0000000..233df7a --- /dev/null +++ b/ginka/sample.py @@ -0,0 +1,144 @@ +import math + +import numpy as np +import torch +import torch.nn.functional as F + +from .model import MASK_TOKEN, MAP_H, MAP_W, SeperatedModels +from .utils import compute_remaining, MAP_SIZE, compute_adjacency_mask + +# MaskGIT 采样函数:通过迭代去掩码从离散隐变量 z 生成地图 + + +def wall_growth_sample( + model: torch.nn.Module, + inp: torch.Tensor, + z: torch.Tensor, + struct: torch.Tensor, + target_density: torch.Tensor, + max_steps: int = 24 +) -> np.ndarray: + # 墙壁生长算法:从 inp 中已有的墙壁出发,逐步向外生长 + # 每步 MASK 位置决策(墙/非墙)后,新邻接面成为下一轮 MASK,逐步外扩 + state = inp.clone() # [B, MAP_SIZE] + # 初始 MASK:邻接已有墙壁的空地 + init_adj = compute_adjacency_mask(state) + state[init_adj & (state == 0)] = MASK_TOKEN + + for step in range(max_steps): + mask_pos = (state == MASK_TOKEN) # [B, MAP_SIZE] + if not mask_pos.any(): + break + + remain = compute_remaining(state, target_density, 1) + logits = model(state, z, struct, remain) + probs = F.softmax(logits, dim=-1) + + wall_prob = probs[:, :, 1] # [B, MAP_SIZE] + hits = (wall_prob > 0.5) & mask_pos + state[hits] = 1 + state[mask_pos & ~hits] = 0 + + adj = compute_adjacency_mask(state) + new_mask = adj & (state == 0) + if not new_mask.any(): + break + state[new_mask] = MASK_TOKEN + + state[state == MASK_TOKEN] = 0 + return state.cpu().numpy().reshape(state.size(0), MAP_H, MAP_W) + + +def maskgit_sample( + model: torch.nn.Module, inp: torch.Tensor, z: torch.Tensor, + struct: torch.Tensor, target_density: torch.Tensor, + stage: int, steps: int, + target_tiles: list[int] +) -> np.ndarray: + # 每步只从预测中选出置信度最高的若干 target_tile 候选揭开, + # 其余已有结构(墙/门等非空地非掩码)原样保留, + # 空地与掩码保持为 MASK,等待后续步骤继续填充。 + current = inp.clone() + target_tensor = torch.tensor(target_tiles, dtype=torch.long, device=inp.device) + + # 迭代去掩码:每步根据置信度分数重新决定掩码位置 + for step in range(steps): + remain = compute_remaining(current, target_density, stage) + logits = model(current, z, struct, remain) + probs = F.softmax(logits, dim=-1) + + dist = torch.distributions.Categorical(probs) + sampled = dist.sample() + + confidences = torch.gather(probs, -1, sampled.unsqueeze(-1)).squeeze(-1) + + # 余弦退火调度:随步数推进,保留掩码的位置数量递减至 0 + ratio = math.cos(((step + 1) / steps) * math.pi / 2) + num_to_mask = math.floor(ratio * MAP_SIZE) + + # 结构位:current 中非空地、非掩码的位置(来自上一阶段,始终保留) + struct_mask = (current[:] != MASK_TOKEN) & (current[:] != 0) + # 候选位:sampled 为目标图块且不覆盖结构位 + candidate_mask = torch.isin(sampled[:], target_tensor) & ~struct_mask + cand_count = candidate_mask.sum() + reveal_count = max(0, int(cand_count.item()) - num_to_mask) + next_state = current[:].clone() + if reveal_count > 0 and cand_count > 0: + cand_indices = candidate_mask.nonzero(as_tuple=False) + cand_conf = confidences[:][cand_indices[:, 0], cand_indices[:, 1]] + top_k = min(reveal_count, cand_conf.size(0)) + _, top_idx = torch.topk(cand_conf, k=top_k, largest=True) + reveal_rows = cand_indices[top_idx, 0] + reveal_cols = cand_indices[top_idx, 1] + next_state[reveal_rows, reveal_cols] = sampled[reveal_rows, reveal_cols] + # 结构位原样保留,其余未揭开的置为 MASK + non_struct_non_revealed = (next_state == current[:]) & ~struct_mask + next_state[non_struct_non_revealed & (next_state != MASK_TOKEN)] = MASK_TOKEN + # 空地也重新标为 MASK(允许下一步继续填充) + next_state[next_state == 0] = MASK_TOKEN + current = next_state + + if (current[:] == MASK_TOKEN).all(): + break + + # 兜底:未被填充的掩码位视为空地(不属于本阶段负责的图块) + still_masked = (current[:] == MASK_TOKEN) + current[still_masked] = 0 + + return current.cpu().numpy().reshape(current.size(0), MAP_H, MAP_W) + +def full_generate( + inp: torch.Tensor, + z1: torch.Tensor, z2: torch.Tensor, z3: torch.Tensor, + struct: torch.Tensor, + target_density: torch.Tensor, + models: SeperatedModels, + steps: int = 18 +) -> tuple[np.ndarray, np.ndarray, np.ndarray]: + # 三阶段级联生成:Stage1 墙壁生长 → Stage2 门/怪/入口 → Stage3 资源 + # 返回 (stage1结果, stage1+2合并, 最终完整地图),形状均为 [B, H, W] + device = inp.device + + pred1_np = wall_growth_sample( + models.mg1, inp, z1, struct, target_density + ) # [B, H, W] + inp2 = torch.tensor(pred1_np.reshape(pred1_np.shape[0], -1), dtype=torch.long, device=device) + inp2[inp2 == 0] = MASK_TOKEN + + pred2_np = maskgit_sample( + models.mg2, inp2, z2, struct, target_density, 2, + steps, target_tiles=[2, 4, 5, 6] + ) # [B, H, W] + merged12 = pred1_np.copy() + merged12[pred2_np != 0] = pred2_np[pred2_np != 0] + inp3 = torch.tensor(merged12.reshape(merged12.shape[0], -1), dtype=torch.long, device=device) + inp3[inp3 == 0] = MASK_TOKEN + + pred3_np = maskgit_sample( + models.mg3, inp3, z3, struct, target_density, 3, + steps, target_tiles=[3] + ) # [B, H, W] + merged123 = merged12.copy() + merged123[pred3_np != 0] = pred3_np[pred3_np != 0] + + return pred1_np, merged12, merged123 diff --git a/ginka/train_seperated.py b/ginka/train_seperated.py index e5e8092..96e2b37 100644 --- a/ginka/train_seperated.py +++ b/ginka/train_seperated.py @@ -1,5 +1,4 @@ import argparse -import math import os import sys import random @@ -10,16 +9,22 @@ import numpy as np import torch import torch.nn as nn import torch.nn.functional as F -import torch.optim as optim from tqdm import tqdm from torch.utils.data import DataLoader -from .vqvae.quantize import VectorQuantizer -from .vqvae.model import GinkaVQVAE, DistFieldEncoder -from .maskGIT.model import GinkaMaskGIT +from .model import SeperatedModels +from .model import ( + MASK_TOKEN, MAP_W, MAP_H, EPOCHS, VQ_GAMMA +) +from .utils import ( + compute_remaining, MAP_SIZE, summarize_codebook_hits +) +from .sample import full_generate from .dataset import GinkaSeperatedDataset from shared.image import matrix_to_image_cv -from shared.distance import DIST_VOCAB, compute_distance_field_tensor + +# 图块 ID 定义: +# 0. 空地 1. 墙壁 2. 普通门 3. 资源 4. 怪物 5. 入口 6. 机关门 7. 掩码(MASK_TOKEN) # 三阶段级联地图生成训练脚本 # @@ -34,120 +39,6 @@ from shared.distance import DIST_VOCAB, compute_distance_field_tensor # stage2 → door / monster / entrance(功能性实体) # stage3 → resource(资源点) -# 图块 ID 定义: -# 0. 空地 1. 墙壁 2. 普通门 3. 资源 4. 怪物 5. 入口 6. 机关门 7. 掩码(MASK_TOKEN) - -# 共用 VQ-VAE 超参(共享的编码维度) -VQ_D_Z = 64 # 码字维度 -VQ_BETA = 1.0 # commit loss 权重(防止编码器输出漂离 codebook) -VQ_GAMMA = 0.0 # entropy loss 权重(当前未启用) - -# 三通道 VQ 各自独立超参(L、K、层数、维度等均独立配置) -# Stage1 墙壁骨架 — 结构最复杂,模型容量最大 -VQ_L1 = 24 -VQ_K1 = 32 -VQ_D_MODEL1 = 384 -VQ_NHEAD1 = 8 -VQ_LAYERS1 = 6 -VQ_DIM_FF1 = 1536 - -# Stage2 功能元素 — 中等复杂度 -VQ_L2 = 12 -VQ_K2 = 16 -VQ_D_MODEL2 = 256 -VQ_NHEAD2 = 4 -VQ_LAYERS2 = 6 -VQ_DIM_FF2 = 1024 - -# Stage3 资源分布 — 最简单,模型容量最小 -VQ_L3 = 8 -VQ_K3 = 16 -VQ_D_MODEL3 = 192 -VQ_NHEAD3 = 4 -VQ_LAYERS3 = 4 -VQ_DIM_FF3 = 768 -L_DIST = 8 # 距离场码字序列长度 -K_DIST = 16 # 距离场 codebook 大小 - -# 距离场编码器超参 -DIST_D_MODEL = 384 # 距离场编码器模型维度 -DIST_LAYERS = 6 # 距离场编码器 Transformer 层数 -DIST_DIM_FF = 1536 # 距离场编码器 FF 维度 -DIST_NHEAD = 8 # 距离场编码器注意力头数 -VQ_BETA_DIST = 0.5 # 距离场 commit loss 权重 - -# 第一阶段 MaskGIT 超参 -STAGE1_MG_DMODEL = 512 -STAGE1_MG_NHEAD = 4 -STAGE1_MG_NUM_LAYERS = 8 -STAGE1_MG_DIM_FF = 2048 - -# 第二阶段 MaskGIT 超参 -STAGE2_MG_DMODEL = 256 -STAGE2_MG_NHEAD = 4 -STAGE2_MG_NUM_LAYERS = 6 -STAGE2_MG_DIM_FF = 1024 - -# 第三阶段 MaskGIT 超参 -STAGE3_MG_DMODEL = 256 -STAGE3_MG_NHEAD = 4 -STAGE3_MG_NUM_LAYERS = 6 -STAGE3_MG_DIM_FF = 1024 - -# 三阶段 Cross Entropy 损失权重(可调节各阶段对总损失的贡献比例) -STAGE1_CE_WEIGHT = 1.0 -STAGE2_CE_WEIGHT = 1.0 -STAGE3_CE_WEIGHT = 1.0 - -# 各阶段 VQ commit loss 权重(当前未单独使用,统一由 VQ_BETA 控制) -STAGE1_VQ_WEIGHT = 0.5 -STAGE2_VQ_WEIGHT = 0.5 -STAGE3_VQ_WEIGHT = 0.5 - -# 全局参数 -NUM_CLASSES = 8 # 图块类型数 -MASK_TOKEN = 7 # 掩码图块 -TOTAL_K = VQ_K1 + VQ_K2 + VQ_K3 # 预计算,供日志输出使用 -MAP_W = 13 # 地图宽度 -MAP_H = 13 # 地图高度 -MAP_SIZE = MAP_W * MAP_H # 地图大小 -DENSITY_DIM = 5 # [wall, door, monster, entrance, resource] -GENERATE_STEP = 18 # MaskGIT 采样步数 -SUBSET_WEIGHTS = (0.5, 0.3, 0.2) # 每个子集的概率 - -WALL_DENSITY_IDX = 0 -DOOR_DENSITY_IDX = 1 -MONSTER_DENSITY_IDX = 2 -ENTRANCE_DENSITY_IDX = 3 -RESOURCE_DENSITY_IDX = 4 - -MG_Z_DROPOUT = 0.1 # z 隐变量 Dropout 概率 -MG_STRUCT_DROPOUT = 0.1 # 结构参量 Dropout 概率 - -# 邻接损失权重(三阶段) -LAMBDA_ADJ1 = 0.2 -LAMBDA_ADJ2 = 0.1 -LAMBDA_ADJ3 = 0.05 - -# Patch 损失权重(三阶段)及核参数 -LAMBDA_PATCH1 = 0.2 -LAMBDA_PATCH2 = 0.2 -LAMBDA_PATCH3 = 0.2 -PATCH_KERNEL_SIZE = 5 -PATCH_SIGMA = 1.2 - -# 损失参数 -VQ_BETA = 0.5 # 承诺损失权重 - -# 训练超参 -BATCH_SIZE = 64 # 每批样本数 -LR = 1e-4 # AdamW 初始学习率 -MIN_LR = 1e-6 # 余弦退火最低学习率 -WEIGHT_DECAY = 1e-4 # L2 正则化系数 -EPOCHS = 400 # 总训练轮数 -CHECKPOINT = 20 # 每隔多少 epoch 保存检查点并执行验证 -REFERENCE_SAMPLE_PROB = 0.2 # 训练时将参考掩码图无梯度自采样 1-3 步的概率 - device = torch.device( "cuda:0" if torch.cuda.is_available() else "mps" if torch.backends.mps.is_available() @@ -156,6 +47,18 @@ device = torch.device( disable_tqdm = not sys.stdout.isatty() +# 训练与推理超参 +VQ_BETA = 0.5 # 承诺损失权重 +STAGE1_CE_WEIGHT = 1.0 # Stage1 CE 损失权重 +STAGE2_CE_WEIGHT = 1.0 # Stage2 CE 损失权重 +STAGE3_CE_WEIGHT = 1.0 # Stage3 CE 损失权重 +MG_Z_DROPOUT = 0.1 # z 隐变量 Dropout 概率 +BATCH_SIZE = 64 # 每批样本数 +CHECKPOINT = 20 # 每隔多少 epoch 保存检查点并执行验证 +GENERATE_STEP = 18 # MaskGIT 采样步数 +SEED_SAMPLE_STEPS = 24 # 墙壁种子生长采样步数 +SUBSET_WEIGHTS = (0.5, 0.3, 0.2) # 每个子集的概率 + def _str2bool(v: str): if isinstance(v, bool): return v if v.lower() in ('true', '1', 'yes'): return True @@ -171,141 +74,13 @@ def parse_arguments(): parser.add_argument("--load_optim", type=_str2bool, default=True) return parser.parse_args() -def build_model(device: torch.device): - # 三组 VQ-VAE 编码器:各自独立编码一个阶段的地图上下文,全部超参独立 - vq1 = GinkaVQVAE( - num_classes=NUM_CLASSES, L=VQ_L1, K=VQ_K1, - d_model=VQ_D_MODEL1, nhead=VQ_NHEAD1, num_layers=VQ_LAYERS1, dim_ff=VQ_DIM_FF1, - map_h=MAP_H, map_w=MAP_W - ).to(device) - vq2 = GinkaVQVAE( - num_classes=NUM_CLASSES, L=VQ_L2, K=VQ_K2, - d_model=VQ_D_MODEL2, nhead=VQ_NHEAD2, num_layers=VQ_LAYERS2, dim_ff=VQ_DIM_FF2, - map_h=MAP_H, map_w=MAP_W - ).to(device) - vq3 = GinkaVQVAE( - num_classes=NUM_CLASSES, L=VQ_L3, K=VQ_K3, - d_model=VQ_D_MODEL3, nhead=VQ_NHEAD3, num_layers=VQ_LAYERS3, dim_ff=VQ_DIM_FF3, - map_h=MAP_H, map_w=MAP_W - ).to(device) - - # 三个独立 MaskGIT 解码器,分别接收各自阶段的 z_q 作为条件 - mg1 = GinkaMaskGIT( - num_classes=NUM_CLASSES, d_model=STAGE1_MG_DMODEL, d_z=VQ_D_Z, dim_ff=STAGE1_MG_DIM_FF, - nhead=STAGE1_MG_NHEAD, num_layers=STAGE1_MG_NUM_LAYERS, map_h=MAP_H, map_w=MAP_W, - z_seq_len=VQ_L1, z_dist_len=L_DIST - ).to(device) - mg2 = GinkaMaskGIT( - num_classes=NUM_CLASSES, d_model=STAGE2_MG_DMODEL, d_z=VQ_D_Z, dim_ff=STAGE2_MG_DIM_FF, - nhead=STAGE2_MG_NHEAD, num_layers=STAGE2_MG_NUM_LAYERS, map_h=MAP_H, map_w=MAP_W, - z_seq_len=VQ_L2, z_dist_len=L_DIST - ).to(device) - mg3 = GinkaMaskGIT( - num_classes=NUM_CLASSES, d_model=STAGE3_MG_DMODEL, d_z=VQ_D_Z, dim_ff=STAGE3_MG_DIM_FF, - nhead=STAGE3_MG_NHEAD, num_layers=STAGE3_MG_NUM_LAYERS, map_h=MAP_H, map_w=MAP_W, - z_seq_len=VQ_L3, z_dist_len=L_DIST - ).to(device) - - # 三个独立 VectorQuantizer:各阶段使用自己的码本大小 - quantizer1 = VectorQuantizer(K=VQ_K1, d_z=VQ_D_Z).to(device) - quantizer2 = VectorQuantizer(K=VQ_K2, d_z=VQ_D_Z).to(device) - quantizer3 = VectorQuantizer(K=VQ_K3, d_z=VQ_D_Z).to(device) - quantizers = (quantizer1, quantizer2, quantizer3) - - # 距离场编码器与量化器:将 L1 距离场编码为离散 latent z_dist - dist_encoder = DistFieldEncoder( - vocab=DIST_VOCAB, L=L_DIST, d_z=VQ_D_Z, d_model=DIST_D_MODEL, - nhead=DIST_NHEAD, num_layers=DIST_LAYERS, dim_ff=DIST_DIM_FF, - map_h=MAP_H, map_w=MAP_W - ).to(device) - dist_quantizer = VectorQuantizer(K=K_DIST, d_z=VQ_D_Z).to(device) - - # latent dropout 用可学习 mask token,各阶段共享 - latent_mask_embedding = nn.Parameter(torch.randn(1, 1, VQ_D_Z, device=device) * 0.02) - - # 所有模块参数合并到同一优化器 - all_params = ( - list(vq1.parameters()) + list(vq2.parameters()) + list(vq3.parameters()) + - list(mg1.parameters()) + list(mg2.parameters()) + list(mg3.parameters()) + - list(quantizer1.parameters()) + list(quantizer2.parameters()) + list(quantizer3.parameters()) + - list(dist_encoder.parameters()) + list(dist_quantizer.parameters()) + - [latent_mask_embedding] - ) - optimizer = optim.AdamW(all_params, lr=LR, weight_decay=1e-4) - # 余弦退火:从 LR 线性衰减至 MIN_LR,周期为全部训练轮数 - scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=EPOCHS, eta_min=MIN_LR) - - return vq1, vq2, vq3, mg1, mg2, mg3, quantizers, optimizer, scheduler, latent_mask_embedding, dist_encoder, dist_quantizer - -def cross_entropy_loss(logits, target): - # logits: [B, L, C],需转为 [B, C, L] 以匹配 cross_entropy 期望格式 - return F.cross_entropy(logits.permute(0, 2, 1), target) - -def adjacency_loss(logits, target): - # 邻接损失:约束相邻两格同时为空地的概率 - # logits: [B, S, C] — MaskGIT 解码器输出 - # target: [B, S] — 目标类别 ID,不含 MASK 标记 - B, S, C = logits.shape - H = 13 - W = 13 - probs = F.softmax(logits, dim=-1) # [B, S, C] - p_floor = probs[:, :, 0].view(B, H, W) # [B, H, W] — 地板概率 - t = target.view(B, H, W) - t_floor = (t == 0).float() # 地板标注为 1,其余为 0 - - # 水平边:左格 × 右格 - joint_h = p_floor[:, :, :-1] * p_floor[:, :, 1:] # [B, H, W-1] - target_h = t_floor[:, :, :-1] * t_floor[:, :, 1:] # [B, H, W-1] - - # 垂直边:上格 × 下格 - joint_v = p_floor[:, :-1, :] * p_floor[:, 1:, :] # [B, H-1, W] - target_v = t_floor[:, :-1, :] * t_floor[:, 1:, :] # [B, H-1, W] - - loss_h = F.binary_cross_entropy(joint_h, target_h, reduction='mean') - loss_v = F.binary_cross_entropy(joint_v, target_v, reduction='mean') - return (loss_h + loss_v) / 2.0 - -def gaussian_kernel(kernel_size, sigma, device): - # 生成归一化二维高斯卷积核 [1, 1, K, K] - k = kernel_size - center = (k - 1) / 2.0 - xs = torch.arange(k, dtype=torch.float32, device=device) - center - gx = torch.exp(-xs ** 2 / (2.0 * sigma ** 2)) - gy = torch.exp(-xs ** 2 / (2.0 * sigma ** 2)) - g2d = gx[:, None] * gy[None, :] # [K, K] - g2d = g2d / g2d.sum() # 归一化 - return g2d.view(1, 1, k, k) - -def patch_loss(logits, target, kernel_size=5, sigma=1.2): - # Patch 损失:高斯核加权的邻域 CE 平滑损失 - # logits: [B, S, C] - # target: [B, S] - B, S, C = logits.shape - H = 13 - W = 13 - - # 逐格 CE(不做 reduction) - ce = F.cross_entropy( - logits.reshape(-1, C), target.reshape(-1), reduction='none' - ).view(B, H, W) # [B, H, W] - - # 高斯核 - kernel = gaussian_kernel(kernel_size, sigma, logits.device) # [1, 1, K, K] - - # replicate 填充后用 unfold 提取邻域 - pad = kernel_size // 2 - ce_padded = F.pad( - ce.view(B, 1, H, W), (pad, pad, pad, pad), mode='replicate' - ) - # patches: [B, K*K, H*W] - patches = F.unfold(ce_padded, kernel_size=(kernel_size, kernel_size)) - patches = patches.view(B, kernel_size * kernel_size, H, W) # [B, K*K, H, W] - - # 加权求和 - k_flat = kernel.view(1, kernel_size * kernel_size, 1, 1) - smoothed = (patches * k_flat).sum(dim=1) # [B, H, W] - - return smoothed.mean() +def cross_entropy_loss(logits, target, mask): + # logits: [B, L, C],target: [B, L],mask: [B, L] bool(True = 参与 loss) + loss = F.cross_entropy(logits.permute(0, 2, 1), target, reduction='none') + masked = loss[mask] + if masked.numel() == 0: + return torch.tensor(0.0, device=logits.device, requires_grad=True) + return masked.mean() def apply_z_dropout( z_q: torch.Tensor, @@ -317,847 +92,159 @@ def apply_z_dropout( mask = torch.rand(z_q.shape[0], z_q.shape[1], 1, device=z_q.device) < drop_prob return z_q * (~mask).float() + mask_embedding * mask.float() -def summarize_codebook_hits(code_hits) -> tuple[float, float, int]: - # code_hits 为 tuple of 3 tensors(各量器不同 K) - combined = torch.cat([h.flatten() for h in code_hits], dim=0) - total_hits = combined.sum() - if total_hits.item() <= 0: - return 0.0, 0.0, 0 - - probs = combined / total_hits - perplexity = torch.exp( - -(probs * torch.log(probs.clamp_min(1e-10))).sum() - ).item() - active_codes = int((combined > 0).sum().item()) - usage_rate = active_codes / combined.numel() - return perplexity, usage_rate, active_codes - def quantize_stage_latents( - quantizers: tuple[VectorQuantizer, VectorQuantizer, VectorQuantizer], + models: SeperatedModels, z_e1: torch.Tensor, z_e2: torch.Tensor, z_e3: torch.Tensor -) -> tuple[tuple[torch.Tensor, torch.Tensor, torch.Tensor], torch.Tensor, torch.Tensor]: - quantizer1, quantizer2, quantizer3 = quantizers - z_q1, _, commit_loss1, _, code_hits1 = quantizer1(z_e1) - z_q2, _, commit_loss2, _, code_hits2 = quantizer2(z_e2) - z_q3, _, commit_loss3, _, code_hits3 = quantizer3(z_e3) +) -> tuple: + z_q1, _, commit_loss1, _, code_hits1, entropy1 = models.quantizer1(z_e1) + z_q2, _, commit_loss2, _, code_hits2, entropy2 = models.quantizer2(z_e2) + z_q3, _, commit_loss3, _, code_hits3, entropy3 = models.quantizer3(z_e3) commit_loss = (commit_loss1 + commit_loss2 + commit_loss3) / 3 + entropy_loss = entropy1 + entropy2 + entropy3 code_hits = (code_hits1, code_hits2, code_hits3) - return (z_q1, z_q2, z_q3), commit_loss, code_hits + return (z_q1, z_q2, z_q3), commit_loss, code_hits, entropy_loss -def build_reference_rollout_steps(prob: float) -> int: - if random.random() >= prob: - return 0 - - return random.randint(1, 3) - -def sample_reference_inputs( - model: torch.nn.Module, - reference: torch.Tensor, - z_q: torch.Tensor, - z_dist: torch.Tensor, - struct: torch.Tensor, - target_density: torch.Tensor, - stage: int, - rollout_steps: int -) -> torch.Tensor: - if rollout_steps <= 0: - return reference - - sampled_reference = reference.clone() - with torch.no_grad(): - current = sampled_reference.clone() - z_q_detached = z_q.detach() - z_dist_detached = z_dist.detach() - - for _ in range(rollout_steps): - masked_positions = current == MASK_TOKEN - masked_counts = masked_positions.sum(dim=1) - if int(masked_counts.sum().item()) <= 0: - break - - remain = compute_remaining(current, target_density, stage) - logits = model(current, z_q_detached, z_dist_detached, struct, remain) - probs = F.softmax(logits, dim=-1) - dist = torch.distributions.Categorical(probs) - predicted = dist.sample() - confidence = torch.gather( - probs, - -1, - predicted.unsqueeze(-1) - ).squeeze(-1) - - for local_idx in range(current.size(0)): - masked_count = int(masked_counts[local_idx].item()) - if masked_count <= 0: - continue - - masked_indices = masked_positions[local_idx].nonzero(as_tuple=True)[0] - reveal_count = max(1, math.ceil(masked_count * 0.1)) - reveal_count = min(reveal_count, masked_indices.numel()) - masked_confidence = confidence[local_idx, masked_indices] - _, reveal_order = torch.topk( - masked_confidence, - k=reveal_count, - largest=True - ) - reveal_indices = masked_indices[reveal_order] - current[local_idx, reveal_indices] = predicted[local_idx, reveal_indices] - - sampled_reference = current - - return sampled_reference - -def compute_remaining( - current: torch.Tensor, - target_density: torch.Tensor, - stage: int -) -> torch.Tensor: - remain = torch.zeros(current.size(0), DENSITY_DIM, device=current.device) - - visible_wall = (current == 1).sum(dim=1).float() / MAP_SIZE - visible_door = ((current == 2) | (current == 6)).sum(dim=1).float() / MAP_SIZE - visible_monster = (current == 4).sum(dim=1).float() / MAP_SIZE - visible_entrance = (current == 5).sum(dim=1).float() / MAP_SIZE - visible_resource = (current == 3).sum(dim=1).float() / MAP_SIZE - - if stage == 1: - remain[:, WALL_DENSITY_IDX] = ( - target_density[:, WALL_DENSITY_IDX] - visible_wall - ).clamp(min=0.0, max=1.0) - elif stage == 2: - remain[:, DOOR_DENSITY_IDX] = ( - target_density[:, DOOR_DENSITY_IDX] - visible_door - ).clamp(min=0.0, max=1.0) - remain[:, MONSTER_DENSITY_IDX] = ( - target_density[:, MONSTER_DENSITY_IDX] - visible_monster - ).clamp(min=0.0, max=1.0) - remain[:, ENTRANCE_DENSITY_IDX] = ( - target_density[:, ENTRANCE_DENSITY_IDX] - visible_entrance - ).clamp(min=0.0, max=1.0) - elif stage == 3: - remain[:, RESOURCE_DENSITY_IDX] = ( - target_density[:, RESOURCE_DENSITY_IDX] - visible_resource - ).clamp(min=0.0, max=1.0) - - return remain - -def rect_mask( - ratio: float, h_range: tuple[int, int] = (2, 7), - w_range: tuple[int, int] = (2, 7) -) -> np.ndarray: - # 纯矩形分块掩码,反复放置随机矩形直到掩码格数达标 - target = int(MAP_SIZE * ratio) - mask = np.zeros((MAP_H, MAP_W), dtype=bool) - while mask.sum() < target: - bh = np.random.randint(h_range[0], h_range[1]) - bw = np.random.randint(w_range[0], w_range[1]) - x = np.random.randint(0, MAP_H - bh + 1) - y = np.random.randint(0, MAP_W - bw + 1) - mask[x:x + bh, y:y + bw] = True - return mask - -def compute_adjacency_mask(flat_state: torch.Tensor) -> torch.Tensor: - # 返回 [MAP_SIZE] bool,True 表示该位置与任意墙壁 4-邻接 - # flat_state: [MAP_SIZE] 整数(可选 batch 维度的话取第 0 行) - if flat_state.dim() > 1: - flat_state = flat_state[0] - state_2d = flat_state.reshape(MAP_H, MAP_W) - wall = (state_2d == 1) - adj = torch.zeros_like(wall, dtype=torch.bool) - adj[:, 1:] |= wall[:, :-1] - adj[:, :-1] |= wall[:, 1:] - adj[1:, :] |= wall[:-1, :] - adj[:-1, :] |= wall[1:, :] - return adj.flatten() - -def wall_growth_sample( - model: torch.nn.Module, - z: torch.Tensor, - z_dist: torch.Tensor, - struct: torch.Tensor, - target_density: torch.Tensor, - max_steps: int = 24 -) -> np.ndarray: - # 生长算法:初始外圈墙壁固定,仅邻接一圈为 MASK,其余为 FLOOR(0) - # 每步 MASK 位置决策(墙/非墙)后,新邻接面成为下一轮 MASK,逐步外扩 - outer_idx = [] - for r in range(MAP_H): - for c in range(MAP_W): - if r == 0 or r == MAP_H - 1 or c == 0 or c == MAP_W - 1: - outer_idx.append(r * MAP_W + c) - outer_idx = torch.tensor(outer_idx, dtype=torch.long, device=z.device) - - state = torch.full((1, MAP_SIZE), 0, dtype=torch.long, device=z.device) - state[0, outer_idx] = 1 - # 初始 MASK:邻接外圈墙壁的一圈 - init_adj = compute_adjacency_mask(state[0]) - state[0, init_adj & (state[0] == 0)] = MASK_TOKEN - - for step in range(max_steps): - mask_pos = state[0] == MASK_TOKEN - if mask_pos.sum() == 0: - break - - remain = compute_remaining(state, target_density, 1) - logits = model(state, z, z_dist, struct, remain) - probs = F.softmax(logits, dim=-1) - - was_mask = mask_pos.clone() - mask_idx = mask_pos.nonzero(as_tuple=True)[0] - wall_prob = probs[0, mask_idx, 1] - hits = wall_prob > 0.5 - if hits.any(): - state[0, mask_idx[hits]] = 1 - # 本轮 MASK 决策完毕:未提交的 MASK → FLOOR(0) - state[0, mask_idx[~hits]] = 0 - - # 下一轮 MASK:邻接墙壁且本轮不是 MASK(即新触及的 FLOOR) - adj = compute_adjacency_mask(state[0]) - new_mask = adj & (state[0] == 0) & (~was_mask) - if new_mask.sum() == 0: - break - state[0, new_mask] = MASK_TOKEN - - state[0, state[0] == MASK_TOKEN] = 0 - return state[0].cpu().numpy().reshape(MAP_H, MAP_W) - -def maskgit_sample( - model: torch.nn.Module, inp: torch.Tensor, z: torch.Tensor, - z_dist: torch.Tensor, struct: torch.Tensor, target_density: torch.Tensor, - stage: int, steps: int, - target_tiles: list[int] | None = None, keep_fixed: bool = True -) -> np.ndarray: - # target_tiles: 本阶段负责生成的图块 ID 列表;None 表示接受所有类别(stage1) - # keep_fixed=True:锁定输入中已有的非掩码/非空地位,使上一阶段结构保持不变 - # keep_fixed=False:结构位保留,但每步结束后空地重新标为 MASK(探索模式) - # - # 有 target_tiles 时的核心逻辑: - # 每步只从预测中选出置信度最高的若干 target_tile 候选揭开, - # 其余已有结构(墙/门等非空地非掩码)原样保留, - # 空地与掩码保持为 MASK,等待后续步骤继续填充。 - current = inp.clone() - has_target = target_tiles is not None - if has_target: - target_tensor = torch.tensor(target_tiles, dtype=torch.long, device=inp.device) - - # 迭代去掩码:每步根据置信度分数重新决定掩码位置 - for step in range(steps): - remain = compute_remaining(current, target_density, stage) - logits = model(current, z, z_dist, struct, remain) - probs = F.softmax(logits, dim=-1) - - dist = torch.distributions.Categorical(probs) - sampled = dist.sample() - - confidences = torch.gather(probs, -1, sampled.unsqueeze(-1)).squeeze(-1) - - # 余弦退火调度:随步数推进,保留掩码的位置数量递减至 0 - ratio = math.cos(((step + 1) / steps) * math.pi / 2) - num_to_mask = math.floor(ratio * MAP_SIZE) - - if not has_target: - # stage1:无 target 约束,仅锁定 fixed 位(若 keep_fixed) - if keep_fixed: - fixed_mask = (current[0] != MASK_TOKEN) - sampled[0, fixed_mask] = current[0, fixed_mask] - confidences[0, fixed_mask] = 1.0 - if num_to_mask > 0: - _, mask_indices = torch.topk(confidences[0], k=num_to_mask, largest=False) - sampled[0].scatter_(0, mask_indices, MASK_TOKEN) - current = sampled - else: - # 有 target_tiles:基于当前 current 构建下一状态 - # 结构位:current 中非空地、非掩码的位置(来自上一阶段,始终保留) - struct_mask = (current[0] != MASK_TOKEN) & (current[0] != 0) - # 候选位:sampled 为目标图块且不覆盖结构位 - candidate_mask = torch.isin(sampled[0], target_tensor) & ~struct_mask - # 对候选位按置信度排序,选出置信度最高的若干位揭开 - cand_count = candidate_mask.sum().item() - reveal_count = max(0, int(cand_count) - num_to_mask) - next_state = current[0].clone() - if reveal_count > 0 and cand_count > 0: - cand_indices = candidate_mask.nonzero(as_tuple=True)[0] - cand_conf = confidences[0][cand_indices] - top_k = min(reveal_count, cand_conf.size(0)) - _, top_idx = torch.topk(cand_conf, k=top_k, largest=True) - reveal_indices = cand_indices[top_idx] - next_state[reveal_indices] = sampled[0][reveal_indices] - # 后处理:进度未超 75% 时,随机将新揭开位的 20%-40% 再次掩码, - # 抑制目标图块过密生成;后期不再压制,确保最终能全部揭开 - if step / steps <= 0.75: - suppress_ratio = random.uniform(0.2, 0.4) - suppress_k = max(1, int(reveal_indices.size(0) * suppress_ratio)) - suppress_perm = torch.randperm( - reveal_indices.size(0), device=inp.device - )[:suppress_k] - next_state[reveal_indices[suppress_perm]] = MASK_TOKEN - # 结构位原样保留,其余未揭开的置为 MASK - non_struct_non_revealed = (next_state == current[0]) & ~struct_mask - next_state[non_struct_non_revealed & (next_state != MASK_TOKEN)] = MASK_TOKEN - # free 模式下,空地也重新标为 MASK(允许下一步继续填充) - if not keep_fixed: - next_state[next_state == 0] = MASK_TOKEN - current = next_state.unsqueeze(0) - - if (current[0] == MASK_TOKEN).sum() == 0: - break - - # 兜底:若仍有残余掩码位,按模式填充 - still_masked = (current[0] == MASK_TOKEN) - if still_masked.any(): - if has_target: - # 目标模式下,未被填充的位置视为空地(不属于本阶段负责的图块) - current[0, still_masked] = 0 - else: - remain = compute_remaining(current, target_density, stage) - logits = model(current, z, z_dist, struct, remain) - current[0, still_masked] = torch.argmax(logits[0, still_masked], dim=-1) - - return current[0].cpu().numpy().reshape(MAP_H, MAP_W) - -def full_generate_specific_z( - input: torch.Tensor, - z_q: tuple[torch.Tensor, torch.Tensor, torch.Tensor], - z_dist: torch.Tensor, - struct: torch.Tensor, - target_density: torch.Tensor, - models: list[torch.nn.Module], - device: torch.device, - keep_fixed: tuple[bool, bool, bool] = (True, True, True) -) -> tuple: - vq1, vq2, vq3, mg1, mg2, mg3, quantizers, _, _, _, _, _ = models - - with torch.no_grad(): - z1, z2, z3 = z_q - - # 三阶段级联生成:Stage1 使用生长算法,Stage2/3 保持 MaskGIT - pred1_np = wall_growth_sample( - mg1, z1, z_dist, struct, target_density - ) - inp2 = torch.tensor(pred1_np.flatten(), dtype=torch.long, device=device).reshape(1, MAP_SIZE) - inp2[inp2 == 0] = MASK_TOKEN - - pred2_np = maskgit_sample( - mg2, inp2, z2, z_dist, struct, target_density, 2, - GENERATE_STEP, target_tiles=[2, 6, 4, 5], keep_fixed=keep_fixed[1] - ) - merged12 = pred1_np.copy() - merged12[pred2_np != 0] = pred2_np[pred2_np != 0] - inp3 = torch.tensor(merged12.flatten(), dtype=torch.long, device=device).reshape(1, MAP_SIZE) - inp3[inp3 == 0] = MASK_TOKEN - - pred3_np = maskgit_sample( - mg3, inp3, z3, z_dist, struct, target_density, 3, - GENERATE_STEP, target_tiles=[3], keep_fixed=keep_fixed[2] - ) - merged123 = merged12.copy() - merged123[pred3_np != 0] = pred3_np[pred3_np != 0] - - return pred1_np, merged12, merged123 - -def inpaint_generate( - raw_map: np.ndarray, - mask: np.ndarray, - z_q: tuple[torch.Tensor, torch.Tensor, torch.Tensor], - z_dist: torch.Tensor, - struct: torch.Tensor, - target_density: torch.Tensor, - models: list[torch.nn.Module], - device: torch.device -): - # 三阶段矩形掩码修补生成 - # - z 来自完整地图 VQ 编码,作为全局先验 - # - 未掩码区域的原始图块逐阶段保留,避免被误当空地重填 - vq1, vq2, vq3, mg1, mg2, mg3, quantizers, _, _, _, _, _ = models - z1, z2, z3 = z_q - - # Stage1: 补全墙壁 — 生长方式,掩码区域内仅邻接圈为 MASK,其余为 FLOOR(0) - inp1 = raw_map.copy() - non_wf = (inp1 != 0) & (inp1 != 1) - inp1[non_wf] = 0 - # 掩码区域先全设为 FLOOR(0),再找出邻接圈设为 MASK - inp1[mask] = 0 - inp1_t = torch.tensor(inp1.flatten(), dtype=torch.long, device=device).reshape(1, MAP_SIZE) - # 初始 MASK:掩码区域内与固定墙壁邻接的一圈 - mask_t = torch.tensor(mask.flatten(), dtype=torch.bool, device=device) - init_state = inp1_t.clone() - init_adj = compute_adjacency_mask(init_state[0]) - init_mask = init_adj & (init_state[0] == 0) & mask_t - init_state[0, init_mask] = MASK_TOKEN - - with torch.no_grad(): - state = init_state - for step in range(GENERATE_STEP): - mask_pos = state[0] == MASK_TOKEN - if mask_pos.sum() == 0: - break - - remain = compute_remaining(state, target_density, 1) - logits = mg1(state, z1, z_dist, struct, remain) - probs = F.softmax(logits, dim=-1) - - was_mask = mask_pos.clone() - mask_idx = mask_pos.nonzero(as_tuple=True)[0] - wall_prob = probs[0, mask_idx, 1] - hits = wall_prob > 0.5 - if hits.any(): - state[0, mask_idx[hits]] = 1 - state[0, mask_idx[~hits]] = 0 - - adj = compute_adjacency_mask(state[0]) - new_mask = adj & (state[0] == 0) & (~was_mask) & mask_t - if new_mask.sum() == 0: - break - state[0, new_mask] = MASK_TOKEN - - state[0, state[0] == MASK_TOKEN] = 0 - pred1_np = state[0].cpu().numpy().reshape(MAP_H, MAP_W) - - # Stage2: 补全门/怪/入口,保留未掩码区的原始非墙壁结构 - merged_s1 = pred1_np.copy() - preserve_s2 = (~mask) & (raw_map != 0) & (raw_map != 1) - merged_s1[preserve_s2] = raw_map[preserve_s2] - - inp2 = merged_s1.copy() - inp2[inp2 == 0] = MASK_TOKEN - inp2_t = torch.tensor(inp2.flatten(), dtype=torch.long, device=device).reshape(1, MAP_SIZE) - - with torch.no_grad(): - pred2_np = maskgit_sample( - mg2, inp2_t, z2, z_dist, struct, target_density, 2, - GENERATE_STEP, target_tiles=[2, 6, 4, 5], keep_fixed=True - ) - - # Stage3: 补全资源,保留未掩码区的原始资源 - merged_s2 = merged_s1.copy() - merged_s2[pred2_np != 0] = pred2_np[pred2_np != 0] - res_preserve = (raw_map == 3) & (~mask) - merged_s2[res_preserve] = 3 - - inp3 = merged_s2.copy() - inp3[inp3 == 0] = MASK_TOKEN - inp3_t = torch.tensor(inp3.flatten(), dtype=torch.long, device=device).reshape(1, MAP_SIZE) - - with torch.no_grad(): - pred3_np = maskgit_sample( - mg3, inp3_t, z3, z_dist, struct, target_density, 3, - GENERATE_STEP, target_tiles=[3], keep_fixed=True - ) - - merged_s3 = merged_s2.copy() - merged_s3[pred3_np != 0] = pred3_np[pred3_np != 0] - - return pred1_np, merged_s1, merged_s2, merged_s3 - -def annotate(img: np.ndarray, text: str, y: int = 14) -> np.ndarray: - # 在图片左上角叠加文字标注(黑色描边 + 白色填充,确保任意背景下可读) - img = img.copy() - cv2.putText(img, text, (2, y), cv2.FONT_HERSHEY_SIMPLEX, 0.4, (0, 0, 0), 2) - cv2.putText(img, text, (2, y), cv2.FONT_HERSHEY_SIMPLEX, 0.4, (255, 255, 255), 1) - return img - -def annotate_labels( - img: np.ndarray, - struct: torch.Tensor, - target_density: torch.Tensor -) -> np.ndarray: - # 三行标注:第一行结构标签,后两行显示五维目标密度 - s = struct.tolist() - d = target_density.tolist() - line1 = f"sym:{s[0]} outer:{s[1]}" - line2 = f"wall:{d[0]:.2f} door:{d[1]:.2f}" - line3 = f"enemy:{d[2]:.2f} ent:{d[3]:.2f} res:{d[4]:.2f}" - img = img.copy() - for text, y in [(line1, 12), (line2, 24), (line3, 36)]: - cv2.putText(img, text, (2, y), cv2.FONT_HERSHEY_SIMPLEX, 0.35, (0, 0, 0), 2) - cv2.putText(img, text, (2, y), cv2.FONT_HERSHEY_SIMPLEX, 0.35, (255, 255, 255), 1) - return img - -def rand_keep() -> tuple[bool, bool, bool]: - b = random.choice([True, False]) - return (b, b, b) - -def keep_label(kf: tuple[bool, bool, bool]) -> str: - return 'fix' if kf[0] else 'free' - -def build_dataset_sample_case( - dataset: GinkaSeperatedDataset, - models: list[torch.nn.Module], - dist_models: tuple, - device: torch.device, - idx: int | None = None -) -> dict: - vq1, vq2, vq3, mg1, mg2, mg3, quantizers, _, _, _, _, _ = models - dist_encoder, dist_quantizer = dist_models - sample = dataset.random_sample_map(idx=idx) - - enc1_t = sample["encoder_stage1"].to(device).reshape(1, MAP_SIZE) - enc2_t = sample["encoder_stage2"].to(device).reshape(1, MAP_SIZE) - enc3_t = sample["encoder_stage3"].to(device).reshape(1, MAP_SIZE) - struct_t = sample["struct_inject"].to(device).reshape(1, -1) - target_density_t = sample["target_density"].to(device).reshape(1, -1) - dist_field_t = sample["distance_field"].to(device).reshape(1, -1) - - with torch.no_grad(): - z_e1 = vq1(enc1_t) - z_e2 = vq2(enc2_t) - z_e3 = vq3(enc3_t) - z_q, commit_loss, code_hits = quantize_stage_latents( - quantizers, z_e1, z_e2, z_e3 - ) - z_e_dist = dist_encoder(dist_field_t) - z_dist, _, _, _, _ = dist_quantizer(z_e_dist) - - return { - "sample": sample, - "struct": struct_t, - "target_density": target_density_t, - "z_q": z_q, - "z_dist": z_dist, - "sample_idx": sample["sample_idx"] - } - -def sample_case_label(case: dict) -> str: - return case["sample"]["map_name"] - -# 验证可视化 part1:3×3 网格;行1=编码器输入,行2=掩码输入,行3=三阶段预测(合并) -def visualize_part1(batch, logits1, logits2, logits3, tile_dict): - SEP = 3 - TILE_SIZE = 32 - img_h = MAP_H * TILE_SIZE - img_w = MAP_W * TILE_SIZE - - def to_img(mat): - return matrix_to_image_cv(mat, tile_dict, TILE_SIZE) - - pred1 = torch.argmax(logits1[0], dim=-1).cpu().numpy().reshape(MAP_H, MAP_W) - pred2 = torch.argmax(logits2[0], dim=-1).cpu().numpy().reshape(MAP_H, MAP_W) - pred3 = torch.argmax(logits3[0], dim=-1).cpu().numpy().reshape(MAP_H, MAP_W) - - enc1_np = batch["encoder_stage1"][0].numpy().reshape(MAP_H, MAP_W) - enc2_np = batch["encoder_stage2"][0].numpy().reshape(MAP_H, MAP_W) - enc3_np = batch["encoder_stage3"][0].numpy().reshape(MAP_H, MAP_W) - inp1_np = batch["input_stage1"][0].numpy().reshape(MAP_H, MAP_W) - inp2_np = batch["input_stage2"][0].numpy().reshape(MAP_H, MAP_W) - inp3_np = batch["input_stage3"][0].numpy().reshape(MAP_H, MAP_W) - - # 将各阶段掩码输入中的 MASK 位用模型预测值填充,保留非掩码位原值 - result1 = inp1_np.copy() - result1[inp1_np == MASK_TOKEN] = pred1[inp1_np == MASK_TOKEN] - result2 = inp2_np.copy() - result2[inp2_np == MASK_TOKEN] = pred2[inp2_np == MASK_TOKEN] - result3 = inp3_np.copy() - result3[inp3_np == MASK_TOKEN] = pred3[inp3_np == MASK_TOKEN] - - rows = [ - [annotate(to_img(enc1_np), batch["map_name"][0]), to_img(enc2_np), to_img(enc3_np)], - [to_img(inp1_np), to_img(inp2_np), to_img(inp3_np)], - [to_img(result1), to_img(result2), to_img(result3)], - ] - grid = np.ones((3 * img_h + 4 * SEP, 3 * img_w + 4 * SEP, 3), dtype=np.uint8) * 255 - for r, row in enumerate(rows): - for c, img in enumerate(row): - y = SEP + r * (img_h + SEP) - x = SEP + c * (img_w + SEP) - grid[y:y + img_h, x:x + img_w] = img - return grid - -# 验证可视化 part2:行1=真实地图三阶段,行2=stage1 输入与使用真实 z 自回归生成的各阶段结果 -def visualize_part2(batch, z_q, z_dist, models, device, tile_dict): - SEP = 3 - TILE_SIZE = 32 - img_h = MAP_H * TILE_SIZE - img_w = MAP_W * TILE_SIZE - - def to_img(mat): - return matrix_to_image_cv(mat, tile_dict, TILE_SIZE) - - inp1_t = batch["input_stage1"][0:1].to(device).reshape(1, MAP_SIZE) - struct_t = batch["struct_inject"][0:1].to(device) - target_density_t = batch["target_density"][0:1].to(device) - z_q_single = (z_q[0][0:1], z_q[1][0:1], z_q[2][0:1]) - kf = rand_keep() - auto_pred1_np, auto_merged12, auto_merged123 = full_generate_specific_z( - inp1_t, z_q_single, z_dist[0:1], struct_t, target_density_t, models, device, keep_fixed=kf - ) - kf_label = 'fix' if kf[0] else 'free' - - enc1_np = batch["encoder_stage1"][0].numpy().reshape(MAP_H, MAP_W) - enc2_np = batch["encoder_stage2"][0].numpy().reshape(MAP_H, MAP_W) - enc3_np = batch["encoder_stage3"][0].numpy().reshape(MAP_H, MAP_W) - inp1_np = batch["input_stage1"][0].numpy().reshape(MAP_H, MAP_W) - - struct_cpu = batch["struct_inject"][0] - target_density_cpu = batch["target_density"][0] - - rows = [ - [annotate(to_img(enc1_np), batch["map_name"][0]), to_img(enc2_np), to_img(enc3_np)], - [ - annotate(to_img(inp1_np), kf_label), - annotate_labels(to_img(auto_pred1_np), struct_cpu, target_density_cpu), - annotate_labels(to_img(auto_merged12), struct_cpu, target_density_cpu), - annotate_labels(to_img(auto_merged123), struct_cpu, target_density_cpu) - ], - ] - grid = np.ones((2 * img_h + 3 * SEP, 4 * img_w + 5 * SEP, 3), dtype=np.uint8) * 255 - for r, row in enumerate(rows): - for c, img in enumerate(row): - y = SEP + r * (img_h + SEP) - x = SEP + c * (img_w + SEP) - grid[y:y + img_h, x:x + img_w] = img - return grid - -# 验证可视化 part4:2×3 网格;保留稀疏墙壁种子,但 z 与标签来自训练集样本 -def visualize_rand( +# 墙壁种子生成:随机放置墙壁后让模型生长 +def visualize_seed( train_dataset: GinkaSeperatedDataset, - models: list[torch.nn.Module], - dist_models: tuple, - device: torch.device, - tile_dict -): - SEP = 3 - TILE_SIZE = 32 - img_h = MAP_H * TILE_SIZE - img_w = MAP_W * TILE_SIZE - - def to_img(mat): - return matrix_to_image_cv(mat, tile_dict, TILE_SIZE) - - vq1, vq2, vq3, mg1, mg2, mg3, quantizers, _, _, _, _, _ = models - dist_encoder, dist_quantizer = dist_models - - samples_data = [] - for _ in range(4): - sample = train_dataset.random_sample_map() - raw_map = sample["raw_map"].cpu().numpy().reshape(MAP_H, MAP_W) - ratio = random.uniform(0.2, 0.8) - mask = rect_mask(ratio) - # 验证时同样修正掩码,避免孤立墙壁分量完全被盖住 - target1 = sample["encoder_stage1"].cpu().numpy().reshape(MAP_H, MAP_W) - from ginka.dataset import ensure_wall_connection - mask = ensure_wall_connection(mask, target1) - - enc1_t = sample["encoder_stage1"].to(device).reshape(1, MAP_SIZE) - enc2_t = sample["encoder_stage2"].to(device).reshape(1, MAP_SIZE) - enc3_t = sample["encoder_stage3"].to(device).reshape(1, MAP_SIZE) - struct_t = sample["struct_inject"].to(device).reshape(1, -1) - target_density_t = sample["target_density"].to(device).reshape(1, -1) - dist_field_t = sample["distance_field"].to(device).reshape(1, -1) - - with torch.no_grad(): - z_e1 = vq1(enc1_t) - z_e2 = vq2(enc2_t) - z_e3 = vq3(enc3_t) - z_q, _, _ = quantize_stage_latents(quantizers, z_e1, z_e2, z_e3) - z_e_dist = dist_encoder(dist_field_t) - z_dist, _, _, _, _ = dist_quantizer(z_e_dist) - - _, _, _, merged_s3 = inpaint_generate( - raw_map, mask, z_q, z_dist, struct_t, target_density_t, - models, device - ) - - masked_display = raw_map.copy() - masked_display[mask] = MASK_TOKEN - - samples_data.append({ - "masked": masked_display, - "result": merged_s3, - "mask": mask, - "label": f"{sample['map_name']} {int(ratio * 100)}%" - }) - - row1 = [ - annotate(to_img(samples_data[0]["masked"]), samples_data[0]["label"]), - to_img(samples_data[0]["result"]), - annotate(to_img(samples_data[1]["masked"]), samples_data[1]["label"]), - to_img(samples_data[1]["result"]), - ] - row2 = [ - annotate(to_img(samples_data[2]["masked"]), samples_data[2]["label"]), - to_img(samples_data[2]["result"]), - annotate(to_img(samples_data[3]["masked"]), samples_data[3]["label"]), - to_img(samples_data[3]["result"]), - ] - grid = np.ones((2 * img_h + 3 * SEP, 4 * img_w + 5 * SEP, 3), dtype=np.uint8) * 255 - for r, row in enumerate([row1, row2]): - for c, img in enumerate(row): - y = SEP + r * (img_h + SEP) - x = SEP + c * (img_w + SEP) - grid[y:y + img_h, x:x + img_w] = img - return grid - -def visualize_validate( - batch, logits1, logits2, logits3, z_q, z_dist, - models: list[torch.nn.Module], dist_models: tuple, device: torch.device, tile_dict, - train_dataset: GinkaSeperatedDataset, epoch: int, batch_idx: int -): - save_dir = f"result/seperated/e{epoch}" - os.makedirs(save_dir, exist_ok=True) - cv2.imwrite(f"{save_dir}/val{batch_idx}.png", visualize_part1(batch, logits1, logits2, logits3, tile_dict)) - cv2.imwrite(f"{save_dir}/full{batch_idx}.png", visualize_part2(batch, z_q, z_dist, models, device, tile_dict)) - cv2.imwrite(f"{save_dir}/rand{batch_idx}.png", visualize_rand(train_dataset, models, dist_models, device, tile_dict)) - -def validate( - dataloader: DataLoader, - models: list[torch.nn.Module], - dist_models: tuple, + models: SeperatedModels, device: torch.device, tile_dict, - train_dataset: GinkaSeperatedDataset, epoch: int ): - vq1, vq2, vq3, mg1, mg2, mg3, quantizers, _, _, _, _, _ = models - dist_encoder, dist_quantizer = dist_models - quantizer1, quantizer2, quantizer3 = quantizers + SEP = 3 + TILE_SIZE = 32 + img_h = MAP_H * TILE_SIZE + img_w = MAP_W * TILE_SIZE - # 切换为推理模式(关闭 Dropout / BatchNorm 统计更新) - for m in [vq1, vq2, vq3, mg1, mg2, mg3, dist_encoder]: - m.eval() + def to_img(mat): + return matrix_to_image_cv(mat, tile_dict, TILE_SIZE) - # 累计各阶段损失(跨所有 batch 求和,最终除以 batch 数得到均值) - loss1_total = torch.Tensor([0]).to(device) - loss2_total = torch.Tensor([0]).to(device) - loss3_total = torch.Tensor([0]).to(device) - commit_total = torch.Tensor([0]).to(device) - adj1_total = torch.Tensor([0]).to(device) - adj2_total = torch.Tensor([0]).to(device) - adj3_total = torch.Tensor([0]).to(device) - patch1_total = torch.Tensor([0]).to(device) - patch2_total = torch.Tensor([0]).to(device) - patch3_total = torch.Tensor([0]).to(device) - code_hits_total = (torch.zeros(quantizer1.K, device=device), torch.zeros(quantizer2.K, device=device), torch.zeros(quantizer3.K, device=device)) # validate + save_dir = f"result/seperated/e{epoch}" + os.makedirs(save_dir, exist_ok=True) - density_metrics = { - 1: {"mae": 0.0, "over": 0.0, "count": 0}, - 2: {"mae": 0.0, "over": 0.0, "count": 0}, - 4: {"mae": 0.0, "over": 0.0, "count": 0}, - 5: {"mae": 0.0, "over": 0.0, "count": 0}, - 3: {"mae": 0.0, "over": 0.0, "count": 0}, - } + for img_i in range(5): + samples_data = [] + for _ in range(4): + sample = train_dataset.random_sample_map() + struct_t = sample["struct_inject"].to(device).reshape(1, -1) + target_density_t = sample["target_density"].to(device).reshape(1, -1) + enc1_t = sample["encoder_stage1"].to(device).reshape(1, MAP_SIZE) + enc2_t = sample["encoder_stage2"].to(device).reshape(1, MAP_SIZE) + enc3_t = sample["encoder_stage3"].to(device).reshape(1, MAP_SIZE) - idx = 0 + with torch.no_grad(): + z_e1 = models.vq1(enc1_t) + z_e2 = models.vq2(enc2_t) + z_e3 = models.vq3(enc3_t) + z_q, _, _, _ = quantize_stage_latents(models, z_e1, z_e2, z_e3) + z1, z2, z3 = z_q - with torch.no_grad(): - for batch in tqdm(dataloader, leave=False, desc="Validate Progress", disable=disable_tqdm): + inp = torch.full((1, MAP_SIZE), 0, dtype=torch.long, device=device) + seed_count = random.randint(7, 14) - # 三阶段各自的掩码输入、预测目标和 VQ 编码器输入 - inp1 = batch["input_stage1"].to(device).reshape(-1, MAP_SIZE) - target1 = batch["target_stage1"].to(device).reshape(-1, MAP_SIZE) - enc1 = batch["encoder_stage1"].to(device).reshape(-1, MAP_SIZE) + # 分层采样:将 13x13 地图划分为 4x4 网格,从不同格中随机取种子点 + grid_h, grid_w = 4, 4 + cell_h = MAP_H // grid_h + 1 + cell_w = MAP_W // grid_w + 1 + all_cells = [(r, c) for r in range(grid_h) for c in range(grid_w)] + chosen_cells = random.sample(all_cells, seed_count) + for cell_r, cell_c in chosen_cells: + r_min = cell_r * cell_h + r_max = min(r_min + cell_h, MAP_H) + c_min = cell_c * cell_w + c_max = min(c_min + cell_w, MAP_W) + sr = random.randint(r_min, max(r_min, r_max - 1)) + sc = random.randint(c_min, max(c_min, c_max - 1)) + flat_idx = sr * MAP_W + sc + inp[0, flat_idx] = 1 - inp2 = batch["input_stage2"].to(device).reshape(-1, MAP_SIZE) - target2 = batch["target_stage2"].to(device).reshape(-1, MAP_SIZE) - enc2 = batch["encoder_stage2"].to(device).reshape(-1, MAP_SIZE) + _, _, merged = full_generate( + inp, z1, z2, z3, + struct_t, target_density_t, models, + steps=SEED_SAMPLE_STEPS + ) - inp3 = batch["input_stage3"].to(device).reshape(-1, MAP_SIZE) - target3 = batch["target_stage3"].to(device).reshape(-1, MAP_SIZE) - enc3 = batch["encoder_stage3"].to(device).reshape(-1, MAP_SIZE) + samples_data.append(merged[0]) - struct = batch["struct_inject"].to(device) - target_density = batch["target_density"].to(device) - dist_field = batch["distance_field"].to(device) + grid = np.ones((2 * img_h + 3 * SEP, 2 * img_w + 3 * SEP, 3), dtype=np.uint8) * 255 + for r in range(2): + for c in range(2): + y = SEP + r * (img_h + SEP) + x = SEP + c * (img_w + SEP) + grid[y:y + img_h, x:x + img_w] = to_img(samples_data[r * 2 + c]) + cv2.imwrite(f"{save_dir}/seed{img_i}.png", grid) - # 距离场编码与量化 - z_e_dist = dist_encoder(dist_field) - z_dist, _, _, _, _ = dist_quantizer(z_e_dist) +# 矩形掩码生成:对真实地图随机掩码后让模型补全 +def visualize_mask( + train_dataset: GinkaSeperatedDataset, + models: SeperatedModels, + device: torch.device, + tile_dict, + epoch: int +): + SEP = 3 + TILE_SIZE = 32 + img_h = MAP_H * TILE_SIZE + img_w = MAP_W * TILE_SIZE - # VQ 编码:各阶段独立编码并分别量化 - z_e1 = vq1(enc1) # [B, L, d_z] - z_e2 = vq2(enc2) - z_e3 = vq3(enc3) + def to_img(mat): + return matrix_to_image_cv(mat, tile_dict, TILE_SIZE) - z_q, commit_loss, code_hits = quantize_stage_latents( - quantizers, z_e1, z_e2, z_e3 - ) - z_q1, z_q2, z_q3 = z_q + save_dir = f"result/seperated/e{epoch}" + os.makedirs(save_dir, exist_ok=True) - remain1 = compute_remaining(inp1, target_density, 1) - remain2 = compute_remaining(inp2, target_density, 2) - remain3 = compute_remaining(inp3, target_density, 3) + for img_i in range(5): + samples_data = [] + for _ in range(4): + sample = train_dataset.random_sample_map() + raw_map = sample["raw_map"].cpu().numpy().reshape(MAP_H, MAP_W) + ratio = random.uniform(0.2, 0.8) - # 三阶段 MaskGIT 推理:各阶段接收自己的 z_q 和共享的 z_dist - logits1 = mg1(inp1, z_q1, z_dist, struct, remain1) - logits2 = mg2(inp2, z_q2, z_dist, struct, remain2) - logits3 = mg3(inp3, z_q3, z_dist, struct, remain3) + enc1_t = sample["encoder_stage1"].to(device).reshape(1, MAP_SIZE) + enc2_t = sample["encoder_stage2"].to(device).reshape(1, MAP_SIZE) + enc3_t = sample["encoder_stage3"].to(device).reshape(1, MAP_SIZE) + struct_t = sample["struct_inject"].to(device).reshape(1, -1) + target_density_t = sample["target_density"].to(device).reshape(1, -1) - loss1_total += cross_entropy_loss(logits1, target1) - loss2_total += cross_entropy_loss(logits2, target2) - loss3_total += cross_entropy_loss(logits3, target3) - commit_total += commit_loss - adj1_total += adjacency_loss(logits1, target1) - adj2_total += adjacency_loss(logits2, target2) - adj3_total += adjacency_loss(logits3, target3) - patch1_total += patch_loss(logits1, target1, PATCH_KERNEL_SIZE, PATCH_SIGMA) - patch2_total += patch_loss(logits2, target2, PATCH_KERNEL_SIZE, PATCH_SIGMA) - patch3_total += patch_loss(logits3, target3, PATCH_KERNEL_SIZE, PATCH_SIGMA) - code_hits_total = (code_hits_total[0] + code_hits[0], code_hits_total[1] + code_hits[1], code_hits_total[2] + code_hits[2]) + with torch.no_grad(): + z_e1 = models.vq1(enc1_t) + z_e2 = models.vq2(enc2_t) + z_e3 = models.vq3(enc3_t) + z_q, _, _, _ = quantize_stage_latents(models, z_e1, z_e2, z_e3) + z1, z2, z3 = z_q - # 计算各目标对象的真实密度误差与过量生成密度 - pred1_map = torch.argmax(logits1, dim=-1).cpu() - pred2_map = torch.argmax(logits2, dim=-1).cpu() # [B, MAP_SIZE] - pred3_map = torch.argmax(logits3, dim=-1).cpu() - true1_map = target1.cpu() # [B, MAP_SIZE] - true2_map = target2.cpu() # [B, MAP_SIZE] - true3_map = target3.cpu() - metric_sources = [ - (1, pred1_map, true1_map), - (2, pred2_map, true2_map), - (4, pred2_map, true2_map), - (5, pred2_map, true2_map), - (3, pred3_map, true3_map), - ] - for tile_id, pred_map_batch, true_map_batch in metric_sources: - for batch_idx in range(pred_map_batch.size(0)): - pred_map = pred_map_batch[batch_idx] - true_map = true_map_batch[batch_idx] - pred_count = float((pred_map == tile_id).sum().item()) - true_count = float((true_map == tile_id).sum().item()) - if tile_id == 2: - pred_count += float((pred_map == 6).sum().item()) - true_count += float((true_map == 6).sum().item()) - density_metrics[tile_id]["mae"] += abs(pred_count - true_count) / MAP_SIZE - density_metrics[tile_id]["over"] += max(pred_count - true_count, 0.0) / MAP_SIZE - density_metrics[tile_id]["count"] += 1 + mask = torch.rand(MAP_SIZE, device=device) < ratio + inp = torch.tensor(raw_map.flatten(), dtype=torch.long, device=device).reshape(1, MAP_SIZE) + inp[0, mask] = MASK_TOKEN - # 每个 batch 生成三种可视化图(val/full/rand) - visualize_validate( - batch, logits1, logits2, logits3, z_q, z_dist, - models, dist_models, device, tile_dict, train_dataset, epoch, idx - ) - idx += 1 + _, _, merged = full_generate( + inp, z1, z2, z3, struct_t, target_density_t, models + ) - tile_names = {1: 'wall', 2: 'door', 4: 'enemy', 5: 'entrance', 3: 'resource'} - for tile_id in [1, 2, 4, 5, 3]: - count = density_metrics[tile_id]["count"] - avg_mae = density_metrics[tile_id]["mae"] / count if count > 0 else 0.0 - avg_over = density_metrics[tile_id]["over"] / count if count > 0 else 0.0 - tqdm.write(f" density {tile_names[tile_id]}: mae={avg_mae:.4f} over={avg_over:.4f}") + samples_data.append(merged[0]) - # 恢复训练模式 - for m in [vq1, vq2, vq3, mg1, mg2, mg3, dist_encoder]: - m.train() - - return loss1_total, loss2_total, loss3_total, adj1_total, adj2_total, adj3_total, patch1_total, patch2_total, patch3_total, commit_total, code_hits_total + grid = np.ones((2 * img_h + 3 * SEP, 2 * img_w + 3 * SEP, 3), dtype=np.uint8) * 255 + for r in range(2): + for c in range(2): + y = SEP + r * (img_h + SEP) + x = SEP + c * (img_w + SEP) + grid[y:y + img_h, x:x + img_w] = to_img(samples_data[r * 2 + c]) + cv2.imwrite(f"{save_dir}/mask{img_i}.png", grid) def train(device: torch.device): args = parse_arguments() - result = build_model(device) - vq1, vq2, vq3, mg1, mg2, mg3, quantizers, optimizer, scheduler, latent_mask_embedding, dist_encoder, dist_quantizer = result - models = [vq1, vq2, vq3, mg1, mg2, mg3, quantizers, optimizer, scheduler, latent_mask_embedding, dist_encoder, dist_quantizer] - dist_models = (dist_encoder, dist_quantizer) - quantizer1, quantizer2, quantizer3 = quantizers + result = SeperatedModels(device) tqdm.write(f"Device: {device}") model_list = [ - ("vq1", vq1), ("vq2", vq2), ("vq3", vq3), - ("mg1", mg1), ("mg2", mg2), ("mg3", mg3), - ("quantizer1", quantizer1), ("quantizer2", quantizer2), ("quantizer3", quantizer3), - ("dist_encoder", dist_encoder), ("dist_quantizer", dist_quantizer) + ("vq1", result.vq1), ("vq2", result.vq2), ("vq3", result.vq3), + ("mg1", result.mg1), ("mg2", result.mg2), ("mg3", result.mg3), + ("quantizer1", result.quantizer1), ("quantizer2", result.quantizer2), ("quantizer3", result.quantizer3) ] total_params = 0 for name, m in model_list: @@ -1170,28 +257,7 @@ def train(device: torch.device): if args.resume: # 从指定检查点恢复:加载所有模型权重及训练状态 - ckpt = torch.load(args.state, map_location=device) - vq1.load_state_dict(ckpt["vq1"]) - vq2.load_state_dict(ckpt["vq2"]) - vq3.load_state_dict(ckpt["vq3"]) - mg1.load_state_dict(ckpt["mg1"]) - mg2.load_state_dict(ckpt["mg2"]) - mg3.load_state_dict(ckpt["mg3"]) - quantizer1.load_state_dict(ckpt["quantizer1"]) - quantizer2.load_state_dict(ckpt["quantizer2"]) - quantizer3.load_state_dict(ckpt["quantizer3"]) - if "dist_encoder" in ckpt: - dist_encoder.load_state_dict(ckpt["dist_encoder"]) - if "dist_quantizer" in ckpt: - dist_quantizer.load_state_dict(ckpt["dist_quantizer"]) - if "latent_mask_embedding" in ckpt: - latent_mask_embedding.data.copy_(ckpt["latent_mask_embedding"]) - # load_optim=False 时可跳过优化器/调度器恢复(适合调整学习率后继续训练) - if args.load_optim and "optimizer" in ckpt: - optimizer.load_state_dict(ckpt["optimizer"]) - if args.load_optim and "scheduler" in ckpt: - scheduler.load_state_dict(ckpt["scheduler"]) - start_epoch = ckpt.get("epoch", 0) # 从上次保存的 epoch 继续 + start_epoch = result.load(args.state, load_optim=args.load_optim, map_location=device) tqdm.write(f"Resumed from epoch {start_epoch}: {args.state}") os.makedirs("result/seperated", exist_ok=True) @@ -1203,14 +269,6 @@ def train(device: torch.device): dataset, batch_size=BATCH_SIZE, shuffle=True ) - dataset_val = GinkaSeperatedDataset( - args.validate, subset_weights=SUBSET_WEIGHTS, - density_stats=dataset.density_stats # 复用训练集统计量,保证归一化语义一致 - ) - dataloader_val = DataLoader( - dataset_val, batch_size=min(BATCH_SIZE, len(dataset_val) // 8), shuffle=True - ) - # 预加载图块图像,键为文件名(不含扩展名),用于可视化时将 ID 映射为像素图 tile_dict = {} for f in os.listdir("tiles"): @@ -1225,13 +283,7 @@ def train(device: torch.device): loss2_total = torch.Tensor([0]).to(device) loss3_total = torch.Tensor([0]).to(device) commit_total = torch.Tensor([0]).to(device) - adj1_total = torch.Tensor([0]).to(device) - adj2_total = torch.Tensor([0]).to(device) - adj3_total = torch.Tensor([0]).to(device) - patch1_total = torch.Tensor([0]).to(device) - patch2_total = torch.Tensor([0]).to(device) - patch3_total = torch.Tensor([0]).to(device) - code_hits_total = (torch.zeros(quantizer1.K, device=device), torch.zeros(quantizer2.K, device=device), torch.zeros(quantizer3.K, device=device)) # validate + code_hits_total = (torch.zeros(result.quantizer1.K, device=device), torch.zeros(result.quantizer2.K, device=device), torch.zeros(result.quantizer3.K, device=device)) # validate for batch in tqdm(dataloader, leave=False, desc="Epoch Progress", disable=disable_tqdm): # 三阶段各自的掩码输入序列、预测目标和编码器上下文 @@ -1250,71 +302,50 @@ def train(device: torch.device): # 结构条件向量:[cond_sym, cond_outer] struct = batch["struct_inject"].to(device) target_density = batch["target_density"].to(device) - dist_field = batch["distance_field"].to(device) - optimizer.zero_grad() # 训练循环 + result.optimizer.zero_grad() # 训练循环 # VQ 编码:各阶段编码器分别处理各自上下文切片 - z_e1 = vq1(enc1) # [B, L, d_z] - z_e2 = vq2(enc2) - z_e3 = vq3(enc3) + z_e1 = result.vq1(enc1) # [B, L, d_z] + z_e2 = result.vq2(enc2) + z_e3 = result.vq3(enc3) # 三阶段分别量化,各自使用独立 codebook - z_q, commit_loss, code_hits = quantize_stage_latents( - quantizers, z_e1, z_e2, z_e3 + z_q, commit_loss, code_hits, entropy_loss = quantize_stage_latents( + result, z_e1, z_e2, z_e3 ) z_q1, z_q2, z_q3 = z_q - # 距离场编码与量化 - z_e_dist = dist_encoder(dist_field) - z_dist_raw, _, commit_loss_dist, _, _ = dist_quantizer(z_e_dist) - z_dist = z_dist_raw - # latent dropout:训练时随机丢弃部分码字,替换为可学习 mask 嵌入 - z_q1 = apply_z_dropout(z_q1, latent_mask_embedding, MG_Z_DROPOUT) - z_q2 = apply_z_dropout(z_q2, latent_mask_embedding, MG_Z_DROPOUT) - z_q3 = apply_z_dropout(z_q3, latent_mask_embedding, MG_Z_DROPOUT) - - rollout_steps = build_reference_rollout_steps(REFERENCE_SAMPLE_PROB) - inp1 = sample_reference_inputs( - mg1, inp1, z_q1, z_dist, struct, target_density, 1, rollout_steps - ) - inp2 = sample_reference_inputs( - mg2, inp2, z_q2, z_dist, struct, target_density, 2, rollout_steps - ) - inp3 = sample_reference_inputs( - mg3, inp3, z_q3, z_dist, struct, target_density, 3, rollout_steps - ) + z_q1 = apply_z_dropout(z_q1, result.latent_mask_embedding, MG_Z_DROPOUT) + z_q2 = apply_z_dropout(z_q2, result.latent_mask_embedding, MG_Z_DROPOUT) + z_q3 = apply_z_dropout(z_q3, result.latent_mask_embedding, MG_Z_DROPOUT) remain1 = compute_remaining(inp1, target_density, 1) remain2 = compute_remaining(inp2, target_density, 2) remain3 = compute_remaining(inp3, target_density, 3) - # 三阶段 MaskGIT 前向:各阶段接收自己的 z_q、z_dist、struct 和动态 remain - logits1 = mg1(inp1, z_q1, z_dist, struct, remain1) - logits2 = mg2(inp2, z_q2, z_dist, struct, remain2) - logits3 = mg3(inp3, z_q3, z_dist, struct, remain3) + # 三阶段 MaskGIT 前向:各阶段接收自己的 z_q、struct 和动态 remain + logits1 = result.mg1(inp1, z_q1, struct, remain1) + logits2 = result.mg2(inp2, z_q2, struct, remain2) + logits3 = result.mg3(inp3, z_q3, struct, remain3) - # 三阶段 Cross Entropy + 邻接损失 + Patch 损失 + VQ commit loss 加权求和 - loss1 = cross_entropy_loss(logits1, target1) - loss2 = cross_entropy_loss(logits2, target2) - loss3 = cross_entropy_loss(logits3, target3) - - adj2 = adjacency_loss(logits2, target2) - adj3 = adjacency_loss(logits3, target3) - patch2 = patch_loss(logits2, target2, PATCH_KERNEL_SIZE, PATCH_SIGMA) - patch3 = patch_loss(logits3, target3, PATCH_KERNEL_SIZE, PATCH_SIGMA) + # 三阶段 Cross Entropy:仅对输入中为 MASK_TOKEN 的位置计算 loss + mask1 = (inp1 == MASK_TOKEN) + mask2 = (inp2 == MASK_TOKEN) + mask3 = (inp3 == MASK_TOKEN) + loss1 = cross_entropy_loss(logits1, target1, mask1) + loss2 = cross_entropy_loss(logits2, target2, mask2) + loss3 = cross_entropy_loss(logits3, target3, mask3) loss1_weighted = STAGE1_CE_WEIGHT * loss1 loss2_weighted = STAGE2_CE_WEIGHT * loss2 loss3_weighted = STAGE3_CE_WEIGHT * loss3 - adj_weighted = LAMBDA_ADJ2 * adj2 + LAMBDA_ADJ3 * adj3 - patch_weighted = LAMBDA_PATCH2 * patch2 + LAMBDA_PATCH3 * patch3 - commit_weighted = VQ_BETA * commit_loss + VQ_BETA_DIST * commit_loss_dist - loss = loss1_weighted + loss2_weighted + loss3_weighted + adj_weighted + patch_weighted + commit_weighted + commit_weighted = VQ_BETA * commit_loss + VQ_GAMMA * entropy_loss + loss = loss1_weighted + loss2_weighted + loss3_weighted + commit_weighted loss.backward() - optimizer.step() + result.optimizer.step() # detach 后累加,避免保留计算图占用显存 loss_total += loss.detach() @@ -1322,97 +353,40 @@ def train(device: torch.device): loss2_total += loss2.detach() loss3_total += loss3.detach() commit_total += commit_loss.detach() - adj2_total += adj2.detach() - adj3_total += adj3.detach() - patch2_total += patch2.detach() - patch3_total += patch3.detach() code_hits_total = (code_hits_total[0] + code_hits[0].detach(), code_hits_total[1] + code_hits[1].detach(), code_hits_total[2] + code_hits[2].detach()) # accumulate train # 每个 epoch 结束后更新学习率 - scheduler.step() + result.scheduler.step() data_length = len(dataloader) - train_perplexity, train_usage_rate, train_active_codes = summarize_codebook_hits(code_hits_total) + stats = summarize_codebook_hits(code_hits_total) + parts = [] + for name in ["q1(stage1)", "q2(stage2)", "q3(stage3)"]: + s = stats[name] + parts.append(f"{s['active']}/{s['K']} ppl={s['ppl']:.1f}") + c = stats["combined"] tqdm.write( f"[{datetime.now().strftime('%Y-%m-%d %H:%M:%S')}] " f"E: {epoch + 1} | " f"Loss: {loss_total.item() / data_length:.4f} | " f"CE: {loss1_total.item() / data_length:.4f}, {loss2_total.item() / data_length:.4f}, {loss3_total.item() / data_length:.4f} | " - f"ADJ: {(LAMBDA_ADJ2 * adj2_total.item() + LAMBDA_ADJ3 * adj3_total.item()) / data_length:.4f}, {adj2_total.item() / data_length:.4f}, {adj3_total.item() / data_length:.4f} | " - f"PAT: {(LAMBDA_PATCH2 * patch2_total.item() + LAMBDA_PATCH3 * patch3_total.item()) / data_length:.4f}, {patch2_total.item() / data_length:.4f}, {patch3_total.item() / data_length:.4f} | " f"VQ: {commit_total.item() / data_length:.4f} | " - f"PPL: {train_perplexity:.4f} | " - f"Usage: {train_usage_rate:.4f} ({train_active_codes}/{TOTAL_K}) | " - f"LR: {scheduler.get_last_lr()[0]:.6f}" + f"VQ: {' | '.join(parts)} | " + f"Total: {c['active']}/{c['K']} ppl={c['ppl']:.1f} | " + f"LR: {result.scheduler.get_last_lr()[0]:.6f}" ) - - # 每 CHECKPOINT 个 epoch 执行一次验证、可视化和检查点保存 + + # 每 CHECKPOINT 个 epoch 执行可视化并保存检查点 if (epoch + 1) % CHECKPOINT == 0: - losses = validate( - dataloader_val, models, dist_models, device, tile_dict, dataset, epoch + 1 - ) - loss1_total, loss2_total, loss3_total, _, adj2_total, adj3_total, _, patch2_total, patch3_total, commit_total, code_hits_total = losses - loss1_weighted = STAGE1_CE_WEIGHT * loss1_total - loss2_weighted = STAGE2_CE_WEIGHT * loss2_total - loss3_weighted = STAGE3_CE_WEIGHT * loss3_total - adj_weighted = LAMBDA_ADJ2 * adj2_total + LAMBDA_ADJ3 * adj3_total - patch_weighted = LAMBDA_PATCH2 * patch2_total + LAMBDA_PATCH3 * patch3_total - commit_weighted = VQ_BETA * commit_total - loss_total = loss1_weighted + loss2_weighted + loss3_weighted + adj_weighted + patch_weighted + commit_weighted - - data_length = len(dataloader_val) - val_perplexity, val_usage_rate, val_active_codes = summarize_codebook_hits(code_hits_total) - tqdm.write( - f"[Validate {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}] " - f"E: {epoch + 1} | " - f"Loss: {loss_total.item() / data_length:.4f} | " - f"CE: {loss1_total.item() / data_length:.4f}, {loss2_total.item() / data_length:.4f}, {loss3_total.item() / data_length:.4f} | " - f"ADJ: {(LAMBDA_ADJ2 * adj2_total.item() + LAMBDA_ADJ3 * adj3_total.item()) / data_length:.4f}, {adj2_total.item() / data_length:.4f}, {adj3_total.item() / data_length:.4f} | " - f"PAT: {(LAMBDA_PATCH2 * patch2_total.item() + LAMBDA_PATCH3 * patch3_total.item()) / data_length:.4f}, {patch2_total.item() / data_length:.4f}, {patch3_total.item() / data_length:.4f} | " - f"VQ: {commit_total.item() / data_length:.4f} | " - f"PPL: {val_perplexity:.4f} | " - f"Usage: {val_usage_rate:.4f} ({val_active_codes}/{TOTAL_K}) | " - ) - + visualize_seed(dataset, result, device, tile_dict, epoch + 1) + visualize_mask(dataset, result, device, tile_dict, epoch + 1) ckpt_path = f"result/seperated/sep-{epoch + 1}.pth" - torch.save({ - "epoch": epoch + 1, - "vq1": vq1.state_dict(), - "vq2": vq2.state_dict(), - "vq3": vq3.state_dict(), - "mg1": mg1.state_dict(), - "mg2": mg2.state_dict(), - "mg3": mg3.state_dict(), - "quantizer1": quantizer1.state_dict(), - "quantizer2": quantizer2.state_dict(), - "quantizer3": quantizer3.state_dict(), - "dist_encoder": dist_encoder.state_dict(), - "dist_quantizer": dist_quantizer.state_dict(), - "latent_mask_embedding": latent_mask_embedding.data, - "optimizer": optimizer.state_dict(), - "scheduler": scheduler.state_dict(), - }, ckpt_path) + result.save(ckpt_path, epoch + 1) tqdm.write(f"Saved checkpoint: {ckpt_path}") # 训练结束后保存最终完整权重(含优化器状态,可用于后续续训或推理) final_path = "result/seperated.pth" - torch.save({ - "epoch": EPOCHS, - "vq1": vq1.state_dict(), - "vq2": vq2.state_dict(), - "vq3": vq3.state_dict(), - "mg1": mg1.state_dict(), - "mg2": mg2.state_dict(), - "mg3": mg3.state_dict(), - "quantizer1": quantizer1.state_dict(), - "quantizer2": quantizer2.state_dict(), - "quantizer3": quantizer3.state_dict(), - "dist_encoder": dist_encoder.state_dict(), - "dist_quantizer": dist_quantizer.state_dict(), - "latent_mask_embedding": latent_mask_embedding.data, - "optimizer": optimizer.state_dict(), - "scheduler": scheduler.state_dict(), - }, final_path) + result.save(final_path, EPOCHS) tqdm.write(f"Training complete. Final model saved: {final_path}") if __name__ == "__main__": diff --git a/ginka/utils.py b/ginka/utils.py index 62d36f0..cd2b001 100644 --- a/ginka/utils.py +++ b/ginka/utils.py @@ -1,93 +1,122 @@ import torch -import torch.nn.functional as F import numpy as np +# 工具函数:密度常量、剩余密度计算、掩码生成、邻接检测 + +MAP_W = 13 # 地图宽度 +MAP_H = 13 # 地图高度 +DENSITY_DIM = 5 # [wall, door, monster, entrance, resource] +MAP_SIZE = MAP_W * MAP_H # 地图大小 + +WALL_DENSITY_IDX = 0 +DOOR_DENSITY_IDX = 1 +MONSTER_DENSITY_IDX = 2 +ENTRANCE_DENSITY_IDX = 3 +RESOURCE_DENSITY_IDX = 4 + def print_memory(device, tag=""): if torch.cuda.is_available(): print(f"{tag} | 当前显存: {torch.cuda.memory_allocated(device) / 1024**2:.2f} MB, 最大显存: {torch.cuda.max_memory_allocated(device) / 1024**2:.2f} MB") else: print("当前设备不支持 cuda.") - -def nms_sampling(noise: np.ndarray, k: int, radius=2): - # noise: [H, W] - noise = noise.copy() - points = [] - for _ in range(k): - idx = np.argmax(noise) - x, y = np.unravel_index(idx, noise.shape) +def compute_remaining( + current: torch.Tensor, + target_density: torch.Tensor, + stage: int +) -> torch.Tensor: + remain = torch.zeros(current.size(0), DENSITY_DIM, device=current.device) - points.append((x, y)) + visible_wall = (current == 1).sum(dim=1).float() / MAP_SIZE + visible_door = ((current == 2) | (current == 6)).sum(dim=1).float() / MAP_SIZE + visible_monster = (current == 4).sum(dim=1).float() / MAP_SIZE + visible_entrance = (current == 5).sum(dim=1).float() / MAP_SIZE + visible_resource = (current == 3).sum(dim=1).float() / MAP_SIZE - # 抑制周围 - x0 = max(0, x - radius) - x1 = min(noise.shape[0], x + radius + 1) - y0 = max(0, y - radius) - y1 = min(noise.shape[1], y + radius + 1) + if stage == 1: + remain[:, WALL_DENSITY_IDX] = ( + target_density[:, WALL_DENSITY_IDX] - visible_wall + ).clamp(min=0.0, max=1.0) + elif stage == 2: + remain[:, DOOR_DENSITY_IDX] = ( + target_density[:, DOOR_DENSITY_IDX] - visible_door + ).clamp(min=0.0, max=1.0) + remain[:, MONSTER_DENSITY_IDX] = ( + target_density[:, MONSTER_DENSITY_IDX] - visible_monster + ).clamp(min=0.0, max=1.0) + remain[:, ENTRANCE_DENSITY_IDX] = ( + target_density[:, ENTRANCE_DENSITY_IDX] - visible_entrance + ).clamp(min=0.0, max=1.0) + elif stage == 3: + remain[:, RESOURCE_DENSITY_IDX] = ( + target_density[:, RESOURCE_DENSITY_IDX] - visible_resource + ).clamp(min=0.0, max=1.0) - noise[x0:x1, y0:y1] = -np.inf + return remain - result = np.zeros_like(noise) - for x, y in points: - result[y, x] = 1 - +def rect_mask( + ratio: float, h_range: tuple[int, int] = (2, 7), + w_range: tuple[int, int] = (2, 7) +) -> np.ndarray: + # 纯矩形分块掩码,反复放置随机矩形直到掩码格数达标 + target = int(MAP_SIZE * ratio) + mask = np.zeros((MAP_H, MAP_W), dtype=bool) + while mask.sum() < target: + bh = np.random.randint(h_range[0], h_range[1]) + bw = np.random.randint(w_range[0], w_range[1]) + x = np.random.randint(0, MAP_H - bh + 1) + y = np.random.randint(0, MAP_W - bw + 1) + mask[x:x + bh, y:y + bw] = True + return mask + +def compute_adjacency_mask(flat_state: torch.Tensor) -> torch.Tensor: + # 返回与输入同形状的 bool tensor,True 表示该位置与任意墙壁 4-邻接 + # 支持 [MAP_SIZE] 和 [B, MAP_SIZE] 两种输入 + was_1d = flat_state.dim() == 1 + if was_1d: + flat_state = flat_state.unsqueeze(0) + state_2d = flat_state.reshape(flat_state.size(0), MAP_H, MAP_W) + wall = (state_2d == 1) + adj = torch.zeros_like(wall, dtype=torch.bool) + adj[:, :, 1:] |= wall[:, :, :-1] + adj[:, :, :-1] |= wall[:, :, 1:] + adj[:, 1:, :] |= wall[:, :-1, :] + adj[:, :-1, :] |= wall[:, 1:, :] + result = adj.reshape(flat_state.size(0), MAP_SIZE) + if was_1d: + result = result.squeeze(0) return result +def summarize_codebook_hits(code_hits): + # code_hits 为 tuple of 3 tensors(各量器不同 K) + # 返回各阶段独立统计 + 汇总统计 + names = ["q1(stage1)", "q2(stage2)", "q3(stage3)"] + result = {} + for hits, name in zip(code_hits, names): + total = hits.sum() + if total.item() <= 0: + result[name] = {"ppl": 0.0, "usage": 0.0, "active": 0, "K": int(hits.numel())} + continue + probs = hits / total + ppl = torch.exp( + -(probs * torch.log(probs.clamp_min(1e-10))).sum() + ).item() + active = int((hits > 0).sum().item()) + usage = active / hits.numel() + result[name] = {"ppl": ppl, "usage": usage, "active": active, "K": int(hits.numel())} -def masked_focal( - logits: torch.Tensor, - target: torch.Tensor, - tile_set: set, - gamma: float = 2.0, - balance: bool = True, -) -> torch.Tensor: - """ - 通道专属 Focal Loss + 逆频类别权重。 - - tile_set 内的位置以真实 tile ID 为目标,tile_set 外的位置以 0(空地)为目标, - 全部位置均参与损失计算。 - - balance=True 时,从 batch 内 corrected 标签的频率自动计算逆频权重, - 消除空地(0)因被大量 non-tile-set 位置填充而主导梯度的问题。 - 权重公式:w[c] = total / (count[c] * C),与 sklearn 'balanced' 一致。 - - Args: - logits: [B, H*W, num_classes] 解码头输出(未经 softmax) - target: [B, H*W] 完整地图 ground truth(整数 tile ID) - tile_set: set of int 本通道专属 tile 集合 - gamma: Focal Loss 聚焦参数 - balance: 是否开启逆频类别权重 - - Returns: - scalar tensor 通道专属加权 Focal Loss(均值) - """ - B, S, C = logits.shape - - # 非专属 tile 位置目标替换为 0(空地) - in_set = torch.zeros(B, S, dtype=torch.bool, device=logits.device) - for t in tile_set: - in_set |= (target == t) - - corrected = target.clone() - corrected[~in_set] = 0 - - # 逆频类别权重:用原始 target 统计频率,避免 corrected 中人工填 0 膨胀 - # count[0],导致 weight[0] 趋近于 0、非专属位置损失被消除的问题 - class_weight = None - if balance: - flat = corrected.view(-1) # [B*S] 原始标签 - counts = torch.bincount(flat, minlength=C).float() # [C] - class_weight = torch.sqrt(flat.numel() / (counts.clamp(min=1.0) * C)) - class_weight[counts == 0] = 0.0 # 未出现类别不参与 - - ce = F.cross_entropy( - logits.view(-1, C), - corrected.view(-1), - weight=class_weight, - reduction='none', - ).view(B, S) # [B, S] - - pt = torch.exp(-ce.detach()) # 正确类预测概率,stop-gradient - fl = (1.0 - pt) ** gamma * ce - - return fl.mean() + # 汇总统计 + combined = torch.cat([h.flatten() for h in code_hits], dim=0) + total_hits = combined.sum() + if total_hits.item() > 0: + probs = combined / total_hits + result["combined"] = { + "ppl": float(torch.exp( + -(probs * torch.log(probs.clamp_min(1e-10))).sum() + ).item()), + "active": int((combined > 0).sum().item()), + "K": int(combined.numel()), + } + else: + result["combined"] = {"ppl": 0.0, "active": 0, "K": int(combined.numel())} + return result diff --git a/ginka/vqvae/model.py b/ginka/vqvae/model.py index 5291075..0fd04d1 100644 --- a/ginka/vqvae/model.py +++ b/ginka/vqvae/model.py @@ -149,56 +149,6 @@ class GinkaVQVAE(nn.Module): return z_e -class DistFieldEncoder(nn.Module): - # 距离场编码器:将 L1 距离场编码为 latent z - # - # 使用与 GinkaVQVAE 相同的 Transformer 架构但更轻量: - # DistEmbedding (vocab=DIST_VOCAB) → + 2D 位置编码 → + summary tokens - # → 浅层 Transformer → Linear 投影 → z_e_dist [B, L, d_z] - - def __init__( - self, vocab: int = 13, L: int = 4, d_z: int = 64, - d_model: int = 128, nhead: int = 4, num_layers: int = 3, - dim_ff: int = 512, map_h: int = 13, map_w: int = 13 - ): - super().__init__() - self.L = L - self.map_h = map_h - self.map_w = map_w - - self.dist_embedding = nn.Embedding(vocab, d_model) - self.row_embedding = nn.Parameter(torch.randn(1, map_h, d_model) * 0.02) - self.col_embedding = nn.Parameter(torch.randn(1, map_w, d_model) * 0.02) - self.summary_tokens = nn.Parameter(torch.randn(1, L, d_model) * 0.02) - - encoder_layer = nn.TransformerEncoderLayer( - d_model=d_model, nhead=nhead, dim_feedforward=dim_ff, batch_first=True, - activation='gelu', norm_first=True, dropout=0.1, - ) - self.transformer = nn.TransformerEncoder(encoder_layer, num_layers=num_layers) - - self.proj = nn.Sequential( - nn.Linear(d_model, d_z), - nn.LayerNorm(d_z), - ) - - def forward(self, dist_field: torch.Tensor) -> torch.Tensor: - # dist_field: [B, H*W] 整数,值域 [0, DIST_MAX_BUCKET] - B, _ = dist_field.shape - - row_idx = torch.arange(self.map_h, device=dist_field.device).repeat_interleave(self.map_w) - col_idx = torch.arange(self.map_w, device=dist_field.device).repeat(self.map_h) - pos = self.row_embedding[0, row_idx] + self.col_embedding[0, col_idx] - - x = self.dist_embedding(dist_field) + pos # [B, H*W, d_model] - summary = self.summary_tokens.expand(B, -1, -1) # [B, L, d_model] - x = torch.cat([summary, x], dim=1) # [B, L+H*W, d_model] - - x = self.transformer(x) - - z_e = self.proj(x[:, :self.L]) # [B, L, d_z] - return z_e - if __name__ == "__main__": device = torch.device("cuda") if torch.cuda.is_available() else torch.device("cpu") diff --git a/ginka/vqvae/quantize.py b/ginka/vqvae/quantize.py index 4966b43..79bc8bf 100644 --- a/ginka/vqvae/quantize.py +++ b/ginka/vqvae/quantize.py @@ -9,13 +9,15 @@ class VectorQuantizer(nn.Module): K: int, d_z: int, decay: float = 0.99, - epsilon: float = 1e-5 + epsilon: float = 1e-5, + dead_threshold: float = 2.0 ): super().__init__() self.K = K self.d_z = d_z self.decay = decay self.epsilon = epsilon + self.dead_threshold = dead_threshold self.codebook = nn.Embedding(K, d_z) nn.init.uniform_(self.codebook.weight, -1.0 / K, 1.0 / K) @@ -30,7 +32,7 @@ class VectorQuantizer(nn.Module): def codebook_stats( self, indices: torch.Tensor - ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: flat_indices = indices.reshape(-1) one_hot = F.one_hot(flat_indices, num_classes=self.K).float() avg_probs = one_hot.mean(dim=0) @@ -39,7 +41,9 @@ class VectorQuantizer(nn.Module): ) usage_rate = (avg_probs > 0).float().mean() usage_count = one_hot.sum(dim=0) - return perplexity, usage_rate, usage_count + # 负熵作为惩罚项:p*log(p),使用越均匀值越负,加到 loss 中鼓励多样使用 + entropy_loss = (avg_probs * torch.log(avg_probs.clamp_min(1e-10))).sum() + return perplexity, usage_rate, usage_count, entropy_loss def ema_update(self, z_flat: torch.Tensor, flat_indices: torch.Tensor): one_hot = F.one_hot(flat_indices, num_classes=self.K).type_as(z_flat) @@ -63,9 +67,19 @@ class VectorQuantizer(nn.Module): normalized_weight = self.ema_weight / normalized_cluster_size.unsqueeze(1) self.codebook.weight.data.copy_(normalized_weight) + # 死码重启:长期未被使用的码字,从当前 batch 中随机选取 z_e 重新初始化 + dead = self.ema_cluster_size < self.dead_threshold + if dead.any(): + n_dead = int(dead.sum()) + src_idx = torch.randint(0, z_flat.size(0), (n_dead,), device=z_flat.device) + src = z_flat[src_idx] + self.codebook.weight.data[dead] = src + self.ema_weight[dead] = src + self.ema_cluster_size[dead] = 1.0 + def forward( self, z_e: torch.Tensor - ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: # z_e: [B, L, d_z] B, L, d_z = z_e.shape @@ -98,8 +112,8 @@ class VectorQuantizer(nn.Module): self.ema_update(z_flat.detach(), flat_indices.detach()) indices = flat_indices.reshape(B, L) - perplexity, usage_rate, usage_count = self.codebook_stats(indices) - return z_q_st, indices, commit_loss, perplexity, usage_count + perplexity, usage_rate, usage_count, entropy_loss = self.codebook_stats(indices) + return z_q_st, indices, commit_loss, perplexity, usage_count, entropy_loss def sample(self, B: int, L: int, device: torch.device) -> torch.Tensor: indices = torch.randint(0, self.K, (B, L), device=device) diff --git a/shared/image.py b/shared/image.py index 8830a09..1251dcb 100644 --- a/shared/image.py +++ b/shared/image.py @@ -1,4 +1,6 @@ +import cv2 import numpy as np +import torch def blend_alpha(bg, fg, alpha): """ 使用 alpha 通道混合前景图块和背景图 """ @@ -37,4 +39,28 @@ def matrix_to_image_cv(map_matrix, tile_set, tile_size=32): canvas[y:y+tile_size, x:x+tile_size], tile_rgb, alpha ) - return canvas \ No newline at end of file + return canvas + +def annotate(img: np.ndarray, text: str, y: int = 14) -> np.ndarray: + # 在图片左上角叠加文字标注(黑色描边 + 白色填充,确保任意背景下可读) + img = img.copy() + cv2.putText(img, text, (2, y), cv2.FONT_HERSHEY_SIMPLEX, 0.4, (0, 0, 0), 2) + cv2.putText(img, text, (2, y), cv2.FONT_HERSHEY_SIMPLEX, 0.4, (255, 255, 255), 1) + return img + +def annotate_labels( + img: np.ndarray, + struct: torch.Tensor, + target_density: torch.Tensor +) -> np.ndarray: + # 三行标注:第一行结构标签,后两行显示五维目标密度 + s = struct.tolist() + d = target_density.tolist() + line1 = f"sym:{s[0]} outer:{s[1]}" + line2 = f"wall:{d[0]:.2f} door:{d[1]:.2f}" + line3 = f"enemy:{d[2]:.2f} ent:{d[3]:.2f} res:{d[4]:.2f}" + img = img.copy() + for text, y in [(line1, 12), (line2, 24), (line3, 36)]: + cv2.putText(img, text, (2, y), cv2.FONT_HERSHEY_SIMPLEX, 0.35, (0, 0, 0), 2) + cv2.putText(img, text, (2, y), cv2.FONT_HERSHEY_SIMPLEX, 0.35, (255, 255, 255), 1) + return img diff --git a/test_refactor.py b/test_refactor.py new file mode 100644 index 0000000..a9357af --- /dev/null +++ b/test_refactor.py @@ -0,0 +1,336 @@ +import os +import sys +import torch +import numpy as np +from datetime import datetime + +os.environ["CUDA_VISIBLE_DEVICES"] = "" +os.makedirs("result/test", exist_ok=True) + +device = torch.device("cpu") + +def log(msg): + print(f" [{datetime.now().strftime('%H:%M:%S')}] {msg}") + +def sep(title): + print(f"\n{'='*60}\n {title}\n{'='*60}") + +sep("1. 模型初始化") + +from ginka.model import ( + SeperatedModels, NUM_CLASSES, MASK_TOKEN, MAP_W, MAP_H, + VQ_L1, VQ_L2, VQ_L3 +) +from ginka.utils import MAP_SIZE + +models = SeperatedModels(device) +log(f"SeperatedModels 创建成功") + +total = sum( + p.numel() + for m in [models.vq1, models.vq2, models.vq3, + models.mg1, models.mg2, models.mg3, + models.quantizer1, models.quantizer2, models.quantizer3] + for p in m.parameters() +) + models.latent_mask_embedding.numel() +log(f"总参数量: {total:,}") + +sep("2. VQ-VAE 编码器前向") + +dummy_map = torch.randint(0, 7, (1, MAP_SIZE), dtype=torch.long) + +with torch.no_grad(): + z_e1 = models.vq1(dummy_map) + z_e2 = models.vq2(dummy_map) + z_e3 = models.vq3(dummy_map) + +log(f"z_e1 shape: {z_e1.shape} (期望 [1, {VQ_L1}, 64])") +log(f"z_e2 shape: {z_e2.shape} (期望 [1, {VQ_L2}, 64])") +log(f"z_e3 shape: {z_e3.shape} (期望 [1, {VQ_L3}, 64])") + +assert z_e1.shape == (1, VQ_L1, 64), f"z_e1 shape mismatch" +assert z_e2.shape == (1, VQ_L2, 64), f"z_e2 shape mismatch" +assert z_e3.shape == (1, VQ_L3, 64), f"z_e3 shape mismatch" +log("VQ-VAE 编码器 [OK]") + +sep("3. VectorQuantizer 量化") + +models.quantizer1.eval() +models.quantizer2.eval() +models.quantizer3.eval() + +with torch.no_grad(): + z_q1, idx1, cl1, ppl1, hits1, ent1 = models.quantizer1(z_e1) + z_q2, idx2, cl2, ppl2, hits2, ent2 = models.quantizer2(z_e2) + z_q3, idx3, cl3, ppl3, hits3, ent3 = models.quantizer3(z_e3) + +log(f"z_q1 shape: {z_q1.shape}") +log(f"indices1 shape: {idx1.shape} max={idx1.max().item()} < K={models.quantizer1.K}") +log(f"z_q2 shape: {z_q2.shape}") +log(f"z_q3 shape: {z_q3.shape}") +log(f"commit_loss: {cl1.item():.4f}, {cl2.item():.4f}, {cl3.item():.4f}") +log(f"entropy_loss: {ent1.item():.4f}, {ent2.item():.4f}, {ent3.item():.4f}") + +assert z_q1.shape == z_e1.shape +assert z_q2.shape == z_e2.shape +assert z_q3.shape == z_e3.shape +log("VectorQuantizer 量化 [OK]") + +# 测试采样 +models.quantizer1.eval() +with torch.no_grad(): + z_sampled = models.quantizer1.sample(1, VQ_L1, device) +log(f"sample 采样 shape: {z_sampled.shape} (期望 [1, {VQ_L1}, 64])") +assert z_sampled.shape == (1, VQ_L1, 64) +log("codebook 采样 [OK]") + +sep("4. MaskGIT 前向") + +dummy_struct = torch.tensor([[3, 1]], dtype=torch.long) +dummy_remain = torch.tensor([[0.2, 0.1, 0.3, 0.1, 0.3]], dtype=torch.float) + +models.mg1.eval() +models.mg2.eval() +models.mg3.eval() + +with torch.no_grad(): + logits1 = models.mg1(dummy_map, z_q1, dummy_struct, dummy_remain) + logits2 = models.mg2(dummy_map, z_q2, dummy_struct, dummy_remain) + logits3 = models.mg3(dummy_map, z_q3, dummy_struct, dummy_remain) + +log(f"mg1 logits shape: {logits1.shape} (期望 [1, {MAP_SIZE}, {NUM_CLASSES}])") +log(f"mg2 logits shape: {logits2.shape}") +log(f"mg3 logits shape: {logits3.shape}") + +assert logits1.shape == (1, MAP_SIZE, NUM_CLASSES) +assert logits2.shape == (1, MAP_SIZE, NUM_CLASSES) +assert logits3.shape == (1, MAP_SIZE, NUM_CLASSES) +log("MaskGIT 前向 [OK]") + +sep("5. 数据集加载") + +from ginka.dataset import GinkaSeperatedDataset + +ds_train = GinkaSeperatedDataset("ginka-dataset.json", subset_weights=(0.5, 0.3, 0.2)) +log(f"训练集大小: {len(ds_train)} 条") +log(f"密度统计: wall [{ds_train.density_stats['wall_min_density']:.3f}, {ds_train.density_stats['wall_max_density']:.3f}]") + +sample = ds_train[0] +log(f"input_stage1 shape: {sample['input_stage1'].shape}") +log(f"target_stage1 shape: {sample['target_stage1'].shape}") +log(f"encoder_stage1 shape: {sample['encoder_stage1'].shape}") +log(f"struct_inject: {sample['struct_inject'].tolist()}") +log(f"target_density: {sample['target_density'].tolist()}") + +for key in ["input_stage1", "input_stage2", "input_stage3", + "target_stage1", "target_stage2", "target_stage3", + "encoder_stage1", "encoder_stage2", "encoder_stage3"]: + assert sample[key].shape == (MAP_H, MAP_W), f"{key} shape wrong: {sample[key].shape}" + +log(f"MASK_TOKEN count in inp1: {(sample['input_stage1'] == MASK_TOKEN).sum().item()}") +log(f"MASK_TOKEN count in inp2: {(sample['input_stage2'] == MASK_TOKEN).sum().item()}") +log(f"MASK_TOKEN count in inp3: {(sample['input_stage3'] == MASK_TOKEN).sum().item()}") +log("数据集加载 [OK]") + +random_sample = ds_train.random_sample_map() +log(f"random_sample keys: {list(random_sample.keys())}") +log(f"raw_map shape: {random_sample['raw_map'].shape}") +log("随机样本采样 [OK]") + +sep("6. compute_remaining") + +from ginka.utils import compute_remaining + +inp_t = sample["input_stage1"].reshape(1, MAP_SIZE) +td_t = sample["target_density"].unsqueeze(0) +r1 = compute_remaining(inp_t, td_t, 1) +r2 = compute_remaining(inp_t, td_t, 2) +r3 = compute_remaining(inp_t, td_t, 3) + +log(f"remain (stage1) shape: {r1.shape} (期望 [1, 5])") +log(f"remain (stage2) shape: {r2.shape}") +log(f"remain (stage3) shape: {r3.shape}") +assert r1.shape == (1, 5) +log(f"stage1 remain values: {r1[0].tolist()}") +log("compute_remaining [OK]") + +sep("7. compute_adjacency_mask") + +from ginka.utils import compute_adjacency_mask + +wall_map = torch.full((1, MAP_SIZE), 0, dtype=torch.long) +wall_map[0, 50] = 1 +adj = compute_adjacency_mask(wall_map) +log(f"adjacency mask shape: {adj.shape}") +log(f"adjacent count: {adj.sum().item()} (期望 4)") + +wall_map_b = torch.stack([wall_map[0], wall_map[0]], dim=0) +adj_b = compute_adjacency_mask(wall_map_b) +log(f"batched adj shape: {adj_b.shape}") +log(f"batched adj[0] sum: {adj_b[0].sum().item()}, adj[1] sum: {adj_b[1].sum().item()}") +log("compute_adjacency_mask [OK]") + +sep("8. wall_growth_sample") + +from ginka.sample import wall_growth_sample + +models.mg1.eval() +with torch.no_grad(): + inp_seed = torch.full((1, MAP_SIZE), 0, dtype=torch.long) + seeds = torch.randperm(MAP_SIZE)[:5] + inp_seed[0, seeds] = 1 + + struct_t = sample["struct_inject"].unsqueeze(0) + td_t = sample["target_density"].unsqueeze(0) + + walls = wall_growth_sample( + models.mg1, inp_seed, z_q1, struct_t, td_t, max_steps=3 + ) + +log(f"wall_growth output shape: {walls.shape} (期望 [1, {MAP_H}, {MAP_W}])") +assert walls.shape == (1, MAP_H, MAP_W) +log(f"wall count: {(walls == 1).sum()}") +log("wall_growth_sample [OK]") + +sep("9. maskgit_sample") + +from ginka.sample import maskgit_sample + +models.mg2.eval() +with torch.no_grad(): + inp2 = torch.tensor(walls.reshape(1, MAP_SIZE), dtype=torch.long) + inp2[inp2 == 0] = MASK_TOKEN + + pred2 = maskgit_sample( + models.mg2, inp2, z_q2, struct_t, td_t, + stage=2, steps=3, target_tiles=[2, 4, 5, 6] + ) + +log(f"maskgit_sample (stage2) output shape: {pred2.shape}") +assert pred2.shape == (1, MAP_H, MAP_W) + +models.mg3.eval() +with torch.no_grad(): + merged12 = walls.copy() + merged12[pred2 != 0] = pred2[pred2 != 0] + inp3 = torch.tensor(merged12.reshape(1, MAP_SIZE), dtype=torch.long) + inp3[inp3 == 0] = MASK_TOKEN + + pred3 = maskgit_sample( + models.mg3, inp3, z_q3, struct_t, td_t, + stage=3, steps=3, target_tiles=[3] + ) + +log(f"maskgit_sample (stage3) output shape: {pred3.shape}") +assert pred3.shape == (1, MAP_H, MAP_W) +log("maskgit_sample [OK]") + +sep("10. full_generate(三阶段级联)") + +from ginka.sample import full_generate + +models.mg1.eval() +models.mg2.eval() +models.mg3.eval() + +with torch.no_grad(): + inp = torch.full((1, MAP_SIZE), 0, dtype=torch.long) + seed_idx = torch.randperm(MAP_SIZE)[:8] + inp[0, seed_idx] = 1 + + pred1, merged12, merged123 = full_generate( + inp, z_q1, z_q2, z_q3, + struct_t, td_t, models, steps=3 + ) + +log(f"pred1 shape: {pred1.shape} (期望 [1, {MAP_H}, {MAP_W}])") +log(f"merged12 shape: {merged12.shape}") +log(f"merged123 shape: {merged123.shape}") +assert pred1.shape == (1, MAP_H, MAP_W) +assert merged12.shape == (1, MAP_H, MAP_W) +assert merged123.shape == (1, MAP_H, MAP_W) +log("full_generate [OK]") + +sep("11. 交叉熵 Loss(仅掩码位置)") + +import torch.nn.functional as F + +def cross_entropy_loss_test(logits, target, mask): + loss = F.cross_entropy(logits.permute(0, 2, 1), target, reduction='none') + masked = loss[mask] + if masked.numel() == 0: + return torch.tensor(0.0, requires_grad=True) + return masked.mean() + +mask = (sample["input_stage1"].reshape(1, MAP_SIZE) == MASK_TOKEN) +target = sample["target_stage1"].reshape(1, MAP_SIZE) + +loss_test = cross_entropy_loss_test(logits1, target, mask) +log(f"masked CE loss: {loss_test.item():.4f}") +log(f"mask 中 True 的数量: {mask.sum().item()} / {MAP_SIZE}") + +# 对比:全量 loss vs 掩码 loss +full_loss = F.cross_entropy(logits1.permute(0, 2, 1), target) +log(f"全量 CE loss: {full_loss.item():.4f} vs 掩码 CE loss: {loss_test.item():.4f}") +log("CE Loss [OK]") + +sep("12. 可视化生成") + +from ginka.dataset import compute_symmetry +from shared.image import matrix_to_image_cv +import cv2 + +tile_dict = {} +for f in os.listdir("tiles"): + name = os.path.splitext(f)[0] + img = cv2.imread(f"tiles/{f}", cv2.IMREAD_UNCHANGED) + if img is not None: + tile_dict[name] = img + +TILE_SIZE = 32 + +def save_map(mat, path, label=""): + if isinstance(mat, torch.Tensor): + mat = mat.cpu().numpy() + if mat.ndim == 3: + mat = mat[0] + img = matrix_to_image_cv(mat, tile_dict, TILE_SIZE) + if label: + img = cv2.copyMakeBorder(img, 0, 24, 0, 0, cv2.BORDER_CONSTANT, value=(255, 255, 255)) + cv2.putText(img, label, (4, TILE_SIZE * MAP_H + 18), + cv2.FONT_HERSHEY_SIMPLEX, 0.5, (0, 0, 0), 1) + cv2.imwrite(path, img) + +save_map(pred1, "result/test/pred1_stage1.png", "Stage1: walls") +save_map(merged12, "result/test/pred2_merged12.png", "Stage1+2: +doors/monsters/entrances") +save_map(merged123, "result/test/pred3_merged123.png", "Full: +resources") + +raw = random_sample["raw_map"].cpu().numpy().reshape(MAP_H, MAP_W) +save_map(raw, "result/test/raw_map.png", "Original ground truth") + +inp1_img = sample["input_stage1"].cpu().numpy().reshape(MAP_H, MAP_W) +save_map(inp1_img, "result/test/inp1.png", "Stage1 input (masked)") + +log("可视化输出到 result/test/ 目录:") +log(" pred1_stage1.png — 仅墙壁") +log(" pred2_merged12.png — 墙壁+功能元素") +log(" pred3_merged123.png — 完整地图") +log(" raw_map.png — 原始真实地图") +log(" inp1.png — Stage1 输入(掩码后)") + +sep("测试结果汇总") + +print("\n 所有 12 项测试通过 [OK]\n") +print(f" - 模型初始化与参数统计") +print(f" - VQ-VAE 编码器 (3 个阶段)") +print(f" - VectorQuantizer 量化与采样") +print(f" - MaskGIT 前向 (3 个阶段)") +print(f" - 数据集加载与掩码策略") +print(f" - compute_remaining") +print(f" - compute_adjacency_mask (含 batch)") +print(f" - wall_growth_sample") +print(f" - maskgit_sample (stage2 + stage3)") +print(f" - full_generate (三阶段级联)") +print(f" - 掩码位置 CE Loss") +print(f" - 可视化输出") +print() diff --git a/需要重训的问题.md b/需要重训的问题.md new file mode 100644 index 0000000..a586b37 --- /dev/null +++ b/需要重训的问题.md @@ -0,0 +1,445 @@ +# 需要重训才能解决的问题 + +`01-sampler-fix.patch` 只改推理采样。下面这些是**训练侧**的问题,改完必须重跑训练才能生效。 + +**证据来源**:checkpoint `sep-320.pth`(epoch 320)、数据集 6731 张、服务器 RTX 5090 实测。 +**代码基准**:GitHub `103b02c`(服务器上的本地版本另有 `LAMBDA_ADJ`/`LAMBDA_PATCH` 归零的未提交改动)。 + +按「值不值得做」排序。 + +--- + +## A. 距离场整条通路是死的 —— 1069 万参数在空转 + +### 现在的表现 + +`dist_encoder` + `dist_quantizer` 从头到尾没起过任何作用。`z_dist` 对所有输入都是**同一个常数向量**。 + +### 实测证据 + +**① 码本只剩 1 个码字** + +``` +dist_quantizer K=16 + perplexity = 1.00 / 16 利用率 6.2% + ema_cluster_size 里只有 14 号非零,占比 100.0% + 16 个码字的范数: [0,0,0,0,0,0,0,0,0,0,0,0,0,0, 7.696, 0] + ↑ 只有这一个活着 +``` + +其余 15 个码字**范数是 0** —— 被 EMA 拉到了原点。 + +**② 编码器输出本身就没有区分度** + +拿 500 张**互不相同**的地图(496/500 距离场各不相同)过一遍编码器: + +| | `dist_encoder` | `vq1`(对照) | 差距 | +|---|---|---|---| +| z_e 每维在样本间的标准差 | **0.007046** | 0.656193 | **93 倍** | +| 不同样本 z_e 之间的距离(中位) | **0.01475** | 2.87652 | **195 倍** | +| z_e 自身的范数 | 7.6967(标准差 0.0002) | — | — | + +500 张不同地图编出来的 z_e 几乎完全相同 —— 相对变化只有 **0.2%**。 + +**③ commit_loss 已经趋零** + +``` +commit_loss = 0.000059 +z_e 到最近码字的距离 : 0.0567 +z_e 到次近码字的距离 : 7.6967 ← 其余码字都在原点,永远选不中 +``` + +### 根本原因(两层) + +**第一层:距离场这个特征在这个数据集上本来就没信息。** + +``` +距离 = 0 51.78% (墙) +距离 = 1 44.78% (紧挨着墙的空地) +距离 = 2 3.37% +距离 = 3 0.07% +距离 = 4 0.00% +``` + +**96.6% 的格子距离值不是 0 就是 1。** 每张图的最大距离均值只有 **2.01**,不同取值种类均值 **3.01**。`DIST_VOCAB` 设了 13(0~12),实际只用到 0~4。 + +原因是真实地图**墙密度 51.6%** —— 墙无处不在,任何空地几乎都紧挨着墙。距离场退化成了"墙 / 非墙"的二值图,跟原图的墙层几乎是同一个东西,提供不了"离墙多远"的层次信息。 + +> 这个特征的设计假设(地图有大片开阔区域)在这个数据集上不成立。 + +**第二层:EMA 死锁,不可逆。** + +``` +所有 z_e 落到同一个码字 + ↓ +其余码字的 ema_cluster_size 每步 × 0.99 衰减 + ↓ +normalized_weight = ema_weight / cluster_size → 0,码字被拉到原点 + ↓ +原点离任何 z_e 都是 7.70,永远不可能再被选中 + ↓ +commit_loss ≈ 0,梯度趋零,编码器停止学习 +``` + +`commit_loss` 这一项本身就在**奖励塌缩** —— 塌缩之后 loss 完美,系统再也不会离开这个状态。 + +### 空转的规模 + +| 模块 | 参数量 | 体积 | +|---|---|---| +| `dist_encoder` | **10,689,600** | 40.8 MiB | +| `dist_quantizer` | 2,064 | — | + +对比 `mg1` 本身是 3492 万。**每个 batch 白跑一次 6 层 Transformer 的前向 + 反向。** + +而且 commit `3cd38b8 feat: 距离场输入` 之后又把它从 `d_model=256 / 3层` 加到 `384 / 6层`、`L_DIST` 从 4 加到 8 —— **加的全是空转的量**。 + +### 改什么 + +**方案 1(最省事):整条删掉。** + +`z_dist` 是常数,经过 `z_dist_proj` 再进 `cond_proj`(线性层)—— **常数输入到线性层等价于一个偏置项**。所以删掉之后**模型行为完全不变**,纯粹省资源。 + +需要动的地方:`build_model` 里去掉 `dist_encoder`/`dist_quantizer`,`GinkaMaskGIT` 的 `z_dist_len` 设 0 或去掉这个输入,`dataset.py` 不再算 `distance_field`。 + +**方案 2:换一个真正有区分度的空间特征。** + +数据集里现成就有没被用的(见问题 G):`roomCount`、`highDegBranchCount`、`val`(16 维,6098 种取值)。 + +**方案 3:保留距离场,但必须同时做问题 B 的防塌缩。** 不做防塌缩,换什么特征都可能重蹈覆辙。 + +### 为什么要重训 + +删掉/替换条件通道会改变模型的输入结构,权重不兼容。 + +### 预期效果 + +- **方案 1**:生成质量**完全不变**(可实测验证),训练每步省一次 1069 万参数的前向反向。**这一条把握度 100%,因为 z_dist 已经证明是常数。** +- 方案 2/3:能不能改善生成质量**未知**,取决于新特征是否真的携带 stage1 需要的信息。 + +--- + +## B. 码本 EMA 死锁正在慢慢吃掉 `quantizer1` + +### 现在的表现 + +`quantizer1` 是 stage1 唯一活着的条件通道(`z_dist` 已死),但它只用了 **37.2%** 的码本容量。 + +### 实测证据 + +| 码本 | 负责 | perplexity | 利用率 | 存活码字 | +|---|---|---|---|---| +| **`quantizer1`** | **stage1 墙** | **11.89 / 32** | **37.2%** | **16 / 32** | +| `quantizer2` | stage2 门/怪/入口 | 15.38 / 16 | **96.1%** | 16 / 16 | +| `quantizer3` | stage3 资源 | 5.98 / 16 | 37.3% | 6 / 16 | +| `dist_quantizer` | 距离场 | 1.00 / 16 | 6.2% | 1 / 16 | + +`quantizer1` 死掉的 16 个码字**范数为 0**,和 A 是同一个 EMA 死锁机制,只是还没吃完。 + +有效信息量:`24 码字 × log₂(11.89) ≈ 86 bit`,设计上限是 `24 × log₂32 = 120 bit`,**用掉 71%**。 + +### 这条和症状的对应关系 + +| 阶段 | 条件通道状态 | 表现 | +|---|---|---| +| **stage1 墙** | z1 用掉 71% + **z_dist 完全死掉** | **差** | +| stage2 门/怪/入口 | z2 **96.1% 健康** | 好 | +| stage3 资源 | z3 37.3% | 好(资源本来不需要结构) | + +**唯一条件通道健康的阶段,正好是效果好的那个。** + +### 根本原因 + +`VectorQuantizer` 只有 EMA 更新,**没有任何防塌缩机制**: + +```python +VQ_GAMMA = 0.0 # entropy loss 权重(当前未启用) ← train_seperated.py:45 +``` + +而 `VQDecodeHead`(`vqvae/model.py:33`)—— 那个用 z 还原原图、逼 z 必须编码地图内容的重建头 —— **在 `train_seperated.py` 里从头到尾没被 import 过**。 + +被删掉的 `train_full.sh` 里"阶段 0 重建预训练 + 阶段 1 冻结 VQ 热身"正是为了防这个,那套流程随脚本一起失效了。 + +### 改什么 + +**① 死码重启(dead-code restart)** —— 最直接。在 `VectorQuantizer.ema_update` 里加: + +```python +# cluster_size 长期低于阈值的码字,重新初始化到某个真实 z_e 附近 +dead = self.ema_cluster_size < DEAD_THRESHOLD # 比如 0.5 +if dead.any(): + src = z_flat[torch.randint(0, z_flat.size(0), (int(dead.sum()),))] + self.codebook.weight.data[dead] = src + self.ema_weight[dead] = src + self.ema_cluster_size[dead] = 1.0 +``` + +**② `VQ_GAMMA` 从 0 调回 0.1**,启用熵损失,鼓励码字使用均匀。 + +**③ 或者恢复 `train_full.sh` 的两段式课程**:先用 `VQDecodeHead` 做重建预训练,再冻结 VQ 训 MaskGIT。 + +### 为什么要重训 + +码本状态是训练出来的,已经死掉的码字在现有 checkpoint 上救不回来(范数为 0,离任何 z_e 都是 7.7)。 + +### 预期效果 + +`quantizer1` 利用率 37% → 期望 70%+,z1 的信息量从 86 bit 往 120 bit 靠。 + +**但对生成质量的实际改善幅度未知** —— 更多的 z 容量能不能变成更好的墙,取决于问题 E(条件注入方式)是不是瓶颈。**这一条我没有实验依据。** + +--- + +## C. 码本监控缺了一半 + +### 现在的表现 + +`dist_quantizer` 塌缩了 320 个 epoch,训练日志一个字都没提。 + +### 实测证据 + +``` +E: 334 | ... | PPL: 31.1291 | Usage: 0.6094 (39/64) +``` + +`64 = VQ_K1 + VQ_K2 + VQ_K3 = 32 + 16 + 16`。**`dist_quantizer` 根本不在统计里。** + +而且合并统计会掩盖问题:`PPL 31.1/64` 看着健康,拆开是 37.2% / 96.1% / 37.3%。 + +### 改什么 + +**① 把 `dist_quantizer` 的命中也统计进去。** + +**② 分开打印每个码本,而不是合并。** 合并的 PPL 没有诊断价值。 + +**③ 加一条离线诊断脚本**,直接从 `.pth` 读 `ema_cluster_size` 和 `codebook.weight`,不用跑前向: + +```python +ck = torch.load("result/seperated/sep-XXX.pth", map_location="cpu") +for n in ["quantizer1","quantizer2","quantizer3","dist_quantizer"]: + cs = ck[n]["ema_cluster_size"]; p = cs/cs.sum() + ppl = float(torch.exp(-(p*torch.log(p.clamp_min(1e-10))).sum())) + print(f"{n} ppl={ppl:.2f}/{cs.numel()} alive={int((cs>1).sum())}") +``` + +(`11-码本权重.pth` 就是这么导出来的,46 KB,可以直接跑这段验证。) + +### 为什么要重训 + +监控本身不用重训 —— **这一条现在就能加**。但它是发现前两个问题的前提。 + +--- + +## D. CE 在全部 169 格上算,一半是「抄写」 + +### 现在的表现 + +日志里的 CE 数值有很大一部分是"把可见 token 原样抄出来"的准确率,不反映生成能力。 + +### 实测证据 + +```python +def cross_entropy_loss(logits, target): + return F.cross_entropy(logits.permute(0, 2, 1), target) # 全部 169 格 +``` + +没有 `ignore_index`,没有按掩码位置筛选。标准 MaskGIT 只在**被掩位置**算 loss。 + +按 `dataset.py` 的掩码分布模拟 2 万次:**mg1 的输入平均有 51.6% 的格子是可见的**。也就是说约一半的 loss 是"看见 `1` 就输出 `1`"的恒等映射任务。 + +**注水比例在三个 stage 之间还不一样** —— mg1 可见格最多,注水最严重。 + +### 改什么 + +```python +def cross_entropy_loss(logits, target, inp=None): + if inp is None: + return F.cross_entropy(logits.permute(0, 2, 1), target) + t = target.clone() + t[inp != MASK_TOKEN] = -100 # 只在掩码位算 + return F.cross_entropy(logits.permute(0, 2, 1), t, ignore_index=-100) +``` + +`adjacency_loss` / `patch_loss` 同理(虽然它们的权重现在是 0)。 + +### 为什么要重训 + +改的是训练目标。 + +### 预期效果 + +梯度不再被恒等映射任务稀释,**而且日志里的 CE 终于能反映真实的生成能力**。 + +对生成质量的改善幅度**未知**。但即使质量不变,这一条也值得做 —— **它让后续所有实验的指标变得可信**。现在没有任何一个数字能反映生成质量。 + +--- + +## E. loss 里没有任何连通性约束 + +### 现在的表现 + +真实数据里 **6731 张地图 100% 地板连通**(无一例外)—— 这是最强的规律。而模型从来没被要求学它。 + +修复采样器之后单次生成的可用率是 **75%**,剩下 25% 靠拒绝采样兜。 + +### 实测证据 + +服务器上的实际配置: + +```python +LAMBDA_ADJ1/2/3 = 0 +LAMBDA_PATCH1/2/3 = 0 +loss = CE1 + CE2 + CE3 + commit +``` + +**只剩逐格交叉熵。** 而 CE 对"墙挪一格把房间封死"和"墙挪一格无所谓"的惩罚完全一样。 + +(顺带:原来的 `patch_loss` 数值上就等于 CE —— 归一化高斯核卷积后再求全局均值 ≈ 直接求均值。我在真实训练中实测过:`PAT 0.2354` vs `CE 0.2351`,**差 0.13%**。所以那个"邻域结构损失"从来没引入过任何结构信息,把权重设成 0 是对的。) + +### 改什么 + +可微的连通性损失很难做。三个方向: + +**① 软连通性代理**:对 floor 概率图做若干轮软膨胀,统计最大连通块的期望占比,作为 loss。 + +**② 保持现状,靠推理侧兜底** —— 就是 `01-sampler-fix.patch` 里的拒绝采样 + `repair_connectivity`。**已经能做到 100%,而且不用重训。** + +**③ 用数据集里现成的拓扑字段做监督**(见问题 G)。 + +### 为什么要重训 + +①③ 改训练目标。②不用重训。 + +### 预期效果 + +**建议先不做。** 推理侧兜底已经解决了这个问题(可用率 100%),投入产出比不高。等其他问题解决之后再看单次可用率能不能自己提上去。 + +--- + +## F. 条件通过全局 AdaLN 注入,传不了空间信息(推测,无实验证据) + +### 现在的表现 + +从零生成时,169 个位置的输入完全相同(全是 `MASK`),只有位置编码不同。要区分出哪一格是墙,全靠一个全局向量。 + +### 代码位置 + +`ginka/maskGIT/model.py:80-81`: + +```python +cond_seq = torch.cat([z_proj, zd_proj, e_struct, e_remain], dim=1) +c = self.cond_proj(cond_seq.reshape(cond_seq.size(0), -1)) # [B, d_model] +``` + +所有条件 token 被 **flatten 成一个 d_model 维向量**,再通过 AdaLN 注入每一层。 + +AdaLN 只能做**全局的 scale / shift**。也就是说 z1 那 86 bit 只能以"全局风格"的形式影响输出,**没有任何机制说"第 42 格应该是墙"**。 + +### 改什么 + +改成 **cross-attention**:z 作为 memory,地图 token 作为 query。这样每个位置能各自去 z 里取自己需要的信息。 + +设计文档 `vqvae-maskgit-design.md` 里写的就是 cross-attention,代码里实现成了 flatten + AdaLN。 + +### 为什么要重训 + +改模型结构,权重完全不兼容。 + +### 把握度 + +**这一条是推测。** 有理论依据(全局向量确实无法编码空间位置),但**我没有实验证据证明它是主要瓶颈**,也不知道改完能改善多少。 + +**建议放到最后** —— 工程量最大(改结构 + 完整重训),而且前面几条做完之后瓶颈可能已经转移。 + +--- + +## G. 数据集里三个字段完全没用上 + +### 实测证据 + +`ginka-dataset.json` 每条记录的字段: + +``` +map, size, val, symmetry, outerWall, roomCount, highDegBranchCount +``` + +`dataset.py` 只读了 `map` 和 `outerWall`。 + +| 字段 | 内容 | 取值种类 | +|---|---|---| +| `val` | 16 维浮点向量 | **6098 种**(6731 张里) | +| `roomCount` | 房间数 | 17 种 | +| `highDegBranchCount` | 高度分支数 | 35 种 | +| `symmetry` | 三位对称性 | 5 种(`dataset.py` 自己重算了一遍) | + +**这些是数据侧已经算好的拓扑信息 —— 恰恰是"墙该怎么排"的直接监督信号,现在全被丢掉了。** + +### 改什么 + +把 `roomCount` / `highDegBranchCount` 加进 `struct_inject`(现在只有对称性 3 bit + outerWall 1 bit),或者作为额外的条件 token。`val` 那 16 维如果是拓扑特征向量,可以直接当条件。 + +**这也可能是问题 A 的替代品** —— 比距离场有信息量得多(6098 种取值 vs 距离场的 3 种)。 + +### 为什么要重训 + +改条件输入。 + +### 预期效果 + +**未知,但这是所有方案里信息量最明确的一个。** 值得优先于 F 尝试。 + +--- + +## H. 两个小问题 + +### H-1 子集 3 是空操作,20% 的数据是重复的 + +```python +def apply_subset3(self, raw): + out = self.apply_subset2(raw) + out[0][out[0] == self.ENTRANCE] = self.MASK_ID # out[0] 是 inp1 + return out +``` + +但 `inp1` 来自 `target1`,而 `create_degreaded`(`dataset.py:212`)已经把 `ENTRANCE(5)` 降级成 `0` 了: + +```python +self.degrade_tile(target1, [DOOR, SPECIAL_DOOR, RESOURCE, MONSTER, ENTRANCE]) +``` + +**`inp1` 的取值只可能是 `{0, 1, 7}`,里面根本没有 `5`。** 已验证 `(inp1 == 5).sum() == 0`。 + +**后果**:设计文档里的"子集 D:入口条件生成"**实际不存在**,`SUBSET_WEIGHTS = (0.5, 0.3, 0.2)` 实际等价于 `(0.5, 0.5)`。 + +**改法**:要么补上真正的入口掩码逻辑(在 `target1` 里保留 ENTRANCE),要么删掉这个子集、把权重改成 `(0.5, 0.5)` 让配置反映实际。 + +### H-2 数据集有 5.23% 重复 + +6731 张里唯一的只有 **6379** 张,重复 352 张。 + +不算严重,但去重是免费的。 + +--- + +## 建议顺序 + +| 顺序 | 做什么 | 要重训 | 把握 | +|---|---|---|---| +| **0** | **加码本监控(问题 C)** | ❌ | 确定 —— 一行的事,而且是发现前两个问题的前提 | +| **1** | **删掉距离场通路(问题 A 方案 1)** | ✅(但只为省资源) | **确定** —— 已证明 `z_dist` 是常数,删掉行为不变 | +| **2** | **CE 只在掩码位算(问题 D)** | ✅ | 中 —— 质量改善未知,**但让所有指标变可信** | +| **3** | **加死码重启 + `VQ_GAMMA=0.1`(问题 B)** | ✅ | 中 —— 码本利用率必涨,质量改善未知 | +| **4** | **用上 `roomCount` / `val`(问题 G)** | ✅ | 中 —— 信息量最明确的一条 | +| 5 | 修子集 3(问题 H-1) | ✅ | 低,但顺手 | +| 6 | 条件注入改 cross-attention(问题 F) | ✅ | 未验证,工程量最大,**建议最后** | +| — | 连通性损失(问题 E) | ✅ | **建议不做** —— 推理侧兜底已达 100% | + +**1~3 可以一起改、跑一次训练。** 4 单独一次好做对比。 + +--- + +## 一句话总结 + +**stage1 的两条条件通道,一条(`z_dist`)完全死了、一条(`z1`)只用了 37%,而 stage2 的那条 96% 健康 —— 这和「墙差、怪好」的症状完全对应。** + +采样器修复能把可用率做到 100%,但生成的墙比真实地图碎(墙连通块 8.2 vs 真实 5.7),那部分要靠上面这些。