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

PyTorch含连续与文本特征的神经网络调试求助

问题分析与修复方案

核心错误点

  • 动态修改全连接层fc1:在forward里重新定义self.fc1会导致模型参数无法被正确优化,且无法适配批量输入
  • 错误的flatten方式:直接x.flatten()会把batch维度也压扁,丢失批量信息
  • 序列特征处理不当:x6、x7是长度为max_seq_length的序列,直接flatten会让维度随序列长度变化,应该用池化(均值/最大值)压缩成固定维度
  • 输入维度计算错误:没有提前计算拼接后的总维度,导致fc1初始化错误

修复后的数据集代码(小优化)

提前解析字符串列表,避免每次__getitem__调用eval,提升效率:

class MyDataset(Dataset):
    def __init__(self, csv_file, max_seq_length=8) -> None:
        super().__init__()
        self.data = pd.read_csv(csv_file)
        self.X = self.data.drop(columns=["label"])
        self.y = self.data["label"]
        self.max_seq_length = max_seq_length
        
        # 提前解析字符串格式的列表
        for col in ['x6', 'x7']:
            self.X[col] = self.X[col].apply(eval)

    def __len__(self):
        return len(self.X)
    
    def __getitem__(self, idx):
        features = {
            'x1': torch.tensor(self.X.iloc[idx]['x1'], dtype=torch.float),
            'x2': torch.tensor(self.X.iloc[idx]['x2'], dtype=torch.float),
            'x3': torch.tensor(self.X.iloc[idx]['x3'], dtype=torch.long),
            'x4': torch.tensor(self.X.iloc[idx]['x4'], dtype=torch.long),
            'x5': torch.tensor(self.X.iloc[idx]['x5'], dtype=torch.long),
            'x6': torch.tensor(self.X.iloc[idx]['x6'], dtype=torch.long),
            'x7': torch.tensor(self.X.iloc[idx]['x7'], dtype=torch.long)
        }

        # 补零确保序列长度一致,处理序列长度小于max_seq_length的情况
        pad_len_x6 = max(0, self.max_seq_length - len(features['x6']))
        features['x6'] = torch.nn.functional.pad(features['x6'], pad=(0, pad_len_x6), mode='constant', value=0)
        
        pad_len_x7 = max(0, self.max_seq_length - len(features['x7']))
        features['x7'] = torch.nn.functional.pad(features['x7'], pad=(0, pad_len_x7), mode='constant', value=0)

        return features, torch.tensor(self.y.iloc[idx], dtype=torch.long)

修复后的模型代码

import torch
import torch.nn as nn
import torch.nn.functional as F

class MyModel(nn.Module):
    def __init__(self, vocab_dicts, embedding_dim, max_seq_length, num_classes):
        super().__init__()
        self.embedding_dim = embedding_dim
        
        # 为每个分类特征单独定义Embedding层,避免顺序出错
        self.emb_x3 = nn.Embedding(vocab_dicts["x3"], embedding_dim)
        self.emb_x4 = nn.Embedding(vocab_dicts["x4"], embedding_dim)
        self.emb_x5 = nn.Embedding(vocab_dicts["x5"], embedding_dim)
        self.emb_x6 = nn.Embedding(vocab_dicts["x6"], embedding_dim)
        self.emb_x7 = nn.Embedding(vocab_dicts["x7"], embedding_dim)
        
        # 提前计算全连接层输入总维度:
        # x3/x4/x5各占embedding_dim维度;x6/x7池化后各占embedding_dim维度;x1/x2各占1维度
        total_input_dim = 3 * embedding_dim + 2 * embedding_dim + 2
        self.fc1 = nn.Linear(total_input_dim, 128)
        self.fc2 = nn.Linear(128, 64)
        self.fc3 = nn.Linear(64, num_classes)

    def forward(self, sample):
        # 处理单值分类特征,输出形状:(batch_size, embedding_dim)
        emb_x3 = self.emb_x3(sample["x3"])
        emb_x4 = self.emb_x4(sample["x4"])
        emb_x5 = self.emb_x5(sample["x5"])
        
        # 处理序列分类特征:先embedding再均值池化,把序列维度压缩为固定维度
        emb_x6 = self.emb_x6(sample["x6"])  # 形状:(batch_size, max_seq_length, embedding_dim)
        emb_x6 = torch.mean(emb_x6, dim=1)  # 形状:(batch_size, embedding_dim)
        
        emb_x7 = self.emb_x7(sample["x7"])  # 形状:(batch_size, max_seq_length, embedding_dim)
        emb_x7 = torch.mean(emb_x7, dim=1)  # 形状:(batch_size, embedding_dim)
        
        # 调整连续特征维度,从(batch_size,)转为(batch_size,1),方便拼接
        x1 = sample["x1"].unsqueeze(1)
        x2 = sample["x2"].unsqueeze(1)
        
        # 拼接所有特征,输出形状:(batch_size, total_input_dim)
        x = torch.cat([emb_x3, emb_x4, emb_x5, emb_x6, emb_x7, x1, x2], dim=1)
        
        # 全连接层前向传播
        x = F.relu(self.fc1(x))
        x = F.relu(self.fc2(x))
        x = self.fc3(x)
        
        return x

if __name__ == "__main__":
    vocab_dicts = {
        "x3": 4,
        "x4": 20000,
        "x5": 20000,
        "x6": 20000,
        "x7": 20000,
    }
    max_seq_length = 8
    embedding_dim = 32
    
    # 测试单样本输入(手动添加batch维度)
    single_features = {
        "x1": torch.tensor([0.0259]),
        "x2": torch.tensor([0.0322]),
        "x3": torch.tensor([1]),
        "x4": torch.tensor([16839]),
        "x5": torch.tensor([7721]),
        "x6": torch.tensor([[6331, 8116, 0, 0, 0, 0, 0, 0]]),
        "x7": torch.tensor([[262, 1930, 2205, 2255, 4022, 4225, 4266, 4274]]),
    }
    model = MyModel(vocab_dicts, embedding_dim, max_seq_length, 3)
    print(model)
    print("单样本输出形状:", model(single_features).shape)  # 预期输出(1, 3)
    
    # 测试批量输入(batch_size=2)
    batch_features = {
        "x1": torch.tensor([0.0259, 0.0123]),
        "x2": torch.tensor([0.0322, 0.0456]),
        "x3": torch.tensor([1, 2]),
        "x4": torch.tensor([16839, 12345]),
        "x5": torch.tensor([7721, 8901]),
        "x6": torch.tensor([[6331, 8116, 0, 0, 0, 0, 0, 0], [1234, 5678, 0, 0, 0, 0, 0, 0]]),
        "x7": torch.tensor([[262, 1930, 2205, 2255, 4022, 4225, 4266, 4274], [9876, 5432, 1111, 2222, 3333, 4444, 5555, 6666]]),
    }
    print("批量输出形状:", model(batch_features).shape)  # 预期输出(2, 3)

关键修复说明

  1. 固定全连接层输入维度:提前计算所有特征拼接后的总维度,避免动态修改层结构,确保模型参数能被正常优化
  2. 保留batch维度:所有特征处理都维持(batch_size, ...)的维度,单样本测试时手动添加batch维度(用unsqueeze或列表包裹)
  3. 序列特征压缩:对x6、x7这类序列特征,用均值池化(或最大值池化)把(batch_size, seq_len, emb_dim)压缩为(batch_size, emb_dim),保证维度固定
  4. 连续特征维度调整:用unsqueeze(1)把连续特征从(batch_size,)转为(batch_size,1),适配拼接要求
  5. 明确Embedding层对应关系:单独定义每个Embedding层,避免ModuleList的顺序歧义,提升代码可读性

批量训练示例

用DataLoader自动处理批量拼接:

from torch.utils.data import DataLoader

dataset = MyDataset("your_data.csv", max_seq_length=8)
dataloader = DataLoader(dataset, batch_size=32, shuffle=True)

for batch_features, batch_labels in dataloader:
    outputs = model(batch_features)
    loss = F.cross_entropy(outputs, batch_labels)
    # 后续执行反向传播、参数更新等步骤

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.08 14:45:54