如何用Lambda函数为Pandas每行执行SQL聚合查询并合并结果
问题描述
我有一个名为final_data的Pandas DataFrame,结构如下:
| cust_id | start_date | end_date |
|---|---|---|
| 10001 | 2022-01-01 | 2022-01-30 |
| 10002 | 2022-02-01 | 2022-02-30 |
| 10003 | 2022-01-01 | 2022-01-30 |
| 10004 | 2022-03-01 | 2022-03-30 |
| 10005 | 2022-02-01 | 2022-02-30 |
SQL数据库中有一个名为penalties的表,结构如下:
| cust_id | level1_pen | level_2_pen | date |
|---|---|---|---|
| 10001 | 1 | 4 | 2022-01-01 |
| 10001 | 1 | 1 | 2022-01-02 |
| 10001 | 0 | 1 | 2022-01-30 |
| 10002 | 1 | 1 | 2022-01-01 |
| 10002 | 5 | 0 | 2022-02-01 |
| 10002 | 4 | 0 | 2022-02-04 |
| 10003 | 1 | 6 | 2022-01-02 |
需要将final_data处理为如下结构,其中total_penalties列是根据每行的cust_id、start_date和end_date,从SQL的penalties表中聚合计算level1_pen + level_2_pen的总和:
| cust_id | start_date | end_date | total_penalties |
|---|---|---|---|
| 10001 | 2022-01-01 | 2022-01-30 | 8 |
| 10002 | 2022-02-01 | 2022-02-30 | 9 |
| 10003 | 2022-01-01 | 2022-01-30 | 7 |
请问如何使用Lambda函数对final_data的每一行,基于该行的cust_id、start_date和end_date变量执行SQL查询并聚合数据,最终合并到原DataFrame中?
解决方案
步骤1:建立数据库连接
首先用SQLAlchemy建立与数据库的连接(替换为你的数据库实际连接信息):
from sqlalchemy import create_engine import pandas as pd # 示例连接字符串,根据数据库类型调整(MySQL、PostgreSQL等) engine = create_engine('mysql+pymysql://username:password@host:port/db_name')
步骤2:用Lambda函数逐行执行查询
通过df.apply()结合Lambda函数,对每一行执行参数化SQL查询,计算总罚金额:
def get_total_penalty(row): # 参数化SQL语句,避免SQL注入风险 sql = """ SELECT SUM(level1_pen + level_2_pen) AS total FROM penalties WHERE cust_id = %s AND date >= %s AND date <= %s """ # 执行查询并提取结果 result = pd.read_sql(sql, engine, params=(row['cust_id'], row['start_date'], row['end_date'])) # 无匹配数据时返回0,避免NaN return result['total'].iloc[0] if not result.empty else 0 # 应用Lambda函数到每一行,生成新列 final_data['total_penalties'] = final_data.apply(lambda row: get_total_penalty(row), axis=1) # 按需过滤无罚则的行 final_data = final_data[final_data['total_penalties'] > 0].reset_index(drop=True)
补充:批量查询优化(针对大数据量)
逐行查询效率较低,数据量大时建议改用批量查询:
# 把final_data的条件传入SQL,一次性拉取所有匹配数据 batch_sql = """ SELECT f.cust_id, f.start_date, f.end_date, SUM(p.level1_pen + p.level2_pen) AS total_penalties FROM ( SELECT cust_id, start_date, end_date FROM final_data_temp ) f LEFT JOIN penalties p ON f.cust_id = p.cust_id AND p.date BETWEEN f.start_date AND f.end_date GROUP BY f.cust_id, f.start_date, f.end_date """ # 先将final_data临时写入数据库(需权限),再执行查询 final_data.to_sql('final_data_temp', engine, if_exists='replace', index=False) penalty_totals = pd.read_sql(batch_sql, engine) # 合并结果到原DataFrame final_data = final_data.merge(penalty_totals, on=['cust_id', 'start_date', 'end_date'], how='left') final_data['total_penalties'] = final_data['total_penalties'].fillna(0)
内容的提问来源于stack exchange,提问作者Tahir Zamaan
相关产品推荐
相关产品推荐

