如何高效生成带权重的无重复DataFrame代表性样本(Pandas/PySpark)
Great question—this is such a common pain point when working with huge aggregated datasets, where expanding rows to match their total counts is totally impractical. Let's walk through efficient solutions for both Pandas (for in-memory use cases) and PySpark (when you're dealing with billion-row distributed data).
Pandas Solution (In-Memory Data)
When your aggregated query | total DataFrame fits in memory, you can leverage Pandas' built-in sampling capabilities to avoid expanding rows entirely. We'll use weighted sampling with no replacement, where each query's selection probability is directly tied to its total value.
Example Code:
import pandas as pd # Sample aggregated DataFrame (matches your 8-query, total=20 example) data = { "query": ["q1", "q2", "q3", "q4", "q5", "q6", "q7", "q8"], "total": [3, 5, 2, 4, 1, 2, 2, 1] } df = pd.DataFrame(data) # Number of unique samples to select sample_size = 5 # Perform weighted, non-replacement sampling sampled_queries = df.sample( n=sample_size, weights="total", replace=False, random_state=42 # Optional, for reproducible results )["query"].tolist() print(sampled_queries)
Key Details:
- Efficiency: Unlike the naive approach of expanding each query into
totalrows (which would blow up memory for large totals), this uses the aggregated weights directly. Time complexity is O(n) where n is the number of unique queries, not the raw row count. - Probability Alignment: The
weights="total"parameter ensures each query's selection chance is proportional to its total occurrence count. - Edge Case Handling: If your requested
sample_sizeexceeds the number of unique queries, Pandas will automatically return all queries without throwing errors.
PySpark Solution (Distributed, Billion-Row Scale)
For truly massive datasets where even the aggregated query-total table is too large for a single machine's memory, PySpark's distributed processing is the way to go. We'll use A-Chao's weighted reservoir sampling algorithm, which efficiently selects k unique elements with probabilities proportional to their weights—no row expansion required.
Example Code:
from pyspark.sql import SparkSession from pyspark.sql import functions as F # Initialize Spark session (tune config for your cluster as needed) spark = SparkSession.builder.appName("WeightedQuerySampling").getOrCreate() # Sample aggregated DataFrame (scaled for distributed workflows) data = [("q1", 3), ("q2", 5), ("q3", 2), ("q4", 4), ("q5", 1), ("q6", 2), ("q7", 2), ("q8", 1)] df = spark.createDataFrame(data, schema=["query", "total"]) # Number of unique samples to select sample_size = 5 # Calculate priority scores (core of A-Chao's algorithm) # Queries with higher total values have lower expected scores, so we sort ascending to pick top k df_with_scores = df.withColumn("random_num", F.rand()) \ .withColumn("priority_score", -F.log(F.col("random_num")) / F.col("total")) # Select top k unique queries sampled_df = df_with_scores.orderBy("priority_score").limit(sample_size).select("query") # View results (avoid collecting large samples to the driver) sampled_df.show()
Key Details:
- Distributed Efficiency: This runs entirely across your Spark cluster—each worker processes its partition of the query-total data, and the final sort/limit uses Spark's optimized shuffle logic. No need to pull all data to a single node.
- Probability Alignment: The
priority_scorecalculation ensures queries with highertotalvalues have a statistically higher chance of being selected, while enforcing no duplicates (since we're selecting unique rows). - Edge Case Handling: Add a quick check to handle cases where you request more samples than there are unique queries:
unique_query_count = df.select("query").distinct().count() if sample_size >= unique_query_count: sampled_df = df.select("query").distinct() else: # Run the sampling logic above pass
内容的提问来源于stack exchange,提问作者Leemosh

