简介:这是一份面向数据分析和数据挖掘初学者的Apriori算法Python实现资源,压缩包共2个文件(1个Python脚本、1个txt数据集),大小仅3KB。Python脚本直接基于apyori库实现经典关联规则挖掘流程,包括构建交易列表、设置最小支持度与置信度、迭代发现频繁项集并输出规则;txt数据集则提供了可直接加载的交易记录,便于运行测试与结果验证。通过该资源,读者可以直观理解Apriori算法的核心原理——利用“非频繁项集的超集也不频繁”的先验性质进行剪枝,以及支持度、置信度两个关键指标的计算方式,甚至进一步迁移到商品捆绑推荐、用户行为分析等真实场景。目前已有9647人学习下载,适合正在学习关联规则、需要可运行代码和配套数据快速上手的读者。
1. 这个算法到底在解决什么问题
先说一个场景。你去超市买东西,收银台打出一张小票,上面有牛奶、面包、黄油。超市老板每天攒下几千张这样的小票,他想知道:买牛奶的人是不是大概率也会买面包?如果这个规律成立,他就可以把牛奶和面包摆在一起,或者给买了牛奶的人发面包优惠券。
Apriori算法就是干这个的。它从大量事务数据(比如购物小票)里找出频繁出现的商品组合,再从中挖掘出“买了A的人也倾向买B”这样的关联规则。这个场景不仅限于零售,电商推荐、医疗诊断记录、Python数据分析里的用户行为分析,都离不开同一套逻辑。
这篇内容我会先用最直白的方式讲清楚Apriori的核心机制,然后给出一份完整的Python实现代码,最后重点说说数据集怎么来、怎么清洗、怎么喂给算法。你不用是算法高手,跟着一步步跑就能出结果。
2. 先搞懂三个核心概念,否则后面代码看不懂
Apriori的原理并不复杂,核心是三个指标:支持度、置信度、提升度。我一个个拆开讲。
2.1 支持度:这个组合到底常不常见
支持度衡量的是“一个项集在全部事务中出现的频率”。公式是:
支持度(A) = 包含A的事务数 / 总事务数比如总共有100张小票,其中20张包含牛奶,那么{牛奶}这个单项集的支持度就是0.2。同理,如果同时包含牛奶和面包的小票有8张,那么项集{牛奶, 面包}的支持度就是0.08。
支持度的作用是过滤掉那些“偶然出现”的组合。你设置一个最小支持度阈值(比如0.05),低于这个值的项集直接扔掉,因为太罕见的东西谈不上规律。
2.2 置信度:条件概率的另一种说法
置信度衡量的是“买了A的人里面,有多少比例也买了B”。公式:
置信度(A → B) = 支持度(A∪B) / 支持度(A)还是上面的例子,{牛奶, 面包}的支持度是0.08,{牛奶}的支持度是0.2,那么置信度(牛奶 → 面包) = 0.08 / 0.2 = 0.4。意思是买了牛奶的顾客中,40%的人也会买面包。
置信度告诉你规则的可信程度,但它有一个陷阱:如果面包本身就是一个畅销品(60%的人都会买),那置信度0.4其实并不算高。这时候需要看提升度。
2.3 提升度:排除“本身就很火”的干扰
提升度是关联规则里最容易被人忽略、但又最重要的指标。公式:
提升度(A → B) = 置信度(A → B) / 支持度(B)如果提升度大于1,说明A的出现对B有正向促进作用;等于1说明A和B相互独立;小于1说明A的出现反而会抑制B。
用数据说话:假设面包的支持度是0.6,那提升度 = 0.4 / 0.6 ≈ 0.67,小于1。这说明牛奶和面包其实存在负相关——买了牛奶的人反而不太买面包。如果仅看置信度0.4就以为找到了强规则,就误判了。
实操经验:我见过不少初学者只盯着置信度筛选规则,结果做出来一堆“买尿布就买啤酒”式的假关联。务必把提升度大于1作为规则的硬性过滤条件。
3. Apriori的核心思想:先频繁,再关联
Apriori算法分两个阶段:第一步找频繁项集,第二步从频繁项集中生成关联规则。第一阶段是算法的精髓,有个著名的“先验性质”作为剪枝依据。
3.1 一个反直觉的剪枝原理
Apriori的先验性质是:如果一个项集是频繁的,那它的所有子集也一定是频繁的。反过来看更有用:如果一个项集不是频繁的,那它的所有超集(包含它的更大项集)一定也不频繁。
举个例子。假设{啤酒}的支持度只有0.02,低于最小支持度阈值0.05。那么在生成2项集时,任何包含啤酒的组合,比如{尿布, 啤酒}、{零食, 啤酒},都不用再计算了,因为它们的支持度一定不会超过0.02。这就是Apriori的剪枝逻辑,它避免了把所有的两两组合都数一遍。
3.2 算法迭代过程,拿3个商品举例
假设只有4种商品:牛奶、面包、鸡蛋、啤酒。最小支持度阈值设为0.5,总事务数4条:
T1: 牛奶, 面包, 鸡蛋 T2: 面包, 鸡蛋, 啤酒 T3: 牛奶, 面包, 鸡蛋 T4: 牛奶, 面包第一步,统计每个单项出现次数。牛奶3次、面包4次、鸡蛋3次、啤酒1次。支持度分别是0.75、1.0、0.75、0.25。啤酒低于阈值,直接淘汰。
第二步,用剩余的{牛奶, 面包, 鸡蛋}生成2项集:{牛奶, 面包}、{牛奶, 鸡蛋}、{面包, 鸡蛋}。统计支持度,分别是3/4、2/4、3/4,全部达标。
第三步,由2项集生成3项集{牛奶, 面包, 鸡蛋},统计支持度2/4,达标。
整个过程自底向上,每一层都在用前一层的频繁项集做拼接。这里有个关键点:生成候选集时只拼接“前k-1项完全相同”的两个项集,这是为了减少重复计算。
3.3 代码里怎么体现这个逻辑
Apriori的Python实现网上有大量版本,但很多写成了教科书风格,装了额外的库才能跑。我这份代码只用pandas和itertools,数据结构直接用DataFrame,贴近实际数据分析场景。
核心函数分三块:候选集生成、支持度计算、频繁项集迭代。先是候选集生成:
import itertools from collections import Counter import pandas as pd def generate_candidates(frequent_itemsets, k): """ 由k-1频繁项集生成k项候选集。 只合并前k-2个元素相同的两个项集。 """ candidates = set() freq_list = list(frequent_itemsets) for i in range(len(freq_list)): for j in range(i + 1, len(freq_list)): set1 = freq_list[i] set2 = freq_list[j] if sorted(set1)[:-1] == sorted(set2)[:-1]: new_candidate = tuple(sorted(set(set1) | set(set2))) if len(new_candidate) == k: candidates.add(new_candidate) return candidates然后是支持度计算。遍历所有事务,用issubset判断候选集是否为事务的子集:
def calculate_support(transactions, candidates, min_support): """ 统计候选集在事务中出现的次数,过滤出频繁项集。 transactions: list of set,每个set是一条事务。 min_support: 最小支持度,0到1之间的浮点数。 """ total = len(transactions) support_count = Counter() for candidate in candidates: candidate_set = set(candidate) for transaction in transactions: if candidate_set.issubset(transaction): support_count[candidate] += 1 frequent = { itemset: count / total for itemset, count in support_count.items() if count / total >= min_support } return frequent注意:如果事务数量很大,这种逐条遍历的写法会非常慢。后面我会讲优化方案。
4. 完整实现:从数据加载到规则输出
光讲片段没意思,我直接给你一份能跑通的完整代码。这份代码我按“数据准备 → 频繁项集 → 关联规则 → 结果展示”四步组织。
4.1 准备一份可直接运行的数据集
标题说了“含数据集”,我就直接把数据集构造好。这里我用随机方式生成一份模拟超市购物篮数据,每条事务包含1到5件商品:
import random import pandas as pd random.seed(42) items = ['牛奶', '面包', '鸡蛋', '啤酒', '尿布', '奶粉', '薯片', '可乐', '洗衣液', '纸巾'] transactions = [] for _ in range(1000): size = random.randint(1, 5) # 提高牛奶、面包、鸡蛋的关联概率,让数据更有规律 transaction = set(random.sample(items, size)) if '牛奶' in transaction and random.random() < 0.7: transaction.add('面包') if '鸡蛋' in transaction and random.random() < 0.5: transaction.add('牛奶') transactions.append(transaction)这里我特意让数据里包含了真实的关联结构:牛奶 → 面包、鸡蛋 → 牛奶。这样跑出来的规则不会全是噪音。
同时我生成一个DataFrame版本,方便你直接查看和导出:
df_transactions = pd.DataFrame({ 'transaction_id': range(1, len(transactions) + 1), 'items': [','.join(sorted(t)) for t in transactions] }) # 顺便导出成CSV,方便其他项目复用它 df_transactions.to_csv('transactions.csv', index=False, encoding='utf-8-sig') print(df_transactions.head(10))用utf-8-sig而不是utf-8,是为了在Excel里打开不乱码,这个细节容易忘。
4.2 频繁项集挖掘全流程
有了数据,开始跑频繁项集迭代:
def find_frequent_itemsets(transactions, min_support): """ 迭代生成所有频繁项集。 """ frequent_itemsets = {} k = 1 # 统计1项集 item_count = Counter() for trans in transactions: for item in trans: item_count[item] += 1 total = len(transactions) current_frequent = { (item,): count / total for item, count in item_count.items() if count / total >= min_support } frequent_itemsets.update(current_frequent) # 迭代生成k项集,直到无法生成新的频繁项集 while current_frequent: k += 1 candidates = generate_candidates(set(current_frequent.keys()), k) if not candidates: break current_frequent = calculate_support(transactions, candidates, min_support) frequent_itemsets.update(current_frequent) return frequent_itemsets min_support = 0.05 frequent_itemsets = find_frequent_itemsets(transactions, min_support) print(f"共发现 {len(frequent_itemsets)} 个频繁项集") for itemset, support in sorted(frequent_itemsets.items(), key=lambda x: x[1], reverse=True)[:10]: print(f"{itemset}: 支持度 = {support:.3f}")跑完之后,你会看到类似这样的输出:
共发现 38 个频繁项集 ('牛奶',): 支持度 = 0.526 ('面包',): 支持度 = 0.521 ('牛奶', '面包'): 支持度 = 0.394 ('鸡蛋', '牛奶'): 支持度 = 0.312 ...我实际跑的时候发现,{牛奶, 面包}的支持度差不多在0.4左右。因为生成数据时加了一句:只要事务里包含牛奶,就有70%概率额外加入面包。所以这个结果是符合预期的,说明代码逻辑正确。
4.3 从频繁项集提取关联规则
频繁项集只是第一步。要得到真正可用的规则,还需要计算置信度和提升度:
def generate_rules(frequent_itemsets, frequent_support, transactions, min_confidence=0.5, min_lift=1.0): """ 从频繁项集生成关联规则,并用置信度和提升度过滤。 """ rules = [] for itemset in frequent_itemsets: # 只有长度大于1的项集才有生成规则的必要 if len(itemset) < 2: continue itemset = set(itemset) # 生成所有非空真子集作为规则的前件 for size in range(1, len(itemset)): for antecedent in itertools.combinations(itemset, size): antecedent = set(antecedent) consequent = itemset - antecedent if not consequent: continue # 支持度(A∪B),其实就是当前频繁项集的支持度 support_ab = frequent_support[tuple(sorted(itemset))] # 支持度(A) support_a = frequent_support[tuple(sorted(antecedent))] # 支持度(B) support_b = frequent_support[tuple(sorted(consequent))] confidence = support_ab / support_a if support_a > 0 else 0 lift = confidence / support_b if support_b > 0 else 0 if confidence >= min_confidence and lift >= min_lift: rules.append({ 'antecedent': antecedent, 'consequent': consequent, 'support': support_ab, 'confidence': confidence, 'lift': lift, }) return sorted(rules, key=lambda x: x['lift'], reverse=True) rules = generate_rules( frequent_itemsets, frequent_itemsets, transactions, min_confidence=0.3, min_lift=1.1 ) result_df = pd.DataFrame(rules) print(f"共发现 {len(result_df)} 条有效规则") print(result_df.head(20))这个函数把每个频繁项集拆成“前件 → 后件”,逐一计算置信度和提升度。比如频繁项集{牛奶, 面包, 鸡蛋}可以拆出{牛奶, 面包} → {鸡蛋},也可以拆出{鸡蛋} → {牛奶, 面包},所以规则数量会远大于频繁项集数量。
提示:如果跑出来的规则太少,优先调低min_confidence,而不是调低min_support。因为支持度过滤的是组合的普遍性,置信度过滤的才是规则的强度。
5. 数据集来源与预处理,这是最容易卡住的地方
标题里特别写了“含数据集”,我就多说说数据集的事。很多人算法代码跑通了,却找不到合适的数据来测试,或者找到数据又不会清洗。这里我分享几条实际路径。
5.1 三个直接可用的数据来源
第一,公开数据集平台。Kaggle上搜“market basket optimization”,能找到几百上千条真实超市小票数据。UCI Machine Learning Repository里面有经典的Online Retail数据集,涵盖英国一家零售商店的在线订单,每行包含订单号和商品描述,非常适合做购物篮分析。
第二,Python库自带的数据。如果你用mlxtend这个库,它自带一个transactional数据集可以直接加载。虽然我这个例子里没依赖它,但你做对比验证时可以试试。
第三,自己构造数据。像我上面那样用random生成,重点在于构造时要模拟真实的关联关系,而不是完全随机。完全随机的数据跑不出任何有意义的规则,这是很多人跑完代码后怀疑人生的原因。
5.2 把原始数据整理成算法需要的格式
Apriori的输入格式有很多种,最常用的有两种:
第一种是“事务ID + 商品”的长表格式。比如:
| transaction_id | item |
|---|---|
| 1 | 牛奶 |
| 1 | 面包 |
| 1 | 鸡蛋 |
| 2 | 啤酒 |
第二种是“事务ID + 商品列表”的单行格式:
| transaction_id | items |
|---|---|
| 1 | 牛奶,面包,鸡蛋 |
| 2 | 啤酒 |
两种格式都需要转换成上面代码中transactions的list-of-set结构。转换逻辑非常简单:
import pandas as pd # 读长表格式 df = pd.read_csv('orders.csv') # 按transaction_id分组,聚合成每一项的商品集合 transactions = ( df.groupby('transaction_id')['item'] .apply(lambda x: frozenset(x)) .tolist() )用frozenset而不是set,是因为set不可哈希,不能作为Counter的键。我这里用了set,但在实际工程里建议转成frozenset。
5.3 数据清洗的一个隐藏注意点
我踩过的一个坑是:商品名的空格和大小写不一致。比如“Milk”和“ milk”会被当成两个不同的商品,导致支持度被稀释。处理方式是统一转成小写、去空格,最好再做一次同义词归并。
df['item'] = df['item'].str.strip().str.lower()还有一点,如果数据里有“退货”或“取消”记录,要提前过滤掉。我以前跑某个数据集,发现不少order_id带“C”前缀表示取消,不删掉的话这些负向记录会污染支持度统计。
6. 调参与效果分析,别让算法白跑
代码跑通了,规则也出来了,但距离“能用”还差两步:参数怎么调、结果怎么解读。这一步不做,你得到的只是一堆数字。
6.1 最小支持度的选取逻辑
min_support设多大?我一般按这个逻辑来:先看商品种类数。如果只有几十种商品,0.05可能合适;如果几千种,0.01甚至0.005才够。原则是保证每个高频商品至少出现在几百条事务里。
更实用的做法是:先设一个较大的阈值,比如0.1,看成生成几条规则。如果规则太少就往下调,每次降一半;如果规则太多(几千条没法看),往上调。我实测下来,1000条购物篮数据,min_support=0.05能跑出几十到上百条规则,这个规模比较适合人工分析。
6.2 怎么解读一条规则的质量
我收到过很多类似的提问:为什么我跑出的规则置信度很高,但实际业务里没用?
这里需要再次强调提升度的意义。置信度高只代表“买了A的人很多也买了B”,但如果B本身就是畅销品,这个规律就没有增量信息。比如上面例子里,如果面包支持度0.6,那“牛奶 → 面包”置信度0.4就没什么价值。真正好的规则应该是:A对B的拉动效果明显,且这种拉动不是因为B本身太火。
再看一个极端情况:置信度80%、提升度1.05的规则,和置信度35%、提升度3.2的规则,哪个更有价值?我的答案是后者。因为它的“发现价值”更大,能带来业务上的新洞察。
6.3 结果可视化,让发现更直观
我一直认为,Python数据分析离不开可视化。关联规则的结果用支持度、置信度、提升度三个维度画散点图非常直观:
import matplotlib.pyplot as plt plt.figure(figsize=(10, 6)) scatter = plt.scatter( result_df['support'], result_df['confidence'], c=result_df['lift'], cmap='viridis', s=result_df['support'] * 800, alpha=0.7 ) plt.colorbar(scatter, label='Lift') plt.xlabel('Support') plt.ylabel('Confidence') plt.title('Association Rules: Support-Confidence-Lift Relationship') plt.show()这个图我经常画。右上方是高支持度高置信度的强规则,颜色偏亮的代表提升度更高,气泡的大小也代表支持度尺度。一眼就能看出哪些规则值得深入挖掘。
7. 性能优化与扩展,这个算法比你想象的慢
Apriori有一个著名的短板:数据量大时速度很慢。它需要反复扫描数据集,每生成一层候选集就全表扫一次。这里我分享几个在真实场景下验证过的优化手段。
7.1 事务压缩与哈希树
频繁项集之外的项,对后续迭代没有任何贡献。比如“啤酒”不是频繁项,那么所有包含啤酒的事务,在下一轮扫描中都可以把啤酒直接删掉,事务变短了,子集判断也就快了。这就是事务压缩的思想。
另外可以用哈希树来存储候选集,判断某个事务包含哪些候选集时,不用每个候选集都做issubset,而是按哈希路径快速定位。Python里可以直接用字典或集合,配合frozenset的切片索引来做类似效果,但代码会复杂不少。
7.2 用mlxtend一行替代手写逻辑
如果你的环境允许装第三方库,mlxtend里的apriori和association_rules封装得相当好:
from mlxtend.frequent_patterns import apriori, association_rules # 要先转成one-hot编码格式 df_encoded = pd.get_dummies(df_transactions['items'].str.get_dummies(',')) frequent = apriori(df_encoded, min_support=0.05, use_colnames=True) rules = association_rules(frequent, metric='lift', min_threshold=1.1) print(rules[['antecedents', 'consequents', 'support', 'confidence', 'lift']])mlxtend内部用了优化过的算法,速度比手写版本快不少。但我还是建议你先跑通手写版,理解每一步在做什么,再切换到现成库。因为实际项目中,你往往需要修改算法逻辑(比如加入商品利润权重),那时候只能自己动手。
7.3 数据量大到内存装不下怎么办
如果事务数量达到百万级,Apriori这类基于内存的算法就不太现实了。两个思路:
一是用FP-Growth算法,它只需要扫描两次数据集,而且不需要生成候选集,速度上比Apriori快一个量级。Python里pyfim库有实现。
二是用Spark的FP-Growth或者分布式关联规则挖掘。大数据量场景下,这是工程上更稳妥的路径。如果只是学习实验,1000到10000条事务,手写Apriori完全够用。
我自己在实际项目中用过很久的Apriori,最大的体会是:这个算法入门不难,但真正用到业务里,数据清洗和参数调优花的精力远超算法本身。遇到“结果看起来挺好但业务不买账”的时候,先别急着改代码,回去看看数据里有没有隐藏的噪声,再检查一遍提升度是不是都大于1。另外一个小技巧,导出CSV时记得用utf-8-sig编码,这样在Windows的Excel里打开才不会乱码,我用的是中文商品名,这个坑踩过不止一次。
本文还有配套的精品资源,点击获取