ginka-generator/ginka/sample.py

293 lines
11 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 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