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

Spark中逻辑回归系数标准误差的计算方法问询

Getting Coefficient Standard Errors & Wald Chi-Square for Spark Logistic Regression

Great question—this is a well-known gap in Spark's built-in BinaryLogisticRegressionSummary since linear models (like LinearRegression) do expose standard errors out of the box. The Statistics.chiSqTest you mentioned is indeed for goodness-of-fit, not coefficient significance, so let's walk through how to calculate standard errors (via the variance-covariance matrix) manually, which will let you compute Wald chi-square stats too.

Why Spark Doesn't Expose This Directly

Spark's Logistic Regression uses gradient-based optimizers (like L-BFGS or SGD) instead of the Newton-Raphson method that traditional stats tools rely on. Newton-Raphson directly computes the Hessian matrix (needed for variance estimates) as part of training, but Spark doesn't expose this matrix in its public API. So we'll calculate it ourselves using the observed Fisher information matrix.

Step-by-Step Implementation

Here's how to extend your existing code to get the variance-covariance matrix and standard errors:

1. Prepare Predictions & Weight Values

First, we need the predicted probabilities for each sample to compute the weight term ( w_i = p_i(1-p_i) ), where ( p_i ) is the predicted probability of the positive class.

import org.apache.spark.ml.linalg.{Matrix, DenseMatrix, Vectors}
import org.apache.spark.sql.functions.{col, udf}

// Get predictions with probabilities
val predictions = binarySummary.predictions

// UDF to calculate w_i = p*(1-p) for the positive class
val computeWeight = udf((probVector: org.apache.spark.ml.linalg.Vector) => {
  val p = probVector(1) // Positive class probability
  p * (1.0 - p)
})

val weightedData = predictions.withColumn("weight", computeWeight(col("probability")))

2. Compute the Observed Hessian Matrix

The Hessian matrix for logistic regression is ( X^T W X ), where ( X ) is the feature matrix (including a constant column for the intercept) and ( W ) is a diagonal matrix of ( w_i ) values. We'll add an intercept column to our features first:

// UDF to add an intercept (constant 1.0) to the feature vector
val addIntercept = udf((features: org.apache.spark.ml.linalg.Vector) => {
  Vectors.dense(1.0 +: features.toDense.values)
})

val dataWithIntercept = weightedData.withColumn(
  "features_with_intercept", 
  addIntercept(col(lr.getFeaturesCol))
)

// Calculate the Hessian matrix by summing weighted outer products of features
val hessian: DenseMatrix = dataWithIntercept.rdd.map { row =>
  val features = row.getAs[org.apache.spark.ml.linalg.Vector]("features_with_intercept")
  val w = row.getAs[Double]("weight")
  // Compute w * features * features^T as a dense matrix
  val outerProduct = features.toDense.values.map(v => v * w).toArray
  new DenseMatrix(features.size, features.size, outerProduct)
}.reduce((matA, matB) => matA.add(matB))

3. Apply Regularization

Since your model uses regParam and elasticNetParam, we need to add the regularization term to the Hessian. Note that Spark doesn't regularize the intercept term by default, so we'll only apply L2 regularization to the feature coefficients:

val regParam = lr.getRegParam
val elasticNetParam = lr.getElasticNetParam
val l2RegStrength = regParam * (1.0 - elasticNetParam) // L2 component of ElasticNet

// Create a regularization matrix: L2 penalty for features, 0 for intercept
val regMatrix = DenseMatrix.zeros(hessian.numRows, hessian.numCols)
for (i <- 0 until hessian.numRows - 1) { // Skip the last row/column (intercept)
  regMatrix.update(i, i, l2RegStrength)
}

// Regularized Hessian matrix
val regularizedHessian = hessian.add(regMatrix)

4. Calculate Variance-Covariance Matrix & Standard Errors

The variance-covariance matrix is the inverse of the regularized Hessian. The standard error of each coefficient is the square root of the corresponding diagonal element in this matrix:

// Compute inverse of the regularized Hessian (variance-covariance matrix)
val covMatrix = regularizedHessian.inverse()

// Extract coefficients and intercept from the model
val coefficients = lrModel.coefficients.toArray
val intercept = lrModel.intercept
val numFeatures = coefficients.length

// Get standard errors for features and intercept
val coeffStandardErrors = (0 until numFeatures).map(i => 
  math.sqrt(covMatrix(i, i))
).toArray
val interceptStandardError = math.sqrt(covMatrix(numFeatures, numFeatures))

5. Build Your Parameter Estimation Table

Now you can assemble the table with coefficients, standard errors, Wald Z-scores, and Wald chi-square values (which are ( (coefficient / SE)^2 )):

// Create parameter table entries for features
val featureParams = coefficients.zip(coeffStandardErrors).zipWithIndex.map { 
  case ((coeff, se), idx) =>
    val zScore = coeff / se
    (s"Feature $idx", coeff, se, zScore, zScore * zScore)
}

// Add intercept entry
val interceptParam = ("Intercept", intercept, interceptStandardError, 
  intercept / interceptStandardError, 
  math.pow(intercept / interceptStandardError, 2)
)

// Full parameter table
val paramEstimationTable = featureParams :+ interceptParam

Key Notes

  • L1 Regularization Caveat: If you set elasticNetParam > 0 (L1 regularization), the Hessian matrix may be singular (irreversible) because some coefficients will be zero. In this case, you can use a pseudo-inverse (regularizedHessian.pinv()) instead of inverse(), or consider re-running the model without L1 regularization if you need reliable variance estimates.
  • Large Sample Assumption: This method uses the observed Fisher information, which is consistent with the expected information matrix in large samples—so it's valid for most practical use cases.
  • Wald Chi-Square: The last column in your table is the Wald chi-square statistic for each coefficient, which you can use to test significance (compare to a chi-square distribution with 1 degree of freedom).

内容的提问来源于stack exchange,提问作者user2129946

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.15 04:53:57