PyTorch自定义模型CUDA适配报错:张量设备不匹配如何修复?
问题:GPU训练时设备不匹配RuntimeError
训练自定义BST模型时,已将模型和输入张量移至CUDA设备,但出现如下错误:
RuntimeError: Expected all tensors to be on the same device, but found at least two devices, cpu and cuda:0! (when checking argument for argument index in method wrapper__index_select)
相关代码片段
自定义模型代码
class EmbeddingLayer(nn.Module): def __init__(self): super(EmbeddingLayer, self).__init__() # other features self.other_features_embedding = [] for feature_name in OTHER_FEATURES: vocabulary = CATEGORICAL_FEATURES_WITH_VOCABULARY[feature_name] embedding_dims = int(math.sqrt(len(vocabulary))) embedding = nn.Embedding(len(vocabulary)+1, embedding_dims) self.other_features_embedding.append(embedding) # transformer features item_vocabulary = CATEGORICAL_FEATURES_WITH_VOCABULARY['item'] self.item_embedding_dims = int(math.sqrt(len(item_vocabulary))) self.item_embedding = nn.Embedding(len(item_vocabulary)+1, self.item_embedding_dims) def forward(self, inputs): # other features encoded_other_features = [] for i, feature_name in enumerate(OTHER_FEATURES): embedding = self.other_features_embedding[i](inputs[feature_name]) encoded_other_features.append(embedding) encoded_other_features = torch.cat(encoded_other_features, -1) # transformer features encoded_sequence_item = self.item_embedding(inputs['sequence_item']) encoded_target_item = self.item_embedding(inputs['target_item']) positions = inputs['target_timestamp'].repeat(sequence_length-1, 1).transpose(0, 1) - inputs['sequence_timestamp'] encoded_positions = positions.repeat(1, self.item_embedding_dims).reshape(-1, self.item_embedding_dims, sequence_length-1).transpose(1,2) encoded_sequence_item_with_position = encoded_sequence_item + encoded_positions encoded_transformer_features = torch.cat((encoded_sequence_item_with_position, encoded_target_item.reshape(-1, 1, self.item_embedding_dims)), 1) return encoded_other_features, encoded_transformer_features class BST(nn.Module): def __init__(self, hidden_units, dropout, num_heads): super(BST, self).__init__() ... self.embedding_layer = EmbeddingLayer() ... def forward(self, inputs): other_features, transformer_features = self.embedding_layer(inputs) ... return self.output(features)
模型初始化与训练代码
model = BST([256, 128], 0.3, 1) model.to(device) def train(model, optimizer, dataloader): model.train() for inputs in tqdm(dataloader, total=len(dataloader)): for k, v in inputs.items(): inputs[k] = v.to(device) model.zero_grad() pred = model(inputs) ...
错误原因
EmbeddingLayer中other_features_embedding是用普通列表存储的nn.Embedding层,这些子模块没有被PyTorch的Module系统注册管理。调用model.to(device)时,只有被注册的子模块(比如self.item_embedding)会被自动移到CUDA设备,而列表中的嵌入层仍留在CPU,导致输入张量(CUDA)与嵌入层(CPU)设备不匹配,触发RuntimeError。
修复方案
将other_features_embedding从普通列表改为nn.ModuleList,让PyTorch自动管理这些子模块,确保调用model.to(device)时所有嵌入层都被移至CUDA设备。
修改后的EmbeddingLayer代码
class EmbeddingLayer(nn.Module): def __init__(self): super(EmbeddingLayer, self).__init__() # other features - 改用nn.ModuleList存储嵌入层 self.other_features_embedding = nn.ModuleList() for feature_name in OTHER_FEATURES: vocabulary = CATEGORICAL_FEATURES_WITH_VOCABULARY[feature_name] embedding_dims = int(math.sqrt(len(vocabulary))) embedding = nn.Embedding(len(vocabulary)+1, embedding_dims) self.other_features_embedding.append(embedding) # transformer features item_vocabulary = CATEGORICAL_FEATURES_WITH_VOCABULARY['item'] self.item_embedding_dims = int(math.sqrt(len(item_vocabulary))) self.item_embedding = nn.Embedding(len(item_vocabulary)+1, self.item_embedding_dims) def forward(self, inputs): # other features - 遍历方式不变,ModuleList支持索引访问 encoded_other_features = [] for i, feature_name in enumerate(OTHER_FEATURES): embedding = self.other_features_embedding[i](inputs[feature_name]) encoded_other_features.append(embedding) encoded_other_features = torch.cat(encoded_other_features, -1) # transformer features 部分代码不变 encoded_sequence_item = self.item_embedding(inputs['sequence_item']) encoded_target_item = self.item_embedding(inputs['target_item']) positions = inputs['target_timestamp'].repeat(sequence_length-1, 1).transpose(0, 1) - inputs['sequence_timestamp'] encoded_positions = positions.repeat(1, self.item_embedding_dims).reshape(-1, self.item_embedding_dims, sequence_length-1).transpose(1,2) encoded_sequence_item_with_position = encoded_sequence_item + encoded_positions encoded_transformer_features = torch.cat((encoded_sequence_item_with_position, encoded_target_item.reshape(-1, 1, self.item_embedding_dims)), 1) return encoded_other_features, encoded_transformer_features
额外说明
nn.ModuleList是PyTorch专门用于存储子模块的容器,会自动注册所有内部的Module,确保model.to(device)、model.parameters()等操作能覆盖到所有子模块。- 原forward方法中的代码无需修改,因为
nn.ModuleList支持通过索引访问子模块,和普通列表的使用方式一致。
内容的提问来源于stack exchange,提问作者tworiver
相关产品推荐
相关产品推荐

