如何用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
相关产品推荐
相关产品推荐

