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