如何用Python为DataFrame添加Momentum列:基于连续possessions判断
为Pandas DataFrame添加Momentum列的解决方案
需要为Pandas DataFrame(简称df)添加momentum列,规则为:当该行的possessions值属于长度≥3的连续整数序列时,取值为True,否则为False。
原始DataFrame
| event | seconds_since_start_of_game | possessions |
|---|---|---|
| FOUL_RECEIVE | 49.0 | 1.0 |
| FOUL_RECEIVE | 73.0 | 1.0 |
| POST_OUT | 86.0 | 1.0 |
| LOST_BALL | 101.0 | 1.0 |
| 2MIN_PROVOKE | 122.0 | 2.0 |
| FOUL_RECEIVE | 137.0 | 2.0 |
| POST_OUT | 148.0 | 2.0 |
| POST_OUT | 306.0 | 4.0 |
| LOST_BALL | 324.0 | 5.0 |
| FOUL_RECEIVE | 362.0 | 6.0 |
| LOST_BALL | 376.0 | 6.0 |
| FOUL_RECEIVE | 399.0 | 7.0 |
期望结果DataFrame
| event | seconds_since_start_of_game | possessions | momentum |
|---|---|---|---|
| FOUL_RECEIVE | 49.0 | 1.0 | False |
| FOUL_RECEIVE | 73.0 | 1.0 | False |
| POST_OUT | 86.0 | 1.0 | False |
| LOST_BALL | 101.0 | 1.0 | False |
| 2MIN_PROVOKE | 122.0 | 2.0 | False |
| FOUL_RECEIVE | 137.0 | 2.0 | False |
| POST_OUT | 148.0 | 2.0 | False |
| POST_OUT | 306.0 | 4.0 | True |
| LOST_BALL | 324.0 | 5.0 | True |
| FOUL_RECEIVE | 362.0 | 6.0 | True |
| LOST_BALL | 376.0 | 6.0 | True |
| FOUL_RECEIVE | 399.0 | 7.0 | True |
解决方案代码
步骤说明
- 提取并排序
possessions的唯一值,找出其中连续的整数序列; - 筛选出长度≥3的连续序列,收集这些序列中的所有数值;
- 基于这些数值判断每行的
momentum值。
代码实现
import pandas as pd import numpy as np # 加载原始数据到df df = pd.DataFrame({ 'event': ['FOUL_RECEIVE', 'FOUL_RECEIVE', 'POST_OUT', 'LOST_BALL', '2MIN_PROVOKE', 'FOUL_RECEIVE', 'POST_OUT', 'POST_OUT', 'LOST_BALL', 'FOUL_RECEIVE', 'LOST_BALL', 'FOUL_RECEIVE'], 'seconds_since_start_of_game': [49.0, 73.0, 86.0, 101.0, 122.0, 137.0, 148.0, 306.0, 324.0, 362.0, 376.0, 399.0], 'possessions': [1.0, 1.0, 1.0, 1.0, 2.0, 2.0, 2.0, 4.0, 5.0, 6.0, 6.0, 7.0] }) # 提取排序后的唯一possessions值 unique_poss = sorted(df['possessions'].unique()) # 找出连续序列的分割点 breaks = np.where(np.diff(unique_poss) != 1)[0] + 1 # 分割为连续序列组 poss_groups = np.split(unique_poss, breaks) # 筛选出长度≥3的连续序列,合并为有效数值集合 valid_poss = set().union(*[group for group in poss_groups if len(group) >= 3]) # 添加momentum列 df['momentum'] = df['possessions'].isin(valid_poss) print(df)
代码解释
np.diff(unique_poss) !=1:判断相邻唯一值是否连续(差值为1),定位不连续的分割位置;np.split:根据分割位置将唯一值拆分为多个连续数值组;set().union(*...):将符合长度要求的连续序列中的数值合并为一个集合,方便后续判断;df['possessions'].isin(valid_poss):快速判断每行的possessions是否属于有效序列,生成布尔值的momentum列。
内容的提问来源于stack exchange,提问作者CypEg
相关产品推荐
相关产品推荐

