Apriori算法输出格式化问题:提取frozenset元素并区分关联对与三元组
解决Apriori关联规则格式化输出与元素提取问题
我看你在Apriori算法的规则输出上遇到了两个核心痛点:没法按置信度降序筛选出前5个二元组、三元组关联规则,还有frozenset里的元素没法整理成易读格式。我帮你调整代码,直接搞定这两个问题。
问题分析
原代码的主要问题出在规则类型判断逻辑和结果格式化上:
- 之前通过
len(rules[i][1])判断二元/三元组,会漏掉{B}→{A}这种标准二元组规则,还会把不同形式的三元组规则混在一起 - 没有把frozenset转换成直观的字符串格式,输出会显得杂乱无章
修改后的完整代码
from numpy import * import pandas as pd from operator import itemgetter def loadDataSet(data=None): return pd.read_csv(data, sep = ' ', error_bad_lines=False) def createCandidateSet(data): C1 = [] for transaction in data: for item in transaction: if not [item] in C1: C1.append([item]) C1.sort() return list(map(frozenset, C1)) def scanData(dataset, Ck, support): ssCount = {} for tID in dataset: for candidate in Ck: if candidate.issubset(tID): if not candidate in ssCount: ssCount[candidate] = 1 else: ssCount[candidate]+=1 res = [] supportData ={} for key in ssCount: currSupport = ssCount[key] if currSupport >= support: res.insert(0, key) supportData[key] = currSupport return res, supportData def aprioriHelper(Lk, k): res = [] freqItemLen = len(Lk) for i in range(freqItemLen): for j in range(i+1, freqItemLen): L1 = list(Lk[i])[:k-2] L2 = list(Lk[j])[:k-2] L1.sort() L2.sort() if L1 == L2: res.append(Lk[i] | Lk[j]) return res def apriori(dataset, minSupport=100): C1 = createCandidateSet(dataset) D = list(map(set, dataset)) L1, supportData = scanData(D, C1, minSupport) L = [L1] k = 2 while (len(L[k-2]) > 0): Ck = aprioriHelper(L[k-2], k) Lk, supportK = scanData(D, Ck, minSupport) supportData.update(supportK) L.append(Lk) k+=1 return L, supportData def generateRules(L, supportData, conf = 0.7): rules = [] for i in range(1, len(L)): for freq in L[i]: association = [frozenset([item]) for item in freq] if i > 1: rulesFromConsequences(freq, association, supportData, rules, conf) else: calculateConfidence(freq, association, supportData, rules, conf) return rules def calculateConfidence(freq, association, supportData, rules, conf=0.7): filteredAssociations = [] for consequence in association: confidence = supportData[freq]/supportData[freq - consequence] if confidence >= conf: rules.append((freq-consequence, consequence, confidence)) filteredAssociations.append(consequence) return filteredAssociations def rulesFromConsequences(freq, association, supportData, rules, conf=0.7): a_len = len(association[0]) if (len(freq) > (a_len+1)): association_p1 = aprioriHelper(association, a_len+1) association_p1 = calculateConfidence(freq, association_p1, supportData, rules, conf) if len(association_p1) > 1: rulesFromConsequences(freq, association_p1, supportData, rules, conf) # 新增:把frozenset转换成易读的字符串格式 def format_frozenset(fs): return ", ".join(sorted(fs)) def main(): dataset = [line.split() for line in open('datatest.txt')] L, supportData = apriori(dataset, minSupport=8) rules = generateRules(L, supportData, conf=0) # 按置信度降序排序规则 rules_sorted = sorted(rules, key=itemgetter(2), reverse=True) # 分离二元组规则(总元素数为2:前件+后件的元素总数) double_rules = [] # 分离三元组规则(总元素数为3:前件+后件的元素总数) triple_rules = [] for rule in rules_sorted: antecedent, consequent, confidence = rule total_items = len(antecedent) + len(consequent) if total_items == 2 and len(double_rules) <5: double_rules.append(rule) elif total_items ==3 and len(triple_rules) <5: triple_rules.append(rule) # 两个列表都凑够5个就提前退出,提升效率 if len(double_rules)>=5 and len(triple_rules)>=5: break # 格式化输出OUTPUT A print("=== OUTPUT A: 前5个频繁关联二元组(置信度降序) ===") for idx, rule in enumerate(double_rules, 1): ant, cons, conf = rule print(f"{idx}. {{{format_frozenset(ant)}}} → {{{format_frozenset(cons)}}} 置信度: {conf:.4f}") # 格式化输出OUTPUT B print("\n=== OUTPUT B: 前5个频繁关联三元组(置信度降序) ===") for idx, rule in enumerate(triple_rules, 1): ant, cons, conf = rule print(f"{idx}. {{{format_frozenset(ant)}}} → {{{format_frozenset(cons)}}} 置信度: {conf:.4f}") if __name__ == '__main__': main()
关键修改说明
- 新增格式化函数:
format_frozenset把frozenset里的元素排序后拼接成逗号分隔的字符串,比如frozenset({'milk','bread'})会变成bread, milk,输出更整洁 - 修正规则类型判断:通过
len(antecedent) + len(consequent)判断规则涉及的总元素数,二元组总元素数为2,三元组为3,不会漏判任何情况 - 优化筛选逻辑:遍历排序后的规则时,一旦两个列表都凑够5个就提前退出,提升效率;同时自然处理规则数量不足5个的边界情况(比如数据集里只有3个有效二元组规则,就只输出3个)
- 清晰的输出结构:分OUTPUT A和OUTPUT B输出,每个规则带序号、格式化的前后件和保留4位小数的置信度
内容的提问来源于stack exchange,提问作者Srikar Murali
相关产品推荐
相关产品推荐

