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