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
f5f7193df8
commit
46d1f1b450
405
docs/adjacency-patch-loss-design.md
Normal file
405
docs/adjacency-patch-loss-design.md
Normal file
@ -0,0 +1,405 @@
|
||||
# 邻接损失与 Patch 损失设计文档
|
||||
|
||||
## 1. 问题分析
|
||||
|
||||
当前三阶段 MaskGIT 训练仅使用逐格 Cross Entropy(CE)损失,每个格子的损失完全独立计算。然而地图生成任务中不同格子之间存在强烈的空间相关性:
|
||||
|
||||
- **邻接关系**:两格是否同为可通行区域(地板)是地图连通性的基础,孤立的地板块意味着该区域无法到达
|
||||
- **局部一致性**:一个格子的类别高度依赖其周围格子的布局,例如门必须出现在墙与地板的交界处、怪物通常成簇分布
|
||||
|
||||
CE 损失将每个格子的预测视为独立的多分类问题,无法显式建模这些空间依赖,导致生成的地图在连通性和局部结构上表现不佳。
|
||||
|
||||
本方案引入两种互补的空间损失函数:
|
||||
|
||||
| 损失类型 | 作用范围 | 核心机制 | 主要受益阶段 |
|
||||
| -------- | -------- | -------- | ------------ |
|
||||
| 邻接损失 | 紧邻四邻域 | BCE 约束相邻两格同时为空地的概率 | 阶段一(墙壁骨架) |
|
||||
| Patch 损失 | K×K 邻域 | 高斯核加权 CE 平滑,使损失梯度在空间上扩散 | 全部三个阶段 |
|
||||
|
||||
## 2. 邻接损失
|
||||
|
||||
### 2.1 设计思路
|
||||
|
||||
对于地图中的每一条相邻边(共 312 条:13×12 水平 + 12×13 垂直),计算两端格子均为空地的联合概率,并与真实标注做 BCE(二元交叉熵)对比。
|
||||
|
||||
$$
|
||||
P(\text{both floor})_{(u,v)} = P(\text{floor} \mid u) \cdot P(\text{floor} \mid v)
|
||||
$$
|
||||
|
||||
其中 $P(\text{floor} \mid u) = \text{softmax}(\text{logits}_u)[\text{class}=0]$,即第 0 类(地板)的 softmax 概率。
|
||||
|
||||
真实标注为二元值:若目标地图中两格均为地板(ID=0),则为 1;否则为 0。
|
||||
|
||||
**直觉解释**:
|
||||
- 若两格都是地板,BCE 推动二者的 floor 概率同时提高,强化连通区域的一致性
|
||||
- 若一格是墙、一格是地板,BCE 推动墙格的 floor 概率降低,与 CE 损失形成梯度叠加
|
||||
- 若两格都是墙,BCE 推动二者的 floor 概率同时降低,与 CE 一致
|
||||
|
||||
邻接损失的核心价值不在于引入新信息,而在于**改变损失曲面的几何形状**——通过将逐格独立损失改为逐边联合损失,使相邻格子的梯度互相耦合。当一个格子的预测出错时,其四个邻居的梯度也会受到影响,从而加速模型学习空间一致性。
|
||||
|
||||
### 2.2 数学定义
|
||||
|
||||
设地图尺寸为 $H \times W = 13 \times 13$,logits 形状为 $[B, S, C]$,目标形状为 $[B, S]$($S = 169$,$C = 8$)。
|
||||
|
||||
对水平边 $(i,j) \leftrightarrow (i,j+1)$(共 $H \times (W-1) = 156$ 条):
|
||||
|
||||
$$
|
||||
\mathcal{L}_{\text{adj}}^h = \frac{1}{B \cdot H \cdot (W-1)} \sum_b \sum_{i=0}^{H-1} \sum_{j=0}^{W-2} \text{BCE}\left(
|
||||
p_{b,i,j}^f \cdot p_{b,i,j+1}^f,\;
|
||||
\mathbb{I}[t_{b,i,j} = 0 \land t_{b,i,j+1} = 0]
|
||||
\right)
|
||||
$$
|
||||
|
||||
对垂直边 $(i,j) \leftrightarrow (i+1,j)$(共 $(H-1) \times W = 156$ 条):
|
||||
|
||||
$$
|
||||
\mathcal{L}_{\text{adj}}^v = \frac{1}{B \cdot (H-1) \cdot W} \sum_b \sum_{i=0}^{H-2} \sum_{j=0}^{W-1} \text{BCE}\left(
|
||||
p_{b,i,j}^f \cdot p_{b,i+1,j}^f,\;
|
||||
\mathbb{I}[t_{b,i,j} = 0 \land t_{b,i+1,j} = 0]
|
||||
\right)
|
||||
$$
|
||||
|
||||
总邻接损失为二者均值:
|
||||
|
||||
$$
|
||||
\mathcal{L}_{\text{adj}} = \frac{1}{2} (\mathcal{L}_{\text{adj}}^h + \mathcal{L}_{\text{adj}}^v)
|
||||
$$
|
||||
|
||||
### 2.3 实现
|
||||
|
||||
```python
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
def adjacency_loss(logits, target):
|
||||
# logits: [B, S, C] — MaskGIT 解码器输出
|
||||
# target: [B, S] — 目标类别 ID,不含 MASK 标记
|
||||
B, S, C = logits.shape
|
||||
H = W = 13
|
||||
probs = F.softmax(logits, dim=-1) # [B, S, C]
|
||||
p_floor = probs[:, :, 0].view(B, H, W) # [B, H, W] — 地板概率
|
||||
t = target.view(B, H, W) # [B, H, W]
|
||||
t_floor = (t == 0).float() # 地板标注为 1,其余为 0
|
||||
|
||||
# 水平边:左格 × 右格
|
||||
joint_h = p_floor[:, :, :-1] * p_floor[:, :, 1:] # [B, H, W-1]
|
||||
target_h = t_floor[:, :, :-1] * t_floor[:, :, 1:] # [B, H, W-1]
|
||||
|
||||
# 垂直边:上格 × 下格
|
||||
joint_v = p_floor[:, :-1, :] * p_floor[:, 1:, :] # [B, H-1, W]
|
||||
target_v = t_floor[:, :-1, :] * t_floor[:, 1:, :] # [B, H-1, W]
|
||||
|
||||
loss_h = F.binary_cross_entropy(joint_h, target_h, reduction='mean')
|
||||
loss_v = F.binary_cross_entropy(joint_v, target_v, reduction='mean')
|
||||
return (loss_h + loss_v) / 2.0
|
||||
```
|
||||
|
||||
**边界处理**:只计算完整存在的边(水平 13×12=156 条,垂直 12×13=156 条),共计 312 条边。边界格子的某些方向没有邻居,自然不产生该方向的边。
|
||||
|
||||
**阶段适用性**:邻接损失直接对地板-地板连通性建模,在三个阶段中的适用程度不同:
|
||||
|
||||
| 阶段 | target 地板占比 | 邻接损失有效性 | 说明 |
|
||||
| ---- | --------------- | -------------- | ---- |
|
||||
| 一 | ~50% | 高 | 地板/墙壁二分类,地板连通性决定地图骨架质量 |
|
||||
| 二 | ~70% | 中 | 墙壁已降级为地板,地板占比更高,但仍有门与怪物放置的地板连通约束 |
|
||||
| 三 | ~90% | 低 | 仅保留地板与资源,大部分区域为地板,BCE 信号趋弱 |
|
||||
|
||||
**建议**:邻接损失默认仅作用于阶段一。若想对所有阶段生效,可为二、三阶段设置更小的权重 `LAMBDA_ADJ_STAGE2`、`LAMBDA_ADJ_STAGE3`。
|
||||
|
||||
### 2.4 超参数
|
||||
|
||||
| 参数 | 建议值 | 说明 |
|
||||
| ---- | ------ | ---- |
|
||||
| `LAMBDA_ADJ` | 0.1 | 邻接损失总权重,作用于阶段一 |
|
||||
| `LAMBDA_ADJ_STAGE2` | 0.05 | 阶段二邻接损失权重(可选,默认与 LAMBDA_ADJ 相同) |
|
||||
| `LAMBDA_ADJ_STAGE3` | 0.02 | 阶段三邻接损失权重(可选) |
|
||||
|
||||
### 2.5 扩展方向
|
||||
|
||||
当前方案仅建模地板-地板成对关系。可扩展为多类别邻接矩阵:
|
||||
|
||||
- 对 7 个有效类别两两组合得到 7×7 = 49 种邻接关系
|
||||
- 计算联合概率 `P(class_a | u) * P(class_b | v)` 得到 49 维向量
|
||||
- 与目标地图的 one-hot 邻接关系做 CE
|
||||
|
||||
多类邻接能显式建模"门必须出现在墙与地板交界""怪物应位于地板上"等约束,但计算量较大(49 × 312 = 15,288 对/样本),建议作为地板块邻接验证有效后的进阶方案。
|
||||
|
||||
## 3. Patch 损失
|
||||
|
||||
### 3.1 设计思路
|
||||
|
||||
逐格 CE 损失的梯度仅作用于单格。Patch 损失通过高斯核对每格的 CE 损失进行空间卷积,使一个格子的损失值等于其 K×K 邻域内所有格子 CE 的加权平均。这带来的效果是:
|
||||
|
||||
- 一个格子的预测错误会通过高斯核"扩散"到其邻居,邻居的梯度也会向纠错方向移动
|
||||
- 等效于增强了局部区域的梯度信号,使模型更难忽视局部上下文
|
||||
- 相比邻接损失关心特定类别关系(地板-地板),Patch 损失是**类别无关**的空间平滑,对所有 tile 类型均有效
|
||||
|
||||
**与邻接损失的区别**:
|
||||
|
||||
| 维度 | 邻接损失 | Patch 损失 |
|
||||
| ---- | -------- | ---------- |
|
||||
| 作用范围 | 仅紧邻四邻域(距离 1) | K×K 邻域(距离 2,可调) |
|
||||
| 类别感知 | 是(地板-地板特定) | 否(CE 已包含类别) |
|
||||
| 权重机制 | 等权(乘积) | 高斯衰减权重 |
|
||||
| 梯度耦合 | 乘性(P_A × P_B) | 加性(加权 CE 求和) |
|
||||
|
||||
### 3.2 数学定义
|
||||
|
||||
设 CE 损失在每格的定义为(不做 reduction):
|
||||
|
||||
$$
|
||||
\text{ce}_{b,i,j} = -\log \frac{\exp(\text{logits}_{b,i,j}[t_{b,i,j}])}{\sum_{c=0}^{C-1} \exp(\text{logits}_{b,i,j}[c])}
|
||||
$$
|
||||
|
||||
高斯核 $G \in \mathbb{R}^{K \times K}$,中心 $(k_c, k_c) = (\lfloor K/2 \rfloor, \lfloor K/2 \rfloor)$:
|
||||
|
||||
$$
|
||||
G_{u,v} = \frac{1}{Z} \exp\left(-\frac{(u - k_c)^2 + (v - k_c)^2}{2\sigma^2}\right), \quad Z = \sum_{u,v} G_{u,v}
|
||||
$$
|
||||
|
||||
Patch 损失定义为逐格平滑后的 CE 的均值:
|
||||
|
||||
$$
|
||||
\mathcal{L}_{\text{patch}} = \frac{1}{B \cdot H \cdot W} \sum_b \sum_{i=0}^{H-1} \sum_{j=0}^{W-1} \sum_{u=0}^{K-1} \sum_{v=0}^{K-1}
|
||||
G_{u,v} \cdot \text{ce}_{b,\; i+u-k_c,\; j+v-k_c}
|
||||
$$
|
||||
|
||||
其中越界的 $\text{ce}$ 用边缘复制(replicate padding)填充。
|
||||
|
||||
### 3.3 实现
|
||||
|
||||
```python
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
def gaussian_kernel(kernel_size, sigma, device):
|
||||
# 生成归一化二维高斯卷积核 [1, 1, K, K]
|
||||
k = kernel_size
|
||||
center = (k - 1) / 2.0
|
||||
xs = torch.arange(k, dtype=torch.float32, device=device) - center
|
||||
gx = torch.exp(-xs ** 2 / (2.0 * sigma ** 2))
|
||||
gy = torch.exp(-xs ** 2 / (2.0 * sigma ** 2))
|
||||
g2d = gx[:, None] * gy[None, :] # [K, K]
|
||||
g2d = g2d / g2d.sum() # 归一化
|
||||
return g2d.view(1, 1, k, k)
|
||||
|
||||
def patch_loss(logits, target, kernel_size=5, sigma=1.2):
|
||||
# logits: [B, S, C]
|
||||
# target: [B, S]
|
||||
B, S, C = logits.shape
|
||||
H = W = 13
|
||||
|
||||
# 逐格 CE(不做 reduction)
|
||||
ce = F.cross_entropy(
|
||||
logits.reshape(-1, C), target.reshape(-1), reduction='none'
|
||||
).view(B, H, W) # [B, H, W]
|
||||
|
||||
# 高斯核
|
||||
kernel = gaussian_kernel(kernel_size, sigma, logits.device) # [1, 1, K, K]
|
||||
|
||||
# replicate 填充后用 unfold 提取邻域
|
||||
pad = kernel_size // 2
|
||||
ce_padded = F.pad(ce.view(B, 1, H, W), (pad, pad, pad, pad), mode='replicate')
|
||||
# patches: [B, K*K, H*W]
|
||||
patches = F.unfold(ce_padded, kernel_size=(kernel_size, kernel_size))
|
||||
patches = patches.view(B, kernel_size * kernel_size, H, W) # [B, K*K, H, W]
|
||||
|
||||
# 加权求和
|
||||
k_flat = kernel.view(1, kernel_size * kernel_size, 1, 1) # [1, K*K, 1, 1]
|
||||
smoothed = (patches * k_flat).sum(dim=1) # [B, H, W]
|
||||
|
||||
return smoothed.mean()
|
||||
```
|
||||
|
||||
**设计要点**:
|
||||
|
||||
- `replicate` 填充:边界格子的邻域超出地图范围时,用最近的边缘值填充,避免引入零值的虚假信号
|
||||
- 核归一化:高斯核归一化到和为 1,确保 Patch 损失的数值量级与原始 CE 大致相当,便于权重调参
|
||||
- `F.unfold` 实现:比手动循环或 `conv2d` 更高效,在 CUDA 上有针对性的算子优化
|
||||
|
||||
### 3.4 超参数
|
||||
|
||||
| 参数 | 建议值 | 说明 |
|
||||
| ---- | ------ | ---- |
|
||||
| `PATCH_KERNEL_SIZE` | 5 | 卷积核尺寸,5×5 覆盖约 15% 的地图面积 |
|
||||
| `PATCH_SIGMA` | 1.2 | 高斯核标准差,控制权重衰减速度。σ 越大平滑越强 |
|
||||
| `LAMBDA_PATCH` | 0.1 | Patch 损失总权重 |
|
||||
| `LAMBDA_PATCH_STAGE1` | 0.1 | 阶段一 Patch 权重(可选,默认同 LAMBDA_PATCH) |
|
||||
| `LAMBDA_PATCH_STAGE2` | 0.1 | 阶段二 Patch 权重 |
|
||||
| `LAMBDA_PATCH_STAGE3` | 0.1 | 阶段三 Patch 权重 |
|
||||
|
||||
**Kernel Size 选择**:13×13 地图上,5×5 是合理的起点(覆盖每个方向 ±2 格)。若效果不明显可尝试 7×7,但需注意边界效应增强。
|
||||
|
||||
**Sigma 选择**:σ = 1.2 在 5×5 内权重分布为:
|
||||
|
||||
| 距离 | 0 | 1 | 2 |
|
||||
| ---- | --- | --- | --- |
|
||||
| 权重 | ~0.273 | ~0.164 | ~0.053 |
|
||||
|
||||
中心权重约为边缘的 5 倍,保证自身 CE 仍占主导,同时给予邻居适度的梯度贡献。
|
||||
|
||||
### 3.5 扩展方向
|
||||
|
||||
**可学习卷积核**:将固定高斯核替换为可训练参数 `nn.Parameter`,在训练过程中自适应学习最优的空间权重分布。实现简单,只需将 `gaussian_kernel` 改为可学习参数即可。
|
||||
|
||||
**多尺度 Patch**:同时使用 3×3、5×5、7×7 三种核,各自的平滑结果取加权和,使模型在不同空间尺度上均受到约束。
|
||||
|
||||
**类感知核**:不同类别使用不同的平滑权重。例如墙壁类不应与怪物类互相平滑,可以按类别加权。实现时需在 CE 中保留类别维度,对每类分别卷积。
|
||||
|
||||
## 4. 总损失组合
|
||||
|
||||
最终训练损失为原有各项与新两项的加权和:
|
||||
|
||||
$$
|
||||
\begin{aligned}
|
||||
\mathcal{L}_{\text{total}} = &
|
||||
\sum_{s=1}^{3} w_{\text{ce}}^s \cdot \mathcal{L}_{\text{ce}}^s \\
|
||||
+ & \sum_{s=1}^{3} \lambda_{\text{adj}}^s \cdot \mathcal{L}_{\text{adj}}^s \\
|
||||
+ & \sum_{s=1}^{3} \lambda_{\text{patch}}^s \cdot \mathcal{L}_{\text{patch}}^s \\
|
||||
+ & \beta \cdot \mathcal{L}_{\text{commit}} + \beta_{\text{dist}} \cdot \mathcal{L}_{\text{commit\_dist}}
|
||||
\end{aligned}
|
||||
$$
|
||||
|
||||
其中 $w_{\text{ce}}^s$ 即现有的 `STAGE1_CE_WEIGHT` / `STAGE2_CE_WEIGHT` / `STAGE3_CE_WEIGHT`(默认均为 1.0)。
|
||||
|
||||
**推荐权重配置**:
|
||||
|
||||
| 损失项 | 阶段一 | 阶段二 | 阶段三 |
|
||||
| ------ | ------ | ------ | ------ |
|
||||
| CE 损失 | 1.0 | 1.0 | 1.0 |
|
||||
| 邻接损失 | 0.1 | 0.05 | 0.02 |
|
||||
| Patch 损失 | 0.1 | 0.1 | 0.1 |
|
||||
|
||||
阶段一地板/墙壁二分类是连通性建模的核心,邻接损失权重最高。阶段二、三地板占比渐增,邻接损失的判别力减弱,权重递减。Patch 损失作为类别无关的平滑约束,三个阶段等权重即可。
|
||||
|
||||
### 与现有超参数的位置关系
|
||||
|
||||
新增超参数应追加在 `train_seperated.py` 现有超参数块(第 42-118 行)之后:
|
||||
|
||||
```python
|
||||
# 邻接损失权重(三阶段)
|
||||
LAMBDA_ADJ1 = 0.1
|
||||
LAMBDA_ADJ2 = 0.05
|
||||
LAMBDA_ADJ3 = 0.02
|
||||
|
||||
# Patch 损失权重(三阶段)及核参数
|
||||
LAMBDA_PATCH1 = 0.1
|
||||
LAMBDA_PATCH2 = 0.1
|
||||
LAMBDA_PATCH3 = 0.1
|
||||
PATCH_KERNEL_SIZE = 5
|
||||
PATCH_SIGMA = 1.2
|
||||
```
|
||||
|
||||
## 5. 训练集成
|
||||
|
||||
### 5.1 代码变更范围
|
||||
|
||||
| 文件 | 变更类型 | 说明 |
|
||||
| ---- | -------- | ---- |
|
||||
| `ginka/train_seperated.py` | 修改 | 新增两个损失函数定义、超参数、训练循环中的损失累加 |
|
||||
|
||||
损失函数直接定义在 `train_seperated.py` 中,与现有的 `cross_entropy_loss`(第 201 行)并列放置,保持代码组织的一致性。
|
||||
|
||||
### 5.2 训练循环修改
|
||||
|
||||
在现有损失计算(第 969-977 行)的基础上,追加邻接损失与 Patch 损失:
|
||||
|
||||
```python
|
||||
# 三阶段 Cross Entropy(不变)
|
||||
loss1 = cross_entropy_loss(logits1, target1)
|
||||
loss2 = cross_entropy_loss(logits2, target2)
|
||||
loss3 = cross_entropy_loss(logits3, target3)
|
||||
|
||||
# 新增:邻接损失
|
||||
adj1 = adjacency_loss(logits1, target1) if LAMBDA_ADJ1 > 0 else 0.0
|
||||
adj2 = adjacency_loss(logits2, target2) if LAMBDA_ADJ2 > 0 else 0.0
|
||||
adj3 = adjacency_loss(logits3, target3) if LAMBDA_ADJ3 > 0 else 0.0
|
||||
|
||||
# 新增:Patch 损失
|
||||
patch1 = patch_loss(
|
||||
logits1, target1, PATCH_KERNEL_SIZE, PATCH_SIGMA
|
||||
) if LAMBDA_PATCH1 > 0 else 0.0
|
||||
patch2 = patch_loss(
|
||||
logits2, target2, PATCH_KERNEL_SIZE, PATCH_SIGMA
|
||||
) if LAMBDA_PATCH2 > 0 else 0.0
|
||||
patch3 = patch_loss(
|
||||
logits3, target3, PATCH_KERNEL_SIZE, PATCH_SIGMA
|
||||
) if LAMBDA_PATCH3 > 0 else 0.0
|
||||
|
||||
# 加权汇总
|
||||
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
|
||||
)
|
||||
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
|
||||
)
|
||||
```
|
||||
|
||||
对应地在日志输出区域追加邻接损失和 Patch 损失的值,便于监控。
|
||||
|
||||
### 5.3 验证循环同步修改
|
||||
|
||||
`validate` 函数(第 720-780 行附近)中同样需要追加邻接损失和 Patch 损失的计算,保持训练/验证损失口径一致。
|
||||
|
||||
### 5.4 性能影响
|
||||
|
||||
| 损失函数 | 额外计算量 | 额外显存 |
|
||||
| -------- | ---------- | -------- |
|
||||
| adjacency_loss | softmax + 4 次逐元素乘法 + 2 次 BCE | 极小(仅 B×H×W 级中间张量) |
|
||||
| patch_loss | 逐格 CE + unfold + 加权求和 | 中等(unfold 产生 B×K²×H×W 中间张量) |
|
||||
|
||||
对于 B=64、K=5、H=W=13,unfold 中间张量大小为 64×25×169 ≈ 270k float,显存开销约 1MB,可以忽略。总额外计算时间预估在 5% 以内。
|
||||
|
||||
## 6. 验证方案
|
||||
|
||||
### 6.1 定量指标
|
||||
|
||||
在验证集上对比以下指标(基线 vs 加邻接损失 vs 加 Patch 损失 vs 两者叠加):
|
||||
|
||||
| 指标 | 含义 | 期望趋势 |
|
||||
| ---- | ---- | -------- |
|
||||
| 阶段一 CE 损失 | 墙壁骨架准确度 | 基本持平或微降 |
|
||||
| 阶段二 CE 损失 | 门/怪物/入口准确度 | 持平或改善 |
|
||||
| 阶段三 CE 损失 | 资源放置准确度 | 持平或改善 |
|
||||
| 邻接损失值 | 地板连通性 | 显著下降 |
|
||||
| 地板连通分量数 | 生成地图中可通行区域的连通分量个数(理想为 1) | 下降 |
|
||||
| 孤立地板块比例 | 四邻域均为墙壁的地板格子占比 | 下降 |
|
||||
| Codebook 使用率 | VQ 码本利用率 | 不变或微升 |
|
||||
|
||||
### 6.2 可视化验证
|
||||
|
||||
对相同随机种子、相同随机采码条件下生成的地图进行可视化对比,重点关注:
|
||||
|
||||
1. **墙壁骨架质量**(阶段一):墙壁是否形成合理的房间轮廓和通道,是否出现孤立墙壁碎片
|
||||
2. **连通性**(全图):从入口出发能否到达所有可通行区域,是否存在被墙壁完全包围的孤岛
|
||||
3. **局部一致性**:门是否出现在墙壁上并连接两个地板区域、怪物是否出现在地板上而非墙壁上
|
||||
4. **资源分布**(阶段三):资源是否均匀分布在可通行区域内
|
||||
|
||||
验证时对不同 z 采样分别生成 4-8 张图,并排对比,使用 `shared/image.py` 的 `matrix_to_image_cv` 输出 PNG 文件。
|
||||
|
||||
### 6.3 消融实验
|
||||
|
||||
| 实验编号 | CE 损失 | 邻接损失 | Patch 损失 | 目的 |
|
||||
| -------- | ------- | -------- | ---------- | ---- |
|
||||
| E0 | ✓ | ✗ | ✗ | 基线(当前模型) |
|
||||
| E1 | ✓ | ✓ | ✗ | 验证邻接损失单独效果 |
|
||||
| E2 | ✓ | ✗ | ✓ | 验证 Patch 损失单独效果 |
|
||||
| E3 | ✓ | ✓ | ✓ | 验证两者叠加效果 |
|
||||
|
||||
### 6.4 超参数搜索建议
|
||||
|
||||
若初值效果不理想,按以下优先级调参:
|
||||
|
||||
1. `LAMBDA_ADJ1`:[0.05, 0.1, 0.2, 0.5] — 邻接损失最关键,先找到合适的强度
|
||||
2. `LAMBDA_PATCH1`:[0.05, 0.1, 0.2] — Patch 损失作为补充平滑
|
||||
3. `PATCH_SIGMA`:[0.8, 1.2, 1.6] — 控制平滑范围,sigma 越大平滑越强
|
||||
4. `PATCH_KERNEL_SIZE`:[3, 5, 7] — 控制邻域大小,注意 7×7 在地图上的覆盖接近 30%
|
||||
@ -12,18 +12,17 @@
|
||||
|
||||
### 1. 掩码方式
|
||||
|
||||
**仅使用矩形分块掩码**(`dataset.py` `std_mask` 的第二种模式),不采用散点随机掩码:
|
||||
**仅使用矩形分块掩码**(`dataset.py` `std_mask` 的第二种模式),不采用散点随机掩码。
|
||||
|
||||
- 在 13×13 网格上,反复放置随机矩形块(高宽 2~6),直到掩码格数达到目标比例
|
||||
- 掩码比例可调,默认沿用 Beta(2,2) 采样(范围 [5%, 100%]),同时支持指定固定比例用于对照实验
|
||||
验证时掩码比例随机采样自 `Uniform(0.2, 0.8)`。
|
||||
|
||||
### 2. 条件输入 z
|
||||
|
||||
不再随机采样 z,改为使用**被掩码地图经 VQ-VAE 编码后得到的真实 z**。这意味着:
|
||||
不再随机采样 z,改为使用**完整地图(掩码前)经 VQ-VAE 编码后得到的真实 z**:
|
||||
|
||||
- 编码器输入可能是部分掩码后的地图(掩码位置填 MASK_ID=7)
|
||||
- VQ-VAE 从可见部分推断整体结构,得到的 z 携带地图的全局语义信息
|
||||
- MaskGIT 以该 z 为条件,在掩码区域补全被遮盖的图块
|
||||
- 对掩码后的地图做 VQ 编码没有意义——掩码区域噪声会污染码本表示
|
||||
- z 携带完整地图的全局结构信息,作为"答案的轮廓"注入 MaskGIT
|
||||
- MaskGIT 以完整 z 为条件,对照可见图块,在掩码区域补全被遮盖的图块
|
||||
|
||||
### 3. 生成流程
|
||||
|
||||
@ -31,45 +30,32 @@
|
||||
|
||||
```
|
||||
输入: 完整地图 gt [13×13]
|
||||
↓ ↓
|
||||
VQ-VAE 编码 生成矩形掩码 mask [13×13 bool]
|
||||
gt → z1,z2,z3 gt[mask] = MASK_TOKEN → inp
|
||||
↓ ↓
|
||||
└──────┬─────────┘
|
||||
↓
|
||||
生成矩形掩码 mask [13×13 bool] (比率可调)
|
||||
↓
|
||||
inp = gt.copy(); inp[mask] = MASK_TOKEN
|
||||
↓
|
||||
VQ-VAE 编码 inp → z1, z2, z3
|
||||
↓
|
||||
Stage1: mg1(inp, z1) → 补全墙壁 (mask ∩ target=1)
|
||||
Stage1: mg1(inp, z1) → 补全墙壁
|
||||
Stage2: mg2(inp1+stage1_out, z2) → 补全门/怪物/入口
|
||||
Stage3: mg3(inp2+stage2_out, z3) → 补全资源
|
||||
↓
|
||||
输出: 最终地图 + 逐阶段可视化对比
|
||||
```
|
||||
|
||||
### 4. 与现有 `maskgit_sample` 的适配
|
||||
### 4. 逐阶段保留未掩码区域
|
||||
|
||||
现有 `maskgit_sample` 已支持 `keep_fixed` 参数,天然适配修补场景:
|
||||
与 `full_generate_specific_z` 不同,修补场景下需要将**未掩码区域的原始图块逐阶段传递**,避免被误当成空地重填:
|
||||
|
||||
- `keep_fixed=True`(默认):输入中已有的非掩码图块在生成过程中锁定不变,MaskGIT 仅填充掩码位
|
||||
- 需要新增:将矩形掩码 `mask` 传入函数,使模型只尝试修改掩码区域
|
||||
- Stage1 后将 raw_map 中未掩码区的非 0/1 图块覆盖到 stage1 输出
|
||||
- Stage3 前将 raw_map 中未掩码区的资源覆盖回去
|
||||
|
||||
## 可调参数
|
||||
### 5. 验证可视化
|
||||
|
||||
| 参数 | 说明 | 默认值 |
|
||||
|------|------|--------|
|
||||
| `mask_ratio` | 掩码比例,可指定固定值或 `None` 使用 Beta(2,2) 采样 | `None`(Beta 采样) |
|
||||
| `block_h_range` | 矩形高度范围 | `(2, 7)` |
|
||||
| `block_w_range` | 矩形宽度范围 | `(2, 7)` |
|
||||
| `keep_fixed` | 三阶段是否锁定已有结构 | `(True, True, True)` |
|
||||
替换现有的验证 `rand{batch_idx}.png`(原为随机墙壁种子 + 随机 z 生成),改为修补方式:
|
||||
|
||||
## 实现要点
|
||||
|
||||
1. 将 `std_mask` 的矩形掩码生成逻辑提取为独立函数 `rect_mask(ratio, h_range, w_range) -> np.ndarray`
|
||||
2. 新增 `inpaint_generate(map, mask, models, device)` 函数,包装完整修补流程
|
||||
3. 验证时对同一张地图使用不同 `mask_ratio`(如 0%,30%,50%,70%)生成多组结果,横向对比
|
||||
4. 可视化输出时同时展示:原始地图、掩码地图、各阶段生成结果、最终结果
|
||||
|
||||
## 预期效果
|
||||
|
||||
- 任务从"开放式生成"变为"约束式修补",难度显著降低
|
||||
- 模型可利用可见区域的上下文信息(如墙面走向、房间布局)推断掩码区域内容
|
||||
- 不同掩码比例可衡量模型对空间结构的理解深度
|
||||
- 不显示真实地图,仅展示掩码输入与生成结果的对比
|
||||
- 每张 `rand` 图展示 4 组样本,每组为"掩码地图 + 生成结果"并排
|
||||
- 2 行 × 4 列布局:每行放 2 组,每组左侧掩码输入、右侧生成结果
|
||||
- 掩码区域由 MASK 图块(ID=7)自身标识,无需额外叠加颜色
|
||||
- 每张子图标注样本名称和掩码比例(如 `map_v1 34%`)
|
||||
|
||||
@ -416,6 +416,21 @@ def compute_remaining(
|
||||
|
||||
return remain
|
||||
|
||||
def rect_mask(
|
||||
ratio: float, h_range: tuple[int, int] = (2, 7),
|
||||
w_range: tuple[int, int] = (2, 7)
|
||||
) -> np.ndarray:
|
||||
# 纯矩形分块掩码,反复放置随机矩形直到掩码格数达标
|
||||
target = int(MAP_SIZE * ratio)
|
||||
mask = np.zeros((MAP_H, MAP_W), dtype=bool)
|
||||
while mask.sum() < target:
|
||||
bh = np.random.randint(h_range[0], h_range[1])
|
||||
bw = np.random.randint(w_range[0], w_range[1])
|
||||
x = np.random.randint(0, MAP_H - bh + 1)
|
||||
y = np.random.randint(0, MAP_W - bw + 1)
|
||||
mask[x:x + bh, y:y + bw] = True
|
||||
return mask
|
||||
|
||||
def maskgit_sample(
|
||||
model: torch.nn.Module, inp: torch.Tensor, z: torch.Tensor,
|
||||
z_dist: torch.Tensor, struct: torch.Tensor, target_density: torch.Tensor,
|
||||
@ -551,6 +566,71 @@ def full_generate_specific_z(
|
||||
|
||||
return pred1_np, merged12, merged123
|
||||
|
||||
def inpaint_generate(
|
||||
raw_map: np.ndarray,
|
||||
mask: np.ndarray,
|
||||
z_q: tuple[torch.Tensor, torch.Tensor, torch.Tensor],
|
||||
z_dist: torch.Tensor,
|
||||
struct: torch.Tensor,
|
||||
target_density: torch.Tensor,
|
||||
models: list[torch.nn.Module],
|
||||
device: torch.device
|
||||
):
|
||||
# 三阶段矩形掩码修补生成
|
||||
# - z 来自完整地图 VQ 编码,作为全局先验
|
||||
# - 未掩码区域的原始图块逐阶段保留,避免被误当空地重填
|
||||
vq1, vq2, vq3, mg1, mg2, mg3, quantizers, optimizer, scheduler = models
|
||||
z1, z2, z3 = z_q
|
||||
|
||||
# Stage1: 补全墙壁
|
||||
inp1 = raw_map.copy()
|
||||
inp1[mask] = MASK_TOKEN
|
||||
non_wf = (inp1 != 0) & (inp1 != 1) & (inp1 != MASK_TOKEN)
|
||||
inp1[non_wf] = 0
|
||||
inp1_t = torch.tensor(inp1.flatten(), dtype=torch.long, device=device).reshape(1, MAP_SIZE)
|
||||
|
||||
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
|
||||
)
|
||||
|
||||
# Stage2: 补全门/怪/入口,保留未掩码区的原始非墙壁结构
|
||||
merged_s1 = pred1_np.copy()
|
||||
preserve_s2 = (~mask) & (raw_map != 0) & (raw_map != 1)
|
||||
merged_s1[preserve_s2] = raw_map[preserve_s2]
|
||||
|
||||
inp2 = merged_s1.copy()
|
||||
inp2[inp2 == 0] = MASK_TOKEN
|
||||
inp2_t = torch.tensor(inp2.flatten(), dtype=torch.long, device=device).reshape(1, MAP_SIZE)
|
||||
|
||||
with torch.no_grad():
|
||||
pred2_np = maskgit_sample(
|
||||
mg2, inp2_t, z2, z_dist, struct, target_density, 2,
|
||||
GENERATE_STEP, target_tiles=[2, 6, 4, 5], keep_fixed=True
|
||||
)
|
||||
|
||||
# Stage3: 补全资源,保留未掩码区的原始资源
|
||||
merged_s2 = merged_s1.copy()
|
||||
merged_s2[pred2_np != 0] = pred2_np[pred2_np != 0]
|
||||
res_preserve = (raw_map == 3) & (~mask)
|
||||
merged_s2[res_preserve] = 3
|
||||
|
||||
inp3 = merged_s2.copy()
|
||||
inp3[inp3 == 0] = MASK_TOKEN
|
||||
inp3_t = torch.tensor(inp3.flatten(), dtype=torch.long, device=device).reshape(1, MAP_SIZE)
|
||||
|
||||
with torch.no_grad():
|
||||
pred3_np = maskgit_sample(
|
||||
mg3, inp3_t, z3, z_dist, struct, target_density, 3,
|
||||
GENERATE_STEP, target_tiles=[3], keep_fixed=True
|
||||
)
|
||||
|
||||
merged_s3 = merged_s2.copy()
|
||||
merged_s3[pred3_np != 0] = pred3_np[pred3_np != 0]
|
||||
|
||||
return pred1_np, merged_s1, merged_s2, merged_s3
|
||||
|
||||
def annotate(img: np.ndarray, text: str, y: int = 14) -> np.ndarray:
|
||||
# 在图片左上角叠加文字标注(黑色描边 + 白色填充,确保任意背景下可读)
|
||||
img = img.copy()
|
||||
@ -710,7 +790,7 @@ def visualize_part2(batch, z_q, z_dist, models, device, tile_dict):
|
||||
return grid
|
||||
|
||||
# 验证可视化 part4:2×3 网格;保留稀疏墙壁种子,但 z 与标签来自训练集样本
|
||||
def visualize_part4(
|
||||
def visualize_rand(
|
||||
train_dataset: GinkaSeperatedDataset,
|
||||
models: list[torch.nn.Module],
|
||||
dist_models: tuple,
|
||||
@ -725,33 +805,60 @@ def visualize_part4(
|
||||
def to_img(mat):
|
||||
return matrix_to_image_cv(mat, tile_dict, TILE_SIZE)
|
||||
|
||||
n_walls = random.randint(math.floor(MAP_SIZE * 0.02), math.floor(MAP_SIZE * 0.06))
|
||||
seed = torch.full((1, MAP_SIZE), MASK_TOKEN, dtype=torch.long, device=device)
|
||||
wall_pos = torch.randperm(MAP_SIZE, device=device)[:n_walls]
|
||||
seed[0, wall_pos] = 1
|
||||
seed_np = seed[0].cpu().numpy().reshape(MAP_H, MAP_W)
|
||||
vq1, vq2, vq3, mg1, mg2, mg3, quantizers, optimizer, scheduler = models
|
||||
dist_encoder, dist_quantizer = dist_models
|
||||
|
||||
results = []
|
||||
for _ in range(5):
|
||||
case = build_dataset_sample_case(train_dataset, models, dist_models, device)
|
||||
kf = rand_keep()
|
||||
sample = case["sample"]
|
||||
_, _, merged123 = full_generate_specific_z(
|
||||
seed, case["z_q"], case["z_dist"], case["struct"], case["target_density"],
|
||||
models, device, keep_fixed=kf
|
||||
)
|
||||
result = annotate_labels(
|
||||
to_img(merged123), sample["struct_inject"], sample["target_density"]
|
||||
)
|
||||
results.append(
|
||||
annotate(result, f"{sample_case_label(case)} {keep_label(kf)}", y=50)
|
||||
samples_data = []
|
||||
for _ in range(4):
|
||||
sample = train_dataset.random_sample_map()
|
||||
raw_map = sample["raw_map"].numpy().reshape(MAP_H, MAP_W)
|
||||
ratio = random.uniform(0.2, 0.8)
|
||||
mask = rect_mask(ratio)
|
||||
|
||||
enc1_t = sample["encoder_stage1"].to(device).reshape(1, MAP_SIZE)
|
||||
enc2_t = sample["encoder_stage2"].to(device).reshape(1, MAP_SIZE)
|
||||
enc3_t = sample["encoder_stage3"].to(device).reshape(1, MAP_SIZE)
|
||||
struct_t = sample["struct_inject"].to(device).reshape(1, -1)
|
||||
target_density_t = sample["target_density"].to(device).reshape(1, -1)
|
||||
dist_field_t = sample["distance_field"].to(device).reshape(1, -1)
|
||||
|
||||
with torch.no_grad():
|
||||
z_e1 = vq1(enc1_t)
|
||||
z_e2 = vq2(enc2_t)
|
||||
z_e3 = vq3(enc3_t)
|
||||
z_q, _, _ = quantize_stage_latents(quantizers, z_e1, z_e2, z_e3)
|
||||
z_e_dist = dist_encoder(dist_field_t)
|
||||
z_dist, _, _, _, _ = dist_quantizer(z_e_dist)
|
||||
|
||||
_, _, _, merged_s3 = inpaint_generate(
|
||||
raw_map, mask, z_q, z_dist, struct_t, target_density_t,
|
||||
models, device
|
||||
)
|
||||
|
||||
row1 = [annotate(to_img(seed_np), 'seed')] + results[:2]
|
||||
row2 = results[2:]
|
||||
rows = [row1, row2]
|
||||
grid = np.ones((2 * img_h + 3 * SEP, 3 * img_w + 4 * SEP, 3), dtype=np.uint8) * 255
|
||||
for r, row in enumerate(rows):
|
||||
masked_display = raw_map.copy()
|
||||
masked_display[mask] = MASK_TOKEN
|
||||
|
||||
samples_data.append({
|
||||
"masked": masked_display,
|
||||
"result": merged_s3,
|
||||
"mask": mask,
|
||||
"label": f"{sample['map_name']} {int(ratio * 100)}%"
|
||||
})
|
||||
|
||||
row1 = [
|
||||
annotate(to_img(samples_data[0]["masked"]), samples_data[0]["label"]),
|
||||
to_img(samples_data[0]["result"]),
|
||||
annotate(to_img(samples_data[1]["masked"]), samples_data[1]["label"]),
|
||||
to_img(samples_data[1]["result"]),
|
||||
]
|
||||
row2 = [
|
||||
annotate(to_img(samples_data[2]["masked"]), samples_data[2]["label"]),
|
||||
to_img(samples_data[2]["result"]),
|
||||
annotate(to_img(samples_data[3]["masked"]), samples_data[3]["label"]),
|
||||
to_img(samples_data[3]["result"]),
|
||||
]
|
||||
grid = np.ones((2 * img_h + 3 * SEP, 4 * img_w + 5 * SEP, 3), dtype=np.uint8) * 255
|
||||
for r, row in enumerate([row1, row2]):
|
||||
for c, img in enumerate(row):
|
||||
y = SEP + r * (img_h + SEP)
|
||||
x = SEP + c * (img_w + SEP)
|
||||
@ -767,7 +874,7 @@ def visualize_validate(
|
||||
os.makedirs(save_dir, exist_ok=True)
|
||||
cv2.imwrite(f"{save_dir}/val{batch_idx}.png", visualize_part1(batch, logits1, logits2, logits3, tile_dict))
|
||||
cv2.imwrite(f"{save_dir}/full{batch_idx}.png", visualize_part2(batch, z_q, z_dist, models, device, tile_dict))
|
||||
cv2.imwrite(f"{save_dir}/rand{batch_idx}.png", visualize_part4(train_dataset, models, dist_models, device, tile_dict))
|
||||
cv2.imwrite(f"{save_dir}/rand{batch_idx}.png", visualize_rand(train_dataset, models, dist_models, device, tile_dict))
|
||||
|
||||
def validate(
|
||||
dataloader: DataLoader,
|
||||
|
||||
Loading…
Reference in New Issue
Block a user