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

如何在PySpark 1.5.1中生成带数量系数的关联规则

Optimizing Quantity Coefficient Calculation for Association Rules in PySpark 1.5.1

Looking at your code, the biggest performance bottleneck is definitely the qty_coef function being called in a loop after collecting association rules to the Driver. This leads to N+1 SQL queries (one per rule pair), which is extremely inefficient—each query triggers a full scan of your table, and processing happens serially on the Driver instead of distributed across Executors. Let's fix this with a distributed approach that precomputes all necessary quantity statistics in one pass.

Key Issues in Original Code

  • Serial Processing: Collecting asso_list to the Driver and looping through it forces all coefficient calculations to run on a single machine, wasting Spark's distributed capabilities.
  • Redundant Queries: Each qty_coef call runs nested SQL queries that scan the same table multiple times, adding massive overhead.
  • Unnecessary Data Movement: Shuffling data back to the Driver for small, repeated operations kills performance for large datasets.

Optimized Approach

Instead of calculating coefficients one pair at a time, we'll precompute all required quantity aggregates for frequent item pairs in a distributed way, then join these aggregates with your association rules to get the final coefficients. Here's how to implement it:

Step 1: Precompute Coefficient Aggregates for Frequent Pairs

First, we'll generate all valid item pairs from transactions that contain both items, then calculate the total quantity of each item across those shared transactions. This avoids repeated table scans.

Step 2: Join with Association Rules

Once we have the precomputed quantity aggregates, we'll join them with your confidence-filtered association rules to calculate the coefficient in a distributed manner.

Full Optimized Code

#!/usr/bin/python
from pyspark import SparkContext, HiveContext
from pyspark.mllib.fpm import FPGrowth
import time

def read_data():
    sql="""select t.orderno_nosplit, t.prod_code, t.item_code, sum(t.item_qty) as item_qty 
           from ioc_fdm.fdm_dwr_ioc_fcs_pk_spu_item_f_chain t 
           group by t.prod_code, t.orderno_nosplit, t.item_code """
    data = sql_context.sql(sql)
    return data.cache()

def train(prod):
    spu = total_spu.filter(total_spu.prod_code == prod)
    print 'data length', spu.count(), time.strftime("%H:%M:%S")
    supp = 0.1
    conf = 0.7
    
    # Register table and cache
    sql_context.registerDataFrameAsTable(spu, 'spu_table')
    sql_context.cacheTable('spu_table')
    print 'table register over', time.strftime("%H:%M:%S")
    
    # Prepare transaction sets for FP-Growth
    trans_sets = spu.rdd.repartition(32) \
                       .map(lambda x: (x[0], x[2])) \
                       .groupByKey() \
                       .mapValues(list) \
                       .values() \
                       .cache()
    print 'trans group over', time.strftime("%H:%M:%S")
    
    # Train FP-Growth model
    model = FPGrowth.train(trans_sets, supp, 10)
    print 'model train over', time.strftime("%H:%M:%S")
    
    # Extract frequent itemsets of size 1 and 2
    model_f1 = model.freqItemsets().filter(lambda x: len(x[0]) == 1)
    model_f2 = model.freqItemsets().filter(lambda x: len(x[0]) == 2)
    
    # Create broadcast map for single-item frequencies (for confidence calculation)
    model_f1_tuple = model_f1.map(lambda (items, freq): (items[0], freq))
    bc_single_freq = sc.broadcast(model_f1_tuple.collectAsMap())
    
    # Generate association rules with confidence (both directions: A->B and B->A)
    def generate_rules(itemset, freq):
        item_a, item_b = itemset
        conf_a_to_b = float(freq) / bc_single_freq.value[item_a]
        conf_b_to_a = float(freq) / bc_single_freq.value[item_b]
        return [(item_a, item_b, conf_a_to_b), (item_b, item_a, conf_b_to_a)]
    
    rules_rdd = model_f2.flatMap(lambda (items, freq): generate_rules(items, freq))
    print 'rule generation over', time.strftime("%H:%M:%S")
    
    # Filter rules by confidence threshold
    filtered_rules = rules_rdd.filter(lambda x: x[2] >= conf)
    print 'rule filtering over', time.strftime("%H:%M:%S")
    
    # ---------------------- Optimized Coefficient Calculation ----------------------
    # Precompute total quantities for item pairs in shared transactions
    # 1. Join spu table with itself on transaction ID to get item pairs in the same transaction
    pair_qty_sql = """
        SELECT t1.item_code as item1, t2.item_code as item2,
               SUM(t1.item_qty) as total_qty1, SUM(t2.item_qty) as total_qty2
        FROM spu_table t1
        JOIN spu_table t2 ON t1.orderno_nosplit = t2.orderno_nosplit
        WHERE t1.item_code != t2.item_code
        GROUP BY t1.item_code, t2.item_code
    """
    pair_qty_df = sql_context.sql(pair_qty_sql).cache()
    print 'pair quantity precomputed', time.strftime("%H:%M:%S")
    
    # Convert pair quantity data to RDD for joining with rules
    pair_qty_rdd = pair_qty_df.rdd.map(lambda x: ((x[0], x[1]), (float(x[3])/x[2])))
    
    # Join filtered rules with precomputed coefficients
    final_rdd = filtered_rules.map(lambda x: ((x[0], x[1]), (x[0], x[1], x[2]))) \
                              .join(pair_qty_rdd) \
                              .map(lambda x: (x[1][0][0], x[1][0][1], x[1][0][2], x[1][1]))
    
    # Convert to DataFrame
    asso_df = sql_context.createDataFrame(final_rdd, ['item1', 'item2', 'conf', 'coef'])
    print 'final dataframe created', time.strftime("%H:%M:%S")
    # ---------------------- End Optimized Coefficient Calculation ----------------------
    
    # Save results
    path = "hdfs:/user/hive/wilber/%s" % (prod)
    asso_df.write.mode('overwrite').parquet(path)
    
    # Cleanup cache
    sql_context.clearCache()

if __name__ == '__main__':
    sc = SparkContext()
    sql_context = HiveContext(sc)
    prod_list = sc.textFile('hdfs:/user/hive/wilber/prod_list').collect()
    total_spu = read_data()
    print 'spu read over', time.strftime("%H:%M:%S")
    for prod in prod_list:
        print 'processing prod', prod
        train(prod)

Why This Is Faster

  • Single Pass Aggregation: The pair_qty_sql query computes all item pair quantity totals in one scan of the table, instead of N separate queries.
  • Distributed Processing: All coefficient calculations happen on Executors, not the Driver—no more serial loops over collected data.
  • Reduced Data Movement: We only shuffle necessary data (item pairs and their aggregates) instead of moving entire transaction datasets multiple times.

Additional Tips for PySpark 1.5.1

  • Adjust Repartitioning: The repartition(32) value should be tuned based on your cluster's core count—aim for 2-4 partitions per core.
  • Cache Strategically: We cache intermediate DataFrames/RDDs that are reused (like pair_qty_df and trans_sets) to avoid recomputing expensive operations.
  • Avoid Nested SQL: The original qty_coef function used nested subqueries, which are poorly optimized in older PySpark versions. The join approach is much more efficient.

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.27 10:03:53