mirror of
https://github.com/unanmed/ginka-generator.git
synced 2026-08-14 18:12:28 +08:00
337 lines
11 KiB
Python
337 lines
11 KiB
Python
import os
|
||
import sys
|
||
import torch
|
||
import numpy as np
|
||
from datetime import datetime
|
||
|
||
os.environ["CUDA_VISIBLE_DEVICES"] = ""
|
||
os.makedirs("result/test", exist_ok=True)
|
||
|
||
device = torch.device("cpu")
|
||
|
||
def log(msg):
|
||
print(f" [{datetime.now().strftime('%H:%M:%S')}] {msg}")
|
||
|
||
def sep(title):
|
||
print(f"\n{'='*60}\n {title}\n{'='*60}")
|
||
|
||
sep("1. 模型初始化")
|
||
|
||
from ginka.model import (
|
||
SeperatedModels, NUM_CLASSES, MASK_TOKEN, MAP_W, MAP_H,
|
||
VQ_L1, VQ_L2, VQ_L3
|
||
)
|
||
from ginka.utils import MAP_SIZE
|
||
|
||
models = SeperatedModels(device)
|
||
log(f"SeperatedModels 创建成功")
|
||
|
||
total = sum(
|
||
p.numel()
|
||
for m in [models.vq1, models.vq2, models.vq3,
|
||
models.mg1, models.mg2, models.mg3,
|
||
models.quantizer1, models.quantizer2, models.quantizer3]
|
||
for p in m.parameters()
|
||
) + models.latent_mask_embedding.numel()
|
||
log(f"总参数量: {total:,}")
|
||
|
||
sep("2. VQ-VAE 编码器前向")
|
||
|
||
dummy_map = torch.randint(0, 7, (1, MAP_SIZE), dtype=torch.long)
|
||
|
||
with torch.no_grad():
|
||
z_e1 = models.vq1(dummy_map)
|
||
z_e2 = models.vq2(dummy_map)
|
||
z_e3 = models.vq3(dummy_map)
|
||
|
||
log(f"z_e1 shape: {z_e1.shape} (期望 [1, {VQ_L1}, 64])")
|
||
log(f"z_e2 shape: {z_e2.shape} (期望 [1, {VQ_L2}, 64])")
|
||
log(f"z_e3 shape: {z_e3.shape} (期望 [1, {VQ_L3}, 64])")
|
||
|
||
assert z_e1.shape == (1, VQ_L1, 64), f"z_e1 shape mismatch"
|
||
assert z_e2.shape == (1, VQ_L2, 64), f"z_e2 shape mismatch"
|
||
assert z_e3.shape == (1, VQ_L3, 64), f"z_e3 shape mismatch"
|
||
log("VQ-VAE 编码器 [OK]")
|
||
|
||
sep("3. VectorQuantizer 量化")
|
||
|
||
models.quantizer1.eval()
|
||
models.quantizer2.eval()
|
||
models.quantizer3.eval()
|
||
|
||
with torch.no_grad():
|
||
z_q1, idx1, cl1, ppl1, hits1, ent1 = models.quantizer1(z_e1)
|
||
z_q2, idx2, cl2, ppl2, hits2, ent2 = models.quantizer2(z_e2)
|
||
z_q3, idx3, cl3, ppl3, hits3, ent3 = models.quantizer3(z_e3)
|
||
|
||
log(f"z_q1 shape: {z_q1.shape}")
|
||
log(f"indices1 shape: {idx1.shape} max={idx1.max().item()} < K={models.quantizer1.K}")
|
||
log(f"z_q2 shape: {z_q2.shape}")
|
||
log(f"z_q3 shape: {z_q3.shape}")
|
||
log(f"commit_loss: {cl1.item():.4f}, {cl2.item():.4f}, {cl3.item():.4f}")
|
||
log(f"entropy_loss: {ent1.item():.4f}, {ent2.item():.4f}, {ent3.item():.4f}")
|
||
|
||
assert z_q1.shape == z_e1.shape
|
||
assert z_q2.shape == z_e2.shape
|
||
assert z_q3.shape == z_e3.shape
|
||
log("VectorQuantizer 量化 [OK]")
|
||
|
||
# 测试采样
|
||
models.quantizer1.eval()
|
||
with torch.no_grad():
|
||
z_sampled = models.quantizer1.sample(1, VQ_L1, device)
|
||
log(f"sample 采样 shape: {z_sampled.shape} (期望 [1, {VQ_L1}, 64])")
|
||
assert z_sampled.shape == (1, VQ_L1, 64)
|
||
log("codebook 采样 [OK]")
|
||
|
||
sep("4. MaskGIT 前向")
|
||
|
||
dummy_struct = torch.tensor([[3, 1]], dtype=torch.long)
|
||
dummy_remain = torch.tensor([[0.2, 0.1, 0.3, 0.1, 0.3]], dtype=torch.float)
|
||
|
||
models.mg1.eval()
|
||
models.mg2.eval()
|
||
models.mg3.eval()
|
||
|
||
with torch.no_grad():
|
||
logits1 = models.mg1(dummy_map, z_q1, dummy_struct, dummy_remain)
|
||
logits2 = models.mg2(dummy_map, z_q2, dummy_struct, dummy_remain)
|
||
logits3 = models.mg3(dummy_map, z_q3, dummy_struct, dummy_remain)
|
||
|
||
log(f"mg1 logits shape: {logits1.shape} (期望 [1, {MAP_SIZE}, {NUM_CLASSES}])")
|
||
log(f"mg2 logits shape: {logits2.shape}")
|
||
log(f"mg3 logits shape: {logits3.shape}")
|
||
|
||
assert logits1.shape == (1, MAP_SIZE, NUM_CLASSES)
|
||
assert logits2.shape == (1, MAP_SIZE, NUM_CLASSES)
|
||
assert logits3.shape == (1, MAP_SIZE, NUM_CLASSES)
|
||
log("MaskGIT 前向 [OK]")
|
||
|
||
sep("5. 数据集加载")
|
||
|
||
from ginka.dataset import GinkaSeperatedDataset
|
||
|
||
ds_train = GinkaSeperatedDataset("ginka-dataset.json", subset_weights=(0.5, 0.3, 0.2))
|
||
log(f"训练集大小: {len(ds_train)} 条")
|
||
log(f"密度统计: wall [{ds_train.density_stats['wall_min_density']:.3f}, {ds_train.density_stats['wall_max_density']:.3f}]")
|
||
|
||
sample = ds_train[0]
|
||
log(f"input_stage1 shape: {sample['input_stage1'].shape}")
|
||
log(f"target_stage1 shape: {sample['target_stage1'].shape}")
|
||
log(f"encoder_stage1 shape: {sample['encoder_stage1'].shape}")
|
||
log(f"struct_inject: {sample['struct_inject'].tolist()}")
|
||
log(f"target_density: {sample['target_density'].tolist()}")
|
||
|
||
for key in ["input_stage1", "input_stage2", "input_stage3",
|
||
"target_stage1", "target_stage2", "target_stage3",
|
||
"encoder_stage1", "encoder_stage2", "encoder_stage3"]:
|
||
assert sample[key].shape == (MAP_H, MAP_W), f"{key} shape wrong: {sample[key].shape}"
|
||
|
||
log(f"MASK_TOKEN count in inp1: {(sample['input_stage1'] == MASK_TOKEN).sum().item()}")
|
||
log(f"MASK_TOKEN count in inp2: {(sample['input_stage2'] == MASK_TOKEN).sum().item()}")
|
||
log(f"MASK_TOKEN count in inp3: {(sample['input_stage3'] == MASK_TOKEN).sum().item()}")
|
||
log("数据集加载 [OK]")
|
||
|
||
random_sample = ds_train.random_sample_map()
|
||
log(f"random_sample keys: {list(random_sample.keys())}")
|
||
log(f"raw_map shape: {random_sample['raw_map'].shape}")
|
||
log("随机样本采样 [OK]")
|
||
|
||
sep("6. compute_remaining")
|
||
|
||
from ginka.utils import compute_remaining
|
||
|
||
inp_t = sample["input_stage1"].reshape(1, MAP_SIZE)
|
||
td_t = sample["target_density"].unsqueeze(0)
|
||
r1 = compute_remaining(inp_t, td_t, 1)
|
||
r2 = compute_remaining(inp_t, td_t, 2)
|
||
r3 = compute_remaining(inp_t, td_t, 3)
|
||
|
||
log(f"remain (stage1) shape: {r1.shape} (期望 [1, 5])")
|
||
log(f"remain (stage2) shape: {r2.shape}")
|
||
log(f"remain (stage3) shape: {r3.shape}")
|
||
assert r1.shape == (1, 5)
|
||
log(f"stage1 remain values: {r1[0].tolist()}")
|
||
log("compute_remaining [OK]")
|
||
|
||
sep("7. compute_adjacency_mask")
|
||
|
||
from ginka.utils import compute_adjacency_mask
|
||
|
||
wall_map = torch.full((1, MAP_SIZE), 0, dtype=torch.long)
|
||
wall_map[0, 50] = 1
|
||
adj = compute_adjacency_mask(wall_map)
|
||
log(f"adjacency mask shape: {adj.shape}")
|
||
log(f"adjacent count: {adj.sum().item()} (期望 4)")
|
||
|
||
wall_map_b = torch.stack([wall_map[0], wall_map[0]], dim=0)
|
||
adj_b = compute_adjacency_mask(wall_map_b)
|
||
log(f"batched adj shape: {adj_b.shape}")
|
||
log(f"batched adj[0] sum: {adj_b[0].sum().item()}, adj[1] sum: {adj_b[1].sum().item()}")
|
||
log("compute_adjacency_mask [OK]")
|
||
|
||
sep("8. wall_growth_sample")
|
||
|
||
from ginka.sample import wall_growth_sample
|
||
|
||
models.mg1.eval()
|
||
with torch.no_grad():
|
||
inp_seed = torch.full((1, MAP_SIZE), 0, dtype=torch.long)
|
||
seeds = torch.randperm(MAP_SIZE)[:5]
|
||
inp_seed[0, seeds] = 1
|
||
|
||
struct_t = sample["struct_inject"].unsqueeze(0)
|
||
td_t = sample["target_density"].unsqueeze(0)
|
||
|
||
walls = wall_growth_sample(
|
||
models.mg1, inp_seed, z_q1, struct_t, td_t, max_steps=3
|
||
)
|
||
|
||
log(f"wall_growth output shape: {walls.shape} (期望 [1, {MAP_H}, {MAP_W}])")
|
||
assert walls.shape == (1, MAP_H, MAP_W)
|
||
log(f"wall count: {(walls == 1).sum()}")
|
||
log("wall_growth_sample [OK]")
|
||
|
||
sep("9. maskgit_sample")
|
||
|
||
from ginka.sample import maskgit_sample
|
||
|
||
models.mg2.eval()
|
||
with torch.no_grad():
|
||
inp2 = torch.tensor(walls.reshape(1, MAP_SIZE), dtype=torch.long)
|
||
inp2[inp2 == 0] = MASK_TOKEN
|
||
|
||
pred2 = maskgit_sample(
|
||
models.mg2, inp2, z_q2, struct_t, td_t,
|
||
stage=2, steps=3, target_tiles=[2, 4, 5, 6]
|
||
)
|
||
|
||
log(f"maskgit_sample (stage2) output shape: {pred2.shape}")
|
||
assert pred2.shape == (1, MAP_H, MAP_W)
|
||
|
||
models.mg3.eval()
|
||
with torch.no_grad():
|
||
merged12 = walls.copy()
|
||
merged12[pred2 != 0] = pred2[pred2 != 0]
|
||
inp3 = torch.tensor(merged12.reshape(1, MAP_SIZE), dtype=torch.long)
|
||
inp3[inp3 == 0] = MASK_TOKEN
|
||
|
||
pred3 = maskgit_sample(
|
||
models.mg3, inp3, z_q3, struct_t, td_t,
|
||
stage=3, steps=3, target_tiles=[3]
|
||
)
|
||
|
||
log(f"maskgit_sample (stage3) output shape: {pred3.shape}")
|
||
assert pred3.shape == (1, MAP_H, MAP_W)
|
||
log("maskgit_sample [OK]")
|
||
|
||
sep("10. full_generate(三阶段级联)")
|
||
|
||
from ginka.sample import full_generate
|
||
|
||
models.mg1.eval()
|
||
models.mg2.eval()
|
||
models.mg3.eval()
|
||
|
||
with torch.no_grad():
|
||
inp = torch.full((1, MAP_SIZE), 0, dtype=torch.long)
|
||
seed_idx = torch.randperm(MAP_SIZE)[:8]
|
||
inp[0, seed_idx] = 1
|
||
|
||
pred1, merged12, merged123 = full_generate(
|
||
inp, z_q1, z_q2, z_q3,
|
||
struct_t, td_t, models, steps=3
|
||
)
|
||
|
||
log(f"pred1 shape: {pred1.shape} (期望 [1, {MAP_H}, {MAP_W}])")
|
||
log(f"merged12 shape: {merged12.shape}")
|
||
log(f"merged123 shape: {merged123.shape}")
|
||
assert pred1.shape == (1, MAP_H, MAP_W)
|
||
assert merged12.shape == (1, MAP_H, MAP_W)
|
||
assert merged123.shape == (1, MAP_H, MAP_W)
|
||
log("full_generate [OK]")
|
||
|
||
sep("11. 交叉熵 Loss(仅掩码位置)")
|
||
|
||
import torch.nn.functional as F
|
||
|
||
def cross_entropy_loss_test(logits, target, mask):
|
||
loss = F.cross_entropy(logits.permute(0, 2, 1), target, reduction='none')
|
||
masked = loss[mask]
|
||
if masked.numel() == 0:
|
||
return torch.tensor(0.0, requires_grad=True)
|
||
return masked.mean()
|
||
|
||
mask = (sample["input_stage1"].reshape(1, MAP_SIZE) == MASK_TOKEN)
|
||
target = sample["target_stage1"].reshape(1, MAP_SIZE)
|
||
|
||
loss_test = cross_entropy_loss_test(logits1, target, mask)
|
||
log(f"masked CE loss: {loss_test.item():.4f}")
|
||
log(f"mask 中 True 的数量: {mask.sum().item()} / {MAP_SIZE}")
|
||
|
||
# 对比:全量 loss vs 掩码 loss
|
||
full_loss = F.cross_entropy(logits1.permute(0, 2, 1), target)
|
||
log(f"全量 CE loss: {full_loss.item():.4f} vs 掩码 CE loss: {loss_test.item():.4f}")
|
||
log("CE Loss [OK]")
|
||
|
||
sep("12. 可视化生成")
|
||
|
||
from ginka.dataset import compute_symmetry
|
||
from shared.image import matrix_to_image_cv
|
||
import cv2
|
||
|
||
tile_dict = {}
|
||
for f in os.listdir("tiles"):
|
||
name = os.path.splitext(f)[0]
|
||
img = cv2.imread(f"tiles/{f}", cv2.IMREAD_UNCHANGED)
|
||
if img is not None:
|
||
tile_dict[name] = img
|
||
|
||
TILE_SIZE = 32
|
||
|
||
def save_map(mat, path, label=""):
|
||
if isinstance(mat, torch.Tensor):
|
||
mat = mat.cpu().numpy()
|
||
if mat.ndim == 3:
|
||
mat = mat[0]
|
||
img = matrix_to_image_cv(mat, tile_dict, TILE_SIZE)
|
||
if label:
|
||
img = cv2.copyMakeBorder(img, 0, 24, 0, 0, cv2.BORDER_CONSTANT, value=(255, 255, 255))
|
||
cv2.putText(img, label, (4, TILE_SIZE * MAP_H + 18),
|
||
cv2.FONT_HERSHEY_SIMPLEX, 0.5, (0, 0, 0), 1)
|
||
cv2.imwrite(path, img)
|
||
|
||
save_map(pred1, "result/test/pred1_stage1.png", "Stage1: walls")
|
||
save_map(merged12, "result/test/pred2_merged12.png", "Stage1+2: +doors/monsters/entrances")
|
||
save_map(merged123, "result/test/pred3_merged123.png", "Full: +resources")
|
||
|
||
raw = random_sample["raw_map"].cpu().numpy().reshape(MAP_H, MAP_W)
|
||
save_map(raw, "result/test/raw_map.png", "Original ground truth")
|
||
|
||
inp1_img = sample["input_stage1"].cpu().numpy().reshape(MAP_H, MAP_W)
|
||
save_map(inp1_img, "result/test/inp1.png", "Stage1 input (masked)")
|
||
|
||
log("可视化输出到 result/test/ 目录:")
|
||
log(" pred1_stage1.png — 仅墙壁")
|
||
log(" pred2_merged12.png — 墙壁+功能元素")
|
||
log(" pred3_merged123.png — 完整地图")
|
||
log(" raw_map.png — 原始真实地图")
|
||
log(" inp1.png — Stage1 输入(掩码后)")
|
||
|
||
sep("测试结果汇总")
|
||
|
||
print("\n 所有 12 项测试通过 [OK]\n")
|
||
print(f" - 模型初始化与参数统计")
|
||
print(f" - VQ-VAE 编码器 (3 个阶段)")
|
||
print(f" - VectorQuantizer 量化与采样")
|
||
print(f" - MaskGIT 前向 (3 个阶段)")
|
||
print(f" - 数据集加载与掩码策略")
|
||
print(f" - compute_remaining")
|
||
print(f" - compute_adjacency_mask (含 batch)")
|
||
print(f" - wall_growth_sample")
|
||
print(f" - maskgit_sample (stage2 + stage3)")
|
||
print(f" - full_generate (三阶段级联)")
|
||
print(f" - 掩码位置 CE Loss")
|
||
print(f" - 可视化输出")
|
||
print()
|