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

如何用torchdiffeq的odeint实现同心圆环二分类?解决维度报错

神经ODE实现同心圆环二分类的错误修复与实现方案

错误根源分析

  • ODE函数输出维度不匹配:你的ODEFunc最后一层输出维度为1,但神经ODE要求导数的维度必须和输入状态y的维度(此处为2)完全一致——因为导数是对输入的每个维度计算变化率,维度不匹配会直接触发形状错误。
  • odeint输出维度未处理:odeint返回的张量形状为[时间步数, 批量大小, 特征维度](此处是[2, 1024, 2]),直接传入后续线性层会导致维度不兼容,必须提取最后一个时间步的结果作为ODE演化后的最终状态。

修正后的完整代码

import torch
import torch.nn as nn
from torchdiffeq import odeint

class ODEFunc(nn.Module):
    def __init__(self):
        super(ODEFunc, self).__init__()
        hdim = 32
        # 最后一层输出维度改为2,与输入状态y的维度保持一致
        self.net = nn.Sequential(
            nn.Linear(2, hdim),
            nn.Tanh(),
            nn.Linear(hdim, hdim),
            nn.Tanh(),
            nn.Linear(hdim, 2)
        )

    def forward(self, t, y):
        return self.net(y)

class Model(nn.Module):
    def __init__(self, odefunc, device="cpu"):
        super(Model, self).__init__()
        self.odefunc = odefunc
        # 输入维度为2(ODE最后一步的输出维度),输出1用于二分类判断
        self.linear_layer = nn.Linear(2, 1)
        self.device = device

    def forward(self, y):
        t_span = torch.linspace(0., 1., 2).to(self.device)
        # 调用odeint完成状态演化
        pred_y = odeint(self.odefunc, y, t_span)
        # 提取最后一个时间步的状态,形状从[2, 1024, 2]转为[1024, 2]
        final_state = pred_y[-1]
        # 映射到二分类输出区间
        yhat = self.linear_layer(final_state)
        # 转为0-1概率值,适配二分类损失计算
        yhat = torch.sigmoid(yhat)
        return yhat

# 测试与训练示例
if __name__ == "__main__":
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    # 生成同心圆环模拟数据
    def generate_concentric_data(n_samples=1024):
        # 第一类:半径1附近的样本
        r1 = torch.randn(n_samples//2) * 0.1 + 1.0
        theta1 = torch.rand(n_samples//2) * 2 * torch.pi
        x1 = r1 * torch.cos(theta1)
        y1 = r1 * torch.sin(theta1)
        # 第二类:半径3附近的样本
        r2 = torch.randn(n_samples//2) * 0.1 + 3.0
        theta2 = torch.rand(n_samples//2) * 2 * torch.pi
        x2 = r2 * torch.cos(theta2)
        y2 = r2 * torch.sin(theta2)
        # 合并数据与标签
        X = torch.cat([torch.stack([x1,y1], dim=1), torch.stack([x2,y2], dim=1)], dim=0)
        y = torch.cat([torch.zeros(n_samples//2), torch.ones(n_samples//2)], dim=0).unsqueeze(1)
        return X.to(device), y.to(device)
    
    # 初始化训练组件
    odefunc = ODEFunc().to(device)
    model = Model(odefunc, device=device).to(device)
    criterion = nn.BCELoss()
    optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)
    
    # 训练循环
    X_train, y_train = generate_concentric_data()
    epochs = 100
    for epoch in range(epochs):
        optimizer.zero_grad()
        y_pred = model(X_train)
        loss = criterion(y_pred, y_train)
        loss.backward()
        optimizer.step()
        if (epoch+1) % 10 == 0:
            acc = ((y_pred > 0.5) == y_train).float().mean().item()
            print(f"Epoch {epoch+1}, Loss: {loss.item():.4f}, Accuracy: {acc:.4f}")

核心修复说明

  • 对齐ODE输出维度:将nn.Linear(hdim, 1)改为nn.Linear(hdim, 2),确保导数维度与输入状态维度一致,这是解决形状错误的核心。
  • 处理ODE演化结果:通过pred_y[-1]提取最后时间步的状态,调整张量形状以适配后续分类层的输入要求。
  • 适配二分类逻辑:添加sigmoid函数将输出转换为0-1概率值,配合BCELoss完成二分类训练。

内容的提问来源于stack exchange,提问作者Skywalker

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.25 01:20:35