import torch import torch.nn as nn import torch.optim as optim from .vqvae.quantize import VectorQuantizer from .vqvae.model import GinkaVQVAE from .maskGIT.model import GinkaMaskGIT # 共用 VQ-VAE 超参(共享的编码维度) VQ_D_Z = 64 # 码字维度 VQ_GAMMA = 0.1 # entropy loss 权重,鼓励码本使用均匀 # 三通道 VQ 各自独立超参(L、K、层数、维度等均独立配置) # Stage1 墙壁骨架 — 结构最复杂,模型容量最大 VQ_L1 = 24 VQ_K1 = 32 VQ_D_MODEL1 = 384 VQ_NHEAD1 = 8 VQ_LAYERS1 = 6 VQ_DIM_FF1 = 1536 # Stage2 功能元素 — 中等复杂度 VQ_L2 = 12 VQ_K2 = 16 VQ_D_MODEL2 = 256 VQ_NHEAD2 = 4 VQ_LAYERS2 = 6 VQ_DIM_FF2 = 1024 # Stage3 资源分布 — 最简单,模型容量最小 VQ_L3 = 8 VQ_K3 = 16 VQ_D_MODEL3 = 192 VQ_NHEAD3 = 4 VQ_LAYERS3 = 4 VQ_DIM_FF3 = 768 # 第一阶段 MaskGIT 超参 STAGE1_MG_DMODEL = 512 STAGE1_MG_NHEAD = 4 STAGE1_MG_NUM_LAYERS = 8 STAGE1_MG_DIM_FF = 2048 # 第二阶段 MaskGIT 超参 STAGE2_MG_DMODEL = 256 STAGE2_MG_NHEAD = 4 STAGE2_MG_NUM_LAYERS = 6 STAGE2_MG_DIM_FF = 1024 # 第三阶段 MaskGIT 超参 STAGE3_MG_DMODEL = 256 STAGE3_MG_NHEAD = 4 STAGE3_MG_NUM_LAYERS = 6 STAGE3_MG_DIM_FF = 1024 # 各阶段 VQ commit loss 权重(当前未单独使用,统一由 VQ_BETA 控制) STAGE1_VQ_WEIGHT = 0.5 STAGE2_VQ_WEIGHT = 0.5 STAGE3_VQ_WEIGHT = 0.5 # 全局参数 NUM_CLASSES = 8 # 图块类型数 MASK_TOKEN = 7 # 掩码图块 TOTAL_K = VQ_K1 + VQ_K2 + VQ_K3 # 预计算,供日志输出使用 MAP_W = 13 # 地图宽度 MAP_H = 13 # 地图高度 LR = 1e-4 # AdamW 初始学习率 MIN_LR = 1e-6 # 余弦退火最低学习率 WEIGHT_DECAY = 1e-4 # L2 正则化系数 EPOCHS = 400 # 总训练轮数 class SeperatedModels: # 三阶段级联模型集合,封装所有子模块、优化器和调度器 vq1: GinkaVQVAE vq2: GinkaVQVAE vq3: GinkaVQVAE mg1: GinkaMaskGIT mg2: GinkaMaskGIT mg3: GinkaMaskGIT quantizers: tuple[VectorQuantizer, VectorQuantizer, VectorQuantizer] quantizer1: VectorQuantizer quantizer2: VectorQuantizer quantizer3: VectorQuantizer optimizer: optim.AdamW scheduler: optim.lr_scheduler.CosineAnnealingLR latent_mask_embedding: nn.Parameter def __init__(self, device: torch.device): # 三组 VQ-VAE 编码器:各自独立编码一个阶段的地图上下文 self.vq1 = GinkaVQVAE( num_classes=NUM_CLASSES, L=VQ_L1, K=VQ_K1, d_model=VQ_D_MODEL1, nhead=VQ_NHEAD1, num_layers=VQ_LAYERS1, dim_ff=VQ_DIM_FF1, map_h=MAP_H, map_w=MAP_W ).to(device) self.vq2 = GinkaVQVAE( num_classes=NUM_CLASSES, L=VQ_L2, K=VQ_K2, d_model=VQ_D_MODEL2, nhead=VQ_NHEAD2, num_layers=VQ_LAYERS2, dim_ff=VQ_DIM_FF2, map_h=MAP_H, map_w=MAP_W ).to(device) self.vq3 = GinkaVQVAE( num_classes=NUM_CLASSES, L=VQ_L3, K=VQ_K3, d_model=VQ_D_MODEL3, nhead=VQ_NHEAD3, num_layers=VQ_LAYERS3, dim_ff=VQ_DIM_FF3, map_h=MAP_H, map_w=MAP_W ).to(device) # 三个独立 MaskGIT 解码器,分别接收各自阶段的 z_q 作为条件 self.mg1 = GinkaMaskGIT( num_classes=NUM_CLASSES, d_model=STAGE1_MG_DMODEL, d_z=VQ_D_Z, dim_ff=STAGE1_MG_DIM_FF, nhead=STAGE1_MG_NHEAD, num_layers=STAGE1_MG_NUM_LAYERS, map_h=MAP_H, map_w=MAP_W, z_seq_len=VQ_L1 ).to(device) self.mg2 = GinkaMaskGIT( num_classes=NUM_CLASSES, d_model=STAGE2_MG_DMODEL, d_z=VQ_D_Z, dim_ff=STAGE2_MG_DIM_FF, nhead=STAGE2_MG_NHEAD, num_layers=STAGE2_MG_NUM_LAYERS, map_h=MAP_H, map_w=MAP_W, z_seq_len=VQ_L2 ).to(device) self.mg3 = GinkaMaskGIT( num_classes=NUM_CLASSES, d_model=STAGE3_MG_DMODEL, d_z=VQ_D_Z, dim_ff=STAGE3_MG_DIM_FF, nhead=STAGE3_MG_NHEAD, num_layers=STAGE3_MG_NUM_LAYERS, map_h=MAP_H, map_w=MAP_W, z_seq_len=VQ_L3 ).to(device) # 三个独立 VectorQuantizer:各阶段使用自己的码本大小 self.quantizer1 = VectorQuantizer(K=VQ_K1, d_z=VQ_D_Z).to(device) self.quantizer2 = VectorQuantizer(K=VQ_K2, d_z=VQ_D_Z).to(device) self.quantizer3 = VectorQuantizer(K=VQ_K3, d_z=VQ_D_Z).to(device) self.quantizers = (self.quantizer1, self.quantizer2, self.quantizer3) # latent dropout 用可学习 mask token,各阶段共享 self.latent_mask_embedding = nn.Parameter( torch.randn(1, 1, VQ_D_Z, device=device) * 0.02 ) # 所有模块参数合并到同一优化器,余弦退火衰减至 MIN_LR all_params = ( list(self.vq1.parameters()) + list(self.vq2.parameters()) + list(self.vq3.parameters()) + list(self.mg1.parameters()) + list(self.mg2.parameters()) + list(self.mg3.parameters()) + list(self.quantizer1.parameters()) + list(self.quantizer2.parameters()) + list(self.quantizer3.parameters()) + [self.latent_mask_embedding] ) self.optimizer = optim.AdamW(all_params, lr=LR, weight_decay=WEIGHT_DECAY) self.scheduler = optim.lr_scheduler.CosineAnnealingLR( self.optimizer, T_max=EPOCHS, eta_min=MIN_LR ) def __iter__(self): # 向后兼容:支持元组解包 return iter(( self.vq1, self.vq2, self.vq3, self.mg1, self.mg2, self.mg3, self.quantizers, self.optimizer, self.scheduler, self.latent_mask_embedding )) def __getitem__(self, idx): return list(self)[idx] def load(self, ckpt_path: str, load_optim: bool = True, map_location: str = "cpu") -> int: # 从检查点加载模型权重和训练状态,返回恢复的 epoch 编号 ckpt = torch.load(ckpt_path, map_location=map_location) self.vq1.load_state_dict(ckpt["vq1"]) self.vq2.load_state_dict(ckpt["vq2"]) self.vq3.load_state_dict(ckpt["vq3"]) self.mg1.load_state_dict(ckpt["mg1"]) self.mg2.load_state_dict(ckpt["mg2"]) self.mg3.load_state_dict(ckpt["mg3"]) self.quantizer1.load_state_dict(ckpt["quantizer1"]) self.quantizer2.load_state_dict(ckpt["quantizer2"]) self.quantizer3.load_state_dict(ckpt["quantizer3"]) if "latent_mask_embedding" in ckpt: self.latent_mask_embedding.data.copy_(ckpt["latent_mask_embedding"]) if load_optim and "optimizer" in ckpt: self.optimizer.load_state_dict(ckpt["optimizer"]) if load_optim and "scheduler" in ckpt: self.scheduler.load_state_dict(ckpt["scheduler"]) return ckpt.get("epoch", 0) def save(self, path: str, epoch: int): # 保存完整检查点(模型权重 + 优化器/调度器状态) torch.save({ "epoch": epoch, "vq1": self.vq1.state_dict(), "vq2": self.vq2.state_dict(), "vq3": self.vq3.state_dict(), "mg1": self.mg1.state_dict(), "mg2": self.mg2.state_dict(), "mg3": self.mg3.state_dict(), "quantizer1": self.quantizer1.state_dict(), "quantizer2": self.quantizer2.state_dict(), "quantizer3": self.quantizer3.state_dict(), "latent_mask_embedding": self.latent_mask_embedding.data, "optimizer": self.optimizer.state_dict(), "scheduler": self.scheduler.state_dict(), }, path)