【Bug已解决】How to construct a network with two inputs in PyTorch 解决方案 【Bug已解决】How to construct a network with two inputs in PyTorch 解决方案问题描述在实际的深度学习项目中很多任务需要处理多种类型的输入数据。例如图像分类任务中输入既包括图像还包括元数据如拍摄位置、时间等推荐系统中输入包括用户特征和物品特征多模态学习中输入包括文本和图像视频理解中输入包括视频帧和音频这些场景都需要构建一个能够同时接受两个或多个输入的网络模型。然而PyTorch 的nn.Module默认的forward方法通常只接受一个输入张量很多开发者在尝试构建多输入网络时会遇到困惑。常见的问题包括如何定义forward方法来接受多个输入如何在DataLoader中处理多输入数据如何将不同模态的特征进行融合如何处理不同输入的预处理差异如何在多输入网络中使用批量训练错误复现以下代码演示了构建多输入网络时常见的错误import torch import torch.nn as nn # 错误1forward 方法只接受一个参数 class WrongNetwork1(nn.Module): def __init__(self): super().__init__() self.fc nn.Linear(100, 10) def forward(self, x): # 只接受一个输入无法处理两个输入 return self.fc(x) model WrongNetwork1() x1 torch.randn(32, 100) x2 torch.randn(32, 50) try: # 尝试传入两个输入 output model(x1, x2) except Exception as e: print(f错误1: {e}) # TypeError: forward() takes 2 positional arguments but 3 were given # 错误2DataLoader 返回的数据格式不匹配 from torch.utils.data import Dataset, DataLoader class WrongDataset(Dataset): def __init__(self, n100): self.x1 torch.randn(n, 100) self.x2 torch.randn(n, 50) self.y torch.randint(0, 10, (n,)) def __len__(self): return len(self.x1) def __getitem__(self, idx): # 返回三个元素但默认 collate_fn 可能无法正确处理 return self.x1[idx], self.x2[idx], self.y[idx] dataset WrongDataset() # DataLoader 可以处理但需要模型 forward 也接受对应参数 dataloader DataLoader(dataset, batch_size4) # 错误3特征维度不匹配 class WrongNetwork2(nn.Module): def __init__(self): super().__init__() self.branch1 nn.Linear(100, 64) self.branch2 nn.Linear(50, 64) self.classifier nn.Linear(64, 10) def forward(self, x1, x2): out1 self.branch1(x1) # (batch, 64) out2 self.branch2(x2) # (batch, 64) # 错误直接相加但没有确保特征维度一致 combined out1 out2 return self.classifier(combined) model WrongNetwork2() # 如果 x1 和 x2 的 batch_size 不一致会报错 x1 torch.randn(32, 100) x2 torch.randn(16, 50) # 不同的 batch_size try: output model(x1, x2) except Exception as e: print(f错误3: {e}) # RuntimeError: The size of tensor a (32) must match the size of tensor b (16)根因分析1.forward方法的设计PyTorch 的nn.Module.forward方法可以接受任意数量的参数。关键在于forward的参数定义要与实际传入的参数匹配。多输入网络只需在forward方法中定义多个参数即可。2. 特征融合策略多输入网络的核心挑战在于如何将不同分支提取的特征进行有效融合。常见的融合策略包括早期融合在输入层直接拼接原始特征中期融合各分支独立提取特征后在中间层拼接晚期融合各分支独立预测最后集成结果交叉注意力使用注意力机制让不同模态的特征交互3. DataLoader 的数据组织DataLoader的默认collate_fn会自动将__getitem__返回的元组中的每个元素分别堆叠。如果__getitem__返回(x1, x2, y)DataLoader会自动将其组织为(batch_x1, batch_x2, batch_y)。4. 不同模态的预处理不同输入可能需要不同的预处理。例如图像需要归一化和 resize文本需要 tokenize 和 padding数值特征需要标准化。这些预处理通常在 Dataset 的__getitem__中完成。解决方案方案一基本的多输入网络import torch import torch.nn as nn class TwoInputNetwork(nn.Module): 基本的双输入网络 def __init__(self, input1_dim, input2_dim, hidden_dim, output_dim): super().__init__() # 分支1处理第一个输入 self.branch1 nn.Sequential( nn.Linear(input1_dim, hidden_dim), nn.ReLU(), nn.Dropout(0.3), nn.Linear(hidden_dim, hidden_dim), nn.ReLU(), ) # 分支2处理第二个输入 self.branch2 nn.Sequential( nn.Linear(input2_dim, hidden_dim), nn.ReLU(), nn.Dropout(0.3), nn.Linear(hidden_dim, hidden_dim), nn.ReLU(), ) # 融合层拼接后的分类器 self.classifier nn.Sequential( nn.Linear(hidden_dim * 2, hidden_dim), nn.ReLU(), nn.Dropout(0.3), nn.Linear(hidden_dim, output_dim), ) def forward(self, x1, x2): 前向传播接受两个输入。 Args: x1: 第一个输入 (batch_size, input1_dim) x2: 第二个输入 (batch_size, input2_dim) Returns: 输出 (batch_size, output_dim) # 各分支独立处理 out1 self.branch1(x1) # (batch, hidden_dim) out2 self.branch2(x2) # (batch, hidden_dim) # 特征拼接 combined torch.cat([out1, out2], dim1) # (batch, hidden_dim * 2) # 分类 output self.classifier(combined) return output # 使用示例 model TwoInputNetwork( input1_dim100, input2_dim50, hidden_dim128, output_dim10, ) x1 torch.randn(32, 100) x2 torch.randn(32, 50) output model(x1, x2) print(f输出形状: {output.shape}) # torch.Size([32, 10])方案二图像 元数据的双输入网络import torch import torch.nn as nn class ImageMetadataNetwork(nn.Module): 图像 元数据的双输入网络 def __init__(self, num_metadata_features, num_classes): super().__init__() # 图像分支CNN self.image_branch nn.Sequential( nn.Conv2d(3, 32, kernel_size3, padding1), nn.BatchNorm2d(32), nn.ReLU(), nn.MaxPool2d(2), nn.Conv2d(32, 64, kernel_size3, padding1), nn.BatchNorm2d(64), nn.ReLU(), nn.MaxPool2d(2), nn.Conv2d(64, 128, kernel_size3, padding1), nn.BatchNorm2d(128), nn.ReLU(), nn.AdaptiveAvgPool2d((1, 1)), # 全局平均池化 ) # 元数据分支MLP self.metadata_branch nn.Sequential( nn.Linear(num_metadata_features, 64), nn.ReLU(), nn.Dropout(0.3), nn.Linear(64, 64), nn.ReLU(), ) # 融合分类器 self.classifier nn.Sequential( nn.Linear(128 64, 128), nn.ReLU(), nn.Dropout(0.3), nn.Linear(128, num_classes), ) def forward(self, image, metadata): Args: image: 图像张量 (batch, 3, H, W) metadata: 元数据特征 (batch, num_metadata_features) # 图像特征提取 img_features self.image_branch(image) img_features img_features.flatten(1) # (batch, 128) # 元数据处理 meta_features self.metadata_branch(metadata) # (batch, 64) # 融合 combined torch.cat([img_features, meta_features], dim1) # 分类 output self.classifier(combined) return output # 使用示例 model ImageMetadataNetwork(num_metadata_features10, num_classes5) images torch.randn(16, 3, 64, 64) metadata torch.randn(16, 10) output model(images, metadata) print(f输出形状: {output.shape}) # torch.Size([16, 5])方案三文本 图像的多模态网络import torch import torch.nn as nn class TextImageNetwork(nn.Module): 文本 图像的多模态网络 def __init__(self, vocab_size, embed_dim, num_classes, pad_idx0): super().__init__() # 文本分支 self.text_branch nn.Sequential( nn.Embedding(vocab_size, embed_dim, padding_idxpad_idx), # 这里简化实际可使用 LSTM/Transformer ) self.text_encoder nn.LSTM(embed_dim, 128, batch_firstTrue, bidirectionalTrue) self.text_fc nn.Sequential( nn.Linear(256, 128), nn.ReLU(), ) # 图像分支 self.image_branch nn.Sequential( nn.Conv2d(3, 32, 3, padding1), nn.ReLU(), nn.MaxPool2d(2), nn.Conv2d(32, 64, 3, padding1), ![配图](https://i-blog.csdnimg.cn/img_convert/02f718d179e1ec0d10c7821ec39dab37.png) nn.ReLU(), nn.AdaptiveAvgPool2d((1, 1)), ) self.image_fc nn.Sequential( nn.Linear(64, 128), nn.ReLU(), ) # 融合层 self.fusion nn.Sequential( nn.Linear(128 128, 128), nn.ReLU(), nn.Dropout(0.3), nn.Linear(128, num_classes), ) def forward(self, text, image): Args: text: 文本索引序列 (batch, seq_len) image: 图像 (batch, 3, H, W) # 文本处理 embedded self.text_branch(text) # (batch, seq_len, embed_dim) lstm_out, (hidden, _) self.text_encoder(embedded) # 拼接最后正向和反向的隐状态 text_features torch.cat([hidden[-2], hidden[-1]], dim1) # (batch, 256) text_features self.text_fc(text_features) # (batch, 128) # 图像处理 img_features self.image_branch(image) img_features img_features.flatten(1) # (batch, 64) img_features self.image_fc(img_features) # (batch, 128) # 融合 combined torch.cat([text_features, img_features], dim1) output self.fusion(combined) return output完整修复代码以下是一个完整的、生产级别的多输入网络实现包含数据处理、模型定义、训练和评估import torch import torch.nn as nn import torch.optim as optim from torch.utils.data import Dataset, DataLoader import numpy as np from typing import Tuple, Dict, Any, List, Optional # 数据集定义 class MultiInputDataset(Dataset): 多输入数据集。 支持图像 数值特征 文本等多种输入组合。 def __init__(self, num_samples1000, image_size32, num_metadata10, seq_len20, vocab_size100): self.num_samples num_samples # 生成模拟数据 self.images torch.randn(num_samples, 3, image_size, image_size) self.metadata torch.randn(num_samples, num_metadata) self.texts torch.randint(1, vocab_size, (num_samples, seq_len)) self.labels torch.randint(0, 5, (num_samples,)) def __len__(self): return self.num_samples def __getitem__(self, idx): return { image: self.images[idx], metadata: self.metadata[idx], text: self.texts[idx], label: self.labels[idx], } def multi_input_collate_fn(batch: List[Dict[str, torch.Tensor]]) - Dict[str, torch.Tensor]: 自定义 collate_fn将字典列表转为批量字典。 result {} for key in batch[0]: result[key] torch.stack([item[key] for item in batch]) return result # 模型定义 class MultiModalNetwork(nn.Module): 多模态网络支持图像、数值特征和文本输入。 使用中期融合策略。 def __init__(self, config: Dict[str, Any]): super().__init__() self.config config # 图像分支 image_channels config.get(image_channels, 3) image_embed_dim config.get(image_embed_dim, 128) self.image_branch nn.Sequential( nn.Conv2d(image_channels, 32, kernel_size3, padding1), nn.BatchNorm2d(32), nn.ReLU(), nn.MaxPool2d(2), nn.Conv2d(32, 64, kernel_size3, padding1), nn.BatchNorm2d(64), nn.ReLU(), nn.MaxPool2d(2), nn.Conv2d(64, 128, kernel_size3, padding1), nn.BatchNorm2d(128), nn.ReLU(), nn.AdaptiveAvgPool2d((1, 1)), ) self.image_fc nn.Sequential( nn.Linear(128, image_embed_dim), nn.ReLU(), nn.Dropout(0.3), ) # 数值特征分支 metadata_dim config.get(metadata_dim, 10) metadata_embed_dim config.get(metadata_embed_dim, 64) self.metadata_branch nn.Sequential( nn.Linear(metadata_dim, 64), nn.ReLU(), nn.BatchNorm1d(64), nn.Dropout(0.3), nn.Linear(64, metadata_embed_dim), nn.ReLU(), ) # 文本分支 vocab_size config.get(vocab_size, 100) text_embed_dim config.get(text_embed_dim, 64) text_hidden_dim config.get(text_hidden_dim, 128) self.text_embedding nn.Embedding(vocab_size, text_embed_dim, padding_idx0) self.text_lstm nn.LSTM( text_embed_dim, text_hidden_dim, batch_firstTrue, bidirectionalTrue, dropout0.3, ) self.text_fc nn.Sequential( nn.Linear(text_hidden_dim * 2, 128), nn.ReLU(), nn.Dropout(0.3), ) # 融合层 fusion_input_dim image_embed_dim metadata_embed_dim 128 fusion_hidden_dim config.get(fusion_hidden_dim, 128) num_classes config.get(num_classes, 5) self.fusion nn.Sequential( nn.Linear(fusion_input_dim, fusion_hidden_dim), nn.ReLU(), nn.BatchNorm1d(fusion_hidden_dim), nn.Dropout(0.4), nn.Linear(fusion_hidden_dim, fusion_hidden_dim // 2), nn.ReLU(), nn.Dropout(0.3), nn.Linear(fusion_hidden_dim // 2, num_classes), ) # 初始化权重 self._init_weights() def _init_weights(self): 初始化权重 for m in self.modules(): if isinstance(m, nn.Linear): nn.init.xavier_uniform_(m.weight) if m.bias is not None: nn.init.zeros_(m.bias) elif isinstance(m, nn.Conv2d): nn.init.kaiming_normal_(m.weight, modefan_out, nonlinearityrelu) def forward(self, image, metadata, text): 多输入前向传播。 Args: image: 图像张量 (batch, 3, H, W) metadata: 数值特征 (batch, metadata_dim) text: 文本索引序列 (batch, seq_len) Returns: 分类 logits (batch, num_classes) # 图像分支 img_feat self.image_branch(image) # (B, 128, 1, 1) img_feat img_feat.flatten(1) # (B, 128) img_feat self.image_fc(img_feat) # (B, image_embed_dim) # 元数据分支 meta_feat self.metadata_branch(metadata) # (B, metadata_embed_dim) # 文本分支 text_emb self.text_embedding(text) # (B, seq_len, text_embed_dim) lstm_out, (hidden, _) self.text_lstm(text_emb) text_feat torch.cat([hidden[-2], hidden[-1]], dim1) # (B, 256) text_feat self.text_fc(text_feat) # (B, 128) # 融合 combined torch.cat([img_feat, meta_feat, text_feat], dim1) # 分类 output self.fusion(combined) return output def extract_features(self, image, metadata, text): 提取融合后的特征用于可视化或进一步分析 img_feat self.image_branch(image).flatten(1) img_feat self.image_fc(img_feat) meta_feat self.metadata_branch(metadata) text_emb self.text_embedding(text) lstm_out, (hidden, _) self.text_lstm(text_emb) text_feat torch.cat([hidden[-2], hidden[-1]], dim1) text_feat self.text_fc(text_feat) combined torch.cat([img_feat, meta_feat, text_feat], dim1) return combined # 训练器 class MultiInputTrainer: 多输入网络训练器 def __init__(self, model, devicecuda, learning_rate1e-3): self.model model.to(device) self.device torch.device(device) self.optimizer optim.AdamW(model.parameters(), lrlearning_rate, weight_decay1e-4) self.criterion nn.CrossEntropyLoss() self.scheduler optim.lr_scheduler.CosineAnnealingLR(self.optimizer, T_max10) self.train_losses [] self.val_losses [] self.train_accs [] self.val_accs [] def train_epoch(self, dataloader): 训练一个 epoch self.model.train() total_loss 0 correct 0 total 0 for batch in dataloader: # 将所有输入移到设备 images batch[image].to(self.device) metadata batch[metadata].to(self.device) texts batch[text].to(self.device) labels batch[label].to(self.device) self.optimizer.zero_grad() # 前向传播多输入 outputs self.model(images, metadata, texts) loss self.criterion(outputs, labels) loss.backward() # 梯度裁剪 torch.nn.utils.clip_grad_norm_(self.model.parameters(), max_norm1.0) self.optimizer.step() total_loss loss.item() pred outputs.argmax(dim1) correct pred.eq(labels).sum().item() total labels.size(0) self.scheduler.step() avg_loss total_loss / len(dataloader) accuracy 100. * correct / total self.train_losses.append(avg_loss) self.train_accs.append(accuracy) return {loss: avg_loss, accuracy: accuracy} torch.no_grad() def validate(self, dataloader): 验证 self.model.eval() total_loss 0 correct 0 total 0 for batch in dataloader: images batch[image].to(self.device) metadata batch[metadata].to(self.device) texts batch[text].to(self.device) labels batch[label].to(self.device) outputs self.model(images, metadata, texts) loss self.criterion(outputs, labels) total_loss loss.item() pred outputs.argmax(dim1) correct pred.eq(labels).sum().item() total labels.size(0) avg_loss total_loss / len(dataloader) accuracy 100. * correct / total self.val_losses.append(avg_loss) self.val_accs.append(accuracy) return {loss: avg_loss, accuracy: accuracy} def train(self, train_loader, val_loader, epochs10): 完整训练 print(f{*60}) print(f开始训练 | 设备: {self.device} | 轮数: {epochs}) print(f{*60}) best_val_acc 0 for epoch in range(1, epochs 1): train_metrics self.train_epoch(train_loader) val_metrics self.validate(val_loader) print(fEpoch {epoch}/{epochs} | fTrain Loss: {train_metrics[loss]:.4f}, Acc: {train_metrics[accuracy]:.2f}% | fVal Loss: {val_metrics[loss]:.4f}, Acc: {val_metrics[accuracy]:.2f}%) if val_metrics[accuracy] best_val_acc: best_val_acc val_metrics[accuracy] torch.save(self.model.state_dict(), best_multi_input_model.pth) print(f - 最佳模型已保存 (Val Acc: {best_val_acc:.2f}%)) print(f\n训练完成最佳验证准确率: {best_val_acc:.2f}%) return self.model # 主程序 if __name__ __main__: # 配置 config { image_channels: 3, image_embed_dim: 128, metadata_dim: 10, metadata_embed_dim: 64, vocab_size: 100, text_embed_dim: 64, text_hidden_dim: 128, fusion_hidden_dim: 128, num_classes: 5, } # 创建数据集 train_dataset MultiInputDataset(num_samples2000, image_size32) val_dataset MultiInputDataset(num_samples400, image_size32) train_loader DataLoader( train_dataset, batch_size32, shuffleTrue, collate_fnmulti_input_collate_fn, ) val_loader DataLoader( val_dataset, batch_size32, shuffleFalse, collate_fnmulti_input_collate_fn, ) # 创建模型 model MultiModalNetwork(config) # 打印模型结构 print(f模型参数数量: {sum(p.numel() for p in model.parameters()):,}) # 训练 device cuda if torch.cuda.is_available() else cpu trainer MultiInputTrainer(model, devicedevice, learning_rate1e-3) trained_model trainer.train(train_loader, val_loader, epochs10) # 推理示例 print(\n--- 推理示例 ---) model.eval() sample { image: torch.randn(1, 3, 32, 32).to(device), metadata: torch.randn(1, 10).to(device), text: torch.randint(1, 100, (1, 20)).to(device), } with torch.no_grad(): output model(sample[image], sample[metadata], sample[text]) pred output.argmax(dim1).item() print(f预测类别: {pred}) print(fLogits: {output.cpu().numpy()}) print(\n完成)常见陷阱与注意事项1. 输入维度的对齐确保所有输入的 batch_size 一致。如果不同输入的 batch_size 不同特征拼接时会报错。2. 特征尺度的归一化不同分支提取的特征可能有不同的尺度值域范围。在融合前使用 BatchNorm 层对每个分支的输出进行归一化可以避免某个分支的特征主导融合结果。3. 融合策略的选择拼接Concatenation最常用保留所有信息但增加了维度相加Addition要求各分支输出维度相同适合残差连接注意力Attention让模型学习不同模态的权重效果最好但计算量大门控Gating使用 sigmoid 门控控制各分支的贡献4. DataLoader 的 collate_fn如果__getitem__返回字典需要自定义collate_fn来正确批处理。默认的collate_fn只支持元组/列表返回值。5. 梯度裁剪多分支网络中不同分支的梯度尺度可能差异很大。使用梯度裁剪clip_grad_norm_可以稳定训练。6. 模态缺失处理在实际应用中某些样本可能缺少某个模态的输入如没有图像。需要设计机制处理模态缺失如使用零向量填充或学习一个默认的模态嵌入。总结构建多输入网络是解决多模态学习问题的关键。核心要点如下forward方法接受多个参数PyTorch 的forward方法天然支持多参数只需定义对应的参数即可。各分支独立处理为每种输入模态设计独立的特征提取分支CNN 处理图像、LSTM 处理文本、MLP 处理数值特征。特征融合策略拼接是最简单有效的融合方式注意力融合效果更好但更复杂。DataLoader 适配使用自定义collate_fn处理多输入数据的批处理。归一化和正则化在融合前使用 BatchNorm 对齐各分支的特征尺度。梯度管理使用梯度裁剪和适当的学习率稳定多分支网络的训练。通过合理设计多输入网络架构可以充分利用不同模态的互补信息提升模型性能。