
3天手写实现关联规则算法,告别复制代码跑不通的坑
刚拿到一段 Apriori 算法的代码,信心满满地粘贴到 PyCharm 里,点运行。结果?报错信息像天书一样,或者更糟糕——程序跑完了,输出的结果全是乱码,支持度置信度根本对不上。你是不是也经历过这种“复制粘贴式学习”的绝望?代码看起来眼熟,变量名也懂,但就是跑不通,改哪都错。
别急,这往往不是环境的问题,而是你根本没搞懂关联规则算法底层的数据流转逻辑。很多教程只给结果,不给过程,导致你像个盲人摸象。今天咱们不整虚的,直接上手手写实现核心逻辑。哪怕你之前只学过基础 Python,只要跟着这篇教程,把数据从“购物小票”变成“规则推荐”的过程拆开揉碎看,保证你彻底明白其中的门道。
概念速懂:关联规则到底在算什么?
在房建工程里,我们常说要“材料配套”,比如买了水泥就得配砂石。在前端开发或电商推荐场景下,关联规则算法(Association Rule Learning)干的就是这事:发现用户行为中隐藏的模式。
最经典的例子是“啤酒与尿布”。超市发现,买尿布的爸爸,经常会顺手买啤酒。这背后有三个核心指标,不懂这三个词,代码写得再溜也是白搭:
支持度 (Support):这个规则有多普遍?比如“买啤酒”的人占所有买过东西的人的 10%,那支持度就是 0.1。
置信度 (Confidence):买了啤酒的人里,有多大比例也会买尿布?如果 100 个买啤酒的人里有 80 个买了尿布,置信度就是 0.8。
提升度 (Lift):这个关联是不是巧合?如果 Lift 1,说明两者正相关,越买越有;如果 Lift ≈ 1,说明两者独立,没啥关系。
关键点:我们要找的是那些支持度和置信度都超过设定阈值的“频繁项集”。Apriori 算法是解决这个问题的经典方案,它的核心思想叫“向下封闭性”——如果一个项集不频繁,它的所有超集一定也不频繁。利用这一点,我们可以剪枝,大幅减少计算量。
环境准备:极简依赖,拒绝玄学报错
很多初学者第一步就卡在环境配置上。为了排除干扰,我们尽量使用原生 Python 库,不依赖复杂的第三方包。
你需要准备:
Python 3.8+:建议用 Anaconda 管理环境,避免版本冲突。
标准库:collections(用于 Counter 统计)、itertools(用于生成组合)。
不需要安装 mlxtend 或 scikit-learn,因为我们要手写实现核心逻辑。依赖越少,出问题的概率越低,而且你能看清每一行代码在干嘛。
打开你的终端,确认 Python 版本:
python --version
如果输出正常,新建一个 apriori_demo.py 文件。我们要从零开始构建数据结构,而不是直接调用黑盒函数。
核心语法:拆解 Apriori 的骨架
Apriori 算法分两步走:生成候选项集 - 剪枝。
1. 生成频繁 1-项集
这是最基础的一步。扫描所有事务(Transaction),统计每个物品出现的次数。
这里有个常见的坑:数据预处理。原始数据通常是列表的列表,比如 [[A, B, C], [A, C]]。我们需要把它转换成集合(Set)或者 Counter,方便快速查找。
from collections import Counter
def get_frequent_1_itemsets(transactions, min_support):
生成频繁1-项集
:param transactions: 原始交易数据,列表的列表
:param min_support: 最小支持度阈值
:return: 字典 {item: support_count}
item_count = Counter()
total_transactions = len(transactions)
# 遍历每一笔交易,累加物品计数
for t in transactions:
# 去重!同一笔交易里买两次A,只算一次
for item in set(t):
item_count[item] += 1
# 过滤出满足最小支持度的物品
frequent_1 = {}
for item, count in item_count.items():
support = count / total_transactions
if support = min_support:
# 存储格式:(物品, 支持度)
frequent_1[item] = support
return frequent_1
注意:这里用了 set(t)。如果一笔交易是 [A, A, B],不转 set 的话,A 会被计数两次,导致支持度虚高。这是新手最容易忽略的细节。
2. 生成候选 k-项集 (k 1)
这是算法最复杂的部分。我们需要从上一轮的频繁 (k-1)-项集,组合出新的 k-项集,并判断它们是否频繁。
假设我们要找 2-项集。我们从频繁 1-项集 {A, B, C} 中两两组合:{A,B}, {A,C}, {B,C}。
然后扫描原始数据,看这些组合出现了多少次。
但如果是 3-项集呢?直接组合会爆炸。Apriori 的精髓在于连接步骤和剪枝步骤。
连接:如果 L2 中有 {A,B} 和 {A,C},且第一个元素相同(都是 A),则可以连接成 {A,B,C}。
剪枝:检查生成的候选集 {A,B,C} 的所有 (k-1) 子集(即 {A,B}, {A,C}, {B,C})是否都在 L2 中。如果有一个不在,直接丢弃。
import itertools
def apriori_generate_candidates(frequent_k_minus_1, k):
生成候选k-项集
:param frequent_k_minus_1: 频繁(k-1)-项集的列表,元素是tuple
:param k: 目标项集大小
:return: 候选k-项集的列表
candidates = set()
# 将频繁项集转为列表,便于索引
freq_list = list(frequent_k_minus_1)
# 双重循环连接
for i in range(len(freq_list)):
for j in range(i + 1, len(freq_list)):
# 取前 k-2 个元素进行比较
# 例如 k=3, 比较前 1 个元素
prefix = freq_list[i][:k-2]
if prefix == freq_list[j][:k-2]:
# 合并
candidate = tuple(sorted(set(freq_list[i]) | set(freq_list[j])))
# 剪枝:检查 candidate 的所有 (k-1) 子集是否频繁
is_frequent = True
for subset in itertools.combinations(candidate, k-1):
if subset not in freq_list:
is_frequent = False
break
if is_frequent:
candidates.add(candidate)
return list(candidates)
这段代码逻辑很密,建议对着注释一步步走。特别是 itertools.combinations 的使用,它能高效生成所有子集,避免手写递归的麻烦。
完整代码示例:跑通一个完整流程
光看片段不够,我们把所有逻辑串起来,写一个完整的 Apriori 类。为了方便演示,我们构造一份模拟的“工地采购数据”。
场景:某工地采购部记录了 100 次采购行为。
数据特征:水泥和砂石经常一起买,电线和开关偶尔一起买。
class Apriori:
def __init__(self, min_support=0.3, min_confidence=0.5):
self.min_support = min_support
self.min_confidence = min_confidence
self.frequent_itemsets = {} # 存储所有频繁项集及其支持度
def fit(self, transactions):
total = len(transactions)
# 1. 生成频繁1-项集
freq_1 = self._get_freq_1(transactions, total)
self.frequent_itemsets.update({(k,): v for k, v in freq_1.items()})
current_freq = list(freq_1.keys())
k = 2
# 2. 迭代生成 k-项集
while current_freq:
candidates = self._generate_candidates(current_freq, k)
if not candidates:
break
new_freq = {}
# 计算候选项集的支持度
for cand in candidates:
count = 0
for t in transactions:
if set(cand).issubset(set(t)):
count += 1
support = count / total
if support = self.min_support:
new_freq[cand] = support
if new_freq:
self.frequent_itemsets.update(new_freq)
current_freq = list(new_freq.keys())
k += 1
else:
break
def _get_freq_1(self, transactions, total):
counts = Counter()
for t in transactions:
for item in set(t):
counts[item] += 1
return {item: count/total for item, count in counts.items() if count/total = self.min_support}
def _generate_candidates(self, prev_freq, k):
# 简化版生成逻辑,实际项目中需优化性能
prev_set = set(prev_freq)
candidates = set()
prev_list = sorted(prev_set)
for i in range(len(prev_list)):
for j in range(i+1, len(prev_list)):
# 检查前 k-2 个元素是否一致
if prev_list[i][:k-2] == prev_list[j][:k-2]:
cand = tuple(sorted(set(prev_list[i]) | set(prev_list[j])))
# 剪枝
valid = True
for sub in itertools.combinations(cand, k-1):
if sub not in prev_set:
valid = False
break
if valid:
candidates.add(cand)
return list(candidates)
def generate_rules(self):
rules = []
for itemset, support in self.frequent_itemsets.items():
if len(itemset) 2:
continue
# 生成规则:A - B, B - A ...
for i in range(len(itemset)):
antecedent = tuple(sorted(itemset[:i] + itemset[i+1:]))
consequent = itemset[i]
# 查找前件的支持度
ant_support = self.frequent_itemsets.get(antecedent, 0)
if ant_support == 0:
continue
confidence = support / ant_support
if confidence = self.min_confidence:
rules.append({
'antecedent': antecedent,
'consequent': consequent,
'support': support,
'confidence': confidence
})
return rules
# --- 运行测试 ---
if __name__ == '__main__':
# 模拟数据:
# 水泥(Cement) 和 砂石(Aggregate) 强关联
# 电线(Wire) 和 开关(Switch) 弱关联
# 砖块(Brick) 独立出现
data = [
['Cement', 'Aggregate', 'Brick'],
['Cement', 'Aggregate'],
['Cement', 'Aggregate', 'Wire'],
['Cement', 'Brick'],
['Aggregate', 'Wire', 'Switch'],
['Aggregate', 'Brick'],
['Cement', 'Aggregate', 'Brick', 'Wire'],
['Cement', 'Aggregate'],
['Aggregate', 'Wire'],
['Cement', 'Brick']
]
apriori = Apriori(min_support=0.4, min_confidence=0.5)
apriori.fit(data)
print(=== 频繁项集 ===)
for itemset, sup in apriori.frequent_itemsets.items():
print(f{itemset}: {sup:.2f})
print(\n=== 关联规则 ===)
for rule in apriori.generate_rules():
print(f{rule['antecedent']} - {rule['consequent']} | Conf: {rule['confidence']:.2f})
运行这段代码,你会看到 ('Cement', 'Aggregate') 的支持度很高,且能生成 Cement - Aggregate 的规则。如果阈值调低,还能看到 Aggregate - Wire 的规则。
调试技巧:如果在 _generate_candidates 里卡住,建议在 candidates.add(cand) 前打印 cand 和 prev_list 的相关部分。很多时候,剪枝逻辑里的 subset not in prev_set 会因为元组顺序不一致而失效,务必确保 tuple(sorted(...)) 的一致性。
常见报错与避坑指南
KeyError: 'Cement'
原因:在计算置信度时,去查找前件的支持度,但前件可能不在 frequent_itemsets 里。
解决:使用 self.frequent_itemsets.get(antecedent, 0),默认值为 0,避免崩溃。
结果为空
原因:min_support 设得太高。
解决:先跑一遍 min_support=0.1,看看有哪些频繁项集,再逐步调整阈值。不要盲目追求高支持度,否则什么都挖不出来。
内存溢出 (MemoryError)
原因:数据量太大,候选集爆炸。
解决:Apriori 不适合超大数据集。如果数据量超过 10 万条,考虑使用 FP-Growth 算法,它构建 FP-Tree,效率远高于 Apriori。但在入门阶段,理解 Apriori 的逻辑更重要。
小结
手写一遍关联规则算法,不是为了替代库函数,而是为了建立对数据结构的直觉。当你明白了支持度、置信度是怎么从原始数据中“数”出来的,你再去看 mlxtend 的官方文档,或者在项目中集成推荐系统时,心里就有底了。
记住,代码跑不通,往往是因为你对数据流的假设错了。下次遇到报错,别急着换库,先打印中间变量,看看数据长什么样。
你在项目里踩过这个坑吗?比如数据预处理时的去重问题,或者阈值设置的纠结?评论区聊聊,咱们一起避坑。