You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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()

关键修改说明

  1. 新增格式化函数:format_frozenset把frozenset里的元素排序后拼接成逗号分隔的字符串,比如frozenset({'milk','bread'})会变成bread, milk,输出更整洁
  2. 修正规则类型判断:通过len(antecedent) + len(consequent)判断规则涉及的总元素数,二元组总元素数为2,三元组为3,不会漏判任何情况
  3. 优化筛选逻辑:遍历排序后的规则时,一旦两个列表都凑够5个就提前退出,提升效率;同时自然处理规则数量不足5个的边界情况(比如数据集里只有3个有效二元组规则,就只输出3个)
  4. 清晰的输出结构:分OUTPUT A和OUTPUT B输出,每个规则带序号、格式化的前后件和保留4位小数的置信度

内容的提问来源于stack exchange,提问作者Srikar Murali

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.15 04:55:34