import math 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 # MaskGIT 采样函数:通过迭代去掩码从离散隐变量 z 生成地图 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] # 初始 MASK:邻接已有墙壁的空地 init_adj = compute_adjacency_mask(state) state[init_adj & (state == 0)] = MASK_TOKEN for step in range(max_steps): mask_pos = (state == MASK_TOKEN) # [B, MAP_SIZE] 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 adj = compute_adjacency_mask(state) new_mask = adj & (state == 0) if not new_mask.any(): break state[new_mask] = MASK_TOKEN state[state == MASK_TOKEN] = 0 return state.cpu().numpy().reshape(state.size(0), MAP_H, MAP_W) 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) # 迭代去掩码:每步根据置信度分数重新决定掩码位置 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) # 余弦退火调度:随步数推进,保留掩码的位置数量递减至 0 ratio = math.cos(((step + 1) / steps) * math.pi / 2) num_to_mask = math.floor(ratio * MAP_SIZE) # 结构位:current 中非空地、非掩码的位置(来自上一阶段,始终保留) struct_mask = (current[:] != MASK_TOKEN) & (current[:] != 0) # 候选位:sampled 为目标图块且不覆盖结构位 candidate_mask = torch.isin(sampled[:], target_tensor) & ~struct_mask cand_count = candidate_mask.sum() reveal_count = max(0, int(cand_count.item()) - num_to_mask) next_state = current[:].clone() if reveal_count > 0 and cand_count > 0: cand_indices = candidate_mask.nonzero(as_tuple=False) cand_conf = confidences[:][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 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 ) -> tuple[np.ndarray, np.ndarray, np.ndarray]: # 三阶段级联生成:Stage1 墙壁生长 → Stage2 门/怪/入口 → Stage3 资源 # 返回 (stage1结果, stage1+2合并, 最终完整地图),形状均为 [B, H, W] device = inp.device pred1_np = wall_growth_sample( models.mg1, inp, z1, struct, target_density ) # [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