Compare commits

..

2 Commits

Author SHA1 Message Date
27c42bb690 chore: rename doc 2026-07-17 16:27:21 +08:00
3cd38b86a9 feat: 距离场输入 2026-07-17 16:26:57 +08:00
6 changed files with 655 additions and 50 deletions

421
docs/latent-improve.md Normal file
View File

@ -0,0 +1,421 @@
# 隐变量增强机制设计文档
## 整体架构概览
```mermaid
flowchart TD
A[原始地图] --> B[数据增强]
A --> C[随机变换]
B --> D[enc1, enc2, enc3]
C --> E[enc1', enc2', enc3']
D --> F[VQ-VAE 编码器]
E --> G[VQ-VAE 编码器]
F --> H["z_e"]
G --> I["z_e'"]
H --> J[量化]
J --> K["z_q"]
K --> L[Latent Dropout]
L --> M["z_q'"]
M --> N[MaskGIT]
H -->|对比损失 MSE| I
O[距离场] --> P[距离编码器]
P --> Q["z_dist"]
Q --> N
```
三个机制分别作用在隐变量的不同环节:
| 机制 | 作用位置 | 作用阶段 | 核心目标 |
| -------------- | ----------------- | --------- | ---------------------------------- |
| Latent Dropout | z_q → z_q' | 训练 | 防止解码器过度依赖特定 latent code |
| 距离场编码 | z_dist 拼接入 z_q | 训练+推理 | 注入连续空间结构信息 |
| 对比学习损失 | z_e vs z_e' | 训练 | 强制编码器学习变换不变表示 |
---
## 机制一Latent Dropout
### 核心思路
训练时以概率 p 随机丢弃(替换为可学习 mask tokenz_q 序列中的部分 latent code迫使 MaskGIT 解码器学会在部分码字缺失的情况下仍能合理生成,从而降低对特定码字组合的过依赖。
### 设计方案
#### Dropout 策略:可学习 Mask Token 替换
在 z_q [B, L, d_z] 上,以概率 p 将某些位置的 latent code 替换为一个可学习的 mask 嵌入向量:
```mermaid
flowchart LR
A["z_q B×L×d_z"] --> B["按 p=0.1 随机选择位置"]
B --> C["选中位置替换为 latent_mask_embedding d_z"]
C --> D["z_q' B×L×d_z → 送入 MaskGIT"]
```
选择 mask token 替换方案(而非置零或随机采码)的理由:
- **置零向量**:改变了 z_q 的范数分布(码字向量模长约为 1零向量模长为 0打破了量化空间的几何结构。
- **随机采码**:过于激进,可能引入语义上完全无关的码字,造成训练不稳定。
- **Mask token**:可学习的嵌入,由模型自适应确定"缺失码字"的最优替代表示,与 BERT / MaskGIT 自身的设计理念一致。
#### Dropout 率 p
- `p = 0.1` — 对于 L=16 的码字序列,平均每 batch 有 1.6 个码字被丢弃,扰动适度。
- 更激进的 `p = 0.2` 可作为消融实验的对比项。
- 推理时 p = 0不使用 dropout
#### 实现位置
Latent dropout 应放在 `GinkaMaskGIT` 内部,在 z 投影之前:
```python
# 在 GinkaMaskGIT.forward 中:
if self.training and self.z_dropout > 0:
z = self.apply_latent_dropout(z) # 以概率 self.z_dropout 替换为 mask token
z_proj = self.z_proj(z)
```
需在 `GinkaMaskGIT.__init__` 中新增参数 `z_dropout: float = 0.0``nn.Parameter` 类型的 `latent_mask_embedding`
#### 与现有 `MG_Z_DROPOUT` 的关系
当前 `train_seperated.py` 第 94 行已定义 `MG_Z_DROPOUT = 0.1`,但未在训练循环中使用。该常量应作为 `GinkaMaskGIT` 构造参数传入,在模型内部完成 dropout 逻辑。
### 预期效果
- codebook 使用率usage rate可能提升解码器不再依赖极少数字就能生成迫使更多码字被"激活"。
- 训练损失可能轻微上升(因为条件信号变弱),但验证时的生成质量应提升,随机采码生成效果应更稳定。
---
## 机制二:距离场编码
### 核心思路
为每个地图位置计算其到最近墙壁的 **L1曼哈顿距离**构成与地图同尺寸13×13的距离场。距离场通过独立编码器编码为额外一组 latent code `z_dist`,与三阶段的 `z_q` 沿序列维拼接后送入 MaskGIT。
距离场提供了 tile 类别标签中不包含的连续空间信息:
- 距离墙壁越近,越适合放置墙壁相关元素(如门、角落资源);
- 距离墙壁越远(中间区域),越适合空旷空间(通路、大面积空地);
- 距离场的梯度方向隐含了"朝向通道/远离墙壁"的结构特征。
### 距离场计算
对于地图 M [13, 13],设墙壁 tile ID 为 1
$$D(i, j) = \min_{(i_w, j_w): M[i_w, j_w] = 1} \left(|i - i_w| + |j - j_w|\right)$$
距离值域为 [0, 24]13×13 网格中最远曼哈顿距离为 12+12=24。计算实现
```python
def compute_distance_field(map_matrix: np.ndarray) -> np.ndarray:
# map_matrix: [13, 13] 整数矩阵
# 返回: [13, 13] 浮点矩阵,每格为到最近墙壁的曼哈顿距离
wall_mask = (map_matrix == 1)
# 直接距离变换
from scipy.ndimage import distance_transform_cdt
return distance_transform_cdt(~wall_mask, metric='taxicab').astype(np.float32)
```
若无 scipy 依赖,也可使用纯 numpy 的 BFS 或多源扩展实现。
#### 距离归一化
距离值需要归一化到合理范围以便编码器处理。推荐**分桶嵌入**(类似 tile embedding
| 方案 | 做法 | 优点 |
| -------- | -------------------------------------- | -------------------------- |
| 连续值 | 归一化到 [0, 1],经 MLP 映射为嵌入 | 保留精确距离 |
| 分桶嵌入 | 距离离散化为 N 个桶,用 Embedding 查表 | 与 tile embedding 风格统一 |
**推荐分桶嵌入**:将距离离散化到 [0, MAX_DIST] 的整数(截断到 MAX_DIST再用 Embedding 映射。MAX_DIST 可取 12超过该距离的格子极少且距离语义趋于饱和共计 13 个桶0-12
### 编码器设计
距离场编码器使用与 `GinkaVQVAE` 相同的 Transformer 架构、但参数更轻量的独立实例:
```mermaid
flowchart TD
A["距离场 B×169分桶后的整数序列"] --> B["DistEmbedding (vocab=13, d_model=128)"]
B --> C["2D 因式分解位置编码"]
C --> D["+ L_dist 个可学习 summary token"]
D --> E["轻量 Transformer Encoder2-3 层, d_model=128, nhead=4, dim_ff=512"]
E --> F["取前 L_dist 个 summary token → Linear 投影 → B×L_dist×d_z=64"]
F --> G["可选:独立 VectorQuantizer 量化K=16"]
G --> H["z_dist B×L_dist×d_z"]
```
#### 设计要点
| 参数 | 建议值 | 说明 |
| ---------- | ------ | ---------------------------------------- |
| L_dist | 4 | 距离场信息量低于 tile 类别4 个码字足够 |
| d_model | 128 | 比 VQ-VAE 的 256 更轻量 |
| num_layers | 3 | 信息密度低,浅层即可编码 |
| K_dist | 16 | 距离场 codebook 容量(可选,见下方讨论) |
#### 是否需要量化距离场 z
距离场是连续信号,原则上不需要 VQ 量化。但考虑以下因素后**推荐量化**
- **推理统一性**:推理时 z 均从 codebook 采样,若 z_dist 连续则无对应采样机制,只能从训练集的距离场编码获得,破坏了"用户无需任何输入"的设计目标。
- **正则化效果**:量化对距离场编码施加信息瓶颈,迫使编码器提取高层结构特征而非记忆具体距离值。
- 实现简单:复用现有 `VectorQuantizer`,只需 K=16 的小型 codebook。
推理时 z_dist 与三阶段 z_q 均可从各自 codebook 独立随机采样。
### 拼接与注入方式
z_dist 直接拼接到 z_q 序列后,通过现有的 `cond_proj` 投影为 AdaLN 条件向量:
```python
# 当前cond_seq = cat([z_proj, e_struct, e_remain]) # [B, L+2+5, d_z]
# 修改后:
z_proj = self.z_proj(z_q) # [B, L, d_z]
zd_proj = self.zd_proj(z_dist) # [B, L_dist, d_z]
cond_seq = torch.cat([z_proj, zd_proj, e_struct, e_remain], dim=1)
# cond_seq: [B, L + L_dist + 2 + 5, d_z]
c = self.cond_proj(cond_seq.reshape(B, -1)) # [B, d_model]
```
需要在 `GinkaMaskGIT` 中:
- 新增 `z_dist_proj``nn.Linear(d_z, d_z)`
- 修改 `cond_proj` 输入维度:`(z_seq_len + L_dist + 2 + 5) * d_z`
- `forward` 签名新增 `z_dist: torch.Tensor` 参数
#### 距离场来源
距离场由**全量地图的墙壁布局**决定,与哪个阶段无关。因此 z_dist 由**一份全量地图的距离场**编码得到,在三阶段间共享使用。实际编码时以 `encoder_stage1`(或 raw_map中的墙壁为基础计算距离场因为只有 stage1 保留了完整的墙壁信息。
```python
# 训练循环中:
dist_field = compute_distance_field(enc1) # enc1 含完整墙壁
z_dist = dist_encoder(dist_field) # 共用一份距离场编码
z_dist = dist_quantizer(z_dist) # 可选量化
# 三阶段 MaskGIT 前向均传入同一 z_dist
logits1 = mg1(inp1, z_q1, z_dist, struct, remain1)
logits2 = mg2(inp2, z_q2, z_dist, struct, remain2)
logits3 = mg3(inp3, z_q3, z_dist, struct, remain3)
```
---
## 机制三:对比学习损失
对同一原始地图施加两种随机等距变换后分别通过编码器,约束两个视图的 latent z_e 尽可能相似,损失使用 MSE
$$\mathcal{L}_{contrast} = \text{MSE}(z_e, z_e')$$
其中 $T_1, T_2$ 从 8 种等距变换中独立随机采样。
### 为何选择 MSE 而非 InfoNCE
| 损失函数 | 适用场景 | 当前适用性 |
| --------------- | ----------------------------- | --------------------------------------------- |
| **MSE推荐** | 正样本对已知且应完全相同 | 变换前后 z_e 应一致MSE 直接约束 |
| InfoNCE | 正负样本对需要从 batch 内挖掘 | batch 中不同地图的 z 理应不同,负样本客观存在 |
考虑到:
- 同一地图的变换版本理应映射到同一 code 序列(或至少相邻 code直接 MSE 约束更精准;
- InfoNCE 将"与 batch 内其他样本拉开距离"作为负约束,但在语义层面上两张不同地图的 z 未必应该远离(结构相似的地图应有相近的 z这可能引入噪声
- MSE 实现更简单,计算量更低。
若实验中发现 MSE 导致 codebook 坍缩(所有 z_e 趋同),可替换为 NT-XentSimCLR 风格),通过温度参数控制分布的熵。
### 变换集合
对 13×13 的二维网格,以下 8 种变换构成等距变换群D4 二面体群):
| 编号 | 变换 |
| ---- | ------------------------------------- |
| T0 | 恒等(无变换) |
| T1 | 旋转 90° 顺时针 |
| T2 | 旋转 180° |
| T3 | 旋转 270° 顺时针(等价于 90° 逆时针) |
| T4 | 水平翻转 |
| T5 | 垂直翻转 |
| T6 | 主对角线翻转(转置) |
| T7 | 副对角线翻转 |
每次构造样本对时,从 8 种变换中独立随机采样两个(允许相同,即 T_i = T_j
### 训练流程
对比损失在**训练循环**中计算,不与数据集的在线增强耦合:
```
1. 从 dataset 获取 batch已经过随机增强用于 MaskGIT 训练目标)
2. 取各阶段的 encoder 输入 enc1, enc2, enc3均保留完整上下文
3. 对 enc1/enc2/enc3 分别施加随机变换 T_a编码得 z_e_a
4. 对 enc1/enc2/enc3 分别施加随机变换 T_b编码得 z_e_b
5. 计算三阶段的 MSE(z_e_a, z_e_b),求和作为总对比损失
```
需注意:
- 当前 dataset 的 `__getitem__` 已对原始地图施加了在线增强(随机旋转/翻转)。对比损失所需的原始 → 变换流程需独立于该增强,在编码阶段重新获取变换前的地图。最简单的方式是在 dataset 中新增返回字段 `raw_map`(未经增强的完整地图),或从 `encoder_stage3`(含所有 tile中取墙壁信息复原。
- 建议在 dataset 中新增 `raw_map` 字段:返回增强前(或逆增强后)的完整地图 [13, 13],用于对比学习的变换和距离场计算。
### 损失权重与调度
$$\mathcal{L}_{total} = \mathcal{L}_{CE} + \beta \cdot \mathcal{L}_{commit} + \lambda_{contrast} \cdot \mathcal{L}_{contrast}$$
- $\lambda_{contrast} = 0.1$ 作为初始值
- 若训练早期对比损失较大导致模型不稳定,可在前 N 个 epoch 线性 warmup $\lambda_{contrast}$ 从 0 到目标值
- 若 codebook 使用率下降(熵减小),需降低 $\lambda_{contrast}$ 或切换为 InfoNCE
### 与距离场编码器的关系
距离场编码器同样受益于对比学习约束——对地图施加等距变换后,距离场也相应变换,但距离场编码器应输出相同的 z_dist。因此对比损失也应作用于距离场编码器
$$\mathcal{L}_{contrast}^{dist} = \text{MSE}(z_{dist}, z_{dist}')$$
可在距离场编码器内部或外部实现,权重与主对比损失相同。
---
## 数据集调整
三个机制均需要额外的数据字段,需调整 `GinkaSeperatedDataset`
### 新增字段
| 字段名 | 形状 | 说明 |
| ---------------- | -------- | ------------------------------------------- |
| `raw_map` | [13, 13] | 未经增强的原始完整地图(所有 tile 类别) |
| `distance_field` | [169] | 从 raw_map 的墙壁计算的 L1 距离场(分桶后) |
### 实现要点
1. `raw_map`:在 `__getitem__` 中,先保存 raw_map再施加在线数据增强得到各阶段输入。
2. `distance_field`:从 `raw_map` 计算,与在线增强无关(距离场由原图墙壁决定,增强不影响墙壁的相对位置关系)。
---
## 模型与训练流程变更
### 新增/修改的模块
| 模块 | 变更类型 | 说明 |
| ----------------------- | -------- | -------------------------------------------------- |
| `GinkaMaskGIT` | 修改 | 新增 z_dropout, latent_mask_embedding, z_dist 处理 |
| `DistFieldEncoder` | 新增 | 距离场编码器(轻量 GinkaVQVAE 变体) |
| `GinkaSeperatedDataset` | 修改 | 新增 raw_map, distance_field 字段 |
| `train_seperated.py` | 修改 | 集成三个机制的训练逻辑 |
### 训练循环伪码
```python
# 每个训练 step
enc1, enc2, enc3, inp1, inp2, inp3, target1, target2, target3,
struct, target_density, raw_map, dist_field = batch
# === VQ 编码 ===
z_e1, z_e2, z_e3 = vq1(enc1), vq2(enc2), vq3(enc3)
z_q1, z_q2, z_q3 = quantize(z_e1, z_e2, z_e3)
# === 距离场编码 ===
z_e_dist = dist_encoder(dist_field)
z_dist = dist_quantizer(z_e_dist)
# === 对比损失 ===
# 对 raw_map 施加两种随机变换,编码两次
enc1_a, enc2_a, enc3_a = apply_transform(raw_map, T_a)
enc1_b, enc2_b, enc3_b = apply_transform(raw_map, T_b)
z_e1_a, z_e2_a, z_e3_a = vq1(enc1_a), vq2(enc2_a), vq3(enc3_a)
z_e1_b, z_e2_b, z_e3_b = vq1(enc1_b), vq2(enc2_b), vq3(enc3_b)
loss_contrast = (
mse(z_e1_a, z_e1_b) + mse(z_e2_a, z_e2_b) + mse(z_e3_a, z_e3_b)
) / 3
# 距离场对比
dist_a = apply_transform(dist_field, T_a)
dist_b = apply_transform(dist_field, T_b)
z_e_dist_a = dist_encoder(dist_a)
z_e_dist_b = dist_encoder(dist_b)
loss_contrast += mse(z_e_dist_a, z_e_dist_b)
# === MaskGIT 前向Latent Dropout 在 mg 内部) ===
logits1 = mg1(inp1, z_q1, z_dist, struct, remain1)
logits2 = mg2(inp2, z_q2, z_dist, struct, remain2)
logits3 = mg3(inp3, z_q3, z_dist, struct, remain3)
# === 总损失 ===
loss = (
STAGE1_CE_WEIGHT * ce(logits1, target1)
+ STAGE2_CE_WEIGHT * ce(logits2, target2)
+ STAGE3_CE_WEIGHT * ce(logits3, target3)
+ VQ_BETA * commit_loss
+ VQ_BETA_DIST * commit_loss_dist
+ LAMBDA_CONTRAST * loss_contrast
)
```
### 新增超参
| 参数 | 建议值 | 说明 |
| ----------------- | ------ | ----------------------------- |
| `Z_DROPOUT` | 0.1 | latent dropout 概率 |
| `L_DIST` | 4 | 距离场码字序列长度 |
| `K_DIST` | 16 | 距离场 codebook 大小 |
| `DIST_D_MODEL` | 128 | 距离场编码器模型维度 |
| `DIST_LAYERS` | 3 | 距离场编码器 Transformer 层数 |
| `DIST_DIM_FF` | 512 | 距离场编码器 FF 维度 |
| `DIST_MAX_BUCKET` | 12 | 距离分桶上限 |
| `LAMBDA_CONTRAST` | 0.1 | 对比损失权重 |
| `VQ_BETA_DIST` | 0.5 | 距离场 commit loss 权重 |
---
## 实施建议
### 阶段一Latent Dropout最低风险最快落地
1. 在 `GinkaMaskGIT` 中新增 `z_dropout` 参数和 `latent_mask_embedding`
2. 在 `build_model` 中将 `MG_Z_DROPOUT` 传入各 mg 实例;
3. 训练并观察 codebook usage rate 和生成质量的变化。
### 阶段二:距离场编码(中等复杂度)
1. 实现距离场计算函数(`shared/distance.py`
2. 实现 `DistFieldEncoder` 类(参考 `GinkaVQVAE` 结构,更轻量);
3. 修改 `GinkaSeperatedDataset`,新增 `raw_map``distance_field`
4. 修改 `GinkaMaskGIT``forward` 签名和 `cond_proj` 维度;
5. 修改训练循环,加入距离场编码和前向链路。
### 阶段三:对比学习损失(最后接入)
1. 实现等距变换工具函数(`shared/transforms.py`
2. 在训练循环中加入对比损失计算;
3. 监控 codebook 熵的变化,按需调整 `LAMBDA_CONTRAST`
4. 与阶段一、二的成果叠加训练。
### 整体叠加
三个阶段不冲突,建议按顺序逐层加入,每次加入后充分训练验证,确保新机制没有破坏已有成果后再推进下一阶段。最终三者叠加后,预期模型在随机采码推理时的生成多样性和结构合理性均应有明显提升。
---
## 消融实验建议
为验证各机制的有效性,建议进行以下消融实验:
| 实验编号 | Latent Dropout | 距离场编码 | 对比损失 | 目的 |
| -------- | -------------- | ---------- | -------- | ----------------------- |
| E0 | ✗ | ✗ | ✗ | 基线(当前模型) |
| E1 | ✓ | ✗ | ✗ | 单独验证 latent dropout |
| E2 | ✗ | ✓ | ✗ | 单独验证距离场编码 |
| E3 | ✗ | ✗ | ✓ | 单独验证对比损失 |
| E4 | ✓ | ✓ | ✗ | latent dropout + 距离场 |
| E5 | ✓ | ✓ | ✓ | 三者全叠加(目标方案) |
每个实验使用相同随机种子、相同训练轮数,对比以下指标:
- 验证集 CE 损失(三阶段分别)
- Codebook usage rate / perplexity
- 随机采码生成地图的墙壁连通性、各类 tile 密度偏差
- 可视化对比(随机采码 5 次生成的地图一致性)

View File

@ -3,6 +3,7 @@ import random
import torch
import numpy as np
from torch.utils.data import Dataset
from shared.distance import compute_distance_field
def load_data(path: str):
with open(path, 'r', encoding="utf-8") as f:
@ -94,6 +95,8 @@ class GinkaSeperatedDataset(Dataset):
return enc1, enc2, enc3
def pack_sample(self, item: dict, map_np: np.ndarray, out: tuple) -> dict:
# out[2] = encoder_stage1含完整墙壁据此计算距离场
dist_field = compute_distance_field(out[2])
return {
"input_stage1": torch.LongTensor(out[0]),
"target_stage1": torch.LongTensor(out[1]),
@ -105,7 +108,8 @@ class GinkaSeperatedDataset(Dataset):
"target_stage3": torch.LongTensor(out[7]),
"encoder_stage3": torch.LongTensor(out[8]),
"struct_inject": self.build_struct_inject(map_np, item['outerWall']),
"target_density": self.build_target_density(item['map'])
"target_density": self.build_target_density(item['map']),
"distance_field": torch.LongTensor(dist_field)
}
def random_sample_map(self, idx: int | None = None) -> dict:
@ -122,7 +126,8 @@ class GinkaSeperatedDataset(Dataset):
"encoder_stage3": torch.LongTensor(enc3),
"struct_inject": self.build_struct_inject(map_np, item['outerWall']),
"target_density": self.build_target_density(item['map']),
"raw_map": torch.LongTensor(map_np)
"raw_map": torch.LongTensor(map_np),
"distance_field": torch.LongTensor(compute_distance_field(enc1))
}
sample['sample_idx'] = idx
return sample

View File

@ -7,12 +7,13 @@ from .maskGIT import Transformer
# 结构标签词表大小
SYM_VOCAB = 8 # symmetryH/V/C 三位组合 0-7
OUTER_VOCAB = 2 # outerWall 0-1
L_DIST = 4 # 距离场码字序列长度
class GinkaMaskGIT(nn.Module):
def __init__(
self, num_classes: int = 16, d_model: int = 192, dim_ff: int = 512,
nhead: int = 8, num_layers: int = 4, map_h: int = 13, map_w: int = 13,
d_z: int = 64, z_seq_len: int = 6
d_z: int = 64, z_seq_len: int = 6, z_dist_len: int = L_DIST
):
super().__init__()
self.map_h = map_h
@ -33,8 +34,11 @@ class GinkaMaskGIT(nn.Module):
# z 投影:逐 token 线性变换,保持序列结构
self.z_proj = nn.Linear(d_z, d_z)
# 条件融合投影z_seq_len 个 z token + 2 个结构 token + 1 个剩余密度 token
self.cond_proj = nn.Linear((z_seq_len + 2 + 5) * d_z, d_model)
# 距离场 z 投影
self.z_dist_proj = nn.Linear(d_z, d_z)
# 条件融合投影z_seq_len 个 z token + z_dist_len 个距离场 token + 2 个结构 token + 5 个剩余密度 token
self.cond_proj = nn.Linear((z_seq_len + z_dist_len + 2 + 5) * d_z, d_model)
# 纯 encoder Transformer条件向量 c 通过 AdaLN 注入每一层
self.transformer = Transformer(
@ -47,11 +51,13 @@ class GinkaMaskGIT(nn.Module):
self,
map: torch.Tensor,
z: torch.Tensor,
z_dist: torch.Tensor,
struct: torch.Tensor,
remain: torch.Tensor
) -> torch.Tensor:
# map: [B, H * W]
# z: [B, z_seq_len, d_z]
# z_dist: [B, z_dist_len, d_z]
# struct: [B, 2] — [cond_sym(0-7), cond_outer(0-1)]
# remain: [B, 5] float — [wall, door, monster, entrance, resource] 剩余密度
@ -61,14 +67,17 @@ class GinkaMaskGIT(nn.Module):
self.outer_embed(struct[:, 1])
], dim=1)
# 剩余密度:连续浮点向量投影为单个 d_z 维 token[B, 1, d_z]
# 剩余密度:连续浮点向量投影为 d_z 维 token[B, 5, d_z]
e_remain = self.remain_proj(remain.unsqueeze(-1))
# z逐 token 投影,保留序列结构 [B, z_seq_len, d_z]
z_proj = self.z_proj(z)
# 拼接所有条件 token → [B, z_seq_len+3, d_z],展平后投影到 d_model
cond_seq = torch.cat([z_proj, e_struct, e_remain], dim=1)
# 距离场 z 投影 [B, z_dist_len, d_z]
zd_proj = self.z_dist_proj(z_dist)
# 拼接所有条件 token → 展平后投影到 d_model
cond_seq = torch.cat([z_proj, zd_proj, e_struct, e_remain], dim=1)
c = self.cond_proj(cond_seq.reshape(cond_seq.size(0), -1)) # [B, d_model]
# tile embedding + 位置编码
@ -113,10 +122,12 @@ if __name__ == "__main__":
z_seq_len=6
).to(device)
z_dist_input = torch.randn(4, L_DIST, 64).to(device) # [4, L_DIST, 64]
print_memory(device, "初始化后")
start = time.perf_counter()
logits = model(map_input, z_input, struct_input, remain_input)
logits = model(map_input, z_input, z_dist_input, struct_input, remain_input)
end = time.perf_counter()
print_memory(device, "前向传播后")

View File

@ -8,16 +8,18 @@ from datetime import datetime
import cv2
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.optim as optim
from tqdm import tqdm
from torch.utils.data import DataLoader
from .vqvae.quantize import VectorQuantizer
from .vqvae.model import GinkaVQVAE
from .vqvae.model import GinkaVQVAE, DistFieldEncoder
from .maskGIT.model import GinkaMaskGIT
from .dataset import GinkaSeperatedDataset
from shared.image import matrix_to_image_cv
from shared.distance import DIST_VOCAB, compute_distance_field_tensor
# 三阶段级联地图生成训练脚本
#
@ -47,6 +49,14 @@ VQ_DIM_FF = 1024 # VQ-VAE 前馈网络隐层维度
VQ_D_MODEL = 256 # VQ-VAE Transformer 模型维度
VQ_NHEAD = 4 # VQ-VAE 多头注意力头数
# 距离场编码器超参
L_DIST = 4 # 距离场码字序列长度
K_DIST = 16 # 距离场 codebook 大小
DIST_D_MODEL = 128 # 距离场编码器模型维度
DIST_LAYERS = 3 # 距离场编码器 Transformer 层数
DIST_DIM_FF = 512 # 距离场编码器 FF 维度
VQ_BETA_DIST = 0.5 # 距离场 commit loss 权重
# 第一阶段 MaskGIT 超参
STAGE1_MG_DMODEL = 512
STAGE1_MG_NHEAD = 4
@ -163,22 +173,45 @@ def build_model(device: torch.device):
quantizer3 = VectorQuantizer(K=VQ_K, d_z=VQ_D_Z).to(device)
quantizers = (quantizer1, quantizer2, quantizer3)
# 九个模块参数合并到同一优化器,端到端联合训练
# 距离场编码器与量化器:将 L1 距离场编码为离散 latent z_dist
dist_encoder = DistFieldEncoder(
vocab=DIST_VOCAB, L=L_DIST, d_z=VQ_D_Z, d_model=DIST_D_MODEL,
nhead=VQ_NHEAD, num_layers=DIST_LAYERS, dim_ff=DIST_DIM_FF,
map_h=MAP_H, map_w=MAP_W
).to(device)
dist_quantizer = VectorQuantizer(K=K_DIST, d_z=VQ_D_Z).to(device)
# latent dropout 用可学习 mask token各阶段共享
latent_mask_embedding = nn.Parameter(torch.randn(1, 1, VQ_D_Z, device=device) * 0.02)
# 所有模块参数合并到同一优化器
all_params = (
list(vq1.parameters()) + list(vq2.parameters()) + list(vq3.parameters()) +
list(mg1.parameters()) + list(mg2.parameters()) + list(mg3.parameters()) +
list(quantizer1.parameters()) + list(quantizer2.parameters()) + list(quantizer3.parameters())
list(quantizer1.parameters()) + list(quantizer2.parameters()) + list(quantizer3.parameters()) +
list(dist_encoder.parameters()) + list(dist_quantizer.parameters()) +
[latent_mask_embedding]
)
optimizer = optim.AdamW(all_params, lr=LR, weight_decay=1e-4)
# 余弦退火:从 LR 线性衰减至 MIN_LR周期为全部训练轮数
scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=EPOCHS, eta_min=MIN_LR)
return vq1, vq2, vq3, mg1, mg2, mg3, quantizers, optimizer, scheduler
return vq1, vq2, vq3, mg1, mg2, mg3, quantizers, optimizer, scheduler, latent_mask_embedding, dist_encoder, dist_quantizer
def cross_entropy_loss(logits, target):
# logits: [B, L, C],需转为 [B, C, L] 以匹配 cross_entropy 期望格式
return F.cross_entropy(logits.permute(0, 2, 1), target)
def apply_z_dropout(
z_q: torch.Tensor,
mask_embedding: nn.Parameter,
drop_prob: float
) -> torch.Tensor:
# 以 drop_prob 概率将 z_q 中的码字替换为可学习 mask 嵌入
# z_q: [B, L, d_z], mask_embedding: [1, 1, d_z]
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()
if total_hits.item() <= 0:
@ -217,6 +250,7 @@ def sample_reference_inputs(
model: torch.nn.Module,
reference: torch.Tensor,
z_q: torch.Tensor,
z_dist: torch.Tensor,
struct: torch.Tensor,
target_density: torch.Tensor,
stage: int,
@ -229,6 +263,7 @@ def sample_reference_inputs(
with torch.no_grad():
current = sampled_reference.clone()
z_q_detached = z_q.detach()
z_dist_detached = z_dist.detach()
for _ in range(rollout_steps):
masked_positions = current == MASK_TOKEN
@ -237,7 +272,7 @@ def sample_reference_inputs(
break
remain = compute_remaining(current, target_density, stage)
logits = model(current, z_q_detached, struct, remain)
logits = model(current, z_q_detached, z_dist_detached, struct, remain)
probs = F.softmax(logits, dim=-1)
dist = torch.distributions.Categorical(probs)
predicted = dist.sample()
@ -304,7 +339,8 @@ def compute_remaining(
def maskgit_sample(
model: torch.nn.Module, inp: torch.Tensor, z: torch.Tensor,
struct: torch.Tensor, target_density: torch.Tensor, stage: int, steps: int,
z_dist: torch.Tensor, struct: torch.Tensor, target_density: torch.Tensor,
stage: int, steps: int,
target_tiles: list[int] | None = None, keep_fixed: bool = True
) -> np.ndarray:
# target_tiles: 本阶段负责生成的图块 ID 列表None 表示接受所有类别stage1
@ -323,7 +359,7 @@ def maskgit_sample(
# 迭代去掩码:每步根据置信度分数重新决定掩码位置
for step in range(steps):
remain = compute_remaining(current, target_density, stage)
logits = model(current, z, struct, remain)
logits = model(current, z, z_dist, struct, remain)
probs = F.softmax(logits, dim=-1)
dist = torch.distributions.Categorical(probs)
@ -390,7 +426,7 @@ def maskgit_sample(
current[0, still_masked] = 0
else:
remain = compute_remaining(current, target_density, stage)
logits = model(current, z, struct, remain)
logits = model(current, z, z_dist, struct, remain)
current[0, still_masked] = torch.argmax(logits[0, still_masked], dim=-1)
return current[0].cpu().numpy().reshape(MAP_H, MAP_W)
@ -398,6 +434,7 @@ def maskgit_sample(
def full_generate_specific_z(
input: torch.Tensor,
z_q: tuple[torch.Tensor, torch.Tensor, torch.Tensor],
z_dist: torch.Tensor,
struct: torch.Tensor,
target_density: torch.Tensor,
models: list[torch.nn.Module],
@ -411,14 +448,14 @@ def full_generate_specific_z(
# 三阶段级联生成,但使用给定的 z
pred1_np = maskgit_sample(
mg1, input.clone(), z1, struct, target_density, 1,
mg1, input.clone(), z1, z_dist, struct, target_density, 1,
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[inp2 == 0] = MASK_TOKEN
pred2_np = maskgit_sample(
mg2, inp2, z2, struct, target_density, 2,
mg2, inp2, z2, z_dist, struct, target_density, 2,
GENERATE_STEP, target_tiles=[2, 6, 4, 5], keep_fixed=keep_fixed[1]
)
merged12 = pred1_np.copy()
@ -427,7 +464,7 @@ def full_generate_specific_z(
inp3[inp3 == 0] = MASK_TOKEN
pred3_np = maskgit_sample(
mg3, inp3, z3, struct, target_density, 3,
mg3, inp3, z3, z_dist, struct, target_density, 3,
GENERATE_STEP, target_tiles=[3], keep_fixed=keep_fixed[2]
)
merged123 = merged12.copy()
@ -469,10 +506,12 @@ def keep_label(kf: tuple[bool, bool, bool]) -> str:
def build_dataset_sample_case(
dataset: GinkaSeperatedDataset,
models: list[torch.nn.Module],
dist_models: tuple,
device: torch.device,
idx: int | None = None
) -> dict:
vq1, vq2, vq3, mg1, mg2, mg3, quantizers, optimizer, scheduler = models
dist_encoder, dist_quantizer = dist_models
sample = dataset.random_sample_map(idx=idx)
enc1_t = sample["encoder_stage1"].to(device).reshape(1, MAP_SIZE)
@ -480,6 +519,7 @@ def build_dataset_sample_case(
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)
@ -488,12 +528,15 @@ def build_dataset_sample_case(
z_q, commit_loss, code_hits = quantize_stage_latents(
quantizers, z_e1, z_e2, z_e3
)
z_e_dist = dist_encoder(dist_field_t)
z_dist, _, _, _, _ = dist_quantizer(z_e_dist)
return {
"sample": sample,
"struct": struct_t,
"target_density": target_density_t,
"z_q": z_q,
"z_dist": z_dist,
"sample_idx": sample["sample_idx"]
}
@ -543,7 +586,7 @@ def visualize_part1(batch, logits1, logits2, logits3, tile_dict):
return grid
# 验证可视化 part2行1=真实地图三阶段行2=stage1 输入与使用真实 z 自回归生成的各阶段结果
def visualize_part2(batch, z_q, models, device, tile_dict):
def visualize_part2(batch, z_q, z_dist, models, device, tile_dict):
SEP = 3
TILE_SIZE = 32
img_h = MAP_H * TILE_SIZE
@ -558,7 +601,7 @@ def visualize_part2(batch, z_q, models, device, tile_dict):
z_q_single = (z_q[0][0:1], z_q[1][0:1], z_q[2][0:1])
kf = rand_keep()
auto_pred1_np, auto_merged12, auto_merged123 = full_generate_specific_z(
inp1_t, z_q_single, struct_t, target_density_t, models, device, keep_fixed=kf
inp1_t, z_q_single, z_dist[0:1], struct_t, target_density_t, models, device, keep_fixed=kf
)
kf_label = 'fix' if kf[0] else 'free'
@ -591,6 +634,7 @@ def visualize_part2(batch, z_q, models, device, tile_dict):
def visualize_part4(
train_dataset: GinkaSeperatedDataset,
models: list[torch.nn.Module],
dist_models: tuple,
device: torch.device,
tile_dict
):
@ -610,11 +654,11 @@ def visualize_part4(
results = []
for _ in range(5):
case = build_dataset_sample_case(train_dataset, models, device)
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["struct"], case["target_density"],
seed, case["z_q"], case["z_dist"], case["struct"], case["target_density"],
models, device, keep_fixed=kf
)
result = annotate_labels(
@ -636,29 +680,31 @@ def visualize_part4(
return grid
def visualize_validate(
batch, logits1, logits2, logits3, z_q,
models: list[torch.nn.Module], device: torch.device, tile_dict,
batch, logits1, logits2, logits3, z_q, z_dist,
models: list[torch.nn.Module], dist_models: tuple, device: torch.device, tile_dict,
train_dataset: GinkaSeperatedDataset, epoch: int, batch_idx: int
):
save_dir = f"result/seperated/e{epoch}"
os.makedirs(save_dir, exist_ok=True)
cv2.imwrite(f"{save_dir}/val{batch_idx}.png", visualize_part1(batch, logits1, logits2, logits3, tile_dict))
cv2.imwrite(f"{save_dir}/full{batch_idx}.png", visualize_part2(batch, z_q, models, device, tile_dict))
cv2.imwrite(f"{save_dir}/rand{batch_idx}.png", visualize_part4(train_dataset, models, device, 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))
def validate(
dataloader: DataLoader,
models: list[torch.nn.Module],
dist_models: tuple,
device: torch.device,
tile_dict,
train_dataset: GinkaSeperatedDataset,
epoch: int
):
vq1, vq2, vq3, mg1, mg2, mg3, quantizers, optimizer, scheduler = models
dist_encoder, dist_quantizer = dist_models
quantizer1, quantizer2, quantizer3 = quantizers
# 切换为推理模式(关闭 Dropout / BatchNorm 统计更新)
for m in [vq1, vq2, vq3, mg1, mg2, mg3]:
for m in [vq1, vq2, vq3, mg1, mg2, mg3, dist_encoder]:
m.eval()
# 累计各阶段损失(跨所有 batch 求和,最终除以 batch 数得到均值)
@ -696,6 +742,11 @@ def validate(
struct = batch["struct_inject"].to(device)
target_density = batch["target_density"].to(device)
dist_field = batch["distance_field"].to(device)
# 距离场编码与量化
z_e_dist = dist_encoder(dist_field)
z_dist, _, _, _, _ = dist_quantizer(z_e_dist)
# VQ 编码:各阶段独立编码并分别量化
z_e1 = vq1(enc1) # [B, L, d_z]
@ -711,10 +762,10 @@ def validate(
remain2 = compute_remaining(inp2, target_density, 2)
remain3 = compute_remaining(inp3, target_density, 3)
# 三阶段 MaskGIT 推理:各阶段接收自己的 z_q
logits1 = mg1(inp1, z_q1, struct, remain1)
logits2 = mg2(inp2, z_q2, struct, remain2)
logits3 = mg3(inp3, z_q3, struct, remain3)
# 三阶段 MaskGIT 推理:各阶段接收自己的 z_q 和共享的 z_dist
logits1 = mg1(inp1, z_q1, z_dist, struct, remain1)
logits2 = mg2(inp2, z_q2, z_dist, struct, remain2)
logits3 = mg3(inp3, z_q3, z_dist, struct, remain3)
loss1_total += cross_entropy_loss(logits1, target1)
loss2_total += cross_entropy_loss(logits2, target2)
@ -751,8 +802,8 @@ def validate(
# 每个 batch 生成三种可视化图val/full/rand
visualize_validate(
batch, logits1, logits2, logits3, z_q,
models, device, tile_dict, train_dataset, epoch, idx
batch, logits1, logits2, logits3, z_q, z_dist,
models, dist_models, device, tile_dict, train_dataset, epoch, idx
)
idx += 1
@ -764,7 +815,7 @@ def validate(
tqdm.write(f" density {tile_names[tile_id]}: mae={avg_mae:.4f} over={avg_over:.4f}")
# 恢复训练模式
for m in [vq1, vq2, vq3, mg1, mg2, mg3]:
for m in [vq1, vq2, vq3, mg1, mg2, mg3, dist_encoder]:
m.train()
return loss1_total, loss2_total, loss3_total, commit_total, code_hits_total
@ -772,15 +823,18 @@ def validate(
def train(device: torch.device):
args = parse_arguments()
models = build_model(device)
vq1, vq2, vq3, mg1, mg2, mg3, quantizers, optimizer, scheduler = models
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]
dist_models = (dist_encoder, dist_quantizer)
quantizer1, quantizer2, quantizer3 = quantizers
tqdm.write(f"Device: {device}")
model_list = [
("vq1", vq1), ("vq2", vq2), ("vq3", vq3),
("mg1", mg1), ("mg2", mg2), ("mg3", mg3),
("quantizer1", quantizer1), ("quantizer2", quantizer2), ("quantizer3", quantizer3)
("quantizer1", quantizer1), ("quantizer2", quantizer2), ("quantizer3", quantizer3),
("dist_encoder", dist_encoder), ("dist_quantizer", dist_quantizer)
]
total_params = 0
for name, m in model_list:
@ -803,6 +857,12 @@ def train(device: torch.device):
quantizer1.load_state_dict(ckpt["quantizer1"])
quantizer2.load_state_dict(ckpt["quantizer2"])
quantizer3.load_state_dict(ckpt["quantizer3"])
if "dist_encoder" in ckpt:
dist_encoder.load_state_dict(ckpt["dist_encoder"])
if "dist_quantizer" in ckpt:
dist_quantizer.load_state_dict(ckpt["dist_quantizer"])
if "latent_mask_embedding" in ckpt:
latent_mask_embedding.data.copy_(ckpt["latent_mask_embedding"])
# load_optim=False 时可跳过优化器/调度器恢复(适合调整学习率后继续训练)
if args.load_optim and "optimizer" in ckpt:
optimizer.load_state_dict(ckpt["optimizer"])
@ -861,6 +921,7 @@ def train(device: torch.device):
# 结构条件向量:[cond_sym, cond_outer]
struct = batch["struct_inject"].to(device)
target_density = batch["target_density"].to(device)
dist_field = batch["distance_field"].to(device)
optimizer.zero_grad()
@ -875,25 +936,35 @@ def train(device: torch.device):
)
z_q1, z_q2, z_q3 = z_q
# 距离场编码与量化
z_e_dist = dist_encoder(dist_field)
z_dist_raw, _, commit_loss_dist, _, _ = dist_quantizer(z_e_dist)
z_dist = z_dist_raw
# latent dropout训练时随机丢弃部分码字替换为可学习 mask 嵌入
z_q1 = apply_z_dropout(z_q1, latent_mask_embedding, MG_Z_DROPOUT)
z_q2 = apply_z_dropout(z_q2, latent_mask_embedding, MG_Z_DROPOUT)
z_q3 = apply_z_dropout(z_q3, latent_mask_embedding, MG_Z_DROPOUT)
rollout_steps = build_reference_rollout_steps(REFERENCE_SAMPLE_PROB)
inp1 = sample_reference_inputs(
mg1, inp1, z_q1, struct, target_density, 1, rollout_steps
mg1, inp1, z_q1, z_dist, struct, target_density, 1, rollout_steps
)
inp2 = sample_reference_inputs(
mg2, inp2, z_q2, struct, target_density, 2, rollout_steps
mg2, inp2, z_q2, z_dist, struct, target_density, 2, rollout_steps
)
inp3 = sample_reference_inputs(
mg3, inp3, z_q3, struct, target_density, 3, rollout_steps
mg3, inp3, z_q3, z_dist, struct, target_density, 3, rollout_steps
)
remain1 = compute_remaining(inp1, target_density, 1)
remain2 = compute_remaining(inp2, target_density, 2)
remain3 = compute_remaining(inp3, target_density, 3)
# 三阶段 MaskGIT 前向:各阶段接收自己的 z_q、struct 和动态 remain 条件
logits1 = mg1(inp1, z_q1, struct, remain1)
logits2 = mg2(inp2, z_q2, struct, remain2)
logits3 = mg3(inp3, z_q3, struct, remain3)
# 三阶段 MaskGIT 前向:各阶段接收自己的 z_q、z_dist、struct 和动态 remain
logits1 = mg1(inp1, z_q1, z_dist, struct, remain1)
logits2 = mg2(inp2, z_q2, z_dist, struct, remain2)
logits3 = mg3(inp3, z_q3, z_dist, struct, remain3)
# 三阶段 Cross Entropy + VQ commit loss 加权求和
loss1 = cross_entropy_loss(logits1, target1)
@ -902,7 +973,7 @@ def train(device: torch.device):
loss1_weighted = STAGE1_CE_WEIGHT * loss1
loss2_weighted = STAGE2_CE_WEIGHT * loss2
loss3_weighted = STAGE3_CE_WEIGHT * loss3
commit_weighted = VQ_BETA * commit_loss
commit_weighted = VQ_BETA * commit_loss + VQ_BETA_DIST * commit_loss_dist
loss = loss1_weighted + loss2_weighted + loss3_weighted + commit_weighted
loss.backward()
@ -936,7 +1007,7 @@ def train(device: torch.device):
# 每 CHECKPOINT 个 epoch 执行一次验证、可视化和检查点保存
if (epoch + 1) % CHECKPOINT == 0:
losses = validate(
dataloader_val, models, device, tile_dict, dataset, epoch + 1
dataloader_val, models, dist_models, device, tile_dict, dataset, epoch + 1
)
loss1_total, loss2_total, loss3_total, commit_total, code_hits_total = losses
loss1_weighted = STAGE1_CE_WEIGHT * loss1_total
@ -970,6 +1041,9 @@ def train(device: torch.device):
"quantizer1": quantizer1.state_dict(),
"quantizer2": quantizer2.state_dict(),
"quantizer3": quantizer3.state_dict(),
"dist_encoder": dist_encoder.state_dict(),
"dist_quantizer": dist_quantizer.state_dict(),
"latent_mask_embedding": latent_mask_embedding.data,
"optimizer": optimizer.state_dict(),
"scheduler": scheduler.state_dict(),
}, ckpt_path)
@ -988,6 +1062,9 @@ def train(device: torch.device):
"quantizer1": quantizer1.state_dict(),
"quantizer2": quantizer2.state_dict(),
"quantizer3": quantizer3.state_dict(),
"dist_encoder": dist_encoder.state_dict(),
"dist_quantizer": dist_quantizer.state_dict(),
"latent_mask_embedding": latent_mask_embedding.data,
"optimizer": optimizer.state_dict(),
"scheduler": scheduler.state_dict(),
}, final_path)

View File

@ -149,6 +149,56 @@ class GinkaVQVAE(nn.Module):
return z_e
class DistFieldEncoder(nn.Module):
# 距离场编码器:将 L1 距离场编码为 latent z
#
# 使用与 GinkaVQVAE 相同的 Transformer 架构但更轻量:
# DistEmbedding (vocab=DIST_VOCAB) → + 2D 位置编码 → + summary tokens
# → 浅层 Transformer → Linear 投影 → z_e_dist [B, L, d_z]
def __init__(
self, vocab: int = 13, L: int = 4, d_z: int = 64,
d_model: int = 128, nhead: int = 4, num_layers: int = 3,
dim_ff: int = 512, map_h: int = 13, map_w: int = 13
):
super().__init__()
self.L = L
self.map_h = map_h
self.map_w = map_w
self.dist_embedding = nn.Embedding(vocab, d_model)
self.row_embedding = nn.Parameter(torch.randn(1, map_h, d_model) * 0.02)
self.col_embedding = nn.Parameter(torch.randn(1, map_w, d_model) * 0.02)
self.summary_tokens = nn.Parameter(torch.randn(1, L, d_model) * 0.02)
encoder_layer = nn.TransformerEncoderLayer(
d_model=d_model, nhead=nhead, dim_feedforward=dim_ff, batch_first=True,
activation='gelu', norm_first=True, dropout=0.1,
)
self.transformer = nn.TransformerEncoder(encoder_layer, num_layers=num_layers)
self.proj = nn.Sequential(
nn.Linear(d_model, d_z),
nn.LayerNorm(d_z),
)
def forward(self, dist_field: torch.Tensor) -> torch.Tensor:
# dist_field: [B, H*W] 整数,值域 [0, DIST_MAX_BUCKET]
B, _ = dist_field.shape
row_idx = torch.arange(self.map_h, device=dist_field.device).repeat_interleave(self.map_w)
col_idx = torch.arange(self.map_w, device=dist_field.device).repeat(self.map_h)
pos = self.row_embedding[0, row_idx] + self.col_embedding[0, col_idx]
x = self.dist_embedding(dist_field) + pos # [B, H*W, d_model]
summary = self.summary_tokens.expand(B, -1, -1) # [B, L, d_model]
x = torch.cat([summary, x], dim=1) # [B, L+H*W, d_model]
x = self.transformer(x)
z_e = self.proj(x[:, :self.L]) # [B, L, d_z]
return z_e
if __name__ == "__main__":
device = torch.device("cuda") if torch.cuda.is_available() else torch.device("cpu")

41
shared/distance.py Normal file
View File

@ -0,0 +1,41 @@
from collections import deque
import numpy as np
import torch
DIST_MAX_BUCKET = 12 # 距离分桶上限
DIST_VOCAB = DIST_MAX_BUCKET + 1 # 0-12 共 13 个桶
def compute_distance_field(map_matrix: np.ndarray) -> np.ndarray:
# map_matrix: [H, W] 整数矩阵
# 返回: [H * W] 分桶后的整数距离,值域 [0, DIST_MAX_BUCKET]
wall_mask = (map_matrix == 1)
h, w = map_matrix.shape
dist = np.full((h, w), h + w, dtype=np.int32)
dist[wall_mask] = 0
queue = deque()
for i in range(h):
for j in range(w):
if wall_mask[i, j]:
queue.append((i, j))
directions = [(0, 1), (0, -1), (1, 0), (-1, 0)]
while queue:
i, j = queue.popleft()
d = dist[i, j]
for di, dj in directions:
ni, nj = i + di, j + dj
if 0 <= ni < h and 0 <= nj < w and dist[ni, nj] > d + 1:
dist[ni, nj] = d + 1
queue.append((ni, nj))
dist = np.clip(dist, 0, DIST_MAX_BUCKET)
return dist.flatten().astype(np.int64)
def compute_distance_field_tensor(map_tensor: torch.Tensor) -> torch.Tensor:
# map_tensor: [B, H, W] 或 [1, H, W]
batch = []
for i in range(map_tensor.shape[0]):
d = compute_distance_field(map_tensor[i].cpu().numpy())
batch.append(torch.from_numpy(d))
return torch.stack(batch, dim=0).to(map_tensor.device)