【Bug已解决】How to construct a network with two inputs in PyTorch 解决方案
问题描述
在实际的深度学习项目中,很多任务需要处理多种类型的输入数据。例如:
- 图像分类任务中,输入既包括图像,还包括元数据(如拍摄位置、时间等)
- 推荐系统中,输入包括用户特征和物品特征
- 多模态学习中,输入包括文本和图像
- 视频理解中,输入包括视频帧和音频
这些场景都需要构建一个能够同时接受两个或多个输入的网络模型。然而,PyTorch 的nn.Module默认的forward方法通常只接受一个输入张量,很多开发者在尝试构建多输入网络时会遇到困惑。
常见的问题包括:
- 如何定义
forward方法来接受多个输入 - 如何在
DataLoader中处理多输入数据 - 如何将不同模态的特征进行融合
- 如何处理不同输入的预处理差异
- 如何在多输入网络中使用批量训练
错误复现
以下代码演示了构建多输入网络时常见的错误:
import torch import torch.nn as nn # ===== 错误1:forward 方法只接受一个参数 ===== 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 # ===== 错误2:DataLoader 返回的数据格式不匹配 ===== from torch.utils.data import Dataset, DataLoader class WrongDataset(Dataset): def __init__(self, n=100): 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_size=4) # ===== 错误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], dim=1) # (batch, hidden_dim * 2) # 分类 output = self.classifier(combined) return output # 使用示例 model = TwoInputNetwork( input1_dim=100, input2_dim=50, hidden_dim=128, output_dim=10, ) 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_size=3, padding=1), nn.BatchNorm2d(32), nn.ReLU(), nn.MaxPool2d(2), nn.Conv2d(32, 64, kernel_size=3, padding=1), nn.BatchNorm2d(64), nn.ReLU(), nn.MaxPool2d(2), nn.Conv2d(64, 128, kernel_size=3, padding=1), 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], dim=1) # 分类 output = self.classifier(combined) return output # 使用示例 model = ImageMetadataNetwork(num_metadata_features=10, num_classes=5) 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_idx=0): super().__init__() # 文本分支 self.text_branch = nn.Sequential( nn.Embedding(vocab_size, embed_dim, padding_idx=pad_idx), # 这里简化,实际可使用 LSTM/Transformer ) self.text_encoder = nn.LSTM(embed_dim, 128, batch_first=True, bidirectional=True) self.text_fc = nn.Sequential( nn.Linear(256, 128), nn.ReLU(), ) # 图像分支 self.image_branch = nn.Sequential( nn.Conv2d(3, 32, 3, padding=1), nn.ReLU(), nn.MaxPool2d(2), nn.Conv2d(32, 64, 3, padding=1),  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]], dim=1) # (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], dim=1) 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_samples=1000, image_size=32, num_metadata=10, seq_len=20, vocab_size=100): 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_size=3, padding=1), nn.BatchNorm2d(32), nn.ReLU(), nn.MaxPool2d(2), nn.Conv2d(32, 64, kernel_size=3, padding=1), nn.BatchNorm2d(64), nn.ReLU(), nn.MaxPool2d(2), nn.Conv2d(64, 128, kernel_size=3, padding=1), 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_idx=0) self.text_lstm = nn.LSTM( text_embed_dim, text_hidden_dim, batch_first=True, bidirectional=True, dropout=0.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, mode='fan_out', nonlinearity='relu') 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]], dim=1) # (B, 256) text_feat = self.text_fc(text_feat) # (B, 128) # 融合 combined = torch.cat([img_feat, meta_feat, text_feat], dim=1) # 分类 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]], dim=1) text_feat = self.text_fc(text_feat) combined = torch.cat([img_feat, meta_feat, text_feat], dim=1) return combined # ===== 训练器 ===== class MultiInputTrainer: """多输入网络训练器""" def __init__(self, model, device='cuda', learning_rate=1e-3): self.model = model.to(device) self.device = torch.device(device) self.optimizer = optim.AdamW(model.parameters(), lr=learning_rate, weight_decay=1e-4) self.criterion = nn.CrossEntropyLoss() self.scheduler = optim.lr_scheduler.CosineAnnealingLR(self.optimizer, T_max=10) 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_norm=1.0) self.optimizer.step() total_loss += loss.item() pred = outputs.argmax(dim=1) 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(dim=1) 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, epochs=10): """完整训练""" 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(f"Epoch {epoch}/{epochs} | " f"Train Loss: {train_metrics['loss']:.4f}, Acc: {train_metrics['accuracy']:.2f}% | " f"Val 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_samples=2000, image_size=32) val_dataset = MultiInputDataset(num_samples=400, image_size=32) train_loader = DataLoader( train_dataset, batch_size=32, shuffle=True, collate_fn=multi_input_collate_fn, ) val_loader = DataLoader( val_dataset, batch_size=32, shuffle=False, collate_fn=multi_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, device=device, learning_rate=1e-3) trained_model = trainer.train(train_loader, val_loader, epochs=10) # 推理示例 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(dim=1).item() print(f"预测类别: {pred}") print(f"Logits: {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 对齐各分支的特征尺度。
- 梯度管理:使用梯度裁剪和适当的学习率,稳定多分支网络的训练。
通过合理设计多输入网络架构,可以充分利用不同模态的互补信息,提升模型性能。