ginka-generator/test_refactor.py

337 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 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()