Spark DataFrame使用Random过滤采样失效问题求助
Let’s break down what’s going wrong with your code, and how to fix it to get the 50% subsampling you want for label 0.0 records.
The Root Cause of Your Issue
Your current code uses a Scala Random instance directly in the filter condition:
val r = scala.util.Random val df2 = df.filter($"label" === 1.0 || r.nextDouble > 0.5)
Here’s the problem: Spark transformations like filter are lazy-executed, and variables referenced in their closures are evaluated once on the Driver node, not for every individual row on Executors. That means r.nextDouble generates a single random value when Spark plans the filter—not for each row in your DataFrame.
- If that single value is ≤ 0.5, your condition simplifies to
label == 1.0 || false—so only label 1.0 rows are kept (which is exactly what you observed). - If it’s > 0.5, the condition becomes
label == 1.0 || true—so all rows are retained.
Either way, you don’t get the per-row 50% sampling you intended for label 0.0.
The Simple Fix: Use Spark’s Built-in rand() Function
Spark provides a rand() function that generates a random double between 0 and 1 for every row, which is exactly what you need for proper subsampling. Here’s how to implement it:
import org.apache.spark.sql.functions.rand // Keep all label 1.0 rows, and ~50% of label 0.0 rows val df2 = df.filter($"label" === 1.0 || rand() > 0.5) // Verify the result df2.groupBy($"label").count.show
This will correctly retain approximately half of your label 0.0 records, alongside all label 1.0 records, as intended.
Alternative: Per-Executor Random Instances (Less Ideal)
If you prefer to use Scala’s Random instead of Spark’s built-in function, you need to ensure a new random instance is created on Executors (to avoid serialization issues and per-row randomness). You can do this with a map transformation:
val df2 = df.map { row => // Create a Random instance inside the map (executed on Executors) val r = new scala.util.Random() val label = row.getAs[Double]("label") if (label == 1.0 || r.nextDouble() > 0.5) Some(row) else None } .filter(_.isDefined) .map(_.get)
Note that this approach is less efficient than using rand() because it requires manual row handling, whereas Spark’s built-in functions are optimized for distributed processing.
内容的提问来源于stack exchange,提问作者Marsellus Wallace

