如何在Pandas分组后向函数传递可变参数并生成新列?
Pandas分组后传递可变参数给自定义函数的实现方法
问题描述
给定如下Pandas DataFrame:
import pandas as pd data = { "Race_ID": [2,2,2,2,2,5,5,5,5,5,5], "Student_ID": [1,2,3,4,5,9,10,2,3,6,5], "theta": [8,9,2,12,4,5,30,3,2,1,50] } df = pd.DataFrame(data)
需要按Race_ID分组后,将函数f(thetai, *theta)应用到每组的theta列,生成新列feature。示例函数逻辑为thetai ** 2 + 同组其他theta值之和,实际使用的函数更为复杂,核心需求是传递当前元素和同组其余所有元素作为可变参数,实际函数代码如下:
import numpy as np from scipy.stats import norm from scipy.integrate import quad def integrand(xi, thetai, *theta): S = 0 for tj in theta: prod = 1 for t in theta: if abs(t - tj) < 1e-10: continue prod = prod * (1 - norm.cdf(xi + thetai - t)) S = S + norm.cdf(xi + thetai - tj) * prod return S * norm.pdf(xi) def f(thetai, *theta): return quad(integrand, -np.inf, np.inf, args=(thetai, *theta))[0]
解决方案
核心思路
通过groupby按赛事分组后,嵌套使用apply:外层分组获取每组的theta序列,内层遍历序列中的每个元素,将当前元素作为第一个参数,组内其余元素打包为可变参数传递给目标函数。
针对示例函数的实现
def f_example(thetai, *theta): return thetai ** 2 + sum(theta) # 生成feature列 df['feature'] = df.groupby('Race_ID')['theta'].apply( lambda group: group.apply(lambda x: f_example(x, *[t for t in group if t != x])) ).reset_index(drop=True)
运行后得到期望输出:
data = { "Race_ID": [2,2,2,2,2,5,5,5,5,5,5], "Student_ID": [1,2,3,4,5,9,10,2,3,6,5], "theta": [8,9,2,12,4,5,30,3,2,1,50], "feature": [91,107,37,167,47,111,961,97,93,91,2541] } result_df = pd.DataFrame(data)
针对实际复杂函数的实现
直接替换函数为自定义的f即可,注意导入所需依赖:
import numpy as np import pandas as pd from scipy.stats import norm from scipy.integrate import quad def integrand(xi, thetai, *theta): S = 0 for tj in theta: prod = 1 for t in theta: if abs(t - tj) < 1e-10: continue prod = prod * (1 - norm.cdf(xi + thetai - t)) S = S + norm.cdf(xi + thetai - tj) * prod return S * norm.pdf(xi) def f(thetai, *theta): return quad(integrand, -np.inf, np.inf, args=(thetai, *theta))[0] # 生成feature列 df['feature'] = df.groupby('Race_ID')['theta'].apply( lambda group: group.apply(lambda x: f(x, *[t for t in group if t != x])) ).reset_index(drop=True)
注意事项
如果组内存在重复的theta值,上述代码会过滤掉所有与当前元素相等的值。若仅需排除当前元素(保留其他重复项),可改用索引过滤:
lambda group: group.apply(lambda x: f(x, *[t for idx, t in enumerate(group) if idx != group.index.get_loc(x.name)]))
内容的提问来源于stack exchange,提问作者Ishigami
相关产品推荐
相关产品推荐

