From 103b02cf08bbcc3a673848a8298d7b0b5f9bef1d Mon Sep 17 00:00:00 2001 From: unanmed <1319491857@qq.com> Date: Sun, 2 Aug 2026 13:35:49 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20=E5=A2=99=E5=A3=81=E7=94=9F=E6=88=90?= =?UTF-8?q?=E7=AD=96=E7=95=A5=E8=B0=83=E6=95=B4?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- docs/wall-growth-generation-design.md | 367 ++++++++++++++++++++++++++ ginka/dataset.py | 64 ++++- ginka/train_seperated.py | 281 ++++++++++++++------ 3 files changed, 627 insertions(+), 85 deletions(-) create mode 100644 docs/wall-growth-generation-design.md diff --git a/docs/wall-growth-generation-design.md b/docs/wall-growth-generation-design.md new file mode 100644 index 0000000..5d5d89c --- /dev/null +++ b/docs/wall-growth-generation-design.md @@ -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 输出 +- [ ] 对比指标:墙壁连通率、孤立墙块数、房间闭合度 diff --git a/ginka/dataset.py b/ginka/dataset.py index 581e23b..6f39afb 100644 --- a/ginka/dataset.py +++ b/ginka/dataset.py @@ -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]) diff --git a/ginka/train_seperated.py b/ginka/train_seperated.py index 3a0943f..e5e8092 100644 --- a/ginka/train_seperated.py +++ b/ginka/train_seperated.py @@ -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"