如何用lambda/np.where为DataFrame按股票分组添加signal列
问题
本人是Python新手,在为DataFrame添加新列signal时遇到问题。现有按Symbol分组的DataFrame,希望按如下规则生成signal列:
- 每组除最后一行外,
signal为NaN - 每组最后一行根据该组前4行的
sig列值设置(如APPL前4行sig为0、1、0、1时设为1,TSLA前4行sig为0、0、1、0时设为0)
寻求用df.apply+lambda或np.where实现的方法。
原始DataFrame:
Symbol open close sig 0 APPL 153.60 152.90 0 1 APPL 152.90 153.55 1 2 APPL 153.55 152.00 0 3 APPL 152.00 153.50 1 4 APPL 153.50 154.10 1 5 TSLA 193.00 192.10 0 6 TSLA 192.10 191.50 0 7 TSLA 191.50 192.90 1 8 TSLA 192.90 192.45 0 9 TSLA 192.45 191.10 0
期望结果:
Symbol open close sig signal 0 APPL 153.60 152.90 0 NaN 1 APPL 152.90 153.55 1 NaN 2 APPL 152.75 152.00 0 NaN 3 APPL 153.00 153.50 1 NaN 4 APPL 153.50 154.10 1 1 5 TSLA 193.00 192.10 0 NaN 6 TSLA 192.10 191.50 0 NaN 7 TSLA 191.50 192.90 1 NaN 8 TSLA 192.90 192.45 0 NaN 9 TSLA 192.45 191.10 0 0
解决方法
方法一:使用groupby.apply+自定义函数
逻辑直观,适合新手理解,每组独立处理:
import pandas as pd import numpy as np # 构造原始DataFrame(已有数据可跳过此步骤) df = pd.DataFrame({ 'Symbol': ['APPL']*5 + ['TSLA']*5, 'open': [153.60, 152.90, 153.55, 152.00, 153.50, 193.00, 192.10, 191.50, 192.90, 192.45], 'close': [152.90, 153.55, 152.00, 153.50, 154.10, 192.10, 191.50, 192.90, 192.45, 191.10], 'sig': [0,1,0,1,1,0,0,1,0,0] }) # 定义sig序列到signal值的映射规则 sig_rule = { '0101': 1, # APPL前4行sig拼接的字符串 '0010': 0 # TSLA前4行sig拼接的字符串 } def handle_group(group): # 先给整组的signal列设为NaN group['signal'] = np.nan # 只处理行数≥5的组(确保有前4行+最后一行) if len(group) >= 5: # 把前4行的sig转成字符串拼接 sig_sequence = ''.join(group['sig'].iloc[:4].astype(str)) # 从映射规则中取对应值,没有匹配则保持NaN group['signal'].iloc[-1] = sig_rule.get(sig_sequence, np.nan) return group # 按Symbol分组处理后重置索引 df = df.groupby('Symbol').apply(handle_group).reset_index(drop=True) print(df)
方法二:使用np.where+分组变换
用transform批量生成标记,再用np.where赋值,效率更高,适合大数据量:
import pandas as pd import numpy as np # 构造原始DataFrame(已有数据可跳过此步骤) df = pd.DataFrame({ 'Symbol': ['APPL']*5 + ['TSLA']*5, 'open': [153.60, 152.90, 153.55, 152.00, 153.50, 193.00, 192.10, 191.50, 192.90, 192.45], 'close': [152.90, 153.55, 152.00, 153.50, 154.10, 192.10, 191.50, 192.90, 192.45, 191.10], 'sig': [0,1,0,1,1,0,0,1,0,0] }) sig_rule = { '0101': 1, '0010': 0 } # 标记每组的最后一行 df['is_last_row'] = df.groupby('Symbol').cumcount(ascending=False) == 0 # 给每组生成前4行的sig拼接字符串 def get_sig_seq(group): return ''.join(group.iloc[:4].astype(str)) if len(group)>=5 else np.nan df['sig_sequence'] = df.groupby('Symbol')['sig'].transform(get_sig_seq) # 用np.where生成signal列:是最后一行就取映射值,否则设为NaN df['signal'] = np.where(df['is_last_row'], df['sig_sequence'].map(sig_rule), np.nan) # 删除临时列 df = df.drop(columns=['is_last_row', 'sig_sequence']) print(df)
内容的提问来源于stack exchange,提问作者B Sailoo
相关产品推荐
相关产品推荐

