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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.24 16:09:25