如何在Java中基于集合重叠过滤Apache Spark数据集的列表列
Solution: Filter Spark Dataset by Overlap with a Specified Set
To solve this problem, we need to filter rows in a Spark Dataset where a comma-separated string column contains at least one value from a given set of primary keys. Below are two robust approaches to implement this in Java:
Approach 1: Use Built-in Spark Functions (Recommended for Performance)
This method leverages Spark's optimized built-in functions to avoid custom UDFs, which is better suited for large datasets.
Step-by-Step Explanation:
- Split the String Column: Convert the comma-separated string into an array of values, handling optional spaces after commas with the regex
",\\s*". - Create a Literal Array of Primary Keys: Convert the primary keys set into a Spark array literal for comparison.
- Check for Intersection: Use
array_intersectto find common values between the split column and primary keys array, then check if the intersection has any elements.
Complete Test Code:
import org.apache.spark.sql.Dataset; import org.apache.spark.sql.Row; import org.apache.spark.sql.RowFactory; import org.apache.spark.sql.SparkSession; import org.apache.spark.sql.functions; import org.apache.spark.sql.types.DataTypes; import org.apache.spark.sql.types.Metadata; import org.apache.spark.sql.types.StructField; import org.apache.spark.sql.types.StructType; import org.junit.Assert; import org.junit.Test; import java.util.Arrays; import java.util.HashSet; import java.util.List; import java.util.Set; import java.util.stream.Collectors; public class SparkFilterTest { private final SparkSession spark = SparkSession.builder() .appName("FilterTest") .master("local[*]") .getOrCreate(); @Test public void testIfColumnHasMentionsInPrimaryKeys() { // Test data List<Row> data = Arrays.asList( RowFactory.create("ID, ID1"), RowFactory.create("ID,COLUMN_UNDERSCORE_1"), RowFactory.create("ID1, ID2") ); // Define schema StructType schema = new StructType(new StructField[]{ new StructField("COLUMN", DataTypes.StringType, false, Metadata.empty()) }); Dataset<Row> rows = spark.createDataFrame(data, schema); // Primary keys to check against Set<String> primaryKeys = new HashSet<>(); primaryKeys.add("ID1"); // 1. Split the column into an array (handle commas with optional spaces) Column splitColumn = functions.split(functions.col("COLUMN"), ",\\s*"); // 2. Convert primary keys to a Spark array literal List<Column> pkLiterals = primaryKeys.stream() .map(functions::lit) .collect(Collectors.toList()); Column pkArray = functions.array(pkLiterals.toArray(new Column[0])); // 3. Check if there's any overlap between the split column and primary keys Column hasMatch = functions.size(functions.array_intersect(splitColumn, pkArray)).gt(0); // Filter the dataset Dataset<Row> filteredRows = rows.filter(hasMatch); // Validate results Assert.assertEquals(2, filteredRows.count()); List<String> resultColumns = filteredRows.select("COLUMN") .as(Encoders.STRING()) .collectAsList(); Assert.assertTrue(resultColumns.contains("ID, ID1")); Assert.assertTrue(resultColumns.contains("ID1, ID2")); } }
Approach 2: Use a Custom UDF (Flexible for Complex Logic)
If you need more custom validation logic (e.g., case-insensitive checks), a User-Defined Function (UDF) is a good choice.
Complete Test Code:
import org.apache.spark.sql.Dataset; import org.apache.spark.sql.Row; import org.apache.spark.sql.RowFactory; import org.apache.spark.sql.SparkSession; import org.apache.spark.sql.Encoders; import org.apache.spark.sql.UserDefinedFunction; import org.apache.spark.sql.functions; import org.apache.spark.sql.types.DataTypes; import org.apache.spark.sql.types.Metadata; import org.apache.spark.sql.types.StructField; import org.apache.spark.sql.types.StructType; import org.junit.Assert; import org.junit.Test; import java.util.Arrays; import java.util.HashSet; import java.util.List; import java.util.Set; public class SparkFilterTest { private final SparkSession spark = SparkSession.builder() .appName("FilterTest") .master("local[*]") .getOrCreate(); @Test public void testIfColumnHasMentionsInPrimaryKeys() { // Test data and schema (same as above) List<Row> data = Arrays.asList( RowFactory.create("ID, ID1"), RowFactory.create("ID,COLUMN_UNDERSCORE_1"), RowFactory.create("ID1, ID2") ); StructType schema = new StructType(new StructField[]{ new StructField("COLUMN", DataTypes.StringType, false, Metadata.empty()) }); Dataset<Row> rows = spark.createDataFrame(data, schema); Set<String> primaryKeys = new HashSet<>(); primaryKeys.add("ID1"); // Define UDF to check for primary key mentions UserDefinedFunction hasMentionInPk = functions.udf( (String columnValue) -> { if (columnValue == null) return false; // Split and trim elements (regex handles spaces) String[] elements = columnValue.split(",\\s*"); for (String elem : elements) { if (primaryKeys.contains(elem)) { return true; } } return false; }, DataTypes.BooleanType ); // Filter using the UDF Dataset<Row> filteredRows = rows.filter(hasMentionInPk.apply(functions.col("COLUMN"))); // Validate results Assert.assertEquals(2, filteredRows.count()); List<String> resultColumns = filteredRows.select("COLUMN") .as(Encoders.STRING()) .collectAsList(); Assert.assertTrue(resultColumns.contains("ID, ID1")); Assert.assertTrue(resultColumns.contains("ID1, ID2")); } }
Key Notes:
- Regex Handling: The
",\\s*"regex ensures we split on commas regardless of whether there's a space after them (e.g.,"ID, ID1"and"ID,ID1"both split correctly). - Performance: The built-in function approach is preferred for large datasets because Spark can optimize these operations better than custom UDFs.
- Null Safety: Both approaches include checks for null values to avoid runtime errors.
内容的提问来源于stack exchange,提问作者Kirill Linnik
相关产品推荐
相关产品推荐

