PyTorch实现2D Neural Cellular Automaton训练无进展问题排查
问题根源与修复方案
核心问题分析
- 不可导操作截断梯度:
torch.argmax是离散且不可导的操作,直接切断了模型输出到损失的梯度传递链,导致模型权重无法更新。 - 计算图完全断开:网格更新过程中混用numpy数组与PyTorch张量,将模型的张量输出转成numpy值赋值,彻底脱离PyTorch计算图,梯度无法回传到模型参数。
- 训练目标逻辑错误:当前损失计算的是更新后网格与当前迭代的输入网格的差异,模型只需保持细胞不变就能维持固定损失,自然不会产生学习行为。
具体修复步骤
1. 替换不可导的离散输出
训练阶段保留模型的连续logit输出,不做argmax离散化;仅在评估/推理阶段使用阈值判断得到最终的0/1网格,确保梯度链完整。
2. 全张量化网格更新流程
用PyTorch张量操作完成邻域提取与网格更新,避免numpy与张量的来回转换,维持计算图的连续性。
3. 修正训练目标
让模型学习从初始网格演化到目标原始图像,而非与当前输入网格对比,给模型明确的学习方向。
修改后的完整代码
import numpy as np import torch from torch import nn import torch.optim as optim class NeuralNetwork(nn.Module): def __init__(self): super().__init__() self.linear_relu_stack = nn.Sequential( nn.Linear(5, 32), nn.ReLU(), nn.Linear(32, 1) # 单通道输出对应二分类logit ) def forward(self, x): # 输入形状为(B, 5),直接传入全连接层 return self.linear_relu_stack(x) # 全张量化的网格更新函数 def update_grid(grid, model, training=True): dsize = grid.shape[0] # 对网格做padding,简化边界邻域提取 padded_grid = torch.nn.functional.pad(grid, (1,1,1,1), mode='replicate') # 提取每个细胞及其上下左右邻域,堆叠为(dsize, dsize, 5) neighbors = torch.stack([ padded_grid[1:-1, 1:-1], # 当前细胞 padded_grid[2:, 1:-1], # 下方细胞 padded_grid[:-2, 1:-1], # 上方细胞 padded_grid[1:-1, 2:], # 右侧细胞 padded_grid[1:-1, :-2], # 左侧细胞 ], dim=-1) # 展平为(dsize*dsize, 5)后输入模型 logits = model(neighbors.reshape(-1, 5)).squeeze(-1) if training: # 训练阶段返回sigmoid处理后的连续值,用于计算损失 return torch.sigmoid(logits).reshape(dsize, dsize) else: # 评估阶段离散化为0/1 return (torch.sigmoid(logits) > 0.5).float().reshape(dsize, dsize) if __name__ == '__main__': dsize = 8 np.random.seed(42) # 定义目标网格(需要学习还原的原始图像) target_domain = torch.tensor(np.random.randint(2, size=(dsize, dsize)), dtype=torch.float32) # 定义NCA的初始状态网格 init_domain = torch.randint(0, 2, size=(dsize, dsize), dtype=torch.float32) model = NeuralNetwork() optimizer = optim.Adam(model.parameters(), lr=0.001) # Adam优化器比SGD更适合这类任务 loss_fn = nn.BCELoss() # 二分类交叉熵损失,匹配0/1目标 # 训练循环 num_train_iter = 1000 current_domain = init_domain.clone() for step in range(num_train_iter): optimizer.zero_grad() # 训练模式下更新网格,返回连续值 updated_domain = update_grid(current_domain, model, training=True) # 计算与目标网格的损失 loss = loss_fn(updated_domain, target_domain) loss.backward() optimizer.step() # 每100步打印损失 if (step + 1) % 100 == 0: print(f'Iteration {step + 1}, Loss: {loss.item():.6f}') # 截断计算图,更新当前网格状态 current_domain = updated_domain.detach() # 评估阶段:从初始状态开始演化 model.eval() eval_domain = init_domain.clone() num_eval_iter = 10 print("\nEvaluation evolution steps:") for step in range(num_eval_iter): eval_domain = update_grid(eval_domain, model, training=False) print(f"\nStep {step+1}:") print(eval_domain.numpy().astype(int))
额外说明
- 模型改为单通道输出,配合sigmoid实现二分类,比原2类输出更简洁高效。
- 使用
BCELoss替代L1损失,更符合二分类任务的损失特性。 - 用padding+张量堆叠提取邻域,比循环操作更高效且能完整保留计算图。
- 训练时对更新后的网格执行
detach,避免计算图无限累积导致内存溢出。
内容的提问来源于stack exchange,提问作者Foundnation
相关产品推荐
相关产品推荐

