feat: 修改损失值

This commit is contained in:
unanmed 2026-07-18 14:57:46 +08:00
parent 27c42bb690
commit f5f7193df8
3 changed files with 224 additions and 27 deletions

View File

@ -0,0 +1,75 @@
# 基于矩形掩码修补的生成策略设计
## 背景
当前生成策略是从训练集采样地图,经 VQ-VAE 编码得到码本 z再用**随机采样的 z**(而非真实 z作为 MaskGIT 的条件进行三阶段级联生成。在当前数据规模下,模型难以从纯随机条件中学习到合理的生成分布,效果不佳。
## 新策略目标
降低任务难度:保持从训练集采样,但不再用随机 z而是**在真实地图上施加随机矩形掩码**,让模型作为"修补inpainting"任务来完成。这样模型既有真实地图的结构先验(通过 z 条件),又有明确的局部填空目标。
## 核心改动
### 1. 掩码方式
**仅使用矩形分块掩码**`dataset.py` `std_mask` 的第二种模式),不采用散点随机掩码:
- 在 13×13 网格上,反复放置随机矩形块(高宽 2~6直到掩码格数达到目标比例
- 掩码比例可调,默认沿用 Beta(2,2) 采样(范围 [5%, 100%]),同时支持指定固定比例用于对照实验
### 2. 条件输入 z
不再随机采样 z改为使用**被掩码地图经 VQ-VAE 编码后得到的真实 z**。这意味着:
- 编码器输入可能是部分掩码后的地图(掩码位置填 MASK_ID=7
- VQ-VAE 从可见部分推断整体结构,得到的 z 携带地图的全局语义信息
- MaskGIT 以该 z 为条件,在掩码区域补全被遮盖的图块
### 3. 生成流程
仍然采用三阶段级联,但每个阶段统一使用同一份掩码布局:
```
输入: 完整地图 gt [13×13]
生成矩形掩码 mask [13×13 bool] (比率可调)
inp = gt.copy(); inp[mask] = MASK_TOKEN
VQ-VAE 编码 inp → z1, z2, z3
Stage1: mg1(inp, z1) → 补全墙壁 (mask ∩ target=1)
Stage2: mg2(inp1+stage1_out, z2) → 补全门/怪物/入口
Stage3: mg3(inp2+stage2_out, z3) → 补全资源
输出: 最终地图 + 逐阶段可视化对比
```
### 4. 与现有 `maskgit_sample` 的适配
现有 `maskgit_sample` 已支持 `keep_fixed` 参数,天然适配修补场景:
- `keep_fixed=True`默认输入中已有的非掩码图块在生成过程中锁定不变MaskGIT 仅填充掩码位
- 需要新增:将矩形掩码 `mask` 传入函数,使模型只尝试修改掩码区域
## 可调参数
| 参数 | 说明 | 默认值 |
|------|------|--------|
| `mask_ratio` | 掩码比例,可指定固定值或 `None` 使用 Beta(2,2) 采样 | `None`Beta 采样) |
| `block_h_range` | 矩形高度范围 | `(2, 7)` |
| `block_w_range` | 矩形宽度范围 | `(2, 7)` |
| `keep_fixed` | 三阶段是否锁定已有结构 | `(True, True, True)` |
## 实现要点
1. 将 `std_mask` 的矩形掩码生成逻辑提取为独立函数 `rect_mask(ratio, h_range, w_range) -> np.ndarray`
2. 新增 `inpaint_generate(map, mask, models, device)` 函数,包装完整修补流程
3. 验证时对同一张地图使用不同 `mask_ratio`(如 0%30%50%70%)生成多组结果,横向对比
4. 可视化输出时同时展示:原始地图、掩码地图、各阶段生成结果、最终结果
## 预期效果
- 任务从"开放式生成"变为"约束式修补",难度显著降低
- 模型可利用可见区域的上下文信息(如墙面走向、房间布局)推断掩码区域内容
- 不同掩码比例可衡量模型对空间结构的理解深度

View File

@ -10,10 +10,12 @@ def load_data(path: str):
data = json.load(f) data = json.load(f)
data_list = [] data_list = []
for value in data["data"].values(): map_names = []
for key, value in data["data"].items():
data_list.append(value) data_list.append(value)
map_names.append(key)
return data_list return data_list, map_names
def compute_symmetry(target_np: np.ndarray) -> tuple: def compute_symmetry(target_np: np.ndarray) -> tuple:
"""从 numpy 地图矩阵中直接计算三种对称性O(H*W)""" """从 numpy 地图矩阵中直接计算三种对称性O(H*W)"""
@ -39,7 +41,7 @@ class GinkaSeperatedDataset(Dataset):
subset_weights: tuple = (0.5, 0.3, 0.2), subset_weights: tuple = (0.5, 0.3, 0.2),
density_stats: dict | None = None density_stats: dict | None = None
): ):
self.data = load_data(data_path) self.data, self.map_names = load_data(data_path)
total = sum(subset_weights) total = sum(subset_weights)
self.subset_cumw = [sum(subset_weights[:i+1]) / total for i in range(len(subset_weights))] self.subset_cumw = [sum(subset_weights[:i+1]) / total for i in range(len(subset_weights))]
@ -130,6 +132,7 @@ class GinkaSeperatedDataset(Dataset):
"distance_field": torch.LongTensor(compute_distance_field(enc1)) "distance_field": torch.LongTensor(compute_distance_field(enc1))
} }
sample['sample_idx'] = idx sample['sample_idx'] = idx
sample['map_name'] = self.map_names[idx]
return sample return sample
def degrade_tile(self, m: np.ndarray, tiles: list) -> np.ndarray: def degrade_tile(self, m: np.ndarray, tiles: list) -> np.ndarray:
@ -240,4 +243,6 @@ class GinkaSeperatedDataset(Dataset):
else: else:
out = self.apply_subset3(map_np) out = self.apply_subset3(map_np)
return self.pack_sample(item, map_np, out) sample = self.pack_sample(item, map_np, out)
sample["map_name"] = self.map_names[idx]
return sample

View File

@ -39,8 +39,8 @@ from shared.distance import DIST_VOCAB, compute_distance_field_tensor
# 共用 VQ-VAE 超参 # 共用 VQ-VAE 超参
# 三组编码器vq1/vq2/vq3共享相同超参分别对三阶段地图上下文独立编码 # 三组编码器vq1/vq2/vq3共享相同超参分别对三阶段地图上下文独立编码
VQ_L = 16 # 码字序列长度(每个编码器输出 L 个码字,量化后合并为 L*3 VQ_L = 8 # 码字序列长度(每个编码器输出 L 个码字,量化后合并为 L*3
VQ_K = 32 # codebook 大小(离散码本条目数) VQ_K = 16 # codebook 大小(离散码本条目数)
VQ_D_Z = 64 # 码字维度 VQ_D_Z = 64 # 码字维度
VQ_BETA = 1.0 # commit loss 权重(防止编码器输出漂离 codebook VQ_BETA = 1.0 # commit loss 权重(防止编码器输出漂离 codebook
VQ_GAMMA = 0.0 # entropy loss 权重(当前未启用) VQ_GAMMA = 0.0 # entropy loss 权重(当前未启用)
@ -52,9 +52,10 @@ VQ_NHEAD = 4 # VQ-VAE 多头注意力头数
# 距离场编码器超参 # 距离场编码器超参
L_DIST = 4 # 距离场码字序列长度 L_DIST = 4 # 距离场码字序列长度
K_DIST = 16 # 距离场 codebook 大小 K_DIST = 16 # 距离场 codebook 大小
DIST_D_MODEL = 128 # 距离场编码器模型维度 DIST_D_MODEL = 256 # 距离场编码器模型维度
DIST_LAYERS = 3 # 距离场编码器 Transformer 层数 DIST_LAYERS = 3 # 距离场编码器 Transformer 层数
DIST_DIM_FF = 512 # 距离场编码器 FF 维度 DIST_DIM_FF = 1024 # 距离场编码器 FF 维度
DIST_NHEAD = 8 # 距离场编码器注意力头数
VQ_BETA_DIST = 0.5 # 距离场 commit loss 权重 VQ_BETA_DIST = 0.5 # 距离场 commit loss 权重
# 第一阶段 MaskGIT 超参 # 第一阶段 MaskGIT 超参
@ -104,6 +105,18 @@ RESOURCE_DENSITY_IDX = 4
MG_Z_DROPOUT = 0.1 # z 隐变量 Dropout 概率 MG_Z_DROPOUT = 0.1 # z 隐变量 Dropout 概率
MG_STRUCT_DROPOUT = 0.1 # 结构参量 Dropout 概率 MG_STRUCT_DROPOUT = 0.1 # 结构参量 Dropout 概率
# 邻接损失权重(三阶段)
LAMBDA_ADJ1 = 0.6
LAMBDA_ADJ2 = 0.3
LAMBDA_ADJ3 = 0.1
# Patch 损失权重(三阶段)及核参数
LAMBDA_PATCH1 = 0.5
LAMBDA_PATCH2 = 0.5
LAMBDA_PATCH3 = 0.5
PATCH_KERNEL_SIZE = 5
PATCH_SIGMA = 1.2
# 损失参数 # 损失参数
VQ_BETA = 0.5 # 承诺损失权重 VQ_BETA = 0.5 # 承诺损失权重
@ -176,7 +189,7 @@ def build_model(device: torch.device):
# 距离场编码器与量化器:将 L1 距离场编码为离散 latent z_dist # 距离场编码器与量化器:将 L1 距离场编码为离散 latent z_dist
dist_encoder = DistFieldEncoder( dist_encoder = DistFieldEncoder(
vocab=DIST_VOCAB, L=L_DIST, d_z=VQ_D_Z, d_model=DIST_D_MODEL, vocab=DIST_VOCAB, L=L_DIST, d_z=VQ_D_Z, d_model=DIST_D_MODEL,
nhead=VQ_NHEAD, num_layers=DIST_LAYERS, dim_ff=DIST_DIM_FF, nhead=DIST_NHEAD, num_layers=DIST_LAYERS, dim_ff=DIST_DIM_FF,
map_h=MAP_H, map_w=MAP_W map_h=MAP_H, map_w=MAP_W
).to(device) ).to(device)
dist_quantizer = VectorQuantizer(K=K_DIST, d_z=VQ_D_Z).to(device) dist_quantizer = VectorQuantizer(K=K_DIST, d_z=VQ_D_Z).to(device)
@ -202,6 +215,72 @@ def cross_entropy_loss(logits, target):
# logits: [B, L, C],需转为 [B, C, L] 以匹配 cross_entropy 期望格式 # logits: [B, L, C],需转为 [B, C, L] 以匹配 cross_entropy 期望格式
return F.cross_entropy(logits.permute(0, 2, 1), target) return F.cross_entropy(logits.permute(0, 2, 1), target)
def adjacency_loss(logits, target):
# 邻接损失:约束相邻两格同时为空地的概率
# logits: [B, S, C] — MaskGIT 解码器输出
# target: [B, S] — 目标类别 ID不含 MASK 标记
B, S, C = logits.shape
H = 13
W = 13
probs = F.softmax(logits, dim=-1) # [B, S, C]
p_floor = probs[:, :, 0].view(B, H, W) # [B, H, W] — 地板概率
t = target.view(B, H, W)
t_floor = (t == 0).float() # 地板标注为 1其余为 0
# 水平边:左格 × 右格
joint_h = p_floor[:, :, :-1] * p_floor[:, :, 1:] # [B, H, W-1]
target_h = t_floor[:, :, :-1] * t_floor[:, :, 1:] # [B, H, W-1]
# 垂直边:上格 × 下格
joint_v = p_floor[:, :-1, :] * p_floor[:, 1:, :] # [B, H-1, W]
target_v = t_floor[:, :-1, :] * t_floor[:, 1:, :] # [B, H-1, W]
loss_h = F.binary_cross_entropy(joint_h, target_h, reduction='mean')
loss_v = F.binary_cross_entropy(joint_v, target_v, reduction='mean')
return (loss_h + loss_v) / 2.0
def gaussian_kernel(kernel_size, sigma, device):
# 生成归一化二维高斯卷积核 [1, 1, K, K]
k = kernel_size
center = (k - 1) / 2.0
xs = torch.arange(k, dtype=torch.float32, device=device) - center
gx = torch.exp(-xs ** 2 / (2.0 * sigma ** 2))
gy = torch.exp(-xs ** 2 / (2.0 * sigma ** 2))
g2d = gx[:, None] * gy[None, :] # [K, K]
g2d = g2d / g2d.sum() # 归一化
return g2d.view(1, 1, k, k)
def patch_loss(logits, target, kernel_size=5, sigma=1.2):
# Patch 损失:高斯核加权的邻域 CE 平滑损失
# logits: [B, S, C]
# target: [B, S]
B, S, C = logits.shape
H = 13
W = 13
# 逐格 CE不做 reduction
ce = F.cross_entropy(
logits.reshape(-1, C), target.reshape(-1), reduction='none'
).view(B, H, W) # [B, H, W]
# 高斯核
kernel = gaussian_kernel(kernel_size, sigma, logits.device) # [1, 1, K, K]
# replicate 填充后用 unfold 提取邻域
pad = kernel_size // 2
ce_padded = F.pad(
ce.view(B, 1, H, W), (pad, pad, pad, pad), mode='replicate'
)
# patches: [B, K*K, H*W]
patches = F.unfold(ce_padded, kernel_size=(kernel_size, kernel_size))
patches = patches.view(B, kernel_size * kernel_size, H, W) # [B, K*K, H, W]
# 加权求和
k_flat = kernel.view(1, kernel_size * kernel_size, 1, 1)
smoothed = (patches * k_flat).sum(dim=1) # [B, H, W]
return smoothed.mean()
def apply_z_dropout( def apply_z_dropout(
z_q: torch.Tensor, z_q: torch.Tensor,
mask_embedding: nn.Parameter, mask_embedding: nn.Parameter,
@ -541,7 +620,7 @@ def build_dataset_sample_case(
} }
def sample_case_label(case: dict) -> str: def sample_case_label(case: dict) -> str:
return f"train#{case['sample_idx']}" return case["sample"]["map_name"]
# 验证可视化 part13×3 网格行1=编码器输入行2=掩码输入行3=三阶段预测(合并) # 验证可视化 part13×3 网格行1=编码器输入行2=掩码输入行3=三阶段预测(合并)
def visualize_part1(batch, logits1, logits2, logits3, tile_dict): def visualize_part1(batch, logits1, logits2, logits3, tile_dict):
@ -573,7 +652,7 @@ def visualize_part1(batch, logits1, logits2, logits3, tile_dict):
result3[inp3_np == MASK_TOKEN] = pred3[inp3_np == MASK_TOKEN] result3[inp3_np == MASK_TOKEN] = pred3[inp3_np == MASK_TOKEN]
rows = [ rows = [
[to_img(enc1_np), to_img(enc2_np), to_img(enc3_np)], [annotate(to_img(enc1_np), batch["map_name"][0]), to_img(enc2_np), to_img(enc3_np)],
[to_img(inp1_np), to_img(inp2_np), to_img(inp3_np)], [to_img(inp1_np), to_img(inp2_np), to_img(inp3_np)],
[to_img(result1), to_img(result2), to_img(result3)], [to_img(result1), to_img(result2), to_img(result3)],
] ]
@ -614,7 +693,7 @@ def visualize_part2(batch, z_q, z_dist, models, device, tile_dict):
target_density_cpu = batch["target_density"][0] target_density_cpu = batch["target_density"][0]
rows = [ rows = [
[to_img(enc1_np), to_img(enc2_np), to_img(enc3_np)], [annotate(to_img(enc1_np), batch["map_name"][0]), to_img(enc2_np), to_img(enc3_np)],
[ [
annotate(to_img(inp1_np), kf_label), annotate(to_img(inp1_np), kf_label),
annotate_labels(to_img(auto_pred1_np), struct_cpu, target_density_cpu), annotate_labels(to_img(auto_pred1_np), struct_cpu, target_density_cpu),
@ -712,6 +791,12 @@ def validate(
loss2_total = torch.Tensor([0]).to(device) loss2_total = torch.Tensor([0]).to(device)
loss3_total = torch.Tensor([0]).to(device) loss3_total = torch.Tensor([0]).to(device)
commit_total = torch.Tensor([0]).to(device) commit_total = torch.Tensor([0]).to(device)
adj1_total = torch.Tensor([0]).to(device)
adj2_total = torch.Tensor([0]).to(device)
adj3_total = torch.Tensor([0]).to(device)
patch1_total = torch.Tensor([0]).to(device)
patch2_total = torch.Tensor([0]).to(device)
patch3_total = torch.Tensor([0]).to(device)
code_hits_total = torch.zeros(3, quantizer1.K, device=device) code_hits_total = torch.zeros(3, quantizer1.K, device=device)
density_metrics = { density_metrics = {
@ -771,6 +856,12 @@ def validate(
loss2_total += cross_entropy_loss(logits2, target2) loss2_total += cross_entropy_loss(logits2, target2)
loss3_total += cross_entropy_loss(logits3, target3) loss3_total += cross_entropy_loss(logits3, target3)
commit_total += commit_loss commit_total += commit_loss
adj1_total += adjacency_loss(logits1, target1)
adj2_total += adjacency_loss(logits2, target2)
adj3_total += adjacency_loss(logits3, target3)
patch1_total += patch_loss(logits1, target1, PATCH_KERNEL_SIZE, PATCH_SIGMA)
patch2_total += patch_loss(logits2, target2, PATCH_KERNEL_SIZE, PATCH_SIGMA)
patch3_total += patch_loss(logits3, target3, PATCH_KERNEL_SIZE, PATCH_SIGMA)
code_hits_total += code_hits code_hits_total += code_hits
# 计算各目标对象的真实密度误差与过量生成密度 # 计算各目标对象的真实密度误差与过量生成密度
@ -818,7 +909,7 @@ def validate(
for m in [vq1, vq2, vq3, mg1, mg2, mg3, dist_encoder]: for m in [vq1, vq2, vq3, mg1, mg2, mg3, dist_encoder]:
m.train() m.train()
return loss1_total, loss2_total, loss3_total, commit_total, code_hits_total return loss1_total, loss2_total, loss3_total, adj1_total, adj2_total, adj3_total, patch1_total, patch2_total, patch3_total, commit_total, code_hits_total
def train(device: torch.device): def train(device: torch.device):
args = parse_arguments() args = parse_arguments()
@ -902,6 +993,12 @@ def train(device: torch.device):
loss2_total = torch.Tensor([0]).to(device) loss2_total = torch.Tensor([0]).to(device)
loss3_total = torch.Tensor([0]).to(device) loss3_total = torch.Tensor([0]).to(device)
commit_total = torch.Tensor([0]).to(device) commit_total = torch.Tensor([0]).to(device)
adj1_total = torch.Tensor([0]).to(device)
adj2_total = torch.Tensor([0]).to(device)
adj3_total = torch.Tensor([0]).to(device)
patch1_total = torch.Tensor([0]).to(device)
patch2_total = torch.Tensor([0]).to(device)
patch3_total = torch.Tensor([0]).to(device)
code_hits_total = torch.zeros(3, quantizer1.K, device=device) code_hits_total = torch.zeros(3, quantizer1.K, device=device)
for batch in tqdm(dataloader, leave=False, desc="Epoch Progress", disable=disable_tqdm): for batch in tqdm(dataloader, leave=False, desc="Epoch Progress", disable=disable_tqdm):
@ -966,15 +1063,25 @@ def train(device: torch.device):
logits2 = mg2(inp2, z_q2, z_dist, struct, remain2) logits2 = mg2(inp2, z_q2, z_dist, struct, remain2)
logits3 = mg3(inp3, z_q3, z_dist, struct, remain3) logits3 = mg3(inp3, z_q3, z_dist, struct, remain3)
# 三阶段 Cross Entropy + VQ commit loss 加权求和 # 三阶段 Cross Entropy + 邻接损失 + Patch 损失 + VQ commit loss 加权求和
loss1 = cross_entropy_loss(logits1, target1) loss1 = cross_entropy_loss(logits1, target1)
loss2 = cross_entropy_loss(logits2, target2) loss2 = cross_entropy_loss(logits2, target2)
loss3 = cross_entropy_loss(logits3, target3) loss3 = cross_entropy_loss(logits3, target3)
adj1 = adjacency_loss(logits1, target1)
adj2 = adjacency_loss(logits2, target2)
adj3 = adjacency_loss(logits3, target3)
patch1 = patch_loss(logits1, target1, PATCH_KERNEL_SIZE, PATCH_SIGMA)
patch2 = patch_loss(logits2, target2, PATCH_KERNEL_SIZE, PATCH_SIGMA)
patch3 = patch_loss(logits3, target3, PATCH_KERNEL_SIZE, PATCH_SIGMA)
loss1_weighted = STAGE1_CE_WEIGHT * loss1 loss1_weighted = STAGE1_CE_WEIGHT * loss1
loss2_weighted = STAGE2_CE_WEIGHT * loss2 loss2_weighted = STAGE2_CE_WEIGHT * loss2
loss3_weighted = STAGE3_CE_WEIGHT * loss3 loss3_weighted = STAGE3_CE_WEIGHT * loss3
adj_weighted = LAMBDA_ADJ1 * adj1 + LAMBDA_ADJ2 * adj2 + LAMBDA_ADJ3 * adj3
patch_weighted = LAMBDA_PATCH1 * patch1 + LAMBDA_PATCH2 * patch2 + LAMBDA_PATCH3 * patch3
commit_weighted = VQ_BETA * commit_loss + VQ_BETA_DIST * commit_loss_dist commit_weighted = VQ_BETA * commit_loss + VQ_BETA_DIST * commit_loss_dist
loss = loss1_weighted + loss2_weighted + loss3_weighted + commit_weighted loss = loss1_weighted + loss2_weighted + loss3_weighted + adj_weighted + patch_weighted + commit_weighted
loss.backward() loss.backward()
optimizer.step() optimizer.step()
@ -985,6 +1092,12 @@ def train(device: torch.device):
loss2_total += loss2.detach() loss2_total += loss2.detach()
loss3_total += loss3.detach() loss3_total += loss3.detach()
commit_total += commit_loss.detach() commit_total += commit_loss.detach()
adj1_total += adj1.detach()
adj2_total += adj2.detach()
adj3_total += adj3.detach()
patch1_total += patch1.detach()
patch2_total += patch2.detach()
patch3_total += patch3.detach()
code_hits_total += code_hits.detach() code_hits_total += code_hits.detach()
# 每个 epoch 结束后更新学习率 # 每个 epoch 结束后更新学习率
@ -994,11 +1107,12 @@ def train(device: torch.device):
train_perplexity, train_usage_rate, train_active_codes = summarize_codebook_hits(code_hits_total) train_perplexity, train_usage_rate, train_active_codes = summarize_codebook_hits(code_hits_total)
tqdm.write( tqdm.write(
f"[{datetime.now().strftime('%Y-%m-%d %H:%M:%S')}] " f"[{datetime.now().strftime('%Y-%m-%d %H:%M:%S')}] "
f"E: {epoch + 1} | Loss: {loss_total.item() / data_length:.6f} | " f"E: {epoch + 1} | "
f"L1: {loss1_total.item() / data_length:.6f} | " f"Loss: {loss_total.item() / data_length:.4f} | "
f"L2: {loss2_total.item() / data_length:.6f} | " f"CE: {loss1_total.item() / data_length:.4f}, {loss2_total.item() / data_length:.4f}, {loss3_total.item() / data_length:.4f} | "
f"L3: {loss3_total.item() / data_length:.6f} | " f"ADJ: {(LAMBDA_ADJ1 * adj1_total.item() + LAMBDA_ADJ2 * adj2_total.item() + LAMBDA_ADJ3 * adj3_total.item()) / data_length:.4f}, {adj1_total.item() / data_length:.4f}, {adj2_total.item() / data_length:.4f}, {adj3_total.item() / data_length:.4f} | "
f"VQ: {commit_total.item() / data_length:.6f} | " f"PAT: {(LAMBDA_PATCH1 * patch1_total.item() + LAMBDA_PATCH2 * patch2_total.item() + LAMBDA_PATCH3 * patch3_total.item()) / data_length:.4f}, {patch1_total.item() / data_length:.4f}, {patch2_total.item() / data_length:.4f}, {patch3_total.item() / data_length:.4f} | "
f"VQ: {commit_total.item() / data_length:.4f} | "
f"PPL: {train_perplexity:.4f} | " f"PPL: {train_perplexity:.4f} | "
f"Usage: {train_usage_rate:.4f} ({train_active_codes}/{code_hits_total.numel()}) | " f"Usage: {train_usage_rate:.4f} ({train_active_codes}/{code_hits_total.numel()}) | "
f"LR: {scheduler.get_last_lr()[0]:.6f}" f"LR: {scheduler.get_last_lr()[0]:.6f}"
@ -1009,22 +1123,25 @@ def train(device: torch.device):
losses = validate( losses = validate(
dataloader_val, models, dist_models, device, tile_dict, dataset, epoch + 1 dataloader_val, models, dist_models, device, tile_dict, dataset, epoch + 1
) )
loss1_total, loss2_total, loss3_total, commit_total, code_hits_total = losses loss1_total, loss2_total, loss3_total, adj1_total, adj2_total, adj3_total, patch1_total, patch2_total, patch3_total, commit_total, code_hits_total = losses
loss1_weighted = STAGE1_CE_WEIGHT * loss1_total loss1_weighted = STAGE1_CE_WEIGHT * loss1_total
loss2_weighted = STAGE2_CE_WEIGHT * loss2_total loss2_weighted = STAGE2_CE_WEIGHT * loss2_total
loss3_weighted = STAGE3_CE_WEIGHT * loss3_total loss3_weighted = STAGE3_CE_WEIGHT * loss3_total
adj_weighted = LAMBDA_ADJ1 * adj1_total + LAMBDA_ADJ2 * adj2_total + LAMBDA_ADJ3 * adj3_total
patch_weighted = LAMBDA_PATCH1 * patch1_total + LAMBDA_PATCH2 * patch2_total + LAMBDA_PATCH3 * patch3_total
commit_weighted = VQ_BETA * commit_total commit_weighted = VQ_BETA * commit_total
loss_total = loss1_weighted + loss2_weighted + loss3_weighted + commit_weighted loss_total = loss1_weighted + loss2_weighted + loss3_weighted + adj_weighted + patch_weighted + commit_weighted
data_length = len(dataloader_val) data_length = len(dataloader_val)
val_perplexity, val_usage_rate, val_active_codes = summarize_codebook_hits(code_hits_total) val_perplexity, val_usage_rate, val_active_codes = summarize_codebook_hits(code_hits_total)
tqdm.write( tqdm.write(
f"[Validate {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}] " f"[Validate {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}] "
f"E: {epoch + 1} | Loss: {loss_total.item() / data_length:.6f} | " f"E: {epoch + 1} | "
f"L1: {loss1_total.item() / data_length:.6f} | " f"Loss: {loss_total.item() / data_length:.4f} | "
f"L2: {loss2_total.item() / data_length:.6f} | " f"CE: {loss1_total.item() / data_length:.4f}, {loss2_total.item() / data_length:.4f}, {loss3_total.item() / data_length:.4f} | "
f"L3: {loss3_total.item() / data_length:.6f} | " f"ADJ: {(LAMBDA_ADJ1 * adj1_total.item() + LAMBDA_ADJ2 * adj2_total.item() + LAMBDA_ADJ3 * adj3_total.item()) / data_length:.4f}, {adj1_total.item() / data_length:.4f}, {adj2_total.item() / data_length:.4f}, {adj3_total.item() / data_length:.4f} | "
f"VQ: {commit_total.item() / data_length:.6f} | " f"PAT: {(LAMBDA_PATCH1 * patch1_total.item() + LAMBDA_PATCH2 * patch2_total.item() + LAMBDA_PATCH3 * patch3_total.item()) / data_length:.4f}, {patch1_total.item() / data_length:.4f}, {patch2_total.item() / data_length:.4f}, {patch3_total.item() / data_length:.4f} | "
f"VQ: {commit_total.item() / data_length:.4f} | "
f"PPL: {val_perplexity:.4f} | " f"PPL: {val_perplexity:.4f} | "
f"Usage: {val_usage_rate:.4f} ({val_active_codes}/{code_hits_total.numel()}) | " f"Usage: {val_usage_rate:.4f} ({val_active_codes}/{code_hits_total.numel()}) | "
) )