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

从Pandas迁移到PySpark后,如何获取Naive Bayes特征重要性?

How to get feature importance from Naive Bayes in PySpark MLlib?

I originally implemented a Naive Bayes classifier using Python and Pandas DataFrames with scikit-learn's BernoulliNB, where I could easily get feature importance via the coef_ attribute (which represents the log probability of each feature in the positive class). Now I'm migrating this to PySpark and can't find an equivalent method in the PySpark MLlib documentation.

Here's my original scikit-learn code for reference:

df = df.fillna(0).toPandas()
X_df = df.drop(['NOT_OPEN', 'unique_id'], axis = 1)
X = X_df.values
Y = df['NOT_OPEN'].values.reshape(-1,1)
mnb = BernoulliNB(fit_prior=True)
y_pred = mnb.fit(X, Y).predict(X)
estimator = mnb.fit(X, Y)
# coef_: 对于二分类问题,这是特征在正类下的估计概率的对数,值越高表示该特征对正类越重要
feature_names = X_df.columns
coefs_with_fns = sorted(zip(estimator.coef_[0], feature_names))

Is there a way to replicate this feature importance functionality using PySpark's NaiveBayes?


Great question! Unfortunately, PySpark's MLlib NaiveBayes doesn't expose a direct coef_ attribute like scikit-learn's BernoulliNB does. But you can replicate this feature importance logic using the model's underlying parameters, since the value you care about (log probability of features in the positive class) is stored in the model's theta attribute.

Background on PySpark NaiveBayes Parameters

PySpark's NaiveBayes stores two key parameters after training:

  • theta: An array where theta[i][j] represents the log-transformed conditional probability log(P(feature j = 1 | class i))
  • pi: An array of log-transformed prior probabilities for each class

For your Bernoulli Naive Bayes setup (matching your original scikit-learn code), theta[positive_class_index] gives you exactly the log probabilities for features in the positive class—this is equivalent to scikit-learn's coef_[0].

Step-by-Step Implementation

Here's how to extract and sort feature importance in PySpark:

from pyspark.ml.classification import NaiveBayes
from pyspark.ml.feature import VectorAssembler
from pyspark.sql import SparkSession

# Initialize Spark session
spark = SparkSession.builder.appName("NaiveBayesFeatureImportance").getOrCreate()

# 1. Prepare your data (replace with your actual DataFrame)
# Fill missing values (equivalent to df.fillna(0) in pandas)
df = df.fillna(0)

# 2. Assemble features into a single vector column (required for PySpark ML)
feature_cols = [col for col in df.columns if col not in ["NOT_OPEN", "unique_id"]]
assembler = VectorAssembler(inputCols=feature_cols, outputCol="features")
df_assembled = assembler.transform(df)

# 3. Train Bernoulli Naive Bayes model (match fit_prior=True from your original code)
nb = NaiveBayes(modelType="bernoulli", fitPrior=True)
model = nb.fit(df_assembled)

# 4. Identify the index of your positive class (e.g., 1 for NOT_OPEN=1)
class_labels = model.labels
positive_class_idx = class_labels.index(1)  # Adjust if your positive class label is different

# 5. Extract log probabilities for the positive class
positive_class_log_probs = model.theta[positive_class_idx]

# 6. Pair feature names with their importance scores and sort
coefs_with_fns = list(zip(feature_cols, positive_class_log_probs))
# Sort descending to get most important features first
coefs_with_fns_sorted = sorted(coefs_with_fns, key=lambda x: x[1], reverse=True)

# Print or use the sorted feature importance
print("Feature Importance (log probability for positive class):")
for feature, score in coefs_with_fns_sorted:
    print(f"{feature}: {score:.4f}")

Notes

  • If you want a different interpretation of feature importance (like the log ratio of positive vs negative class probabilities), you can calculate theta[positive_class_idx] - theta[negative_class_idx] instead.
  • Always verify model.labels to ensure you're targeting the correct class index, as PySpark sorts class labels alphabetically/numerically by default.

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.11 08:54:42