保留重复项拼接DataFrame:将多序列预测合并至原始数据并统一导出
解决方案
核心思路
先把所有预测数据整合为统一数据集,再针对每个序列长度,从原始数据中提取对应行数的记录并关联预测值,最后把所有结果拼接成一个DataFrame,导出为单个CSV文件。
具体实现代码
import pandas as pd # 假设原始数据已加载到 df_batch 中 # 把分散的预测字典整理成统一列表 predictions_list = [ {'product_id': 'mp000000000001321', 'sequence_length': 1, 'prediction': 5.75}, {'product_id': 'mp000000000001321', 'sequence_length': 3, 'prediction': 5.88} ] # 转换为DataFrame方便批量处理 predictions_df = pd.DataFrame(predictions_list) # 初始化空DataFrame存放最终结果 final_df = pd.DataFrame() # 遍历每个预测条目,生成对应结果 for _, pred_row in predictions_df.iterrows(): # 提取当前预测的商品ID、序列长度、预测值 product_id = pred_row['product_id'] seq_len = pred_row['sequence_length'] pred_val = pred_row['prediction'] # 从原始数据中筛选当前商品的前seq_len行记录 filtered_data = df_batch[df_batch['product_id'] == product_id].head(seq_len) # 添加预测相关字段 filtered_data['prediction'] = pred_val filtered_data['sequence_length'] = seq_len # 将当前结果追加到最终数据集 final_df = pd.concat([final_df, filtered_data], ignore_index=True) # 导出为单个CSV文件 final_df.to_csv('merged_predictions.csv', index=False)
代码说明
- 整合预测数据:把分散的单序列预测字典转换成统一DataFrame,避免重复处理逻辑。
- 批量处理每个预测项:针对每个商品+序列长度的组合,精准提取原始数据中对应行数的记录,关联预测值后合并到结果集。
- 统一导出:所有处理完成后一次性导出,解决了原代码生成多个独立CSV的问题,同时覆盖所有商品的预测结果。
内容的提问来源于stack exchange,提问作者Prateek Singh
相关产品推荐
相关产品推荐

