You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.09 13:20:12