feat: 墙壁生成策略调整

This commit is contained in:
unanmed 2026-08-02 13:35:49 +08:00
parent 46d1f1b450
commit 103b02cf08
3 changed files with 627 additions and 85 deletions

View File

@ -0,0 +1,367 @@
# Stage1 墙壁生长算法生成设计文档
## 背景与问题诊断
### 当前 Stage1 的 MaskGIT 采样机制
Stage1 负责生成地图的墙壁/地面骨架floor/wall使用标准 MaskGIT 迭代解码器:输入全 MASK 地图,通过余弦退火调度逐步揭开置信度最高的 token共 18 步完成。每步流程如下:
```
全 MASK 地图 → 模型预测所有位置 → 选出 Top-K 高置信度位置 →
揭开这些位置 → 其余位置重新 MASK → 下一步
```
### 问题现象
当前 Stage1 生成的地图墙壁质量显著低于预期,具体表现为:
- 墙壁分布破碎出现孤立墙块block 不与外围墙壁连接)
- 房间边界不完整,未能形成闭合区域
- 走廊断断续续,不能形成连通通道
- 整体拓扑结构与训练集分布偏差大
### 根本原因MaskGIT 缺乏空间连续约束
MaskGIT 的迭代解掩码unmasking机制本质上是**逐 token 独立决策**:每一步根据余弦退火比例选出全局置信度最高的 N 个位置进行揭开,选取标准仅依赖于该位置的类别预测置信度,不包含任何空间结构约束。
在 Stage1 墙壁生成中,墙壁的核心特性是**连通性**
1. 所有墙壁必须与外围墙壁(地图边框)连通
2. 墙壁必须形成闭合或半闭合的房间边界
3. 墙壁的合理性完全由其邻域关系决定——一个孤立的墙块在游戏中毫无意义
MaskGIT 的全局置信度排序无视了这些空间邻接约束。一个远离外围墙壁的内部位置,即使模型以 99% 置信度预测为墙壁,也不应该被单独采出——因为游戏规则要求墙壁必须连通。当前机制却会欣然接受这个"高置信度的错误"。这本质上是全局独立采样与局部连续性要求之间的结构性矛盾。
类比:
| 机制 | 类比 | 问题 |
| ------------------ | ------------------------------------ | ---------------- |
| MaskGIT 全局 Top-K | 在所有空白处随机填色,先填最"确信"的 | 可能出现孤立色块 |
| 墙壁生长算法 | 从已有色块边缘逐步向外扩展 | 保证连通性 |
---
## 核心思路:生长式墙壁生成
### 基本思想
将 Stage1 从"全局置信度排序 + 逐步解掩码"改为**边界生长**Frontier Growth
1. 初始状态外围一圈固定为墙壁tile=1内部全部为 MASK
2. 每步:模型预测所有位置,但**只在已存在墙壁的邻接位置**接受墙壁预测
3. 新产生的墙壁成为下一轮的"种子",继续向外生长
4. 直到没有新的邻接候选位置或达到最大步数
```
Step 0 Step 1 Step 2 Step 3
####### ####### ####### #######
#.....# #.....# #...#.# #.###.#
#.....# #..#..# #.###.# #####.#
#.....# #.....# #.....# #.#.#.#
####### ####### ####### #######
# = 墙壁 . = MASK/空地
```
### 与原 MaskGIT 的对比
| 维度 | 原 MaskGIT全局 Top-K | 生长算法(邻接约束) |
| ---------- | ------------------------ | ------------------------------------------ |
| 选取机制 | 全局置信度排序取 Top-K | 只在邻接已存在墙壁的位置候选中按置信度选取 |
| 空间约束 | 无,独立决策 | 强制连通:新墙壁必须"挨着"已有墙壁 |
| 墙壁连通性 | 不保证,依赖模型隐式学习 | 结构性保证 |
| 生成过程 | 并行:一步可揭开任意位置 | 扩散式:从边界逐步向内推进 |
| 收敛步数 | 固定 18 步 | 动态,取决于地图墙壁密度 |
---
## 算法详细设计
### 数据结构
```
地图状态矩阵 state[13][13],每格取值:
0: FLOOR空地/地面)
1: WALL墙壁
6: MASK待生成
外墙固化state 中位于最外圈row=0, row=12, col=0, col=12的位置
在初始化时设为 tile=1WALL并在整个过程中保持不可修改
```
### 主算法
```python
def wall_growth_sample(
model, z, z_dist, struct, target_density,
max_steps: int = 24
) -> np.ndarray:
# 1. 初始化:外圈墙壁 + 内部 MASK
state = torch.full((1, MAP_SIZE), MASK_TOKEN, dtype=torch.long)
state[0, outer_wall_indices] = 1 # 固化外圈
for step in range(max_steps):
# 2. 模型前向:预测所有位置的 logits
remain = compute_remaining(state, target_density, stage=1)
logits = model(state, z, z_dist, struct, remain)
probs = F.softmax(logits, dim=-1) # [1, MAP_SIZE, NUM_CLASSES]
wall_probs = probs[0, :, 1] # 每个位置预测为墙壁的概率
# 3. 寻找生长前沿:哪些 MASK 位置与已有墙壁四连通相邻
mask_positions = (state[0] == MASK_TOKEN)
adjacent_to_wall = compute_adjacency_mask(state[0]) # [MAP_SIZE] bool
candidates = mask_positions & adjacent_to_wall
if candidates.sum() == 0:
break # 没有可生长的位置,终止
# 4. 在候选位置中选出置信度最高的墙壁预测
cand_conf = wall_probs[candidates]
num_to_reveal = compute_reveal_count(step, max_steps, candidates.sum())
_, top_indices = torch.topk(cand_conf, k=num_to_reveal)
cand_positions = candidates.nonzero(as_tuple=True)[0][top_indices]
state[0, cand_positions] = 1 # 设为墙壁
# 5. 剩余 MASK 位置 → 空地
state[0, state[0] == MASK_TOKEN] = 0
return state[0].cpu().numpy().reshape(MAP_H, MAP_W)
```
### 邻接计算
```python
def compute_adjacency_mask(flat_state: torch.Tensor) -> torch.Tensor:
# 返回 [MAP_SIZE] 布尔张量True 表示该位置与任意墙壁 4-邻接
state_2d = flat_state.reshape(MAP_H, MAP_W)
wall = (state_2d == 1)
# 上/下/左/右 各 shift 一位取 OR
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()
```
### 每步揭开数量调度
与标准 MaskGIT 的余弦退火调度不同,生长算法在前期应积极揭开(扩大生长前沿),后期逐渐放缓(精细收束):
```python
def compute_reveal_count(step: int, max_steps: int, num_candidates: int) -> int:
# 前期多揭(扩大前沿),后期少揭(精细控制)
t = step / max_steps
ratio = 1.0 - t ** 0.5 # 从 1.0 降至 0.0,前期保持高位
# 不能超过当前候选数量,至少揭 1 个(若有候选)
return max(1, min(num_candidates, math.ceil(num_candidates * ratio)))
```
备选:沿用在生长前沿上的余弦退火——效果需通过实验验证。
### 候选类型扩展FLOOR 的处理
上述算法仅在候选位置接受墙壁tile=1预测。另一种思路是同时接受 FLOORtile=0预测让"空地"也能与墙壁一同向外推进:
| 策略 | 描述 | 优缺点 |
| ------------------ | -------------------------------------------------- | ------------------------------------------ |
| 仅接受墙壁 | 候选位置只接受 tile=1 预测 | 简单,墙壁生长可控;但后期可能挤压空地空间 |
| 同时接受墙壁和空地 | 候选位置按模型预测的最高置信度类别决定是墙还是空地 | 更接近自然生长;但对模型预测质量要求更高 |
推荐先使用"仅接受墙壁"策略,原因是:空地本身没有连通性要求,一旦墙壁骨架确定,空地自然就是剩余位置。实验验证阶段可对比两种策略。
### 终止条件
生长在以下条件之一满足时终止:
1. **无候选位置**:所有 MASK 位置均不邻接现有墙壁——墙壁已完全生长,剩余位置就是空地
2. **达到最大步数**`max_steps` 上限防止极端情况无限循环参考值24 步,对应 13×13 地图)
---
## Stage1 训练方案:矩形掩码 + 连通损失
### 设计动机
标准 MaskGIT 训练使用 `std_mask()`50% 随机散布掩码 / 50% 随机分块掩码),在整图上生成破碎的 MASK 图案。这种训练方式与生长式推理**完全不匹配**
- **训练**:模型看到散落各处的孤立 MASK token学习从全局上下文独立预测每个位置
- **推理**MASK 区域是一个连续的边界生长前沿,模型需要在已有墙壁的基础上向外延伸
训练与推理的分布不一致("train-test mismatch")导致模型在推理时表现不佳:模型从未在训练中见过"全连通空白区域从边界开始生长"的场景。
**核心改造**:将 Stage1 训练改为矩形掩码rectangular masking且只在掩码区域内与已有墙壁连通的位置计算损失。这使训练任务与推理任务完全对齐——都是"从已有墙壁向外延伸填充一个连续区域"。
### 矩形掩码生成
直接使用 `train_seperated.py` 中已有的 `rect_mask` 函数(第 419 行),无需修改。调用方式:
```python
from ginka.train_seperated import rect_mask
mask = rect_mask(ratio=0.5) # 反复放置随机矩形直到覆盖约 50% 面积
```
现有实现通过反复叠加随机小矩形2×2 ~ 7×7来达到目标遮盖比例形状不规则但保证连续性。
### 连通损失计算
仅对掩码区域内**与外部墙壁 4-连通**的位置计算 CE Loss。原因生长算法推理时只有能从已有墙壁"触及"的位置才有机会被生长覆盖——模型不需要学习预测生长前沿触及不到的位置。
```
原图: MASK阴影区: 连通位置(★):
########### ########### ###########
#.........# #.........# #.........#
#.###.....# #.###.....# #.###.....#
#.#.#.....# → #▓▓▓▓▓▓▓..# → #.★▓▓▓▓..# ← MASK 内左下角为墙,
#.....###.# #▓▓▓▓▓▓▓##.# #▓▓▓▓▓▓★#.# 因与外部墙连通故参与 loss
####### ####### ####### MASK 内右上方为地板,不参与 loss
# = 墙 . = 地板 ▓ = MASK 区域 ★ = 参与 loss 的位置
```
具体算法:
```python
def compute_stage1_loss_mask(
rect_mask: np.ndarray, target1: np.ndarray, inp1: np.ndarray
) -> np.ndarray:
# rect_mask: [H, W] boolTrue=被矩形掩码覆盖
# target1: [H, W] intGT 墙壁骨架(仅含 0=floor, 1=wall
# inp1: [H, W] int输入地图MASK 区域为 MASK_ID其余保留原值
# 返回: [H, W] boolTrue=该位置需要计算 loss
# 策略:在 rect_mask 内,仅对从"外部墙壁"可 4-连通触及的 WALL 位置计算 loss
H, W = rect_mask.shape
MASK_ID = 7
loss_mask = np.zeros((H, W), dtype=bool)
# 1. 找到外部墙壁(不在 rect_mask 内且值为 WALL=1 的位置)
external_wall = (~rect_mask) & (inp1 == 1)
# 2. 从外部墙壁出发 BFS沿 target1==1 传播,找到所有连通墙壁
visited = np.zeros((H, W), dtype=bool)
# BFS/DFS 种子:外部墙壁位置
seeds = np.argwhere(external_wall)
for r, c in seeds:
if visited[r, c]:
continue
stack = [(r, c)]
visited[r, c] = True
while stack:
cr, cc = stack.pop()
for nr, nc in [(cr-1, cc), (cr+1, cc), (cr, cc-1), (cr, cc+1)]:
if 0 <= nr < H and 0 <= nc < W:
if not visited[nr, nc] and target1[nr, nc] == 1:
visited[nr, nc] = True
stack.append((nr, nc))
# 3. 连通位置 = rect_mask 内 & 被 BFS 访问到的位置(即与外部墙连通的墙壁)
loss_mask = rect_mask & visited
return loss_mask
```
关键逻辑:
1. **外部墙壁** = 不在掩码区域内且值为 WALL 的位置
2. **BFS 传播** = 从外部墙壁出发,沿 GT 中值为 WALL(1) 的格子 4-连通扩散,标记所有可达的墙壁位置
3. **Loss 掩码** = 掩码区域内、且被 BFS 标记为可达的位置
注意BFS 使用 GTtarget1而非输入inp1因为需要标注的是"GT 中与外部墙壁连通的墙壁分布",而非输入中已存在的墙壁。即使掩码区域内的某个墙壁位置在输入中被 MASK 覆盖(值为 MASK_ID只要它在 GT 中与外部墙壁连通,就应该参与 loss。
loss_mask 中:
- **True** = 掩码区域内、与外部墙壁连通的墙壁位置 → 模型应该预测 WALL
- **False** = 掩码区域内的其他位置(地板、或与外部墙壁不连通的孤立墙壁) → 不参与 loss
### 对比:新旧训练方案
| 维度 | 旧方案MaskGIT std_mask | 新方案(矩形掩码 + 连通损失) |
| ---- | ------------------------- | ---------------------------- |
| 掩码形状 | 随机散布 + 随机方块叠加 | 连续矩形块(复用 `rect_mask` |
| MASK 覆盖比例 | Beta(2,2),约 5%~95% | 可由 `ratio` 参数控制 |
| loss 计算范围 | 所有 MASK 位置 | 仅 MASK 区域中与外部墙壁连通的位置 |
| 训练-推理一致性 | 低:散点 MASK vs 连续生长 | 高:矩形空缺 vs 连续生长边界 |
| 模型学习目标 | 从破碎上下文预测单个位置 | 从已有墙壁向外延伸填补连续区域 |
### 生长式推断
生长算法仅影响**Stage1 的推理采样方式**。推理时:
1. 初始状态外围一圈固定为墙壁tile=1内部全部为 MASK
2. 每步:模型预测所有位置,但只在已存在墙壁的邻接位置接受墙壁预测
3. 新产生的墙壁成为下一轮的"种子",继续向外生长
4. 直到没有新的邻接候选位置或达到最大步数
训练时使用矩形掩码 + 连通损失,推理时使用生长算法——两者在概念上完全一致:都是"从已有墙壁向外延伸"。训练为推理做了准备的准备teacher forcing推理则自回归地展开autoregressive rollout
### 推理管线影响
```
当前推理管线:
Stage1 (MaskGIT Top-K) → floor/wall 骨架
Stage2 (MaskGIT) → +门/怪物/入口
Stage3 (MaskGIT) → +资源
修改后推理管线:
Stage1 (生长算法) → floor/wall 骨架 [改动点]
Stage2 (MaskGIT) → +门/怪物/入口 [不变]
Stage3 (MaskGIT) → +资源 [不变]
```
Stage2 和 Stage3 的训练与推理**保持不变**——它们继续使用标准 MaskGIT 方法。这是因为门/怪物/资源是稀疏离散元素不涉及全局连通性约束MaskGIT 的独立 token 预测对此类任务足够有效。
### 模型参数
生长算法**不改变模型结构**GinkaMaskGIT 的 Transformer 架构不变),仅修改:
- `dataset.py`Stage1 的 `apply_subset*` 方法,将 MaskGIT 掩码策略改为矩形掩码 + 连通损失掩码
- `train_seperated.py`:新增 `wall_growth_sample` 函数替代 `maskgit_sample` 的 Stage1 分支
- 训练 loss 计算中Stage1 的 CE Loss 需要乘以 `loss_mask` 进行筛选
模型检查点格式与 Stage2/3 训练流程不受影响。
---
## 潜在风险与应对
| 风险 | 描述 | 应对策略 |
| -------------- | -------------------------------------------------------------------------- | ------------------------------------------------------------------------ |
| 墙壁密度过高 | 生长不退让,可能填满所有可达区域 | 增加 target_density 的约束权重;在后期步数中降低揭开速率 |
| 生长形状单一 | 从外圈均匀向内心扩散,产生"洋葱圈"形墙壁 | 引入随机扰动(每步以一定概率跳过某些候选位置);调整 target_density 分布 |
| 屋顶封闭 | 生长过早将空地包围封闭,导致内部无法再生长墙壁 | 在揭开判断中加入"是否会导致孤岛"的后处理检查 |
| 模型能力不匹配 | 生长算法依赖模型在邻接位置的预测质量,若模型本身不能捕捉邻接关系则效果有限 | 可结合已有的二维因式位置嵌入,确保模型具备足够的局部感知能力 |
| 步数不均匀 | 不同地图墙壁密度不同,生长步数的固定上限可能不够或过早收敛 | 将 max_steps 设为动态参数,或改为"无新墙即终止"的收敛策略 |
### 墙壁"洋葱圈"问题专项说明
若每步以固定速率在所有邻接生长前沿均匀揭开,且模型倾向于预测墙壁,则墙壁可能从外围逐层向内均匀扩散,形成洋葱圈状结构。应对策略:
1. **target_density 约束**:靶向墙壁密度作为硬约束,当已生成的墙壁数量接近目标密度时,大幅降低每步揭开速率
2. **非均匀揭开**:每步仅在候选位置中随机选取子集,而非全部可选中的最高置信度位置,增加生长路径的随机性
3. **后处理修剪**:生成完成后,对墙壁进行一次 BFS 修剪,移除不改变连通性的"冗余墙壁"分支
---
## 实施计划
### 数据集修改
- [ ] Stage1 训练中直接调用已有 `rect_mask` 替代 `std_mask`
- [ ] 新增 `compute_stage1_loss_mask` 函数BFS 计算连通损失掩码
- [ ] 在 `__getitem__` 返回的 sample 中新增 `"loss_mask_stage1"` 字段
### 推理修改
- [ ] 在 `train_seperated.py` 中实现 `wall_growth_sample` 函数
- [ ] 修改 `full_generate_specific_z` 中 Stage1 的调用,替换为生长算法
- [ ] 实现 `compute_adjacency_mask` 工具函数(四邻接检测)
- [ ] 实现生长步数调度 `compute_reveal_count`
### 训练脚本修改
- [ ] Stage1 的 CE Loss 改为 `(F.cross_entropy(..., reduction='none') * loss_mask_stage1.float()).sum() / loss_mask_stage1.sum()`
- [ ] Stage2/3 的训练与原有逻辑保持一致,不做修改
### 验证
- [ ] 单步测试:加载已有检查点,运行生长算法一次,输出墙壁骨架地图
- [ ] 添加可视化对比:同一组 z 下MaskGIT Top-K vs 生长算法的 Stage1 输出
- [ ] 对比指标:墙壁连通率、孤立墙块数、房间闭合度

View File

@ -5,6 +5,53 @@ import numpy as np
from torch.utils.data import Dataset from torch.utils.data import Dataset
from shared.distance import compute_distance_field from shared.distance import compute_distance_field
def rect_mask(ratio: float, map_size: int = 169) -> np.ndarray:
# 连续矩形分块掩码,反复放置随机矩形直到掩码格数达标
# 复用 train_seperated.py:419 的同名逻辑dtype 为 bool 以与 std_mask 区分
target = int(map_size * ratio)
mask = np.zeros((13, 13), dtype=bool)
while mask.sum() < target:
bh = np.random.randint(2, 7)
bw = np.random.randint(2, 7)
x = np.random.randint(0, 14 - bh)
y = np.random.randint(0, 14 - bw)
mask[x:x + bh, y:y + bw] = True
return mask
def ensure_wall_connection(mask: np.ndarray, target1: np.ndarray) -> np.ndarray:
# BFS 找到所有墙壁连通分量,对完全被掩码覆盖的分量随机保留若干种子点
H, W = mask.shape
result = mask.copy()
visited = np.zeros((H, W), dtype=bool)
directions = [(-1, 0), (1, 0), (0, -1), (0, 1)]
wall_positions = target1 == 1
seeds = np.argwhere(wall_positions & (~visited))
for r, c in seeds:
r, c = int(r), int(c)
if visited[r, c]:
continue
comp = []
stack = [(r, c)]
visited[r, c] = True
while stack:
cr, cc = stack.pop()
comp.append((cr, cc))
for dr, dc in directions:
nr, nc = cr + dr, cc + dc
if 0 <= nr < H and 0 <= nc < W and not visited[nr, nc]:
if target1[nr, nc] == 1:
visited[nr, nc] = True
stack.append((nr, nc))
# 检查该分量是否完全被掩码覆盖
all_masked = all(result[pos] for pos in comp)
if all_masked:
# 随机保留若干种子点(至少 1 个,最多分量长度的一半)
keep = max(1, np.random.randint(1, max(2, len(comp) // 2 + 1)))
chosen = [comp[i] for i in np.random.choice(len(comp), keep, replace=False)]
for pos in chosen:
result[pos] = False
return result
def load_data(path: str): def load_data(path: str):
with open(path, 'r', encoding="utf-8") as f: with open(path, 'r', encoding="utf-8") as f:
data = json.load(f) data = json.load(f)
@ -179,7 +226,7 @@ class GinkaSeperatedDataset(Dataset):
return target1, inp1, target2, inp2, target3, inp3 return target1, inp1, target2, inp2, target3, inp3
def apply_subset1(self, raw: np.ndarray): def apply_subset1(self, raw: np.ndarray):
# 子集 1std_mask 随机掩码 # 子集 1Stage1 使用 rect_maskStage2/3 保持 std_mask
target1, inp1, target2, inp2, target3, inp3 = self.create_degreaded(raw) target1, inp1, target2, inp2, target3, inp3 = self.create_degreaded(raw)
@ -187,8 +234,11 @@ class GinkaSeperatedDataset(Dataset):
enc2 = inp2.copy() enc2 = inp2.copy()
enc3 = raw.copy() enc3 = raw.copy()
# stage1对整图 std_mask # stage1rect_mask连续矩形掩码替代 std_mask
inp1[self.std_mask()] = self.MASK_ID ratio = float(np.random.beta(2, 2)) * 0.95 + 0.05
rmask = rect_mask(ratio)
rmask = ensure_wall_connection(rmask, target1) # 确保墙壁分量有种子点
inp1[rmask] = self.MASK_ID
# stage2对 floor+功能元素区域 std_mask # stage2对 floor+功能元素区域 std_mask
need_mask = np.isin(inp2, [self.FLOOR, self.DOOR, self.SPECIAL_DOOR, self.MONSTER, self.ENTRANCE]) need_mask = np.isin(inp2, [self.FLOOR, self.DOOR, self.SPECIAL_DOOR, self.MONSTER, self.ENTRANCE])
@ -208,8 +258,12 @@ class GinkaSeperatedDataset(Dataset):
enc2 = inp2.copy() enc2 = inp2.copy()
enc3 = raw.copy() enc3 = raw.copy()
need_mask = np.isin(inp2, [self.FLOOR, self.WALL]) ratio = float(np.random.beta(2, 2)) * 0.95 + 0.05
inp1[need_mask & self.std_mask()] = self.MASK_ID rmask1 = rect_mask(ratio)
need_mask = np.isin(inp1, [self.FLOOR, self.WALL])
rmask_final = ensure_wall_connection(rmask1 & need_mask, target1)
inp1[rmask_final] = self.MASK_ID
need_mask = np.isin(inp2, [self.FLOOR, self.DOOR, self.SPECIAL_DOOR, self.MONSTER, self.ENTRANCE]) need_mask = np.isin(inp2, [self.FLOOR, self.DOOR, self.SPECIAL_DOOR, self.MONSTER, self.ENTRANCE])
inp2[need_mask] = self.MASK_ID inp2[need_mask] = self.MASK_ID
need_mask = np.isin(inp3, [self.FLOOR, self.RESOURCE]) need_mask = np.isin(inp3, [self.FLOOR, self.RESOURCE])

View File

@ -37,31 +37,49 @@ from shared.distance import DIST_VOCAB, compute_distance_field_tensor
# 图块 ID 定义: # 图块 ID 定义:
# 0. 空地 1. 墙壁 2. 普通门 3. 资源 4. 怪物 5. 入口 6. 机关门 7. 掩码MASK_TOKEN # 0. 空地 1. 墙壁 2. 普通门 3. 资源 4. 怪物 5. 入口 6. 机关门 7. 掩码MASK_TOKEN
# 共用 VQ-VAE 超参 # 共用 VQ-VAE 超参(共享的编码维度)
# 三组编码器vq1/vq2/vq3共享相同超参分别对三阶段地图上下文独立编码
VQ_L = 8 # 码字序列长度(每个编码器输出 L 个码字,量化后合并为 L*3
VQ_K = 16 # codebook 大小(离散码本条目数)
VQ_D_Z = 64 # 码字维度 VQ_D_Z = 64 # 码字维度
VQ_BETA = 1.0 # commit loss 权重(防止编码器输出漂离 codebook VQ_BETA = 1.0 # commit loss 权重(防止编码器输出漂离 codebook
VQ_GAMMA = 0.0 # entropy loss 权重(当前未启用) VQ_GAMMA = 0.0 # entropy loss 权重(当前未启用)
VQ_LAYERS = 6 # VQ-VAE Transformer 层数
VQ_DIM_FF = 1024 # VQ-VAE 前馈网络隐层维度 # 三通道 VQ 各自独立超参L、K、层数、维度等均独立配置
VQ_D_MODEL = 256 # VQ-VAE Transformer 模型维度 # Stage1 墙壁骨架 — 结构最复杂,模型容量最大
VQ_NHEAD = 4 # VQ-VAE 多头注意力头数 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 大小
# 距离场编码器超参 # 距离场编码器超参
L_DIST = 4 # 距离场码字序列长度 DIST_D_MODEL = 384 # 距离场编码器模型维度
K_DIST = 16 # 距离场 codebook 大小 DIST_LAYERS = 6 # 距离场编码器 Transformer 层数
DIST_D_MODEL = 256 # 距离场编码器模型维度 DIST_DIM_FF = 1536 # 距离场编码器 FF 维度
DIST_LAYERS = 3 # 距离场编码器 Transformer 层数
DIST_DIM_FF = 1024 # 距离场编码器 FF 维度
DIST_NHEAD = 8 # 距离场编码器注意力头数 DIST_NHEAD = 8 # 距离场编码器注意力头数
VQ_BETA_DIST = 0.5 # 距离场 commit loss 权重 VQ_BETA_DIST = 0.5 # 距离场 commit loss 权重
# 第一阶段 MaskGIT 超参 # 第一阶段 MaskGIT 超参
STAGE1_MG_DMODEL = 512 STAGE1_MG_DMODEL = 512
STAGE1_MG_NHEAD = 4 STAGE1_MG_NHEAD = 4
STAGE1_MG_NUM_LAYERS = 6 STAGE1_MG_NUM_LAYERS = 8
STAGE1_MG_DIM_FF = 2048 STAGE1_MG_DIM_FF = 2048
# 第二阶段 MaskGIT 超参 # 第二阶段 MaskGIT 超参
@ -89,6 +107,7 @@ STAGE3_VQ_WEIGHT = 0.5
# 全局参数 # 全局参数
NUM_CLASSES = 8 # 图块类型数 NUM_CLASSES = 8 # 图块类型数
MASK_TOKEN = 7 # 掩码图块 MASK_TOKEN = 7 # 掩码图块
TOTAL_K = VQ_K1 + VQ_K2 + VQ_K3 # 预计算,供日志输出使用
MAP_W = 13 # 地图宽度 MAP_W = 13 # 地图宽度
MAP_H = 13 # 地图高度 MAP_H = 13 # 地图高度
MAP_SIZE = MAP_W * MAP_H # 地图大小 MAP_SIZE = MAP_W * MAP_H # 地图大小
@ -106,14 +125,14 @@ MG_Z_DROPOUT = 0.1 # z 隐变量 Dropout 概率
MG_STRUCT_DROPOUT = 0.1 # 结构参量 Dropout 概率 MG_STRUCT_DROPOUT = 0.1 # 结构参量 Dropout 概率
# 邻接损失权重(三阶段) # 邻接损失权重(三阶段)
LAMBDA_ADJ1 = 0.6 LAMBDA_ADJ1 = 0.2
LAMBDA_ADJ2 = 0.3 LAMBDA_ADJ2 = 0.1
LAMBDA_ADJ3 = 0.1 LAMBDA_ADJ3 = 0.05
# Patch 损失权重(三阶段)及核参数 # Patch 损失权重(三阶段)及核参数
LAMBDA_PATCH1 = 0.5 LAMBDA_PATCH1 = 0.2
LAMBDA_PATCH2 = 0.5 LAMBDA_PATCH2 = 0.2
LAMBDA_PATCH3 = 0.5 LAMBDA_PATCH3 = 0.2
PATCH_KERNEL_SIZE = 5 PATCH_KERNEL_SIZE = 5
PATCH_SIGMA = 1.2 PATCH_SIGMA = 1.2
@ -153,37 +172,44 @@ def parse_arguments():
return parser.parse_args() return parser.parse_args()
def build_model(device: torch.device): def build_model(device: torch.device):
# 三组 VQ-VAE 编码器各自独立编码一个阶段的地图上下文encoder_stage1/2/3 # 三组 VQ-VAE 编码器:各自独立编码一个阶段的地图上下文,全部超参独立
# 输出形状均为 [B, L, d_z],分别送入各自阶段的 quantizer vq1 = GinkaVQVAE(
vq_kwargs = dict( num_classes=NUM_CLASSES, L=VQ_L1, K=VQ_K1,
num_classes=NUM_CLASSES, L=VQ_L, K=VQ_K, d_model=VQ_D_MODEL, d_model=VQ_D_MODEL1, nhead=VQ_NHEAD1, num_layers=VQ_LAYERS1, dim_ff=VQ_DIM_FF1,
nhead=VQ_NHEAD, num_layers=VQ_LAYERS, dim_ff=VQ_DIM_FF, map_h=MAP_H, map_w=MAP_W map_h=MAP_H, map_w=MAP_W
) ).to(device)
vq1 = GinkaVQVAE(**vq_kwargs).to(device) # 编码 stage1 上下文floor/wall vq2 = GinkaVQVAE(
vq2 = GinkaVQVAE(**vq_kwargs).to(device) # 编码 stage2 上下文door/monster/entrance num_classes=NUM_CLASSES, L=VQ_L2, K=VQ_K2,
vq3 = GinkaVQVAE(**vq_kwargs).to(device) # 编码 stage3 上下文resource 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 作为条件 # 三个独立 MaskGIT 解码器,分别接收各自阶段的 z_q 作为条件
mg1 = GinkaMaskGIT( mg1 = GinkaMaskGIT(
num_classes=NUM_CLASSES, d_model=STAGE1_MG_DMODEL, d_z=VQ_D_Z, dim_ff=STAGE1_MG_DIM_FF, 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, nhead=STAGE1_MG_NHEAD, num_layers=STAGE1_MG_NUM_LAYERS, map_h=MAP_H, map_w=MAP_W,
z_seq_len=VQ_L z_seq_len=VQ_L1, z_dist_len=L_DIST
).to(device) ).to(device)
mg2 = GinkaMaskGIT( mg2 = GinkaMaskGIT(
num_classes=NUM_CLASSES, d_model=STAGE2_MG_DMODEL, d_z=VQ_D_Z, dim_ff=STAGE2_MG_DIM_FF, 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, nhead=STAGE2_MG_NHEAD, num_layers=STAGE2_MG_NUM_LAYERS, map_h=MAP_H, map_w=MAP_W,
z_seq_len=VQ_L z_seq_len=VQ_L2, z_dist_len=L_DIST
).to(device) ).to(device)
mg3 = GinkaMaskGIT( mg3 = GinkaMaskGIT(
num_classes=NUM_CLASSES, d_model=STAGE3_MG_DMODEL, d_z=VQ_D_Z, dim_ff=STAGE3_MG_DIM_FF, 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, nhead=STAGE3_MG_NHEAD, num_layers=STAGE3_MG_NUM_LAYERS, map_h=MAP_H, map_w=MAP_W,
z_seq_len=VQ_L z_seq_len=VQ_L3, z_dist_len=L_DIST
).to(device) ).to(device)
# 三个独立 VectorQuantizer各阶段使用自己的码本,避免语义空间相互干扰 # 三个独立 VectorQuantizer各阶段使用自己的码本大小
quantizer1 = VectorQuantizer(K=VQ_K, d_z=VQ_D_Z).to(device) quantizer1 = VectorQuantizer(K=VQ_K1, d_z=VQ_D_Z).to(device)
quantizer2 = VectorQuantizer(K=VQ_K, d_z=VQ_D_Z).to(device) quantizer2 = VectorQuantizer(K=VQ_K2, d_z=VQ_D_Z).to(device)
quantizer3 = VectorQuantizer(K=VQ_K, d_z=VQ_D_Z).to(device) quantizer3 = VectorQuantizer(K=VQ_K3, d_z=VQ_D_Z).to(device)
quantizers = (quantizer1, quantizer2, quantizer3) quantizers = (quantizer1, quantizer2, quantizer3)
# 距离场编码器与量化器:将 L1 距离场编码为离散 latent z_dist # 距离场编码器与量化器:将 L1 距离场编码为离散 latent z_dist
@ -291,17 +317,19 @@ def apply_z_dropout(
mask = torch.rand(z_q.shape[0], z_q.shape[1], 1, device=z_q.device) < drop_prob 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() return z_q * (~mask).float() + mask_embedding * mask.float()
def summarize_codebook_hits(code_hits: torch.Tensor) -> tuple[float, float, int]: def summarize_codebook_hits(code_hits) -> tuple[float, float, int]:
total_hits = code_hits.sum() # 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: if total_hits.item() <= 0:
return 0.0, 0.0, 0 return 0.0, 0.0, 0
probs = code_hits / total_hits probs = combined / total_hits
perplexity = torch.exp( perplexity = torch.exp(
-(probs * torch.log(probs.clamp_min(1e-10))).sum() -(probs * torch.log(probs.clamp_min(1e-10))).sum()
).item() ).item()
active_codes = int((code_hits > 0).sum().item()) active_codes = int((combined > 0).sum().item())
usage_rate = active_codes / code_hits.numel() usage_rate = active_codes / combined.numel()
return perplexity, usage_rate, active_codes return perplexity, usage_rate, active_codes
def quantize_stage_latents( def quantize_stage_latents(
@ -316,7 +344,7 @@ def quantize_stage_latents(
z_q3, _, commit_loss3, _, code_hits3 = quantizer3(z_e3) z_q3, _, commit_loss3, _, code_hits3 = quantizer3(z_e3)
commit_loss = (commit_loss1 + commit_loss2 + commit_loss3) / 3 commit_loss = (commit_loss1 + commit_loss2 + commit_loss3) / 3
code_hits = torch.stack([code_hits1, code_hits2, code_hits3], dim=0) 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
def build_reference_rollout_steps(prob: float) -> int: def build_reference_rollout_steps(prob: float) -> int:
@ -431,6 +459,71 @@ def rect_mask(
mask[x:x + bh, y:y + bw] = True mask[x:x + bh, y:y + bw] = True
return mask return mask
def compute_adjacency_mask(flat_state: torch.Tensor) -> torch.Tensor:
# 返回 [MAP_SIZE] boolTrue 表示该位置与任意墙壁 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( def maskgit_sample(
model: torch.nn.Module, inp: torch.Tensor, z: torch.Tensor, model: torch.nn.Module, inp: torch.Tensor, z: torch.Tensor,
z_dist: torch.Tensor, struct: torch.Tensor, target_density: torch.Tensor, z_dist: torch.Tensor, struct: torch.Tensor, target_density: torch.Tensor,
@ -535,15 +628,14 @@ def full_generate_specific_z(
device: torch.device, device: torch.device,
keep_fixed: tuple[bool, bool, bool] = (True, True, True) keep_fixed: tuple[bool, bool, bool] = (True, True, True)
) -> tuple: ) -> tuple:
vq1, vq2, vq3, mg1, mg2, mg3, quantizers, optimizer, scheduler = models vq1, vq2, vq3, mg1, mg2, mg3, quantizers, _, _, _, _, _ = models
with torch.no_grad(): with torch.no_grad():
z1, z2, z3 = z_q z1, z2, z3 = z_q
# 三阶段级联生成,但使用给定的 z # 三阶段级联生成Stage1 使用生长算法Stage2/3 保持 MaskGIT
pred1_np = maskgit_sample( pred1_np = wall_growth_sample(
mg1, input.clone(), z1, z_dist, struct, target_density, 1, mg1, z1, z_dist, struct, target_density
GENERATE_STEP, target_tiles=[1], keep_fixed=keep_fixed[0]
) )
inp2 = torch.tensor(pred1_np.flatten(), dtype=torch.long, device=device).reshape(1, MAP_SIZE) inp2 = torch.tensor(pred1_np.flatten(), dtype=torch.long, device=device).reshape(1, MAP_SIZE)
inp2[inp2 == 0] = MASK_TOKEN inp2[inp2 == 0] = MASK_TOKEN
@ -579,21 +671,50 @@ def inpaint_generate(
# 三阶段矩形掩码修补生成 # 三阶段矩形掩码修补生成
# - z 来自完整地图 VQ 编码,作为全局先验 # - z 来自完整地图 VQ 编码,作为全局先验
# - 未掩码区域的原始图块逐阶段保留,避免被误当空地重填 # - 未掩码区域的原始图块逐阶段保留,避免被误当空地重填
vq1, vq2, vq3, mg1, mg2, mg3, quantizers, optimizer, scheduler = models vq1, vq2, vq3, mg1, mg2, mg3, quantizers, _, _, _, _, _ = models
z1, z2, z3 = z_q z1, z2, z3 = z_q
# Stage1: 补全墙壁 # Stage1: 补全墙壁 — 生长方式,掩码区域内仅邻接圈为 MASK其余为 FLOOR(0)
inp1 = raw_map.copy() inp1 = raw_map.copy()
inp1[mask] = MASK_TOKEN non_wf = (inp1 != 0) & (inp1 != 1)
non_wf = (inp1 != 0) & (inp1 != 1) & (inp1 != MASK_TOKEN)
inp1[non_wf] = 0 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) 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(): with torch.no_grad():
pred1_np = maskgit_sample( state = init_state
mg1, inp1_t, z1, z_dist, struct, target_density, 1, for step in range(GENERATE_STEP):
GENERATE_STEP, target_tiles=[1], keep_fixed=True 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: 补全门/怪/入口,保留未掩码区的原始非墙壁结构 # Stage2: 补全门/怪/入口,保留未掩码区的原始非墙壁结构
merged_s1 = pred1_np.copy() merged_s1 = pred1_np.copy()
@ -669,7 +790,7 @@ def build_dataset_sample_case(
device: torch.device, device: torch.device,
idx: int | None = None idx: int | None = None
) -> dict: ) -> dict:
vq1, vq2, vq3, mg1, mg2, mg3, quantizers, optimizer, scheduler = models vq1, vq2, vq3, mg1, mg2, mg3, quantizers, _, _, _, _, _ = models
dist_encoder, dist_quantizer = dist_models dist_encoder, dist_quantizer = dist_models
sample = dataset.random_sample_map(idx=idx) sample = dataset.random_sample_map(idx=idx)
@ -805,15 +926,19 @@ def visualize_rand(
def to_img(mat): def to_img(mat):
return matrix_to_image_cv(mat, tile_dict, TILE_SIZE) return matrix_to_image_cv(mat, tile_dict, TILE_SIZE)
vq1, vq2, vq3, mg1, mg2, mg3, quantizers, optimizer, scheduler = models vq1, vq2, vq3, mg1, mg2, mg3, quantizers, _, _, _, _, _ = models
dist_encoder, dist_quantizer = dist_models dist_encoder, dist_quantizer = dist_models
samples_data = [] samples_data = []
for _ in range(4): for _ in range(4):
sample = train_dataset.random_sample_map() sample = train_dataset.random_sample_map()
raw_map = sample["raw_map"].numpy().reshape(MAP_H, MAP_W) raw_map = sample["raw_map"].cpu().numpy().reshape(MAP_H, MAP_W)
ratio = random.uniform(0.2, 0.8) ratio = random.uniform(0.2, 0.8)
mask = rect_mask(ratio) 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) enc1_t = sample["encoder_stage1"].to(device).reshape(1, MAP_SIZE)
enc2_t = sample["encoder_stage2"].to(device).reshape(1, MAP_SIZE) enc2_t = sample["encoder_stage2"].to(device).reshape(1, MAP_SIZE)
@ -885,7 +1010,7 @@ def validate(
train_dataset: GinkaSeperatedDataset, train_dataset: GinkaSeperatedDataset,
epoch: int epoch: int
): ):
vq1, vq2, vq3, mg1, mg2, mg3, quantizers, optimizer, scheduler = models vq1, vq2, vq3, mg1, mg2, mg3, quantizers, _, _, _, _, _ = models
dist_encoder, dist_quantizer = dist_models dist_encoder, dist_quantizer = dist_models
quantizer1, quantizer2, quantizer3 = quantizers quantizer1, quantizer2, quantizer3 = quantizers
@ -904,7 +1029,7 @@ def validate(
patch1_total = torch.Tensor([0]).to(device) patch1_total = torch.Tensor([0]).to(device)
patch2_total = torch.Tensor([0]).to(device) patch2_total = torch.Tensor([0]).to(device)
patch3_total = torch.Tensor([0]).to(device) patch3_total = torch.Tensor([0]).to(device)
code_hits_total = torch.zeros(3, quantizer1.K, device=device) code_hits_total = (torch.zeros(quantizer1.K, device=device), torch.zeros(quantizer2.K, device=device), torch.zeros(quantizer3.K, device=device)) # validate
density_metrics = { density_metrics = {
1: {"mae": 0.0, "over": 0.0, "count": 0}, 1: {"mae": 0.0, "over": 0.0, "count": 0},
@ -969,7 +1094,7 @@ def validate(
patch1_total += patch_loss(logits1, target1, PATCH_KERNEL_SIZE, PATCH_SIGMA) patch1_total += patch_loss(logits1, target1, PATCH_KERNEL_SIZE, PATCH_SIGMA)
patch2_total += patch_loss(logits2, target2, 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) patch3_total += patch_loss(logits3, target3, PATCH_KERNEL_SIZE, PATCH_SIGMA)
code_hits_total += code_hits code_hits_total = (code_hits_total[0] + code_hits[0], code_hits_total[1] + code_hits[1], code_hits_total[2] + code_hits[2])
# 计算各目标对象的真实密度误差与过量生成密度 # 计算各目标对象的真实密度误差与过量生成密度
pred1_map = torch.argmax(logits1, dim=-1).cpu() pred1_map = torch.argmax(logits1, dim=-1).cpu()
@ -1023,7 +1148,7 @@ def train(device: torch.device):
result = build_model(device) result = build_model(device)
vq1, vq2, vq3, mg1, mg2, mg3, quantizers, optimizer, scheduler, latent_mask_embedding, dist_encoder, dist_quantizer = result 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] models = [vq1, vq2, vq3, mg1, mg2, mg3, quantizers, optimizer, scheduler, latent_mask_embedding, dist_encoder, dist_quantizer]
dist_models = (dist_encoder, dist_quantizer) dist_models = (dist_encoder, dist_quantizer)
quantizer1, quantizer2, quantizer3 = quantizers quantizer1, quantizer2, quantizer3 = quantizers
@ -1106,7 +1231,7 @@ def train(device: torch.device):
patch1_total = torch.Tensor([0]).to(device) patch1_total = torch.Tensor([0]).to(device)
patch2_total = torch.Tensor([0]).to(device) patch2_total = torch.Tensor([0]).to(device)
patch3_total = torch.Tensor([0]).to(device) patch3_total = torch.Tensor([0]).to(device)
code_hits_total = torch.zeros(3, quantizer1.K, device=device) code_hits_total = (torch.zeros(quantizer1.K, device=device), torch.zeros(quantizer2.K, device=device), torch.zeros(quantizer3.K, device=device)) # validate
for batch in tqdm(dataloader, leave=False, desc="Epoch Progress", disable=disable_tqdm): for batch in tqdm(dataloader, leave=False, desc="Epoch Progress", disable=disable_tqdm):
# 三阶段各自的掩码输入序列、预测目标和编码器上下文 # 三阶段各自的掩码输入序列、预测目标和编码器上下文
@ -1127,7 +1252,7 @@ def train(device: torch.device):
target_density = batch["target_density"].to(device) target_density = batch["target_density"].to(device)
dist_field = batch["distance_field"].to(device) dist_field = batch["distance_field"].to(device)
optimizer.zero_grad() optimizer.zero_grad() # 训练循环
# VQ 编码:各阶段编码器分别处理各自上下文切片 # VQ 编码:各阶段编码器分别处理各自上下文切片
z_e1 = vq1(enc1) # [B, L, d_z] z_e1 = vq1(enc1) # [B, L, d_z]
@ -1175,18 +1300,16 @@ def train(device: torch.device):
loss2 = cross_entropy_loss(logits2, target2) loss2 = cross_entropy_loss(logits2, target2)
loss3 = cross_entropy_loss(logits3, target3) loss3 = cross_entropy_loss(logits3, target3)
adj1 = adjacency_loss(logits1, target1)
adj2 = adjacency_loss(logits2, target2) adj2 = adjacency_loss(logits2, target2)
adj3 = adjacency_loss(logits3, target3) adj3 = adjacency_loss(logits3, target3)
patch1 = patch_loss(logits1, target1, PATCH_KERNEL_SIZE, PATCH_SIGMA)
patch2 = patch_loss(logits2, target2, PATCH_KERNEL_SIZE, PATCH_SIGMA) patch2 = patch_loss(logits2, target2, PATCH_KERNEL_SIZE, PATCH_SIGMA)
patch3 = patch_loss(logits3, target3, PATCH_KERNEL_SIZE, PATCH_SIGMA) patch3 = patch_loss(logits3, target3, PATCH_KERNEL_SIZE, PATCH_SIGMA)
loss1_weighted = STAGE1_CE_WEIGHT * loss1 loss1_weighted = STAGE1_CE_WEIGHT * loss1
loss2_weighted = STAGE2_CE_WEIGHT * loss2 loss2_weighted = STAGE2_CE_WEIGHT * loss2
loss3_weighted = STAGE3_CE_WEIGHT * loss3 loss3_weighted = STAGE3_CE_WEIGHT * loss3
adj_weighted = LAMBDA_ADJ1 * adj1 + LAMBDA_ADJ2 * adj2 + LAMBDA_ADJ3 * adj3 adj_weighted = LAMBDA_ADJ2 * adj2 + LAMBDA_ADJ3 * adj3
patch_weighted = LAMBDA_PATCH1 * patch1 + LAMBDA_PATCH2 * patch2 + LAMBDA_PATCH3 * patch3 patch_weighted = LAMBDA_PATCH2 * patch2 + LAMBDA_PATCH3 * patch3
commit_weighted = VQ_BETA * commit_loss + VQ_BETA_DIST * commit_loss_dist 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 loss = loss1_weighted + loss2_weighted + loss3_weighted + adj_weighted + patch_weighted + commit_weighted
@ -1199,13 +1322,11 @@ def train(device: torch.device):
loss2_total += loss2.detach() loss2_total += loss2.detach()
loss3_total += loss3.detach() loss3_total += loss3.detach()
commit_total += commit_loss.detach() commit_total += commit_loss.detach()
adj1_total += adj1.detach()
adj2_total += adj2.detach() adj2_total += adj2.detach()
adj3_total += adj3.detach() adj3_total += adj3.detach()
patch1_total += patch1.detach()
patch2_total += patch2.detach() patch2_total += patch2.detach()
patch3_total += patch3.detach() patch3_total += patch3.detach()
code_hits_total += code_hits.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 结束后更新学习率 # 每个 epoch 结束后更新学习率
scheduler.step() scheduler.step()
@ -1217,11 +1338,11 @@ def train(device: torch.device):
f"E: {epoch + 1} | " f"E: {epoch + 1} | "
f"Loss: {loss_total.item() / data_length:.4f} | " 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"CE: {loss1_total.item() / data_length:.4f}, {loss2_total.item() / data_length:.4f}, {loss3_total.item() / data_length:.4f} | "
f"ADJ: {(LAMBDA_ADJ1 * adj1_total.item() + LAMBDA_ADJ2 * adj2_total.item() + LAMBDA_ADJ3 * adj3_total.item()) / data_length:.4f}, {adj1_total.item() / data_length:.4f}, {adj2_total.item() / data_length:.4f}, {adj3_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_PATCH1 * patch1_total.item() + LAMBDA_PATCH2 * patch2_total.item() + LAMBDA_PATCH3 * patch3_total.item()) / data_length:.4f}, {patch1_total.item() / data_length:.4f}, {patch2_total.item() / data_length:.4f}, {patch3_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"VQ: {commit_total.item() / data_length:.4f} | "
f"PPL: {train_perplexity:.4f} | " f"PPL: {train_perplexity:.4f} | "
f"Usage: {train_usage_rate:.4f} ({train_active_codes}/{code_hits_total.numel()}) | " f"Usage: {train_usage_rate:.4f} ({train_active_codes}/{TOTAL_K}) | "
f"LR: {scheduler.get_last_lr()[0]:.6f}" f"LR: {scheduler.get_last_lr()[0]:.6f}"
) )
@ -1230,12 +1351,12 @@ def train(device: torch.device):
losses = validate( losses = validate(
dataloader_val, models, dist_models, device, tile_dict, dataset, epoch + 1 dataloader_val, models, dist_models, device, tile_dict, dataset, epoch + 1
) )
loss1_total, loss2_total, loss3_total, adj1_total, adj2_total, adj3_total, patch1_total, patch2_total, patch3_total, commit_total, code_hits_total = losses 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 loss1_weighted = STAGE1_CE_WEIGHT * loss1_total
loss2_weighted = STAGE2_CE_WEIGHT * loss2_total loss2_weighted = STAGE2_CE_WEIGHT * loss2_total
loss3_weighted = STAGE3_CE_WEIGHT * loss3_total loss3_weighted = STAGE3_CE_WEIGHT * loss3_total
adj_weighted = LAMBDA_ADJ1 * adj1_total + LAMBDA_ADJ2 * adj2_total + LAMBDA_ADJ3 * adj3_total adj_weighted = LAMBDA_ADJ2 * adj2_total + LAMBDA_ADJ3 * adj3_total
patch_weighted = LAMBDA_PATCH1 * patch1_total + LAMBDA_PATCH2 * patch2_total + LAMBDA_PATCH3 * patch3_total patch_weighted = LAMBDA_PATCH2 * patch2_total + LAMBDA_PATCH3 * patch3_total
commit_weighted = VQ_BETA * commit_total commit_weighted = VQ_BETA * commit_total
loss_total = loss1_weighted + loss2_weighted + loss3_weighted + adj_weighted + patch_weighted + commit_weighted loss_total = loss1_weighted + loss2_weighted + loss3_weighted + adj_weighted + patch_weighted + commit_weighted
@ -1246,11 +1367,11 @@ def train(device: torch.device):
f"E: {epoch + 1} | " f"E: {epoch + 1} | "
f"Loss: {loss_total.item() / data_length:.4f} | " 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"CE: {loss1_total.item() / data_length:.4f}, {loss2_total.item() / data_length:.4f}, {loss3_total.item() / data_length:.4f} | "
f"ADJ: {(LAMBDA_ADJ1 * adj1_total.item() + LAMBDA_ADJ2 * adj2_total.item() + LAMBDA_ADJ3 * adj3_total.item()) / data_length:.4f}, {adj1_total.item() / data_length:.4f}, {adj2_total.item() / data_length:.4f}, {adj3_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_PATCH1 * patch1_total.item() + LAMBDA_PATCH2 * patch2_total.item() + LAMBDA_PATCH3 * patch3_total.item()) / data_length:.4f}, {patch1_total.item() / data_length:.4f}, {patch2_total.item() / data_length:.4f}, {patch3_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"VQ: {commit_total.item() / data_length:.4f} | "
f"PPL: {val_perplexity:.4f} | " f"PPL: {val_perplexity:.4f} | "
f"Usage: {val_usage_rate:.4f} ({val_active_codes}/{code_hits_total.numel()}) | " f"Usage: {val_usage_rate:.4f} ({val_active_codes}/{TOTAL_K}) | "
) )
ckpt_path = f"result/seperated/sep-{epoch + 1}.pth" ckpt_path = f"result/seperated/sep-{epoch + 1}.pth"