如何将CSV文件最后一列序列转为字符级one-hot编码矩阵适配PyTorch模型
问题修复与正确实现
原有代码核心错误
- 读取数据时
pd.read_csv('database.csv', usecols=[4])返回的是二维DataFrame对象,直接遍历data得到的是列名,而非列中存储的氨基酸序列值 - 代码逻辑没有遍历每个序列内部的单个氨基酸字符,也没有为每个序列单独生成对应的one-hot编码矩阵,完全不符合字符级编码的需求
正确实现代码
import pandas as pd import torch from torch.nn.utils.rnn import pad_sequence # 1. 正确读取最后一列,squeeze参数将单列DataFrame转为Series,每个元素对应一个氨基酸序列 data = pd.read_csv('database.csv', usecols=[4]).squeeze().tolist() # 2. 氨基酸映射表保持不变 alphabet = ['A', 'C', 'D', 'E', 'F', 'G','H', 'I', 'K', 'L', 'M', 'N', 'P', 'Q', 'R', 'S', 'T', 'V', 'W', 'Y'] char_to_idx = {c:i for i,c in enumerate(alphabet)} num_alphabet = len(alphabet) onehot_list = [] for seq in data: # 3. 遍历单个序列的每个字符生成编码 seq_onehot = [] for char in seq: vec = [0]*num_alphabet vec[char_to_idx[char]] = 1 seq_onehot.append(vec) # 转为tensor方便后续PyTorch使用,形状为[序列长度, 20] onehot_list.append(torch.tensor(seq_onehot, dtype=torch.float32)) # 4. 适配PyTorch批次训练可选:如果序列长度不一致,可使用pad_sequence对齐 # 对齐后形状为 [batch_size, max_len, 20],batch_size为序列数量,max_len为最长序列长度,20为氨基酸类别数 padded_onehot = pad_sequence(onehot_list, batch_first=True)
输出说明
- 单个序列的one-hot编码形状为
[序列长度, 20],20对应20种标准氨基酸 - 对齐后的批次数据可直接输入PyTorch的CNN、RNN等模型使用
- 若序列中存在不在
alphabet里的异常字符,可在映射时增加异常处理逻辑,比如指定默认索引或者过滤序列
内容的提问来源于stack exchange,提问作者khashayar ehteshami
相关产品推荐
相关产品推荐

