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

基于Sklearn自定义TransformerMixin实现SMOTE处理不平衡数据的技术问询

Hey there! Let's walk through fixing your custom SMOTE transformer to align with scikit-learn's API standards and make it functional for your imbalanced data project.

Key Issues in Your Current Implementation

First, let's break down the problems in your code:

  • fit() returns the wrong value: Scikit-learn estimators require the fit() method to return self (the instance of the transformer), not the resampled X. Returning X will break pipeline workflows since pipelines expect to chain fit() → transform() on the same estimator instance.
  • Class name conflicts & PEP8 violation: Naming your class smote (lowercase) clashes with the SMOTE class you're importing, and violates Python's PEP8 class naming convention (classes should use PascalCase).
  • Bug in print statement: print(len(y), ' ', type) will print the built-in type function object instead of the type of y—you need type(y) here.
  • No persistent SMOTE instance: You're creating a SMOTE object inside fit() but not storing it, which means you can't reuse it in the transform() method for new data (like your test set).
  • Unfinished transform() method: The method is incomplete, which is critical for applying the resampling logic consistently.

Fixed Custom SMOTE Transformer

Here's a corrected implementation that follows scikit-learn's API rules and works reliably:

from imblearn.over_sampling import SMOTE
from sklearn.base import BaseEstimator, TransformerMixin

class CustomSMOTE(BaseEstimator, TransformerMixin):
    def __init__(self, kind='regular', n_jobs=-1):
        # Make SMOTE parameters configurable during initialization
        self.kind = kind
        self.n_jobs = n_jobs
        self.smote_instance = None

    def fit(self, X, y=None):
        # Fit the SMOTE instance to learn minority class distribution
        self.smote_instance = SMOTE(kind=self.kind, n_jobs=self.n_jobs)
        self.smote_instance.fit(X, y)
        
        # Print statements for debugging (optional, can remove later)
        print(f"Fit on data shape: {X.shape}, type: {type(X)}")
        print(f"Target length: {len(y)}, type: {type(y)}")
        
        # Required: Return self to comply with scikit-learn's estimator API
        return self

    def transform(self, X, y=None):
        # Ensure the transformer has been fitted first
        if self.smote_instance is None:
            raise ValueError("Transformer not fitted! Call 'fit()' before 'transform()'.")
        
        # SMOTE needs the target variable to identify minority classes
        if y is None:
            raise ValueError("SMOTE requires the target variable 'y' for resampling.")
        
        # Generate resampled balanced dataset
        X_resampled, y_resampled = self.smote_instance.fit_sample(X, y)
        return X_resampled, y_resampled

What's Changed & Why

  1. Configurable Initialization: The __init__ method lets you pass SMOTE parameters (like kind or n_jobs) when creating an instance, making the transformer flexible for different use cases.
  2. Persistent SMOTE Instance: We store the fitted SMOTE object as self.smote_instance, so transform() can reuse the learned distribution to resample new data consistently.
  3. Compliant fit() Method: Returns self so the transformer works seamlessly in scikit-learn pipelines.
  4. Robust transform() Method: Includes checks for unfitted state and missing y, plus properly returns the resampled X and y.
  5. Fixed Debug Prints: Corrected the type(y) call to show the actual type of your target variable.

Using in a Pipeline

If you want to use this transformer in a pipeline (e.g., with a classifier), use imblearn.pipeline.Pipeline instead of scikit-learn's default pipeline—this handles transformers that return both X and y:

from imblearn.pipeline import Pipeline
from sklearn.linear_model import LogisticRegression

# Create a pipeline with resampling + classification
pipeline = Pipeline([
    ('smote_resample', CustomSMOTE(kind='borderline1')),
    ('classifier', LogisticRegression())
])

# Train on imbalanced data
pipeline.fit(X_train, y_train)

# Predict on test data
y_pred = pipeline.predict(X_test)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.22 09:41:51