mirror of
https://github.com/unanmed/ginka-generator.git
synced 2026-08-14 18:12:28 +08:00
145 lines
5.6 KiB
Python
145 lines
5.6 KiB
Python
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
|