如何将PyTorch模型预测结果合并至原始DataFrame?
问题
我从PyTorch模型得到形状为torch.Size([2958, 96])的预测张量。原始数据集包含2958个qid,每个qid对应的文档数最多96条、最少47条,预测张量中缺失位置用-1填充。原始DataFrame形状为(221567, 7)。
需要按qid将PyTorch预测结果合并到该DataFrame中:张量每行对应一个qid,每列对应该qid下对应顺序文档的排名,且文档顺序与DataFrame完全一致。
简化示例如下(张量已转为DataFrame):
import pandas as pd # 预测张量转成的DataFrame,每行对应一个qid,列是该qid下各文档的预测排名 tensor = {'0': ['3', '1','2'],'1': ['2', '1','2'],'2': ['2', '1','-1']} y_pred = pd.DataFrame(tensor) # 原始数据集,qid对应多行文档,顺序与张量列顺序一致 data = {'qid': ['0', '0','0','1', '1','1','2', '2'],'irrelevant_col': ['foo', 'foo','foo','foo', 'bar','bar','bar', 'bar']} original_df = pd.DataFrame(data)
注意:qid==2仅对应2行数据,因此张量第2行第3列的值为-1。目标输出如下:
target = {'qid': ['0', '0','0','1', '1','1','2', '2'],'irrelevant_col': ['foo', 'foo','foo','foo', 'bar','bar','bar', 'bar'],'y_pred': ['3', '1','2','2', '1','2','2', '1']} target_df = pd.DataFrame(target)
编辑说明:已修正示例中目标输出的y_pred值,使其与张量对应一致,qid=2的最后一个y_pred应为1,原张量第2行第3列的-1无需取用
解决方案
可以通过以下步骤实现需求:
处理预测张量,转换为长格式
先将预测张量的DataFrame转置,让索引对应qid,每行的列是该qid下的文档排名,再用melt方法把宽格式转成长格式,同时过滤掉-1的无效值(对应不存在的文档)。为原始DataFrame添加组内序号
对原始DataFrame按qid分组,为每个qid下的文档添加组内顺序编号(与张量的列索引对应)。合并两个数据集
将处理后的预测数据和原始DataFrame按qid和组内序号合并,得到最终结果。
具体代码实现:
import pandas as pd # 1. 处理预测张量 # 转置y_pred,让索引为qid,列是文档的顺序位置 y_pred_transposed = y_pred.T.reset_index() y_pred_transposed.columns = ['qid'] + [f'rank_{i}' for i in range(y_pred.shape[0])] # 转成长格式,过滤-1的无效值 y_pred_long = y_pred_transposed.melt( id_vars='qid', var_name='doc_order', value_name='y_pred' ).query("y_pred != '-1'") # 提取文档的顺序编号(去掉rank_前缀,转成整数) y_pred_long['doc_order'] = y_pred_long['doc_order'].str.replace('rank_', '').astype(int) # 2. 为原始DataFrame添加组内序号 original_df['doc_order'] = original_df.groupby('qid').cumcount() # 3. 合并数据集 final_df = original_df.merge( y_pred_long, on=['qid', 'doc_order'], how='left' ).drop('doc_order', axis=1) # 查看结果 print(final_df)
运行上述代码后,final_df将与目标输出一致。
如果是直接处理PyTorch张量,可先转为numpy数组再构造DataFrame:
import torch # 假设pred_tensor是你的PyTorch预测张量 pred_tensor = torch.tensor([[3,1,2], [2,1,2], [2,1,-1]]) # 转成numpy数组再构造DataFrame,索引对应qid y_pred = pd.DataFrame(pred_tensor.numpy().astype(str))
内容的提问来源于stack exchange,提问作者Tartaglia
相关产品推荐
相关产品推荐

