如何在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
相关产品推荐
相关产品推荐

