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

如何在PyTorch中对字典输入使用CUDA?含字符串特征处理疑问

PyTorch字典格式输入的CUDA调用时机

首先明确:只有张量(Tensor)类型的数据需要移动到CUDA,字符串这类非张量特征本身不参与GPU计算,无需直接调用.cuda(),要先在模型中编码为张量后再处理。

以下是两种实用的处理时机方案:

1. 在DataLoader的collate_fn中批量处理(推荐)

自定义批量处理函数,在数据加载阶段就把所有张量移到CUDA,效率更高:

def collate_fn(batch):
    # 将批量数据整理为key对应列表的字典结构
    batch_dict = {key: [item[key] for item in batch] for key in batch[0].keys()}
    
    # 针对张量类型的字段执行cuda()
    tensor_keys = ['sequence_timestamp', 'target', 'target_timestamp']
    for key in tensor_keys:
        batch_dict[key] = torch.stack(batch_dict[key]).cuda()
    
    # 字符串特征保持原格式,留到模型中编码
    return batch_dict

创建DataLoader时传入该函数:

train_loader = DataLoader(train_dataset, batch_size=32, collate_fn=collate_fn)

2. 在模型forward方法中处理

若不想修改DataLoader,可在模型前向传播时,先迁移张量再处理字符串特征:

class CustomModel(nn.Module):
    def __init__(self, sex_vocab_size, embed_dim):
        super().__init__()
        self.sex_embedding = nn.Embedding(sex_vocab_size, embed_dim)
        self.device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
    
    def encode_sex(self, sex_list):
        # 自定义字符串编码逻辑,返回索引张量
        vocab = {'여성':0, '남성':1}
        return torch.tensor([vocab[sex] for sex in sex_list])
    
    def forward(self, input_dict):
        # 迁移张量到设备
        sequence_ts = input_dict['sequence_timestamp'].to(self.device)
        target = input_dict['target'].to(self.device)
        target_ts = input_dict['target_timestamp'].to(self.device)
        
        # 处理字符串特征:编码为张量后再迁移到设备
        sex_labels = self.encode_sex(input_dict['sex']).to(self.device)
        sex_emb = self.sex_embedding(sex_labels)
        
        # 后续模型计算逻辑...
        return output

核心注意事项

  • 字符串特征不能直接调用.cuda(),必须先转成索引张量,再将张量移到CUDA后输入嵌入层。
  • 批量迁移张量的效率远高于单个迁移,优先选择在collate_fn中处理。

内容的提问来源于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 19:54:27