使用PySpark计算CSV文件中相同值两次出现间的不同值数量
Hey there! Let's work through this problem together—handling 50GB+ CSV files in PySpark definitely requires efficient window operations, so let's break down the solution step by step to get your desired output.
Understanding the Requirement
First, let's clarify what we need: for each occurrence of a value in col1, we want to count the number of distinct values that appear between its previous occurrence and the current row (excluding both the previous and current occurrence itself). For rows where the value appears for the first time, we return null.
Step-by-Step Solution
Here's how to implement this in PySpark, with performance considerations for large datasets:
1. Add a Continuous Row ID
Spark DataFrames are unordered by default, so we first need to assign a unique, continuous row ID to preserve the original order of your CSV data. We'll use row_number() with a window ordered by monotonically_increasing_id() (this ensures we maintain the input order):
from pyspark.sql import functions as F from pyspark.sql.window import Window # Assume your original DataFrame is named `raw_df` window_row_order = Window.orderBy(F.monotonically_increasing_id()) df = raw_df.withColumn("row_id", F.row_number().over(window_row_order))
2. Track the Previous Occurrence of Each Value
Next, we'll use a window partitioned by col1 to get the row ID of the last time each value appeared:
window_prev_occurrence = Window.partitionBy("col1").orderBy("row_id") df = df.withColumn("prev_row_id", F.lag("row_id").over(window_prev_occurrence))
3. Calculate Distinct Values Between Occurrences
Now we need to count distinct values in the range between prev_row_id + 1 and row_id - 1. We'll use a global window to collect all prior rows, filter for the relevant range, extract the values, deduplicate, and count:
# Global window to collect all rows up to the current one window_global = Window.orderBy("row_id") # Filter rows in the target range, extract `col1` values, deduplicate, and count df = df.withColumn( "col2", F.when( F.col("prev_row_id").isNull(), F.lit(None) # First occurrence, return null ).otherwise( F.size( F.array_distinct( F.transform( # Filter rows that fall between the previous occurrence and current row F.filter( F.collect_list(F.struct("row_id", "col1")).over(window_global), lambda x: (x["row_id"] > F.col("prev_row_id")) & (x["row_id"] < F.col("row_id")) ), # Extract just the `col1` value from the filtered rows lambda x: x["col1"] ) ) ) ) ) # Clean up temporary columns final_df = df.drop("row_id", "prev_row_id")
Verifying the Result
If you run this code on your sample input, you'll get exactly the output you're looking for:
+----+-----+ |col1|col2 | +----+-----+ | a| null| | b| null| | c| null| | c| 0| | a| 2| | b| 2| | a| 1| +----+-----+
Performance Tips for Large Datasets
Since you're working with a 50GB+ file, keep these in mind to avoid memory issues:
- Adjust Spark Configurations: Increase executor memory (
spark.executor.memory) and executor cores (spark.executor.cores) to handle large window operations. - Partition Your Data: If your CSV has a natural partition key (e.g., a date column), partition the DataFrame upfront to reduce the data each executor needs to process.
- Avoid Shuffles: The window operations here minimize shuffles compared to join-based approaches, but monitor the Spark UI for shuffle stages and optimize if needed.
内容的提问来源于stack exchange,提问作者GokulaKannan

