Pandas:在分组内查找并标记匹配行序列模式的所有行
在分组DataFrame中标记分类变量的完整序列模式
我完全懂你的需求——要在分组后的DataFrame里,针对分类变量查找特定的序列模式,并且把模式里的每一行都标记出来,而不是只标记模式的起始行对吧?之前我处理用户行为序列分析时也遇到过一模一样的场景,下面给你分享两种可行的解决方案,从直观到高效都有:
先准备示例数据
先创建一个带分组列和分类列的测试DataFrame,方便你跟着实操:
import pandas as pd # 模拟分组后的分类序列数据 data = { 'group': ['A', 'A', 'A', 'A', 'A', 'B', 'B', 'B', 'B', 'B'], 'category': ['X', 'Y', 'Z', 'Y', 'Z', 'Y', 'Z', 'X', 'Y', 'Z'] } df = pd.DataFrame(data)
假设我们要查找的目标模式是 ['Y', 'Z']。
方法1:直观遍历法(适合小数据集)
这种方法逻辑直白,容易理解,适合数据量不大的场景:
- 定义一个处理单个分组的函数,遍历分组内的每个可能起始位置
- 检查当前起始位置的序列是否匹配目标模式
- 如果匹配,就把该序列覆盖的所有行标记为
True
def mark_pattern_in_group(group, target_pattern): pattern_length = len(target_pattern) # 初始化标记列 group['pattern_match'] = False # 遍历所有可能的模式起始位置 for start_idx in range(len(group) - pattern_length + 1): # 提取当前起始位置的序列 current_seq = group['category'].iloc[start_idx:start_idx+pattern_length].tolist() if current_seq == target_pattern: # 标记该模式覆盖的所有行 group.loc[group.index[start_idx:start_idx+pattern_length], 'pattern_match'] = True return group # 应用到每个分组 target_pattern = ['Y', 'Z'] df_marked = df.groupby('group').apply(mark_pattern_in_group, target_pattern=target_pattern).reset_index(drop=True)
运行后得到的结果:
| group | category | pattern_match |
|---|---|---|
| A | X | False |
| A | Y | True |
| A | Z | True |
| A | Y | True |
| A | Z | True |
| B | Y | True |
| B | Z | True |
| B | X | False |
| B | Y | True |
| B | Z | True |
方法2:向量化优化法(适合大数据集)
如果你的数据集很大,遍历法效率会偏低,这时候可以用shift()结合布尔运算实现向量化匹配,速度会快很多:
def mark_pattern_in_group_fast(group, target_pattern): pattern_length = len(target_pattern) group['pattern_match'] = False # 生成一个布尔序列,标记所有模式的起始位置 match_starts = pd.Series(True, index=group.index) for i in range(pattern_length): # 对比当前行、下一行...第n行是否匹配模式的对应位置 match_starts &= (group['category'].shift(-i) == target_pattern[i]) # 把每个匹配起始位置对应的所有行标记为True for start_idx in match_starts[match_starts].index: end_idx = start_idx + pattern_length - 1 # 避免超出分组索引范围 if end_idx <= group.index.max(): group.loc[start_idx:end_idx, 'pattern_match'] = True return group # 应用到每个分组 df_marked_fast = df.groupby('group').apply(mark_pattern_in_group_fast, target_pattern=target_pattern).reset_index(drop=True)
这个方法利用了pandas的向量化运算,避免了逐行遍历序列,在处理十万级以上数据时优势明显。
注意事项
- 如果你的目标模式长度大于分组的行数,函数会自动跳过该分组,不会报错
- 可以灵活修改
target_pattern为任意长度的序列,比如['X', 'Y', 'Z'] - 标记列的名称可以根据你的需求修改,比如改成
is_in_target_sequence
内容的提问来源于stack exchange,提问作者Randall Goodwin
相关产品推荐
相关产品推荐

