ginka-generator/ginka/train_seperated.py

486 lines
20 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 argparse
import os
import sys
import random
from datetime import datetime
import cv2
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from tqdm import tqdm
from torch.utils.data import DataLoader
from .model import SeperatedModels
from .model import (
MASK_TOKEN, MAP_W, MAP_H, EPOCHS, VQ_GAMMA
)
from .utils import (
compute_remaining, MAP_SIZE, summarize_codebook_hits
)
from .sample import full_generate
from .dataset import GinkaSeperatedDataset, rect_mask, ensure_wall_connection
from shared.image import matrix_to_image_cv, annotate
# 图块 ID 定义:
# 0. 空地 1. 墙壁 2. 普通门 3. 资源 4. 怪物 5. 入口 6. 机关门 7. 掩码MASK_TOKEN
# 三阶段级联地图生成训练脚本
#
# 整体架构:
# VQ-VAE三组独立编码器 vq1/vq2/vq3将三阶段地图上下文分别编码为离散潜变量
# 再由三个独立 VectorQuantizer 分别量化为 z_q1/z_q2/z_q3
# 三个独立 MaskGITmg1/mg2/mg3分别以各自阶段 z_q 和 struct_inject 为条件,
# 逐阶段迭代解码地图图块序列。
#
# 三阶段生成目标:
# stage1 → floor / wall地图骨架
# stage2 → door / monster / entrance功能性实体
# stage3 → resource资源点
device = torch.device(
"cuda:0" if torch.cuda.is_available()
else "mps" if torch.backends.mps.is_available()
else "cpu"
)
disable_tqdm = not sys.stdout.isatty()
# 训练与推理超参
VQ_BETA = 0.5 # 承诺损失权重
STAGE1_CE_WEIGHT = 1.0 # Stage1 CE 损失权重
STAGE2_CE_WEIGHT = 1.0 # Stage2 CE 损失权重
STAGE3_CE_WEIGHT = 1.0 # Stage3 CE 损失权重
MG_Z_DROPOUT = 0.1 # z 隐变量 Dropout 概率
BATCH_SIZE = 64 # 每批样本数
CHECKPOINT = 20 # 每隔多少 epoch 保存检查点并执行验证
GENERATE_STEP = 18 # MaskGIT 采样步数
SEED_SAMPLE_STEPS = 24 # 墙壁种子生长采样步数
SUBSET_WEIGHTS = (0.5, 0.3, 0.2) # 每个子集的概率
def _str2bool(v: str):
if isinstance(v, bool): return v
if v.lower() in ('true', '1', 'yes'): return True
if v.lower() in ('false', '0', 'no'): return False
raise argparse.ArgumentTypeError(f"布尔值应为 True/False收到: {v!r}")
def parse_arguments():
parser = argparse.ArgumentParser(description="三阶段级联训练")
parser.add_argument("--resume", type=_str2bool, default=False)
parser.add_argument("--state", type=str, default="", help="续训时检查点路径")
parser.add_argument("--train", type=str, default="ginka-dataset.json")
parser.add_argument("--validate", type=str, default="ginka-eval.json")
parser.add_argument("--load_optim", type=_str2bool, default=True)
return parser.parse_args()
def cross_entropy_loss(logits, target, mask):
# logits: [B, L, C]target: [B, L]mask: [B, L] boolTrue = 参与 loss
loss = F.cross_entropy(logits.permute(0, 2, 1), target, reduction='none')
masked = loss[mask]
if masked.numel() == 0:
return torch.tensor(0.0, device=logits.device, requires_grad=True)
return masked.mean()
def apply_z_dropout(
z_q: torch.Tensor,
mask_embedding: nn.Parameter,
drop_prob: float
) -> torch.Tensor:
# 以 drop_prob 概率将 z_q 中的码字替换为可学习 mask 嵌入
# z_q: [B, L, d_z], mask_embedding: [1, 1, d_z]
mask = torch.rand(z_q.shape[0], z_q.shape[1], 1, device=z_q.device) < drop_prob
return z_q * (~mask).float() + mask_embedding * mask.float()
def quantize_stage_latents(
models: SeperatedModels,
z_e1: torch.Tensor,
z_e2: torch.Tensor,
z_e3: torch.Tensor
) -> tuple:
z_q1, _, commit_loss1, _, code_hits1, entropy1 = models.quantizer1(z_e1)
z_q2, _, commit_loss2, _, code_hits2, entropy2 = models.quantizer2(z_e2)
z_q3, _, commit_loss3, _, code_hits3, entropy3 = models.quantizer3(z_e3)
commit_loss = (commit_loss1 + commit_loss2 + commit_loss3) / 3
entropy_loss = entropy1 + entropy2 + entropy3
code_hits = (code_hits1, code_hits2, code_hits3)
return (z_q1, z_q2, z_q3), commit_loss, code_hits, entropy_loss
# 每张图 2 行 x 4 列MASKED | PREDICTED | MASKED | PREDICTED
def make_compare_grid(samples_src, samples_pred, map_keys, tile_dict, TILE_SIZE):
SEP = 3
img_h = MAP_H * TILE_SIZE
img_w = MAP_W * TILE_SIZE
cols = 4
rows = 2
grid = np.ones((rows * img_h + (rows + 1) * SEP, cols * img_w + (cols + 1) * SEP, 3), dtype=np.uint8) * 255
for r in range(rows):
row_y = SEP + r * (img_h + SEP)
for c in range(cols):
is_src = (c % 2 == 0)
idx = r * 2 + c // 2
col_x = SEP + c * (img_w + SEP)
data = samples_src[idx] if is_src else samples_pred[idx]
grid[row_y:row_y + img_h, col_x:col_x + img_w] = annotate(
matrix_to_image_cv(data, tile_dict, TILE_SIZE), map_keys[idx], y=14
)
# 每两张图片之间的分隔线
for c in range(cols - 1):
line_x = (c + 1) * (img_w + SEP) + SEP // 2
grid[:, line_x:line_x + SEP, :] = 180
return grid
def visualize_seed_growth(
train_dataset: GinkaSeperatedDataset,
models: SeperatedModels,
device: torch.device,
tile_dict,
epoch: int
):
save_dir = f"result/seperated/e{epoch}"
os.makedirs(save_dir, exist_ok=True)
samples_src = []
samples_pred = []
map_keys = []
for _ in range(20):
sample = train_dataset.random_sample_map()
struct_t = sample["struct_inject"].to(device).reshape(1, -1)
target_density_t = sample["target_density"].to(device).reshape(1, -1)
enc1_t = sample["encoder_stage1"].to(device).reshape(1, MAP_SIZE)
enc2_t = sample["encoder_stage2"].to(device).reshape(1, MAP_SIZE)
enc3_t = sample["encoder_stage3"].to(device).reshape(1, MAP_SIZE)
enc1_np = enc1_t.cpu().numpy().reshape(MAP_H, MAP_W)
with torch.no_grad():
z_e1 = models.vq1(enc1_t)
z_e2 = models.vq2(enc2_t)
z_e3 = models.vq3(enc3_t)
z_q, _, _, _ = quantize_stage_latents(models, z_e1, z_e2, z_e3)
z1, z2, z3 = z_q
inp = torch.full((1, MAP_SIZE), 0, dtype=torch.long, device=device)
seed_count = random.randint(7, 14)
seed_idx = torch.randperm(MAP_SIZE, device=device)[:seed_count]
inp[0, seed_idx] = 1
inp_vis = inp.clone()
inp_vis[inp_vis == 0] = MASK_TOKEN
_, _, merged = full_generate(
inp, z1, z2, z3, struct_t, target_density_t, models,
steps=SEED_SAMPLE_STEPS, seed_mode=True, stage1_method="growth"
)
samples_src.append(inp_vis.cpu().numpy().reshape(MAP_H, MAP_W))
samples_pred.append(merged[0])
map_keys.append(sample["map_name"])
for i in range(5):
grid = make_compare_grid(
samples_src[i*4:i*4+4], samples_pred[i*4:i*4+4],
map_keys[i*4:i*4+4], tile_dict, 32
)
cv2.imwrite(f"{save_dir}/seed_growth_{i}.png", grid)
def visualize_seed_maskgit(
train_dataset: GinkaSeperatedDataset,
models: SeperatedModels,
device: torch.device,
tile_dict,
epoch: int
):
save_dir = f"result/seperated/e{epoch}"
os.makedirs(save_dir, exist_ok=True)
samples_src = []
samples_pred = []
map_keys = []
for _ in range(20):
sample = train_dataset.random_sample_map()
struct_t = sample["struct_inject"].to(device).reshape(1, -1)
target_density_t = sample["target_density"].to(device).reshape(1, -1)
enc1_t = sample["encoder_stage1"].to(device).reshape(1, MAP_SIZE)
enc2_t = sample["encoder_stage2"].to(device).reshape(1, MAP_SIZE)
enc3_t = sample["encoder_stage3"].to(device).reshape(1, MAP_SIZE)
enc1_np = enc1_t.cpu().numpy().reshape(MAP_H, MAP_W)
with torch.no_grad():
z_e1 = models.vq1(enc1_t)
z_e2 = models.vq2(enc2_t)
z_e3 = models.vq3(enc3_t)
z_q, _, _, _ = quantize_stage_latents(models, z_e1, z_e2, z_e3)
z1, z2, z3 = z_q
inp = torch.full((1, MAP_SIZE), MASK_TOKEN, dtype=torch.long, device=device)
_, _, merged = full_generate(
inp, z1, z2, z3, struct_t, target_density_t, models,
steps=GENERATE_STEP, seed_mode=True, stage1_method="maskgit"
)
samples_src.append(inp.cpu().numpy().reshape(MAP_H, MAP_W))
samples_pred.append(merged[0])
map_keys.append(sample["map_name"])
for i in range(5):
grid = make_compare_grid(
samples_src[i*4:i*4+4], samples_pred[i*4:i*4+4],
map_keys[i*4:i*4+4], tile_dict, 32
)
cv2.imwrite(f"{save_dir}/seed_maskgit_{i}.png", grid)
def visualize_mask_growth(
train_dataset: GinkaSeperatedDataset,
models: SeperatedModels,
device: torch.device,
tile_dict,
epoch: int
):
save_dir = f"result/seperated/e{epoch}"
os.makedirs(save_dir, exist_ok=True)
samples_src = []
samples_pred = []
map_keys = []
for _ in range(20):
sample = train_dataset.random_sample_map()
raw_map = sample["raw_map"].cpu().numpy().reshape(MAP_H, MAP_W)
enc1_np = sample["encoder_stage1"].cpu().numpy().reshape(MAP_H, MAP_W)
enc1_t = sample["encoder_stage1"].to(device).reshape(1, MAP_SIZE)
enc2_t = sample["encoder_stage2"].to(device).reshape(1, MAP_SIZE)
enc3_t = sample["encoder_stage3"].to(device).reshape(1, MAP_SIZE)
struct_t = sample["struct_inject"].to(device).reshape(1, -1)
target_density_t = sample["target_density"].to(device).reshape(1, -1)
# 矩形分块掩码 + 确保墙壁连通性,与训练时一致
ratio = random.uniform(0.3, 0.9)
rmask = rect_mask(ratio).reshape(MAP_H, MAP_W)
rmask = ensure_wall_connection(rmask, enc1_np)
inp = torch.tensor(raw_map.flatten(), dtype=torch.long, device=device).reshape(1, MAP_SIZE)
inp[0, rmask.flatten()] = MASK_TOKEN
masked_np = inp.cpu().numpy().reshape(MAP_H, MAP_W)
with torch.no_grad():
z_e1 = models.vq1(enc1_t)
z_e2 = models.vq2(enc2_t)
z_e3 = models.vq3(enc3_t)
z_q, _, _, _ = quantize_stage_latents(models, z_e1, z_e2, z_e3)
z1, z2, z3 = z_q
_, _, merged = full_generate(
inp, z1, z2, z3, struct_t, target_density_t, models,
stage1_method="growth"
)
samples_src.append(masked_np)
samples_pred.append(merged[0])
map_keys.append(sample["map_name"])
for i in range(5):
grid = make_compare_grid(
samples_src[i*4:i*4+4], samples_pred[i*4:i*4+4],
map_keys[i*4:i*4+4], tile_dict, 32
)
cv2.imwrite(f"{save_dir}/mask_growth_{i}.png", grid)
def visualize_mask_maskgit(
train_dataset: GinkaSeperatedDataset,
models: SeperatedModels,
device: torch.device,
tile_dict,
epoch: int
):
save_dir = f"result/seperated/e{epoch}"
os.makedirs(save_dir, exist_ok=True)
samples_src = []
samples_pred = []
map_keys = []
for _ in range(20):
sample = train_dataset.random_sample_map()
raw_map = sample["raw_map"].cpu().numpy().reshape(MAP_H, MAP_W)
enc1_np = sample["encoder_stage1"].cpu().numpy().reshape(MAP_H, MAP_W)
enc1_t = sample["encoder_stage1"].to(device).reshape(1, MAP_SIZE)
enc2_t = sample["encoder_stage2"].to(device).reshape(1, MAP_SIZE)
enc3_t = sample["encoder_stage3"].to(device).reshape(1, MAP_SIZE)
struct_t = sample["struct_inject"].to(device).reshape(1, -1)
target_density_t = sample["target_density"].to(device).reshape(1, -1)
# 矩形分块掩码 + 确保墙壁连通性,与训练时一致
ratio = random.uniform(0.3, 0.9)
rmask = rect_mask(ratio).reshape(MAP_H, MAP_W)
rmask = ensure_wall_connection(rmask, enc1_np)
inp = torch.tensor(raw_map.flatten(), dtype=torch.long, device=device).reshape(1, MAP_SIZE)
inp[0, rmask.flatten()] = MASK_TOKEN
masked_np = inp.cpu().numpy().reshape(MAP_H, MAP_W)
with torch.no_grad():
z_e1 = models.vq1(enc1_t)
z_e2 = models.vq2(enc2_t)
z_e3 = models.vq3(enc3_t)
z_q, _, _, _ = quantize_stage_latents(models, z_e1, z_e2, z_e3)
z1, z2, z3 = z_q
_, _, merged = full_generate(
inp, z1, z2, z3, struct_t, target_density_t, models,
stage1_method="maskgit"
)
samples_src.append(masked_np)
samples_pred.append(merged[0])
map_keys.append(sample["map_name"])
for i in range(5):
grid = make_compare_grid(
samples_src[i*4:i*4+4], samples_pred[i*4:i*4+4],
map_keys[i*4:i*4+4], tile_dict, 32
)
cv2.imwrite(f"{save_dir}/mask_maskgit_{i}.png", grid)
def train(device: torch.device):
args = parse_arguments()
result = SeperatedModels(device)
tqdm.write(f"Device: {device}")
model_list = [
("vq1", result.vq1), ("vq2", result.vq2), ("vq3", result.vq3),
("mg1", result.mg1), ("mg2", result.mg2), ("mg3", result.mg3),
("quantizer1", result.quantizer1), ("quantizer2", result.quantizer2), ("quantizer3", result.quantizer3)
]
total_params = 0
for name, m in model_list:
n = sum(p.numel() for p in m.parameters())
total_params += n
tqdm.write(f"{name}: {n:,} params")
tqdm.write(f"Total: {total_params:,} params")
start_epoch = 0
if args.resume:
# 从指定检查点恢复:加载所有模型权重及训练状态
start_epoch = result.load(args.state, load_optim=args.load_optim, map_location=device)
tqdm.write(f"Resumed from epoch {start_epoch}: {args.state}")
os.makedirs("result/seperated", exist_ok=True)
dataset = GinkaSeperatedDataset(
args.train, subset_weights=SUBSET_WEIGHTS
)
dataloader = DataLoader(
dataset, batch_size=BATCH_SIZE, shuffle=True
)
# 预加载图块图像,键为文件名(不含扩展名),用于可视化时将 ID 映射为像素图
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
for epoch in tqdm(range(start_epoch, EPOCHS), desc="Seperated Training", disable=disable_tqdm):
loss_total = torch.Tensor([0]).to(device)
loss1_total = torch.Tensor([0]).to(device)
loss2_total = torch.Tensor([0]).to(device)
loss3_total = torch.Tensor([0]).to(device)
commit_total = torch.Tensor([0]).to(device)
entropy_total = torch.Tensor([0]).to(device)
code_hits_total = (torch.zeros(result.quantizer1.K, device=device), torch.zeros(result.quantizer2.K, device=device), torch.zeros(result.quantizer3.K, device=device)) # validate
for batch in tqdm(dataloader, leave=False, desc="Epoch Progress", disable=disable_tqdm):
# 三阶段各自的掩码输入序列、预测目标和编码器上下文
inp1 = batch["input_stage1"].to(device).reshape(-1, MAP_SIZE)
target1 = batch["target_stage1"].to(device).reshape(-1, MAP_SIZE)
enc1 = batch["encoder_stage1"].to(device).reshape(-1, MAP_SIZE)
inp2 = batch["input_stage2"].to(device).reshape(-1, MAP_SIZE)
target2 = batch["target_stage2"].to(device).reshape(-1, MAP_SIZE)
enc2 = batch["encoder_stage2"].to(device).reshape(-1, MAP_SIZE)
inp3 = batch["input_stage3"].to(device).reshape(-1, MAP_SIZE)
target3 = batch["target_stage3"].to(device).reshape(-1, MAP_SIZE)
enc3 = batch["encoder_stage3"].to(device).reshape(-1, MAP_SIZE)
# 结构条件向量:[cond_sym, cond_outer]
struct = batch["struct_inject"].to(device)
target_density = batch["target_density"].to(device)
result.optimizer.zero_grad() # 训练循环
# VQ 编码:各阶段编码器分别处理各自上下文切片
z_e1 = result.vq1(enc1) # [B, L, d_z]
z_e2 = result.vq2(enc2)
z_e3 = result.vq3(enc3)
# 三阶段分别量化,各自使用独立 codebook
z_q, commit_loss, code_hits, entropy_loss = quantize_stage_latents(
result, z_e1, z_e2, z_e3
)
z_q1, z_q2, z_q3 = z_q
# latent dropout训练时随机丢弃部分码字替换为可学习 mask 嵌入
z_q1 = apply_z_dropout(z_q1, result.latent_mask_embedding, MG_Z_DROPOUT)
z_q2 = apply_z_dropout(z_q2, result.latent_mask_embedding, MG_Z_DROPOUT)
z_q3 = apply_z_dropout(z_q3, result.latent_mask_embedding, MG_Z_DROPOUT)
remain1 = compute_remaining(inp1, target_density, 1)
remain2 = compute_remaining(inp2, target_density, 2)
remain3 = compute_remaining(inp3, target_density, 3)
# 三阶段 MaskGIT 前向:各阶段接收自己的 z_q、struct 和动态 remain
logits1 = result.mg1(inp1, z_q1, struct, remain1)
logits2 = result.mg2(inp2, z_q2, struct, remain2)
logits3 = result.mg3(inp3, z_q3, struct, remain3)
# 三阶段 Cross Entropy
# Stage1: 仅对掩码位置计算 loss
# Stage2/3: 掩码 + 空地位置均计算 loss模型需学习哪里不能放东西
mask1 = (inp1 == MASK_TOKEN)
mask2 = (inp2 == MASK_TOKEN) | (inp2 == 0)
mask3 = (inp3 == MASK_TOKEN) | (inp3 == 0)
loss1 = cross_entropy_loss(logits1, target1, mask1)
loss2 = cross_entropy_loss(logits2, target2, mask2)
loss3 = cross_entropy_loss(logits3, target3, mask3)
loss1_weighted = STAGE1_CE_WEIGHT * loss1
loss2_weighted = STAGE2_CE_WEIGHT * loss2
loss3_weighted = STAGE3_CE_WEIGHT * loss3
commit_weighted = VQ_BETA * commit_loss + VQ_GAMMA * entropy_loss
loss = loss1_weighted + loss2_weighted + loss3_weighted + commit_weighted
loss.backward()
result.optimizer.step()
# detach 后累加,避免保留计算图占用显存
loss_total += loss.detach()
loss1_total += loss1.detach()
loss2_total += loss2.detach()
loss3_total += loss3.detach()
commit_total += commit_loss.detach()
entropy_total += entropy_loss.detach()
code_hits_total = (code_hits_total[0] + code_hits[0].detach(), code_hits_total[1] + code_hits[1].detach(), code_hits_total[2] + code_hits[2].detach()) # accumulate train
# 每个 epoch 结束后更新学习率
result.scheduler.step()
data_length = len(dataloader)
# 总 loss 去掉 entropy_loss 以便与 CE + VQ 对齐
loss_display = (loss_total.item() - VQ_GAMMA * entropy_total.item()) / data_length
entropy_display = entropy_total.item() / data_length
stats = summarize_codebook_hits(code_hits_total)
parts = []
for name in ["q1(stage1)", "q2(stage2)", "q3(stage3)"]:
s = stats[name]
parts.append(f"{s['active']}/{s['K']} ppl={s['ppl']:.1f}")
c = stats["combined"]
tqdm.write(
f"[{datetime.now().strftime('%Y-%m-%d %H:%M:%S')}] "
f"E: {epoch + 1} | "
f"Loss: {loss_display:.4f} | "
f"CE: {loss1_total.item() / data_length:.4f}, {loss2_total.item() / data_length:.4f}, {loss3_total.item() / data_length:.4f} | "
f"cmt: {commit_total.item() / data_length:.4f} ent: {entropy_display:+.4f} | "
f"VQ: {' | '.join(parts)} | "
f"Total: {c['active']}/{c['K']} ppl={c['ppl']:.1f} | "
f"LR: {result.scheduler.get_last_lr()[0]:.6f}"
)
# 每 CHECKPOINT 个 epoch 执行可视化并保存检查点
if (epoch + 1) % CHECKPOINT == 0:
visualize_seed_growth(dataset, result, device, tile_dict, epoch + 1)
visualize_seed_maskgit(dataset, result, device, tile_dict, epoch + 1)
visualize_mask_growth(dataset, result, device, tile_dict, epoch + 1)
visualize_mask_maskgit(dataset, result, device, tile_dict, epoch + 1)
ckpt_path = f"result/seperated/sep-{epoch + 1}.pth"
result.save(ckpt_path, epoch + 1, save_optim=True)
tqdm.write(f"Saved checkpoint: {ckpt_path}")
# 训练结束后保存最终完整权重
final_path = "result/seperated.pth"
result.save(final_path, EPOCHS, save_optim=False)
tqdm.write(f"Training complete. Final model saved: {final_path}")
if __name__ == "__main__":
train(device)