PySpark DataFrame去重:移除id下互为逆序的value配对
Hey there! Your Pandas solution for eliminating symmetric pairs is clever, and we can translate that logic smoothly to PySpark using its built-in functions. Let's break down how to solve this:
Core Idea
Just like your Pandas approach, we'll create a standardized key for each pair of value.x and value.y—this key will be identical for both the original pair and its reverse. Then, we can deduplicate rows based on id and this key, keeping only one instance per unique (id, standardized key) combination.
Implementation Options
Option 1: Using array_sort (Spark 2.4+, More Elegant)
Spark 2.4 introduced array_sort, which simplifies creating the standardized key without messy string concatenation:
from pyspark.sql import SparkSession from pyspark.sql.functions import array, array_sort, concat_ws # Initialize Spark session spark = SparkSession.builder.appName("RemoveSymmetricPairs").getOrCreate() # Sample input DataFrame data = [ (100057, 38953993985, 38993095846), (100057, 38993095845, 38953993985), (100057, 38993095845, 38993095846), (100057, 38993095846, 38953993985), (100011, 38989281716, 38996868028), (100011, 38996868028, 38989281716), (100019, 38916115350, 38994231881), (100019, 38994231881, 38916115350) ] df = spark.createDataFrame(data, ["id", "value.x", "value.y"]) # Create a standardized key: sorted array of value.x and value.y, converted to string df_with_key = df.withColumn( "standardized_pair", concat_ws("_", array_sort(array("value.x", "value.y"))) ) # Deduplicate by id and standardized_pair, keep one instance per group result = df_with_key.dropDuplicates(["id", "standardized_pair"]).drop("standardized_pair") # Show the sorted result result.orderBy("id").show(truncate=False)
Option 2: String Concatenation (Mirror Your Pandas Logic)
If you want to stick closely to your original Pandas code, use when/otherwise to compare values and build the key:
from pyspark.sql.functions import when, concat, col df_with_key = df.withColumn( "standardized_pair", when(col("value.x") > col("value.y"), concat(col("value.x"), col("value.y"))) .otherwise(concat(col("value.y"), col("value.x"))) ) result = df_with_key.dropDuplicates(["id", "standardized_pair"]).drop("standardized_pair") result.orderBy("id").show(truncate=False)
Expected Output
Either option will produce your desired result:
+-------+-------------+-------------+ |id |value.x |value.y | +-------+-------------+-------------+ |100011 |38989281716 |38996868028 | |100019 |38916115350 |38994231881 | |100057 |38953993985 |38993095846 | |100057 |38993095845 |38953993985 | |100057 |38993095845 |38993095846 | +-------+-------------+-------------+
Bonus: Retain Specific Rows
If you need to keep the first occurrence of each pair (instead of a random row from dropDuplicates), use groupBy with first() aggregation:
from pyspark.sql.functions import first result = df_with_key.groupBy("id", "standardized_pair") \ .agg( first("value.x").alias("value.x"), first("value.y").alias("value.y") ) \ .drop("standardized_pair")
内容的提问来源于stack exchange,提问作者Tilo

