PySpark新手求助:优化100GB数据集单列去重计数效率
Hey there! Let's work through this slow distinct count issue you're hitting with your 100GB, 400-column dataset. I've dealt with similar large-scale PySpark tasks before, so here are some practical, actionable fixes to speed things up and get you the correct per-column results:
First: Fix the DataFrame Aggregation Misuse
Your initial df.agg(countDistinct('C1', 'C2', ..., 'CN')) call gives combined distinct counts because that's how the multi-column countDistinct works—it counts unique rows across all specified columns, not per-column uniqueness. To get individual column counts, you need to generate a separate countDistinct expression for each column:
from pyspark.sql.functions import countDistinct # Build aggregation expressions for every column agg_exprs = [countDistinct(col).alias(col) for col in df.columns] # Run the aggregation and convert results to a dictionary for easy access distinct_counts = df.agg(*agg_exprs).collect()[0].asDict()
This gives you the same result as your SQL query but in cleaner, more maintainable code.
Key Optimizations to Cut Runtime from 16 Hours
The main culprit here is the massive shuffle operation triggered by countDistinct—data has to be reorganized across nodes to find unique values. Here's how to tackle that:
1. Use Approximate Counts (If Precision Allows)
If your use case can tolerate a small margin of error (usually 1-5%), swap countDistinct for approx_count_distinct. This uses the HyperLogLog algorithm, which drastically reduces shuffle data and speeds up computation:
from pyspark.sql.functions import approx_count_distinct # You can adjust the relative standard error (rsd) for better precision (default is 0.05) agg_exprs = [approx_count_distinct(col, rsd=0.01).alias(col) for col in df.columns] distinct_counts = df.agg(*agg_exprs).collect()[0].asDict()
This alone can cut your runtime to a fraction of the original—often hours down to minutes for large datasets.
2. Tune Shuffle Partitions
Spark's default shuffle partition count (200) is way too low for a dataset this size with your cluster resources. Adjust it to match your total CPU cores (aim for 2-4 partitions per core):
# Your cluster has 8 worker nodes × 4 cores = 32 cores; 32 × 60 = 1920 partitions is a solid starting point spark.conf.set("spark.sql.shuffle.partitions", 1920)
This ensures data is split into manageable chunks, preventing single tasks from bogging down the cluster.
3. Cache Data to Leverage Local SSDs
Your workers have local SSDs—use them to cache the dataset so subsequent operations don't re-read from remote storage:
from pyspark.sql import StorageLevel # Cache with memory + disk serialization (uses less memory than raw storage) df = df.persist(StorageLevel.MEMORY_AND_DISK_SER) # Trigger caching by running a quick action like count() df.count()
Once cached, all aggregation operations will pull data from local SSDs instead of remote storage, drastically reducing IO time.
4. Process Columns in Batches
Trying to compute 400 distinct counts at once can overload individual tasks. Split columns into smaller batches (e.g., 20 columns per batch) to distribute the load:
batch_size = 20 column_batches = [df.columns[i:i+batch_size] for i in range(0, len(df.columns), batch_size)] all_counts = {} for batch in column_batches: batch_exprs = [countDistinct(col).alias(col) for col in batch] batch_results = df.agg(*batch_exprs).collect()[0].asDict() all_counts.update(batch_results)
This breaks the work into smaller, more manageable chunks that the cluster can process in parallel more efficiently.
5. Check for Data Skew
If some columns have extremely skewed data (e.g., one value makes up 90% of the column), that can cause a single shuffle task to run forever. You can check for skew with:
# Example: Check top 10 most frequent values in column C1 df.groupBy("C1").count().orderBy("count", ascending=False).show(10)
If you find skew, you can fix it by "salting" the skewed column—adding a random prefix to values to split them into multiple partitions, then aggregating twice (once with salt, once without). This is more complex, but worth it if skew is a bottleneck.
Optimizing the SQL Approach (If You Prefer SQL)
If you want to stick with SQL, apply the same optimizations:
- Use
approx_count_distinctinstead ofCOUNT(DISTINCT) - Adjust shuffle partitions before running your query:
SET spark.sql.shuffle.partitions = 1920; SELECT approx_count_distinct(C1) AS C1, approx_count_distinct(C2) AS C2, ... FROM myTable;
内容的提问来源于stack exchange,提问作者breakingduck

