You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

Spark DataFrame使用Random过滤采样失效问题求助

Why Isn't My Label 0.0 Subsampling Logic Working in Spark?

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.13 03:42:50