ginka-generator/ginka/model.py

194 lines
7.3 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 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 = 16
VQ_K1 = 32
VQ_D_MODEL1 = 384
VQ_NHEAD1 = 8
VQ_LAYERS1 = 6
VQ_DIM_FF1 = 1536
# Stage2 功能元素 — 中等复杂度
VQ_L2 = 8
VQ_K2 = 64
VQ_D_MODEL2 = 256
VQ_NHEAD2 = 4
VQ_LAYERS2 = 4
VQ_DIM_FF2 = 1024
# Stage3 资源分布 — 最简单,模型容量最小
VQ_L3 = 24
VQ_K3 = 24
VQ_D_MODEL3 = 256
VQ_NHEAD3 = 8
VQ_LAYERS3 = 4
VQ_DIM_FF3 = 1024
# 第一阶段 MaskGIT 超参
STAGE1_MG_DMODEL = 512
STAGE1_MG_NHEAD = 8
STAGE1_MG_NUM_LAYERS = 8
STAGE1_MG_DIM_FF = 2048
# 第二阶段 MaskGIT 超参
STAGE2_MG_DMODEL = 384
STAGE2_MG_NHEAD = 8
STAGE2_MG_NUM_LAYERS = 6
STAGE2_MG_DIM_FF = 1536
# 第三阶段 MaskGIT 超参
STAGE3_MG_DMODEL = 256
STAGE3_MG_NHEAD = 8
STAGE3_MG_NUM_LAYERS = 6
STAGE3_MG_DIM_FF = 1024
# 各阶段 VQ commit loss 权重(当前未单独使用,统一由 VQ_BETA 控制)
STAGE1_VQ_WEIGHT = 0.2
STAGE2_VQ_WEIGHT = 0.2
STAGE3_VQ_WEIGHT = 0.2
# 全局参数
NUM_CLASSES = 8 # 图块类型数
MASK_TOKEN = 7 # 掩码图块
TOTAL_K = VQ_K1 + VQ_K2 + VQ_K3 # 预计算,供日志输出使用
MAP_W = 13 # 地图宽度
MAP_H = 13 # 地图高度
LR = 2e-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)