如何在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_listto 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_coefcall 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_sqlquery 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_dfandtrans_sets) to avoid recomputing expensive operations. - Avoid Nested SQL: The original
qty_coeffunction used nested subqueries, which are poorly optimized in older PySpark versions. The join approach is much more efficient.
内容的提问来源于stack exchange,提问作者Wilber

