mirror of
https://github.com/unanmed/ginka-generator.git
synced 2026-08-15 02:22:28 +08:00
refactor: 重构代码结构
This commit is contained in:
parent
103b02cf08
commit
763f6258d2
@ -78,3 +78,7 @@
|
|||||||
|
|
||||||
- 编写验证代码时,优先输出可视化结果(图片文件),使用 `shared/image.py` 中的工具
|
- 编写验证代码时,优先输出可视化结果(图片文件),使用 `shared/image.py` 中的工具
|
||||||
- 验证阶段应对不同条件(不同 z 采样)分别生成图片,便于直观对比模型效果
|
- 验证阶段应对不同条件(不同 z 采样)分别生成图片,便于直观对比模型效果
|
||||||
|
|
||||||
|
## 其他
|
||||||
|
|
||||||
|
`app` 目录下的内容不用管,目前尚在训练阶段,还未到达推理发布阶段。
|
||||||
126
docs/重训问题跟踪.md
Normal file
126
docs/重训问题跟踪.md
Normal file
@ -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 不用重训,随时能做。**
|
||||||
@ -3,7 +3,6 @@ import random
|
|||||||
import torch
|
import torch
|
||||||
import numpy as np
|
import numpy as np
|
||||||
from torch.utils.data import Dataset
|
from torch.utils.data import Dataset
|
||||||
from shared.distance import compute_distance_field
|
|
||||||
|
|
||||||
def rect_mask(ratio: float, map_size: int = 169) -> np.ndarray:
|
def rect_mask(ratio: float, map_size: int = 169) -> np.ndarray:
|
||||||
# 连续矩形分块掩码,反复放置随机矩形直到掩码格数达标
|
# 连续矩形分块掩码,反复放置随机矩形直到掩码格数达标
|
||||||
@ -144,8 +143,6 @@ class GinkaSeperatedDataset(Dataset):
|
|||||||
return enc1, enc2, enc3
|
return enc1, enc2, enc3
|
||||||
|
|
||||||
def pack_sample(self, item: dict, map_np: np.ndarray, out: tuple) -> dict:
|
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 {
|
return {
|
||||||
"input_stage1": torch.LongTensor(out[0]),
|
"input_stage1": torch.LongTensor(out[0]),
|
||||||
"target_stage1": torch.LongTensor(out[1]),
|
"target_stage1": torch.LongTensor(out[1]),
|
||||||
@ -158,7 +155,6 @@ class GinkaSeperatedDataset(Dataset):
|
|||||||
"encoder_stage3": torch.LongTensor(out[8]),
|
"encoder_stage3": torch.LongTensor(out[8]),
|
||||||
"struct_inject": self.build_struct_inject(map_np, item['outerWall']),
|
"struct_inject": self.build_struct_inject(map_np, item['outerWall']),
|
||||||
"target_density": self.build_target_density(item['map']),
|
"target_density": self.build_target_density(item['map']),
|
||||||
"distance_field": torch.LongTensor(dist_field)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
def random_sample_map(self, idx: int | None = None) -> dict:
|
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']),
|
"struct_inject": self.build_struct_inject(map_np, item['outerWall']),
|
||||||
"target_density": self.build_target_density(item['map']),
|
"target_density": self.build_target_density(item['map']),
|
||||||
"raw_map": torch.LongTensor(map_np),
|
"raw_map": torch.LongTensor(map_np),
|
||||||
"distance_field": torch.LongTensor(compute_distance_field(enc1))
|
|
||||||
}
|
}
|
||||||
sample['sample_idx'] = idx
|
sample['sample_idx'] = idx
|
||||||
sample['map_name'] = self.map_names[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
|
return inp1, target1, enc1, inp2, target2, enc2, inp3, target3, enc3
|
||||||
|
|
||||||
def apply_subset3(self, raw: np.ndarray):
|
def apply_subset3(self, raw: np.ndarray):
|
||||||
# 子集 3:在 2 的基础上掩码入口
|
# 子集 3:与子集 2 相同(entry 已在 stage2 全掩码中覆盖)
|
||||||
out = self.apply_subset2(raw)
|
out = self.apply_subset2(raw)
|
||||||
out[0][out[0] == self.ENTRANCE] = self.MASK_ID
|
|
||||||
return out
|
return out
|
||||||
|
|
||||||
def __getitem__(self, idx):
|
def __getitem__(self, idx):
|
||||||
|
|||||||
55
ginka/diagnose_codebook.py
Normal file
55
ginka/diagnose_codebook.py
Normal file
@ -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}: 全部码字有非零范数")
|
||||||
@ -7,13 +7,12 @@ from .maskGIT import Transformer
|
|||||||
# 结构标签词表大小
|
# 结构标签词表大小
|
||||||
SYM_VOCAB = 8 # symmetryH/V/C 三位组合 0-7
|
SYM_VOCAB = 8 # symmetryH/V/C 三位组合 0-7
|
||||||
OUTER_VOCAB = 2 # outerWall 0-1
|
OUTER_VOCAB = 2 # outerWall 0-1
|
||||||
L_DIST = 4 # 距离场码字序列长度
|
|
||||||
|
|
||||||
class GinkaMaskGIT(nn.Module):
|
class GinkaMaskGIT(nn.Module):
|
||||||
def __init__(
|
def __init__(
|
||||||
self, num_classes: int = 16, d_model: int = 192, dim_ff: int = 512,
|
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,
|
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__()
|
super().__init__()
|
||||||
self.map_h = map_h
|
self.map_h = map_h
|
||||||
@ -34,11 +33,8 @@ class GinkaMaskGIT(nn.Module):
|
|||||||
# z 投影:逐 token 线性变换,保持序列结构
|
# z 投影:逐 token 线性变换,保持序列结构
|
||||||
self.z_proj = nn.Linear(d_z, d_z)
|
self.z_proj = nn.Linear(d_z, d_z)
|
||||||
|
|
||||||
# 距离场 z 投影
|
# 条件融合投影:z_seq_len 个 z token + 2 个结构 token + 5 个剩余密度 token
|
||||||
self.z_dist_proj = nn.Linear(d_z, d_z)
|
self.cond_proj = nn.Linear((z_seq_len + 2 + 5) * d_z, d_model)
|
||||||
|
|
||||||
# 条件融合投影: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)
|
|
||||||
|
|
||||||
# 纯 encoder Transformer,条件向量 c 通过 AdaLN 注入每一层
|
# 纯 encoder Transformer,条件向量 c 通过 AdaLN 注入每一层
|
||||||
self.transformer = Transformer(
|
self.transformer = Transformer(
|
||||||
@ -51,13 +47,11 @@ class GinkaMaskGIT(nn.Module):
|
|||||||
self,
|
self,
|
||||||
map: torch.Tensor,
|
map: torch.Tensor,
|
||||||
z: torch.Tensor,
|
z: torch.Tensor,
|
||||||
z_dist: torch.Tensor,
|
|
||||||
struct: torch.Tensor,
|
struct: torch.Tensor,
|
||||||
remain: torch.Tensor
|
remain: torch.Tensor
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
# map: [B, H * W]
|
# map: [B, H * W]
|
||||||
# z: [B, z_seq_len, d_z]
|
# 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)]
|
# struct: [B, 2] — [cond_sym(0-7), cond_outer(0-1)]
|
||||||
# remain: [B, 5] float — [wall, door, monster, entrance, resource] 剩余密度
|
# 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:逐 token 投影,保留序列结构 [B, z_seq_len, d_z]
|
||||||
z_proj = self.z_proj(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
|
# 拼接所有条件 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]
|
c = self.cond_proj(cond_seq.reshape(cond_seq.size(0), -1)) # [B, d_model]
|
||||||
|
|
||||||
# tile embedding + 位置编码
|
# tile embedding + 位置编码
|
||||||
@ -96,7 +87,7 @@ if __name__ == "__main__":
|
|||||||
device = torch.device("cuda") if torch.cuda.is_available() else torch.device("cpu")
|
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]
|
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([
|
struct_input = torch.tensor([
|
||||||
[3, 1],
|
[3, 1],
|
||||||
[0, 0],
|
[0, 0],
|
||||||
@ -122,12 +113,10 @@ if __name__ == "__main__":
|
|||||||
z_seq_len=6
|
z_seq_len=6
|
||||||
).to(device)
|
).to(device)
|
||||||
|
|
||||||
z_dist_input = torch.randn(4, L_DIST, 64).to(device) # [4, L_DIST, 64]
|
|
||||||
|
|
||||||
print_memory(device, "初始化后")
|
print_memory(device, "初始化后")
|
||||||
|
|
||||||
start = time.perf_counter()
|
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()
|
end = time.perf_counter()
|
||||||
|
|
||||||
print_memory(device, "前向传播后")
|
print_memory(device, "前向传播后")
|
||||||
|
|||||||
193
ginka/model.py
Normal file
193
ginka/model.py
Normal file
@ -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)
|
||||||
144
ginka/sample.py
Normal file
144
ginka/sample.py
Normal file
@ -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
|
||||||
File diff suppressed because it is too large
Load Diff
181
ginka/utils.py
181
ginka/utils.py
@ -1,93 +1,122 @@
|
|||||||
import torch
|
import torch
|
||||||
import torch.nn.functional as F
|
|
||||||
import numpy as np
|
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=""):
|
def print_memory(device, tag=""):
|
||||||
if torch.cuda.is_available():
|
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")
|
print(f"{tag} | 当前显存: {torch.cuda.memory_allocated(device) / 1024**2:.2f} MB, 最大显存: {torch.cuda.max_memory_allocated(device) / 1024**2:.2f} MB")
|
||||||
else:
|
else:
|
||||||
print("当前设备不支持 cuda.")
|
print("当前设备不支持 cuda.")
|
||||||
|
|
||||||
def nms_sampling(noise: np.ndarray, k: int, radius=2):
|
def compute_remaining(
|
||||||
# noise: [H, W]
|
current: torch.Tensor,
|
||||||
noise = noise.copy()
|
target_density: torch.Tensor,
|
||||||
points = []
|
stage: int
|
||||||
|
) -> torch.Tensor:
|
||||||
|
remain = torch.zeros(current.size(0), DENSITY_DIM, device=current.device)
|
||||||
|
|
||||||
for _ in range(k):
|
visible_wall = (current == 1).sum(dim=1).float() / MAP_SIZE
|
||||||
idx = np.argmax(noise)
|
visible_door = ((current == 2) | (current == 6)).sum(dim=1).float() / MAP_SIZE
|
||||||
x, y = np.unravel_index(idx, noise.shape)
|
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
|
||||||
|
|
||||||
points.append((x, y))
|
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
|
||||||
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)
|
|
||||||
|
|
||||||
noise[x0:x1, y0:y1] = -np.inf
|
def rect_mask(
|
||||||
|
ratio: float, h_range: tuple[int, int] = (2, 7),
|
||||||
result = np.zeros_like(noise)
|
w_range: tuple[int, int] = (2, 7)
|
||||||
for x, y in points:
|
) -> np.ndarray:
|
||||||
result[y, x] = 1
|
# 纯矩形分块掩码,反复放置随机矩形直到掩码格数达标
|
||||||
|
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
|
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,
|
combined = torch.cat([h.flatten() for h in code_hits], dim=0)
|
||||||
target: torch.Tensor,
|
total_hits = combined.sum()
|
||||||
tile_set: set,
|
if total_hits.item() > 0:
|
||||||
gamma: float = 2.0,
|
probs = combined / total_hits
|
||||||
balance: bool = True,
|
result["combined"] = {
|
||||||
) -> torch.Tensor:
|
"ppl": float(torch.exp(
|
||||||
"""
|
-(probs * torch.log(probs.clamp_min(1e-10))).sum()
|
||||||
通道专属 Focal Loss + 逆频类别权重。
|
).item()),
|
||||||
|
"active": int((combined > 0).sum().item()),
|
||||||
tile_set 内的位置以真实 tile ID 为目标,tile_set 外的位置以 0(空地)为目标,
|
"K": int(combined.numel()),
|
||||||
全部位置均参与损失计算。
|
}
|
||||||
|
else:
|
||||||
balance=True 时,从 batch 内 corrected 标签的频率自动计算逆频权重,
|
result["combined"] = {"ppl": 0.0, "active": 0, "K": int(combined.numel())}
|
||||||
消除空地(0)因被大量 non-tile-set 位置填充而主导梯度的问题。
|
return result
|
||||||
权重公式: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()
|
|
||||||
|
|||||||
@ -149,56 +149,6 @@ class GinkaVQVAE(nn.Module):
|
|||||||
|
|
||||||
return z_e
|
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__":
|
if __name__ == "__main__":
|
||||||
device = torch.device("cuda") if torch.cuda.is_available() else torch.device("cpu")
|
device = torch.device("cuda") if torch.cuda.is_available() else torch.device("cpu")
|
||||||
|
|
||||||
|
|||||||
@ -9,13 +9,15 @@ class VectorQuantizer(nn.Module):
|
|||||||
K: int,
|
K: int,
|
||||||
d_z: int,
|
d_z: int,
|
||||||
decay: float = 0.99,
|
decay: float = 0.99,
|
||||||
epsilon: float = 1e-5
|
epsilon: float = 1e-5,
|
||||||
|
dead_threshold: float = 2.0
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.K = K
|
self.K = K
|
||||||
self.d_z = d_z
|
self.d_z = d_z
|
||||||
self.decay = decay
|
self.decay = decay
|
||||||
self.epsilon = epsilon
|
self.epsilon = epsilon
|
||||||
|
self.dead_threshold = dead_threshold
|
||||||
|
|
||||||
self.codebook = nn.Embedding(K, d_z)
|
self.codebook = nn.Embedding(K, d_z)
|
||||||
nn.init.uniform_(self.codebook.weight, -1.0 / K, 1.0 / K)
|
nn.init.uniform_(self.codebook.weight, -1.0 / K, 1.0 / K)
|
||||||
@ -30,7 +32,7 @@ class VectorQuantizer(nn.Module):
|
|||||||
|
|
||||||
def codebook_stats(
|
def codebook_stats(
|
||||||
self, indices: torch.Tensor
|
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)
|
flat_indices = indices.reshape(-1)
|
||||||
one_hot = F.one_hot(flat_indices, num_classes=self.K).float()
|
one_hot = F.one_hot(flat_indices, num_classes=self.K).float()
|
||||||
avg_probs = one_hot.mean(dim=0)
|
avg_probs = one_hot.mean(dim=0)
|
||||||
@ -39,7 +41,9 @@ class VectorQuantizer(nn.Module):
|
|||||||
)
|
)
|
||||||
usage_rate = (avg_probs > 0).float().mean()
|
usage_rate = (avg_probs > 0).float().mean()
|
||||||
usage_count = one_hot.sum(dim=0)
|
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):
|
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)
|
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)
|
normalized_weight = self.ema_weight / normalized_cluster_size.unsqueeze(1)
|
||||||
self.codebook.weight.data.copy_(normalized_weight)
|
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(
|
def forward(
|
||||||
self, z_e: torch.Tensor
|
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]
|
# z_e: [B, L, d_z]
|
||||||
B, L, d_z = z_e.shape
|
B, L, d_z = z_e.shape
|
||||||
|
|
||||||
@ -98,8 +112,8 @@ class VectorQuantizer(nn.Module):
|
|||||||
self.ema_update(z_flat.detach(), flat_indices.detach())
|
self.ema_update(z_flat.detach(), flat_indices.detach())
|
||||||
|
|
||||||
indices = flat_indices.reshape(B, L)
|
indices = flat_indices.reshape(B, L)
|
||||||
perplexity, usage_rate, usage_count = self.codebook_stats(indices)
|
perplexity, usage_rate, usage_count, entropy_loss = self.codebook_stats(indices)
|
||||||
return z_q_st, indices, commit_loss, perplexity, usage_count
|
return z_q_st, indices, commit_loss, perplexity, usage_count, entropy_loss
|
||||||
|
|
||||||
def sample(self, B: int, L: int, device: torch.device) -> torch.Tensor:
|
def sample(self, B: int, L: int, device: torch.device) -> torch.Tensor:
|
||||||
indices = torch.randint(0, self.K, (B, L), device=device)
|
indices = torch.randint(0, self.K, (B, L), device=device)
|
||||||
|
|||||||
@ -1,4 +1,6 @@
|
|||||||
|
import cv2
|
||||||
import numpy as np
|
import numpy as np
|
||||||
|
import torch
|
||||||
|
|
||||||
def blend_alpha(bg, fg, alpha):
|
def blend_alpha(bg, fg, alpha):
|
||||||
""" 使用 alpha 通道混合前景图块和背景图 """
|
""" 使用 alpha 通道混合前景图块和背景图 """
|
||||||
@ -38,3 +40,27 @@ def matrix_to_image_cv(map_matrix, tile_set, tile_size=32):
|
|||||||
)
|
)
|
||||||
|
|
||||||
return canvas
|
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
|
||||||
|
|||||||
336
test_refactor.py
Normal file
336
test_refactor.py
Normal file
@ -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()
|
||||||
445
需要重训的问题.md
Normal file
445
需要重训的问题.md
Normal file
@ -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),那部分要靠上面这些。
|
||||||
Loading…
Reference in New Issue
Block a user