mirror of
https://github.com/unanmed/ginka-generator.git
synced 2026-08-14 18:12:28 +08:00
194 lines
7.3 KiB
Python
194 lines
7.3 KiB
Python
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)
|