mirror of
https://github.com/unanmed/ginka-generator.git
synced 2026-08-14 18:12:28 +08:00
feat: 墙壁生成策略调整
This commit is contained in:
parent
46d1f1b450
commit
103b02cf08
367
docs/wall-growth-generation-design.md
Normal file
367
docs/wall-growth-generation-design.md
Normal 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=1(WALL),并在整个过程中保持不可修改
|
||||
```
|
||||
|
||||
### 主算法
|
||||
|
||||
```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)预测。另一种思路是同时接受 FLOOR(tile=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] bool,True=被矩形掩码覆盖
|
||||
# target1: [H, W] int,GT 墙壁骨架(仅含 0=floor, 1=wall)
|
||||
# inp1: [H, W] int,输入地图(MASK 区域为 MASK_ID,其余保留原值)
|
||||
# 返回: [H, W] bool,True=该位置需要计算 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 使用 GT(target1)而非输入(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 输出
|
||||
- [ ] 对比指标:墙壁连通率、孤立墙块数、房间闭合度
|
||||
@ -5,6 +5,53 @@ import numpy as np
|
||||
from torch.utils.data import Dataset
|
||||
from shared.distance import compute_distance_field
|
||||
|
||||
def rect_mask(ratio: float, map_size: int = 169) -> np.ndarray:
|
||||
# 连续矩形分块掩码,反复放置随机矩形直到掩码格数达标
|
||||
# 复用 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):
|
||||
with open(path, 'r', encoding="utf-8") as f:
|
||||
data = json.load(f)
|
||||
@ -179,7 +226,7 @@ class GinkaSeperatedDataset(Dataset):
|
||||
return target1, inp1, target2, inp2, target3, inp3
|
||||
|
||||
def apply_subset1(self, raw: np.ndarray):
|
||||
# 子集 1:std_mask 随机掩码
|
||||
# 子集 1:Stage1 使用 rect_mask,Stage2/3 保持 std_mask
|
||||
|
||||
target1, inp1, target2, inp2, target3, inp3 = self.create_degreaded(raw)
|
||||
|
||||
@ -187,8 +234,11 @@ class GinkaSeperatedDataset(Dataset):
|
||||
enc2 = inp2.copy()
|
||||
enc3 = raw.copy()
|
||||
|
||||
# stage1:对整图 std_mask
|
||||
inp1[self.std_mask()] = self.MASK_ID
|
||||
# stage1:rect_mask(连续矩形掩码)替代 std_mask
|
||||
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
|
||||
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()
|
||||
enc3 = raw.copy()
|
||||
|
||||
need_mask = np.isin(inp2, [self.FLOOR, self.WALL])
|
||||
inp1[need_mask & self.std_mask()] = self.MASK_ID
|
||||
ratio = float(np.random.beta(2, 2)) * 0.95 + 0.05
|
||||
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])
|
||||
inp2[need_mask] = self.MASK_ID
|
||||
need_mask = np.isin(inp3, [self.FLOOR, self.RESOURCE])
|
||||
|
||||
@ -37,31 +37,49 @@ from shared.distance import DIST_VOCAB, compute_distance_field_tensor
|
||||
# 图块 ID 定义:
|
||||
# 0. 空地 1. 墙壁 2. 普通门 3. 资源 4. 怪物 5. 入口 6. 机关门 7. 掩码(MASK_TOKEN)
|
||||
|
||||
# 共用 VQ-VAE 超参
|
||||
# 三组编码器(vq1/vq2/vq3)共享相同超参,分别对三阶段地图上下文独立编码
|
||||
VQ_L = 8 # 码字序列长度(每个编码器输出 L 个码字,量化后合并为 L*3)
|
||||
VQ_K = 16 # codebook 大小(离散码本条目数)
|
||||
# 共用 VQ-VAE 超参(共享的编码维度)
|
||||
VQ_D_Z = 64 # 码字维度
|
||||
VQ_BETA = 1.0 # commit loss 权重(防止编码器输出漂离 codebook)
|
||||
VQ_GAMMA = 0.0 # entropy loss 权重(当前未启用)
|
||||
VQ_LAYERS = 6 # VQ-VAE Transformer 层数
|
||||
VQ_DIM_FF = 1024 # VQ-VAE 前馈网络隐层维度
|
||||
VQ_D_MODEL = 256 # VQ-VAE Transformer 模型维度
|
||||
VQ_NHEAD = 4 # VQ-VAE 多头注意力头数
|
||||
|
||||
# 三通道 VQ 各自独立超参(L、K、层数、维度等均独立配置)
|
||||
# Stage1 墙壁骨架 — 结构最复杂,模型容量最大
|
||||
VQ_L1 = 24
|
||||
VQ_K1 = 32
|
||||
VQ_D_MODEL1 = 384
|
||||
VQ_NHEAD1 = 8
|
||||
VQ_LAYERS1 = 6
|
||||
VQ_DIM_FF1 = 1536
|
||||
|
||||
# Stage2 功能元素 — 中等复杂度
|
||||
VQ_L2 = 12
|
||||
VQ_K2 = 16
|
||||
VQ_D_MODEL2 = 256
|
||||
VQ_NHEAD2 = 4
|
||||
VQ_LAYERS2 = 6
|
||||
VQ_DIM_FF2 = 1024
|
||||
|
||||
# Stage3 资源分布 — 最简单,模型容量最小
|
||||
VQ_L3 = 8
|
||||
VQ_K3 = 16
|
||||
VQ_D_MODEL3 = 192
|
||||
VQ_NHEAD3 = 4
|
||||
VQ_LAYERS3 = 4
|
||||
VQ_DIM_FF3 = 768
|
||||
L_DIST = 8 # 距离场码字序列长度
|
||||
K_DIST = 16 # 距离场 codebook 大小
|
||||
|
||||
# 距离场编码器超参
|
||||
L_DIST = 4 # 距离场码字序列长度
|
||||
K_DIST = 16 # 距离场 codebook 大小
|
||||
DIST_D_MODEL = 256 # 距离场编码器模型维度
|
||||
DIST_LAYERS = 3 # 距离场编码器 Transformer 层数
|
||||
DIST_DIM_FF = 1024 # 距离场编码器 FF 维度
|
||||
DIST_D_MODEL = 384 # 距离场编码器模型维度
|
||||
DIST_LAYERS = 6 # 距离场编码器 Transformer 层数
|
||||
DIST_DIM_FF = 1536 # 距离场编码器 FF 维度
|
||||
DIST_NHEAD = 8 # 距离场编码器注意力头数
|
||||
VQ_BETA_DIST = 0.5 # 距离场 commit loss 权重
|
||||
|
||||
# 第一阶段 MaskGIT 超参
|
||||
STAGE1_MG_DMODEL = 512
|
||||
STAGE1_MG_NHEAD = 4
|
||||
STAGE1_MG_NUM_LAYERS = 6
|
||||
STAGE1_MG_NUM_LAYERS = 8
|
||||
STAGE1_MG_DIM_FF = 2048
|
||||
|
||||
# 第二阶段 MaskGIT 超参
|
||||
@ -89,6 +107,7 @@ STAGE3_VQ_WEIGHT = 0.5
|
||||
# 全局参数
|
||||
NUM_CLASSES = 8 # 图块类型数
|
||||
MASK_TOKEN = 7 # 掩码图块
|
||||
TOTAL_K = VQ_K1 + VQ_K2 + VQ_K3 # 预计算,供日志输出使用
|
||||
MAP_W = 13 # 地图宽度
|
||||
MAP_H = 13 # 地图高度
|
||||
MAP_SIZE = MAP_W * MAP_H # 地图大小
|
||||
@ -106,14 +125,14 @@ MG_Z_DROPOUT = 0.1 # z 隐变量 Dropout 概率
|
||||
MG_STRUCT_DROPOUT = 0.1 # 结构参量 Dropout 概率
|
||||
|
||||
# 邻接损失权重(三阶段)
|
||||
LAMBDA_ADJ1 = 0.6
|
||||
LAMBDA_ADJ2 = 0.3
|
||||
LAMBDA_ADJ3 = 0.1
|
||||
LAMBDA_ADJ1 = 0.2
|
||||
LAMBDA_ADJ2 = 0.1
|
||||
LAMBDA_ADJ3 = 0.05
|
||||
|
||||
# Patch 损失权重(三阶段)及核参数
|
||||
LAMBDA_PATCH1 = 0.5
|
||||
LAMBDA_PATCH2 = 0.5
|
||||
LAMBDA_PATCH3 = 0.5
|
||||
LAMBDA_PATCH1 = 0.2
|
||||
LAMBDA_PATCH2 = 0.2
|
||||
LAMBDA_PATCH3 = 0.2
|
||||
PATCH_KERNEL_SIZE = 5
|
||||
PATCH_SIGMA = 1.2
|
||||
|
||||
@ -153,37 +172,44 @@ def parse_arguments():
|
||||
return parser.parse_args()
|
||||
|
||||
def build_model(device: torch.device):
|
||||
# 三组 VQ-VAE 编码器:各自独立编码一个阶段的地图上下文(encoder_stage1/2/3)
|
||||
# 输出形状均为 [B, L, d_z],分别送入各自阶段的 quantizer
|
||||
vq_kwargs = dict(
|
||||
num_classes=NUM_CLASSES, L=VQ_L, K=VQ_K, d_model=VQ_D_MODEL,
|
||||
nhead=VQ_NHEAD, num_layers=VQ_LAYERS, dim_ff=VQ_DIM_FF, map_h=MAP_H, map_w=MAP_W
|
||||
)
|
||||
vq1 = GinkaVQVAE(**vq_kwargs).to(device) # 编码 stage1 上下文(floor/wall)
|
||||
vq2 = GinkaVQVAE(**vq_kwargs).to(device) # 编码 stage2 上下文(door/monster/entrance)
|
||||
vq3 = GinkaVQVAE(**vq_kwargs).to(device) # 编码 stage3 上下文(resource)
|
||||
# 三组 VQ-VAE 编码器:各自独立编码一个阶段的地图上下文,全部超参独立
|
||||
vq1 = GinkaVQVAE(
|
||||
num_classes=NUM_CLASSES, L=VQ_L1, K=VQ_K1,
|
||||
d_model=VQ_D_MODEL1, nhead=VQ_NHEAD1, num_layers=VQ_LAYERS1, dim_ff=VQ_DIM_FF1,
|
||||
map_h=MAP_H, map_w=MAP_W
|
||||
).to(device)
|
||||
vq2 = GinkaVQVAE(
|
||||
num_classes=NUM_CLASSES, L=VQ_L2, K=VQ_K2,
|
||||
d_model=VQ_D_MODEL2, nhead=VQ_NHEAD2, num_layers=VQ_LAYERS2, dim_ff=VQ_DIM_FF2,
|
||||
map_h=MAP_H, map_w=MAP_W
|
||||
).to(device)
|
||||
vq3 = GinkaVQVAE(
|
||||
num_classes=NUM_CLASSES, L=VQ_L3, K=VQ_K3,
|
||||
d_model=VQ_D_MODEL3, nhead=VQ_NHEAD3, num_layers=VQ_LAYERS3, dim_ff=VQ_DIM_FF3,
|
||||
map_h=MAP_H, map_w=MAP_W
|
||||
).to(device)
|
||||
|
||||
# 三个独立 MaskGIT 解码器,分别接收各自阶段的 z_q 作为条件
|
||||
mg1 = GinkaMaskGIT(
|
||||
num_classes=NUM_CLASSES, d_model=STAGE1_MG_DMODEL, d_z=VQ_D_Z, dim_ff=STAGE1_MG_DIM_FF,
|
||||
nhead=STAGE1_MG_NHEAD, num_layers=STAGE1_MG_NUM_LAYERS, map_h=MAP_H, map_w=MAP_W,
|
||||
z_seq_len=VQ_L
|
||||
z_seq_len=VQ_L1, z_dist_len=L_DIST
|
||||
).to(device)
|
||||
mg2 = GinkaMaskGIT(
|
||||
num_classes=NUM_CLASSES, d_model=STAGE2_MG_DMODEL, d_z=VQ_D_Z, dim_ff=STAGE2_MG_DIM_FF,
|
||||
nhead=STAGE2_MG_NHEAD, num_layers=STAGE2_MG_NUM_LAYERS, map_h=MAP_H, map_w=MAP_W,
|
||||
z_seq_len=VQ_L
|
||||
z_seq_len=VQ_L2, z_dist_len=L_DIST
|
||||
).to(device)
|
||||
mg3 = GinkaMaskGIT(
|
||||
num_classes=NUM_CLASSES, d_model=STAGE3_MG_DMODEL, d_z=VQ_D_Z, dim_ff=STAGE3_MG_DIM_FF,
|
||||
nhead=STAGE3_MG_NHEAD, num_layers=STAGE3_MG_NUM_LAYERS, map_h=MAP_H, map_w=MAP_W,
|
||||
z_seq_len=VQ_L
|
||||
z_seq_len=VQ_L3, z_dist_len=L_DIST
|
||||
).to(device)
|
||||
|
||||
# 三个独立 VectorQuantizer:各阶段使用自己的码本,避免语义空间相互干扰
|
||||
quantizer1 = VectorQuantizer(K=VQ_K, d_z=VQ_D_Z).to(device)
|
||||
quantizer2 = VectorQuantizer(K=VQ_K, d_z=VQ_D_Z).to(device)
|
||||
quantizer3 = VectorQuantizer(K=VQ_K, d_z=VQ_D_Z).to(device)
|
||||
# 三个独立 VectorQuantizer:各阶段使用自己的码本大小
|
||||
quantizer1 = VectorQuantizer(K=VQ_K1, d_z=VQ_D_Z).to(device)
|
||||
quantizer2 = VectorQuantizer(K=VQ_K2, d_z=VQ_D_Z).to(device)
|
||||
quantizer3 = VectorQuantizer(K=VQ_K3, d_z=VQ_D_Z).to(device)
|
||||
quantizers = (quantizer1, quantizer2, quantizer3)
|
||||
|
||||
# 距离场编码器与量化器:将 L1 距离场编码为离散 latent z_dist
|
||||
@ -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
|
||||
return z_q * (~mask).float() + mask_embedding * mask.float()
|
||||
|
||||
def summarize_codebook_hits(code_hits: torch.Tensor) -> tuple[float, float, int]:
|
||||
total_hits = code_hits.sum()
|
||||
def summarize_codebook_hits(code_hits) -> tuple[float, float, int]:
|
||||
# code_hits 为 tuple of 3 tensors(各量器不同 K)
|
||||
combined = torch.cat([h.flatten() for h in code_hits], dim=0)
|
||||
total_hits = combined.sum()
|
||||
if total_hits.item() <= 0:
|
||||
return 0.0, 0.0, 0
|
||||
|
||||
probs = code_hits / total_hits
|
||||
probs = combined / total_hits
|
||||
perplexity = torch.exp(
|
||||
-(probs * torch.log(probs.clamp_min(1e-10))).sum()
|
||||
).item()
|
||||
active_codes = int((code_hits > 0).sum().item())
|
||||
usage_rate = active_codes / code_hits.numel()
|
||||
active_codes = int((combined > 0).sum().item())
|
||||
usage_rate = active_codes / combined.numel()
|
||||
return perplexity, usage_rate, active_codes
|
||||
|
||||
def quantize_stage_latents(
|
||||
@ -316,7 +344,7 @@ def quantize_stage_latents(
|
||||
z_q3, _, commit_loss3, _, code_hits3 = quantizer3(z_e3)
|
||||
|
||||
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
|
||||
|
||||
def build_reference_rollout_steps(prob: float) -> int:
|
||||
@ -431,6 +459,71 @@ def rect_mask(
|
||||
mask[x:x + bh, y:y + bw] = True
|
||||
return mask
|
||||
|
||||
def compute_adjacency_mask(flat_state: torch.Tensor) -> torch.Tensor:
|
||||
# 返回 [MAP_SIZE] bool,True 表示该位置与任意墙壁 4-邻接
|
||||
# flat_state: [MAP_SIZE] 整数(可选 batch 维度的话取第 0 行)
|
||||
if flat_state.dim() > 1:
|
||||
flat_state = flat_state[0]
|
||||
state_2d = flat_state.reshape(MAP_H, MAP_W)
|
||||
wall = (state_2d == 1)
|
||||
adj = torch.zeros_like(wall, dtype=torch.bool)
|
||||
adj[:, 1:] |= wall[:, :-1]
|
||||
adj[:, :-1] |= wall[:, 1:]
|
||||
adj[1:, :] |= wall[:-1, :]
|
||||
adj[:-1, :] |= wall[1:, :]
|
||||
return adj.flatten()
|
||||
|
||||
def wall_growth_sample(
|
||||
model: torch.nn.Module,
|
||||
z: torch.Tensor,
|
||||
z_dist: torch.Tensor,
|
||||
struct: torch.Tensor,
|
||||
target_density: torch.Tensor,
|
||||
max_steps: int = 24
|
||||
) -> np.ndarray:
|
||||
# 生长算法:初始外圈墙壁固定,仅邻接一圈为 MASK,其余为 FLOOR(0)
|
||||
# 每步 MASK 位置决策(墙/非墙)后,新邻接面成为下一轮 MASK,逐步外扩
|
||||
outer_idx = []
|
||||
for r in range(MAP_H):
|
||||
for c in range(MAP_W):
|
||||
if r == 0 or r == MAP_H - 1 or c == 0 or c == MAP_W - 1:
|
||||
outer_idx.append(r * MAP_W + c)
|
||||
outer_idx = torch.tensor(outer_idx, dtype=torch.long, device=z.device)
|
||||
|
||||
state = torch.full((1, MAP_SIZE), 0, dtype=torch.long, device=z.device)
|
||||
state[0, outer_idx] = 1
|
||||
# 初始 MASK:邻接外圈墙壁的一圈
|
||||
init_adj = compute_adjacency_mask(state[0])
|
||||
state[0, init_adj & (state[0] == 0)] = MASK_TOKEN
|
||||
|
||||
for step in range(max_steps):
|
||||
mask_pos = state[0] == MASK_TOKEN
|
||||
if mask_pos.sum() == 0:
|
||||
break
|
||||
|
||||
remain = compute_remaining(state, target_density, 1)
|
||||
logits = model(state, z, z_dist, struct, remain)
|
||||
probs = F.softmax(logits, dim=-1)
|
||||
|
||||
was_mask = mask_pos.clone()
|
||||
mask_idx = mask_pos.nonzero(as_tuple=True)[0]
|
||||
wall_prob = probs[0, mask_idx, 1]
|
||||
hits = wall_prob > 0.5
|
||||
if hits.any():
|
||||
state[0, mask_idx[hits]] = 1
|
||||
# 本轮 MASK 决策完毕:未提交的 MASK → FLOOR(0)
|
||||
state[0, mask_idx[~hits]] = 0
|
||||
|
||||
# 下一轮 MASK:邻接墙壁且本轮不是 MASK(即新触及的 FLOOR)
|
||||
adj = compute_adjacency_mask(state[0])
|
||||
new_mask = adj & (state[0] == 0) & (~was_mask)
|
||||
if new_mask.sum() == 0:
|
||||
break
|
||||
state[0, new_mask] = MASK_TOKEN
|
||||
|
||||
state[0, state[0] == MASK_TOKEN] = 0
|
||||
return state[0].cpu().numpy().reshape(MAP_H, MAP_W)
|
||||
|
||||
def maskgit_sample(
|
||||
model: torch.nn.Module, inp: torch.Tensor, z: torch.Tensor,
|
||||
z_dist: torch.Tensor, struct: torch.Tensor, target_density: torch.Tensor,
|
||||
@ -535,15 +628,14 @@ def full_generate_specific_z(
|
||||
device: torch.device,
|
||||
keep_fixed: tuple[bool, bool, bool] = (True, True, True)
|
||||
) -> tuple:
|
||||
vq1, vq2, vq3, mg1, mg2, mg3, quantizers, optimizer, scheduler = models
|
||||
vq1, vq2, vq3, mg1, mg2, mg3, quantizers, _, _, _, _, _ = models
|
||||
|
||||
with torch.no_grad():
|
||||
z1, z2, z3 = z_q
|
||||
|
||||
# 三阶段级联生成,但使用给定的 z
|
||||
pred1_np = maskgit_sample(
|
||||
mg1, input.clone(), z1, z_dist, struct, target_density, 1,
|
||||
GENERATE_STEP, target_tiles=[1], keep_fixed=keep_fixed[0]
|
||||
# 三阶段级联生成:Stage1 使用生长算法,Stage2/3 保持 MaskGIT
|
||||
pred1_np = wall_growth_sample(
|
||||
mg1, z1, z_dist, struct, target_density
|
||||
)
|
||||
inp2 = torch.tensor(pred1_np.flatten(), dtype=torch.long, device=device).reshape(1, MAP_SIZE)
|
||||
inp2[inp2 == 0] = MASK_TOKEN
|
||||
@ -579,21 +671,50 @@ def inpaint_generate(
|
||||
# 三阶段矩形掩码修补生成
|
||||
# - 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
|
||||
|
||||
# Stage1: 补全墙壁
|
||||
# Stage1: 补全墙壁 — 生长方式,掩码区域内仅邻接圈为 MASK,其余为 FLOOR(0)
|
||||
inp1 = raw_map.copy()
|
||||
inp1[mask] = MASK_TOKEN
|
||||
non_wf = (inp1 != 0) & (inp1 != 1) & (inp1 != MASK_TOKEN)
|
||||
non_wf = (inp1 != 0) & (inp1 != 1)
|
||||
inp1[non_wf] = 0
|
||||
# 掩码区域先全设为 FLOOR(0),再找出邻接圈设为 MASK
|
||||
inp1[mask] = 0
|
||||
inp1_t = torch.tensor(inp1.flatten(), dtype=torch.long, device=device).reshape(1, MAP_SIZE)
|
||||
# 初始 MASK:掩码区域内与固定墙壁邻接的一圈
|
||||
mask_t = torch.tensor(mask.flatten(), dtype=torch.bool, device=device)
|
||||
init_state = inp1_t.clone()
|
||||
init_adj = compute_adjacency_mask(init_state[0])
|
||||
init_mask = init_adj & (init_state[0] == 0) & mask_t
|
||||
init_state[0, init_mask] = MASK_TOKEN
|
||||
|
||||
with torch.no_grad():
|
||||
pred1_np = maskgit_sample(
|
||||
mg1, inp1_t, z1, z_dist, struct, target_density, 1,
|
||||
GENERATE_STEP, target_tiles=[1], keep_fixed=True
|
||||
)
|
||||
state = init_state
|
||||
for step in range(GENERATE_STEP):
|
||||
mask_pos = state[0] == MASK_TOKEN
|
||||
if mask_pos.sum() == 0:
|
||||
break
|
||||
|
||||
remain = compute_remaining(state, target_density, 1)
|
||||
logits = mg1(state, z1, z_dist, struct, remain)
|
||||
probs = F.softmax(logits, dim=-1)
|
||||
|
||||
was_mask = mask_pos.clone()
|
||||
mask_idx = mask_pos.nonzero(as_tuple=True)[0]
|
||||
wall_prob = probs[0, mask_idx, 1]
|
||||
hits = wall_prob > 0.5
|
||||
if hits.any():
|
||||
state[0, mask_idx[hits]] = 1
|
||||
state[0, mask_idx[~hits]] = 0
|
||||
|
||||
adj = compute_adjacency_mask(state[0])
|
||||
new_mask = adj & (state[0] == 0) & (~was_mask) & mask_t
|
||||
if new_mask.sum() == 0:
|
||||
break
|
||||
state[0, new_mask] = MASK_TOKEN
|
||||
|
||||
state[0, state[0] == MASK_TOKEN] = 0
|
||||
pred1_np = state[0].cpu().numpy().reshape(MAP_H, MAP_W)
|
||||
|
||||
# Stage2: 补全门/怪/入口,保留未掩码区的原始非墙壁结构
|
||||
merged_s1 = pred1_np.copy()
|
||||
@ -669,7 +790,7 @@ def build_dataset_sample_case(
|
||||
device: torch.device,
|
||||
idx: int | None = None
|
||||
) -> 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
|
||||
sample = dataset.random_sample_map(idx=idx)
|
||||
|
||||
@ -805,15 +926,19 @@ def visualize_rand(
|
||||
def to_img(mat):
|
||||
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
|
||||
|
||||
samples_data = []
|
||||
for _ in range(4):
|
||||
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)
|
||||
mask = rect_mask(ratio)
|
||||
# 验证时同样修正掩码,避免孤立墙壁分量完全被盖住
|
||||
target1 = sample["encoder_stage1"].cpu().numpy().reshape(MAP_H, MAP_W)
|
||||
from ginka.dataset import ensure_wall_connection
|
||||
mask = ensure_wall_connection(mask, target1)
|
||||
|
||||
enc1_t = sample["encoder_stage1"].to(device).reshape(1, MAP_SIZE)
|
||||
enc2_t = sample["encoder_stage2"].to(device).reshape(1, MAP_SIZE)
|
||||
@ -885,7 +1010,7 @@ def validate(
|
||||
train_dataset: GinkaSeperatedDataset,
|
||||
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
|
||||
quantizer1, quantizer2, quantizer3 = quantizers
|
||||
|
||||
@ -904,7 +1029,7 @@ def validate(
|
||||
patch1_total = torch.Tensor([0]).to(device)
|
||||
patch2_total = torch.Tensor([0]).to(device)
|
||||
patch3_total = torch.Tensor([0]).to(device)
|
||||
code_hits_total = torch.zeros(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 = {
|
||||
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)
|
||||
patch2_total += patch_loss(logits2, target2, PATCH_KERNEL_SIZE, PATCH_SIGMA)
|
||||
patch3_total += patch_loss(logits3, target3, PATCH_KERNEL_SIZE, PATCH_SIGMA)
|
||||
code_hits_total += code_hits
|
||||
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()
|
||||
@ -1023,7 +1148,7 @@ def train(device: torch.device):
|
||||
|
||||
result = build_model(device)
|
||||
vq1, vq2, vq3, mg1, mg2, mg3, quantizers, optimizer, scheduler, latent_mask_embedding, dist_encoder, dist_quantizer = result
|
||||
models = [vq1, vq2, vq3, mg1, mg2, mg3, quantizers, optimizer, scheduler]
|
||||
models = [vq1, vq2, vq3, mg1, mg2, mg3, quantizers, optimizer, scheduler, latent_mask_embedding, dist_encoder, dist_quantizer]
|
||||
dist_models = (dist_encoder, dist_quantizer)
|
||||
quantizer1, quantizer2, quantizer3 = quantizers
|
||||
|
||||
@ -1106,7 +1231,7 @@ def train(device: torch.device):
|
||||
patch1_total = torch.Tensor([0]).to(device)
|
||||
patch2_total = torch.Tensor([0]).to(device)
|
||||
patch3_total = torch.Tensor([0]).to(device)
|
||||
code_hits_total = torch.zeros(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):
|
||||
# 三阶段各自的掩码输入序列、预测目标和编码器上下文
|
||||
@ -1127,7 +1252,7 @@ def train(device: torch.device):
|
||||
target_density = batch["target_density"].to(device)
|
||||
dist_field = batch["distance_field"].to(device)
|
||||
|
||||
optimizer.zero_grad()
|
||||
optimizer.zero_grad() # 训练循环
|
||||
|
||||
# VQ 编码:各阶段编码器分别处理各自上下文切片
|
||||
z_e1 = vq1(enc1) # [B, L, d_z]
|
||||
@ -1175,18 +1300,16 @@ def train(device: torch.device):
|
||||
loss2 = cross_entropy_loss(logits2, target2)
|
||||
loss3 = cross_entropy_loss(logits3, target3)
|
||||
|
||||
adj1 = adjacency_loss(logits1, target1)
|
||||
adj2 = adjacency_loss(logits2, target2)
|
||||
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)
|
||||
patch3 = patch_loss(logits3, target3, PATCH_KERNEL_SIZE, PATCH_SIGMA)
|
||||
|
||||
loss1_weighted = STAGE1_CE_WEIGHT * loss1
|
||||
loss2_weighted = STAGE2_CE_WEIGHT * loss2
|
||||
loss3_weighted = STAGE3_CE_WEIGHT * loss3
|
||||
adj_weighted = LAMBDA_ADJ1 * adj1 + LAMBDA_ADJ2 * adj2 + LAMBDA_ADJ3 * adj3
|
||||
patch_weighted = LAMBDA_PATCH1 * patch1 + LAMBDA_PATCH2 * patch2 + LAMBDA_PATCH3 * patch3
|
||||
adj_weighted = LAMBDA_ADJ2 * adj2 + LAMBDA_ADJ3 * adj3
|
||||
patch_weighted = LAMBDA_PATCH2 * patch2 + LAMBDA_PATCH3 * patch3
|
||||
commit_weighted = VQ_BETA * commit_loss + VQ_BETA_DIST * commit_loss_dist
|
||||
loss = loss1_weighted + loss2_weighted + loss3_weighted + adj_weighted + patch_weighted + commit_weighted
|
||||
|
||||
@ -1199,13 +1322,11 @@ def train(device: torch.device):
|
||||
loss2_total += loss2.detach()
|
||||
loss3_total += loss3.detach()
|
||||
commit_total += commit_loss.detach()
|
||||
adj1_total += adj1.detach()
|
||||
adj2_total += adj2.detach()
|
||||
adj3_total += adj3.detach()
|
||||
patch1_total += patch1.detach()
|
||||
patch2_total += patch2.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 结束后更新学习率
|
||||
scheduler.step()
|
||||
@ -1217,11 +1338,11 @@ def train(device: torch.device):
|
||||
f"E: {epoch + 1} | "
|
||||
f"Loss: {loss_total.item() / data_length:.4f} | "
|
||||
f"CE: {loss1_total.item() / data_length:.4f}, {loss2_total.item() / data_length:.4f}, {loss3_total.item() / data_length:.4f} | "
|
||||
f"ADJ: {(LAMBDA_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"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"ADJ: {(LAMBDA_ADJ2 * adj2_total.item() + LAMBDA_ADJ3 * adj3_total.item()) / data_length:.4f}, {adj2_total.item() / data_length:.4f}, {adj3_total.item() / data_length:.4f} | "
|
||||
f"PAT: {(LAMBDA_PATCH2 * patch2_total.item() + LAMBDA_PATCH3 * patch3_total.item()) / data_length:.4f}, {patch2_total.item() / data_length:.4f}, {patch3_total.item() / data_length:.4f} | "
|
||||
f"VQ: {commit_total.item() / data_length:.4f} | "
|
||||
f"PPL: {train_perplexity:.4f} | "
|
||||
f"Usage: {train_usage_rate:.4f} ({train_active_codes}/{code_hits_total.numel()}) | "
|
||||
f"Usage: {train_usage_rate:.4f} ({train_active_codes}/{TOTAL_K}) | "
|
||||
f"LR: {scheduler.get_last_lr()[0]:.6f}"
|
||||
)
|
||||
|
||||
@ -1230,12 +1351,12 @@ def train(device: torch.device):
|
||||
losses = validate(
|
||||
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
|
||||
loss2_weighted = STAGE2_CE_WEIGHT * loss2_total
|
||||
loss3_weighted = STAGE3_CE_WEIGHT * loss3_total
|
||||
adj_weighted = LAMBDA_ADJ1 * adj1_total + LAMBDA_ADJ2 * adj2_total + LAMBDA_ADJ3 * adj3_total
|
||||
patch_weighted = LAMBDA_PATCH1 * patch1_total + LAMBDA_PATCH2 * patch2_total + LAMBDA_PATCH3 * patch3_total
|
||||
adj_weighted = LAMBDA_ADJ2 * adj2_total + LAMBDA_ADJ3 * adj3_total
|
||||
patch_weighted = LAMBDA_PATCH2 * patch2_total + LAMBDA_PATCH3 * patch3_total
|
||||
commit_weighted = VQ_BETA * commit_total
|
||||
loss_total = loss1_weighted + loss2_weighted + loss3_weighted + adj_weighted + patch_weighted + commit_weighted
|
||||
|
||||
@ -1246,11 +1367,11 @@ def train(device: torch.device):
|
||||
f"E: {epoch + 1} | "
|
||||
f"Loss: {loss_total.item() / data_length:.4f} | "
|
||||
f"CE: {loss1_total.item() / data_length:.4f}, {loss2_total.item() / data_length:.4f}, {loss3_total.item() / data_length:.4f} | "
|
||||
f"ADJ: {(LAMBDA_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"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"ADJ: {(LAMBDA_ADJ2 * adj2_total.item() + LAMBDA_ADJ3 * adj3_total.item()) / data_length:.4f}, {adj2_total.item() / data_length:.4f}, {adj3_total.item() / data_length:.4f} | "
|
||||
f"PAT: {(LAMBDA_PATCH2 * patch2_total.item() + LAMBDA_PATCH3 * patch3_total.item()) / data_length:.4f}, {patch2_total.item() / data_length:.4f}, {patch3_total.item() / data_length:.4f} | "
|
||||
f"VQ: {commit_total.item() / data_length:.4f} | "
|
||||
f"PPL: {val_perplexity:.4f} | "
|
||||
f"Usage: {val_usage_rate:.4f} ({val_active_codes}/{code_hits_total.numel()}) | "
|
||||
f"Usage: {val_usage_rate:.4f} ({val_active_codes}/{TOTAL_K}) | "
|
||||
)
|
||||
|
||||
ckpt_path = f"result/seperated/sep-{epoch + 1}.pth"
|
||||
|
||||
Loading…
Reference in New Issue
Block a user