ginka-generator/ginka/sample.py

145 lines
5.6 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

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