import heapq import math import random from collections import deque import numpy as np import torch import torch.nn.functional as F from .model import MASK_TOKEN, MAP_H, MAP_W, SeperatedModels from .utils import ( compute_remaining, MAP_SIZE, compute_adjacency_mask, WALL_DENSITY_IDX, DOOR_DENSITY_IDX, MONSTER_DENSITY_IDX, ENTRANCE_DENSITY_IDX, RESOURCE_DENSITY_IDX ) # 采样器超参 DBC_GUMBEL = 2.0 # 揭开顺序的 Gumbel 噪声强度,0 关闭 DBC_JITTER = 0.08 # 图块预算的随机浮动,0 关闭 DBC_TRIES = 6 # 拒绝采样次数 # 每个阶段负责的图块,在 target_density 里对应的下标 STAGE_DENSITY_IDX = { 1: (WALL_DENSITY_IDX,), 2: (DOOR_DENSITY_IDX, MONSTER_DENSITY_IDX, ENTRANCE_DENSITY_IDX), 3: (RESOURCE_DENSITY_IDX,), } _NB4 = ((0, 1), (0, -1), (1, 0), (-1, 0)) def floor_label(m: np.ndarray): mask = (m != 1) lab = -np.ones(m.shape, int) n = 0 for i in range(MAP_H): for j in range(MAP_W): if mask[i, j] and lab[i, j] < 0: q = deque([(i, j)]) lab[i, j] = n while q: a, b = q.popleft() for da, db in _NB4: x, y = a + da, b + db if 0 <= x < MAP_H and 0 <= y < MAP_W and mask[x, y] and lab[x, y] < 0: lab[x, y] = n q.append((x, y)) n += 1 return lab, n def repair_connectivity(m: np.ndarray) -> np.ndarray: # 0-1 Dijkstra 打通最少的墙,使所有非墙格连成一片 m = m.copy() while True: lab, n = floor_label(m) if n <= 1: return m sizes = [int((lab == k).sum()) for k in range(n)] main = int(np.argmax(sizes)) tgt = next(k for k in range(n) if k != main) dist = np.full(m.shape, 1 << 30) pq = [] prev = {} end = None for i in range(MAP_H): for j in range(MAP_W): if lab[i, j] == main: dist[i, j] = 0 heapq.heappush(pq, (0, i, j)) while pq: d, i, j = heapq.heappop(pq) if d > dist[i, j]: continue if lab[i, j] == tgt: end = (i, j) break for da, db in _NB4: x, y = i + da, j + db if not (0 <= x < MAP_H and 0 <= y < MAP_W): continue nd = d + (1 if m[x, y] == 1 else 0) if nd < dist[x, y]: dist[x, y] = nd prev[(x, y)] = (i, j) heapq.heappush(pq, (nd, x, y)) if end is None: return m cur = end while cur in prev: if m[cur] == 1: m[cur] = 0 cur = prev[cur] def maskgit_sample( model: torch.nn.Module, inp: torch.Tensor, z: torch.Tensor, struct: torch.Tensor, target_density: torch.Tensor, stage: int, steps: int, target_tiles: list[int] ) -> np.ndarray: # 每步只从预测中选出置信度最高的若干 target_tile 候选揭开, # 其余已有结构(墙/门等非空地非掩码)原样保留, # 空地与掩码保持为 MASK,等待后续步骤继续填充。 current = inp.clone() target_tensor = torch.tensor(target_tiles, dtype=torch.long, device=inp.device) # 退火对象改成「本阶段还需要放多少个目标图块」 base = float(sum(target_density[0, i] for i in STAGE_DENSITY_IDX[stage])) * MAP_SIZE if DBC_JITTER > 0: base *= 1.0 + random.uniform(-DBC_JITTER, DBC_JITTER) budget = int(round(base)) need0 = max(0, budget - int(torch.isin(current[:], target_tensor).sum())) # 迭代去掩码:每步根据置信度分数重新决定掩码位置 for step in range(steps): remain = compute_remaining(current, target_density, stage) logits = model(current, z, struct, remain) probs = F.softmax(logits, dim=-1) dist = torch.distributions.Categorical(probs) sampled = dist.sample() confidences = torch.gather(probs, -1, sampled.unsqueeze(-1)).squeeze(-1) # Gumbel 噪声退火:提高揭开顺序的多样性(MaskGIT 原论文做法) if DBC_GUMBEL > 0: u = torch.rand_like(confidences).clamp_(1e-9, 1 - 1e-9) score = (torch.log(confidences.clamp_min(1e-9)) + DBC_GUMBEL * (1 - step / steps) * (-torch.log(-torch.log(u)))) else: score = confidences # 结构位:current 中非空地、非掩码的位置(来自上一阶段,始终保留) struct_mask = (current[:] != MASK_TOKEN) & (current[:] != 0) # 候选位:sampled 为目标图块且不覆盖结构位 candidate_mask = torch.isin(sampled[:], target_tensor) & ~struct_mask cand_count = candidate_mask.sum() # 预算感知揭开:按本阶段还需放置的图块数做余弦退火 ratio = math.cos(((step + 1) / steps) * math.pi / 2) still = math.floor(ratio * need0) now = max(0, budget - int(torch.isin(current[:], target_tensor).sum())) reveal_count = min(max(0, now - still), int(cand_count.item())) next_state = current[:].clone() if reveal_count > 0 and cand_count > 0: cand_indices = candidate_mask.nonzero(as_tuple=False) cand_conf = score[:][cand_indices[:, 0], cand_indices[:, 1]] top_k = min(reveal_count, cand_conf.size(0)) _, top_idx = torch.topk(cand_conf, k=top_k, largest=True) reveal_rows = cand_indices[top_idx, 0] reveal_cols = cand_indices[top_idx, 1] next_state[reveal_rows, reveal_cols] = sampled[reveal_rows, reveal_cols] # 结构位原样保留,其余未揭开的置为 MASK non_struct_non_revealed = (next_state == current[:]) & ~struct_mask next_state[non_struct_non_revealed & (next_state != MASK_TOKEN)] = MASK_TOKEN # 空地也重新标为 MASK(允许下一步继续填充) next_state[next_state == 0] = MASK_TOKEN current = next_state if (current[:] == MASK_TOKEN).all(): break # 兜底:未被填充的掩码位视为空地(不属于本阶段负责的图块) still_masked = (current[:] == MASK_TOKEN) current[still_masked] = 0 return current.cpu().numpy().reshape(current.size(0), MAP_H, MAP_W) def sample_with_retry( model: torch.nn.Module, inp: torch.Tensor, z: torch.Tensor, struct: torch.Tensor, target_density: torch.Tensor, stage: int, steps: int, target_tiles: list[int], tries: int = DBC_TRIES ) -> np.ndarray: # 拒绝采样:多次采样取地板连通的那张,都不连通则修复块数最少的 best = None best_n = 1 << 30 for _ in range(tries): m = maskgit_sample( model, inp.clone(), z, struct, target_density, stage, steps, target_tiles=target_tiles ) m_2d = m.reshape(MAP_H, MAP_W) _, n = floor_label(m_2d) if n == 1: return m if n < best_n: best = m best_n = n if len(best.shape) == 3: repaired = repair_connectivity(best[0]) repaired = repaired[np.newaxis, ...] else: repaired = repair_connectivity(best.reshape(MAP_H, MAP_W)) repaired = repaired.reshape(best.shape) return repaired def wall_growth_sample( model: torch.nn.Module, inp: torch.Tensor, z: torch.Tensor, struct: torch.Tensor, target_density: torch.Tensor, max_steps: int = 24 ) -> np.ndarray: # 墙壁生长算法:从 inp 中已有的墙壁出发,逐步向外生长 # 所有未固定位置初始为 MASK,只有与墙壁邻接的 MASK 才参与评估 state = inp.clone() # [B, MAP_SIZE] state[state == 0] = MASK_TOKEN for step in range(max_steps): adj = compute_adjacency_mask(state) mask_pos = adj & (state == MASK_TOKEN) if not mask_pos.any(): break remain = compute_remaining(state, target_density, 1) logits = model(state, z, struct, remain) probs = F.softmax(logits, dim=-1) wall_prob = probs[:, :, 1] # [B, MAP_SIZE] hits = (wall_prob > 0.5) & mask_pos state[hits] = 1 state[mask_pos & ~hits] = 0 state[state == MASK_TOKEN] = 0 return state.cpu().numpy().reshape(state.size(0), MAP_H, MAP_W) def full_generate( inp: torch.Tensor, z1: torch.Tensor, z2: torch.Tensor, z3: torch.Tensor, struct: torch.Tensor, target_density: torch.Tensor, models: SeperatedModels, steps: int = 18, seed_mode: bool = False, stage1_method: str = "maskgit" ) -> tuple[np.ndarray, np.ndarray, np.ndarray]: # 三阶段级联生成 # seed_mode=True 从空白生成墙壁(自由生成) # seed_mode=False 从 inp 中的已有结构出发(掩码补全) # stage1_method: "maskgit"=MaskGIT+拒绝采样, "growth"=墙壁生长算法 # 返回 (stage1结果, stage1+2合并, 最终完整地图),形状均为 [B, H, W] device = inp.device if seed_mode: if stage1_method == "growth": pred1_np = wall_growth_sample( models.mg1, inp, z1, struct, target_density ) else: stage1_inp = torch.full((inp.size(0), MAP_SIZE), MASK_TOKEN, dtype=torch.long, device=device) pred1_np = sample_with_retry( models.mg1, stage1_inp, z1, struct, target_density, 1, steps, target_tiles=[1] ) else: if stage1_method == "growth": pred1_np = wall_growth_sample( models.mg1, inp, z1, struct, target_density ) else: pred1_np = sample_with_retry( models.mg1, inp, z1, struct, target_density, 1, steps, target_tiles=[1] ) # [B, H, W] inp2 = torch.tensor(pred1_np.reshape(pred1_np.shape[0], -1), dtype=torch.long, device=device) inp2[inp2 == 0] = MASK_TOKEN pred2_np = maskgit_sample( models.mg2, inp2, z2, struct, target_density, 2, steps, target_tiles=[2, 4, 5, 6] ) # [B, H, W] merged12 = pred1_np.copy() merged12[pred2_np != 0] = pred2_np[pred2_np != 0] inp3 = torch.tensor(merged12.reshape(merged12.shape[0], -1), dtype=torch.long, device=device) inp3[inp3 == 0] = MASK_TOKEN pred3_np = maskgit_sample( models.mg3, inp3, z3, struct, target_density, 3, steps, target_tiles=[3] ) # [B, H, W] merged123 = merged12.copy() merged123[pred3_np != 0] = pred3_np[pred3_np != 0] return pred1_np, merged12, merged123