本文目录导读:

我来为您详细解释如何编写数据拆分分类脚本,这是机器学习和数据分析中非常常见且重要的步骤。
基础数据拆分脚本
使用 train_test_split
# 安装: pip install scikit-learn pandas numpy
import numpy as np
import pandas as pd
from sklearn.model_selection import train_test_split
# 加载数据
df = pd.read_csv('your_data.csv')
# 分离特征和目标变量
X = df.drop('target', axis=1) # 特征
y = df['target'] # 目标变量
# 基础拆分(训练集80%,测试集20%)
X_train, X_test, y_train, y_test = train_test_split(
X, y,
test_size=0.2, # 测试集比例
random_state=42, # 随机种子,确保结果可重复
stratify=y # 保持类别分布平衡(分类问题推荐)
)
print(f"训练集大小: {X_train.shape}")
print(f"测试集大小: {X_test.shape}")
print(f"训练集类别分布:\n{y_train.value_counts(normalize=True)}")
多类别分类的完整脚本
import pandas as pd
import numpy as np
from sklearn.model_selection import train_test_split, StratifiedKFold
from sklearn.preprocessing import LabelEncoder, StandardScaler
from sklearn.datasets import make_classification
import warnings
warnings.filterwarnings('ignore')
class DataSplitter:
"""数据拆分分类器"""
def __init__(self, test_size=0.2, val_size=0.1, random_state=42):
self.test_size = test_size
self.val_size = val_size
self.random_state = random_state
self.label_encoder = LabelEncoder()
self.scaler = StandardScaler()
def load_and_preprocess(self, data_path, target_column):
"""加载并预处理数据"""
# 加载数据
df = pd.read_csv(data_path)
# 分离特征和目标
X = df.drop(target_column, axis=1)
y = df[target_column]
# 编码目标变量(如果是分类问题)
if y.dtype == 'object':
y = self.label_encoder.fit_transform(y)
print(f"类别映射: {dict(zip(self.label_encoder.classes_,
range(len(self.label_encoder.classes_))))}")
return X, y
def stratified_split(self, X, y):
"""分层抽样拆分"""
# 先拆分为训练+验证集 和 测试集
X_temp, X_test, y_temp, y_test = train_test_split(
X, y,
test_size=self.test_size,
random_state=self.random_state,
stratify=y
)
# 再从训练+验证集中拆分出验证集
val_ratio = self.val_size / (1 - self.test_size)
X_train, X_val, y_train, y_val = train_test_split(
X_temp, y_temp,
test_size=val_ratio,
random_state=self.random_state,
stratify=y_temp
)
return {
'train': (X_train, y_train),
'val': (X_val, y_val),
'test': (X_test, y_test)
}
def scale_features(self, splits):
"""标准化特征"""
X_train = splits['train'][0]
# 拟合scaler
self.scaler.fit(X_train)
scaled_splits = {}
for split_name, (X, y) in splits.items():
X_scaled = self.scaler.transform(X)
scaled_splits[split_name] = (X_scaled, y)
return scaled_splits
def cross_validation_split(self, X, y, n_folds=5):
"""K折交叉验证拆分"""
skf = StratifiedKFold(
n_splits=n_folds,
shuffle=True,
random_state=self.random_state
)
folds = []
for train_idx, val_idx in skf.split(X, y):
X_train, X_val = X.iloc[train_idx], X.iloc[val_idx]
y_train, y_val = y.iloc[train_idx], y.iloc[val_idx]
folds.append({
'train': (X_train.values, y_train.values),
'val': (X_val.values, y_val.values)
})
return folds
def split_by_category(self, df, category_column, target_column):
"""按类别拆分数据"""
categories = df[category_column].unique()
category_splits = {}
for category in categories:
category_data = df[df[category_column] == category]
X = category_data.drop(target_column, axis=1)
y = category_data[target_column]
category_splits[category] = (X, y)
return category_splits
# 使用示例
def main():
# 创建示例数据
X, y = make_classification(
n_samples=1000,
n_features=20,
n_classes=3,
n_informative=15,
random_state=42
)
# 转换为DataFrame
df = pd.DataFrame(X, columns=[f'feature_{i}' for i in range(20)])
df['target'] = y
df.to_csv('sample_data.csv', index=False)
# 初始化拆分器
splitter = DataSplitter(test_size=0.2, val_size=0.1)
# 加载数据
X, y = splitter.load_and_preprocess('sample_data.csv', 'target')
# 进行分层拆分
splits = splitter.stratified_split(X, y)
# 标准化
scaled_splits = splitter.scale_features(splits)
# 输出结果
for split_name, (X_data, y_data) in scaled_splits.items():
print(f"\n{split_name.upper()} 集合:")
print(f" 样本数: {len(X_data)}")
print(f" 特征形状: {X_data.shape}")
print(f" 类别分布: {np.bincount(y_data.astype(int))}")
# 交叉验证
print("\n5折交叉验证:")
folds = splitter.cross_validation_split(X, y, n_folds=5)
for i, fold in enumerate(folds):
train_size = len(fold['train'][0])
val_size = len(fold['val'][0])
print(f" 折 {i+1}: 训练集={train_size}, 验证集={val_size}")
if __name__ == "__main__":
main()
高级数据拆分功能
import pandas as pd
import numpy as np
from sklearn.model_selection import (
train_test_split,
StratifiedKFold,
TimeSeriesSplit,
GroupKFold
)
from datetime import datetime
class AdvancedDataSplitter:
"""高级数据拆分器"""
@staticmethod
def temporal_split(df, date_column, train_end_date, test_start_date):
"""时间序列拆分"""
train = df[df[date_column] < train_end_date]
test = df[df[date_column] >= test_start_date]
# 验证没有数据泄漏
assert train[date_column].max() < test[date_column].min(), \
"训练集和测试集有重叠!"
return train, test
@staticmethod
def group_split(X, y, groups, test_size=0.2):
"""按组拆分(防止同组数据泄露)"""
unique_groups = np.unique(groups)
n_groups = len(unique_groups)
# 随机选择组
n_test_groups = int(n_groups * test_size)
np.random.seed(42)
test_groups = np.random.choice(
unique_groups,
size=n_test_groups,
replace=False
)
# 创建掩码
test_mask = np.isin(groups, test_groups)
train_mask = ~test_mask
return X[train_mask], X[test_mask], y[train_mask], y[test_mask]
@staticmethod
def balanced_undersample(X, y, target_samples_per_class=None):
"""平衡下采样"""
from sklearn.utils import resample
df = pd.DataFrame(X)
df['target'] = y
# 获取最小类别样本数
class_counts = y.value_counts()
if target_samples_per_class is None:
target_samples_per_class = class_counts.min()
# 对每个类别进行下采样
balanced_dfs = []
for cls in class_counts.index:
cls_df = df[df['target'] == cls]
sampled = resample(
cls_df,
replace=False,
n_samples=min(target_samples_per_class, len(cls_df)),
random_state=42
)
balanced_dfs.append(sampled)
balanced_df = pd.concat(balanced_dfs)
return (
balanced_df.drop('target', axis=1).values,
balanced_df['target'].values
)
@staticmethod
def save_splits(splits, output_dir='./data_splits'):
"""保存拆分的文件"""
import os
os.makedirs(output_dir, exist_ok=True)
for split_name, (X, y) in splits.items():
# 保存为CSV
df_split = pd.DataFrame(
np.column_stack([X, y]),
columns=[f'feature_{i}' for i in range(X.shape[1])] + ['target']
)
df_split.to_csv(
f"{output_dir}/{split_name}.csv",
index=False
)
print(f"已保存: {output_dir}/{split_name}.csv")
# 使用示例
def advanced_usage_example():
# 创建包含时间的数据
dates = pd.date_range('2023-01-01', periods=365, freq='D')
data = {
'date': dates,
'value': np.random.randn(365),
'target': np.random.randint(0, 3, 365),
'group': np.random.choice(['A', 'B', 'C', 'D'], 365)
}
df = pd.DataFrame(data)
# 时间序列拆分
splitter = AdvancedDataSplitter()
train, test = splitter.temporal_split(
df,
date_column='date',
train_end_date='2023-09-30',
test_start_date='2023-10-01'
)
print(f"时间序列拆分 - 训练: {len(train)}, 测试: {len(test)}")
# 组拆分
X = df[['value']].values
y = df['target'].values
groups = df['group'].values
X_train, X_test, y_train, y_test = splitter.group_split(X, y, groups)
print(f"组拆分 - 训练: {len(X_train)}, 测试: {len(X_test)}")
# 运行示例
advanced_usage_example()
数据质量检查脚本
def data_quality_report(df, target_column):
"""数据质量报告"""
report = {
'总样本数': len(df),
'特征数量': df.shape[1] - 1,
'目标变量': target_column,
'类别分布': df[target_column].value_counts().to_dict(),
'类别比例': df[target_column].value_counts(normalize=True).to_dict(),
'缺失值': df.isnull().sum().to_dict(),
'重复行': df.duplicated().sum(),
'数据类型': df.dtypes.to_dict()
}
# 检查类别平衡
min_class = min(df[target_column].value_counts())
max_class = max(df[target_column].value_counts())
imbalance_ratio = max_class / min_class
report['类别不平衡比率'] = imbalance_ratio
report['类别是否平衡'] = imbalance_ratio < 2.0
return report
# 使用数据质量检查
df = pd.read_csv('your_data.csv')
report = data_quality_report(df, 'target')
for key, value in report.items():
print(f"{key}: {value}")
实际应用脚本书写建议
最佳实践清单:
- 始终设置 random_state 确保结果可重现
- 使用分层抽样
stratify=y保持类别分布 - 处理数据泄漏 拆分后再进行标准化
- 验证拆分结果 检查类别分布是否一致
- 保存拆分信息 便于后续使用
错误处理模板:
def safe_split_data(df, target_col, test_size=0.2):
"""安全的数据拆分函数"""
try:
# 验证输入
if target_col not in df.columns:
raise ValueError(f"目标列 '{target_col}' 不存在")
if df[target_col].isnull().any():
raise ValueError("目标变量包含缺失值")
if len(df) < 10:
raise ValueError("数据集太小,无法进行有意义的拆分")
# 执行拆分
X = df.drop(target_col, axis=1)
y = df[target_col]
X_train, X_test, y_train, y_test = train_test_split(
X, y,
test_size=test_size,
random_state=42,
stratify=y if len(np.unique(y)) > 1 else None
)
return X_train, X_test, y_train, y_test
except Exception as e:
print(f"数据拆分失败: {str(e)}")
return None, None, None, None
这些脚本涵盖了大部分数据拆分场景,根据您的具体需求选择合适的方法,并根据数据类型(图像、文本、时间序列等)进行调整。