Python中利用空间密度权重实现训练测试集空间一致拆分的方法
基于空间密度权重的训练/测试集拆分方案
可用工具
直接用numpy和pandas就能实现,不需要额外复杂库;如果偏好现成工具组合,scikit-learn的GroupShuffleSplit可以结合权重分箱后的组来完成拆分,但自定义实现会更灵活适配你的空间密度权重场景。
自定义实现思路
方法1:权重分箱分层拆分(稳定优先)
通过将权重分组,保证每组内按比例拆分,最大化保留原数据的空间密度分布,适合权重差异较大的场景:
- 对空间密度权重做等频分箱(建议10-20个箱),把样本按权重区间归类
- 每个箱内独立按预设比例拆分训练/测试样本
- 合并所有箱的拆分结果
import pandas as pd import numpy as np # 假设数据集为df,'weight'列是你的空间密度权重,'features'是特征列,'nitrate_concentration'是目标列 df = pd.read_csv('bavaria_stream_data.csv') split_ratio = 0.8 random_seed = 42 np.random.seed(random_seed) # 步骤1:等频分箱,避免极端权重样本过度集中 df['weight_bin'] = pd.qcut(df['weight'], q=10, labels=False) # 步骤2:按箱拆分样本 train_idx = [] test_idx = [] for bin_id in df['weight_bin'].unique(): bin_samples = df[df['weight_bin'] == bin_id].index.to_list() np.random.shuffle(bin_samples) split_pos = int(len(bin_samples) * split_ratio) train_idx.extend(bin_samples[:split_pos]) test_idx.extend(bin_samples[split_pos:]) # 提取最终训练/测试集 X_train, y_train = df.loc[train_idx, 'features'], df.loc[train_idx, 'nitrate_concentration'] X_test, y_test = df.loc[test_idx, 'features'], df.loc[test_idx, 'nitrate_concentration']
方法2:直接加权随机采样(简洁优先)
如果权重分布相对均匀,可直接用归一化后的权重做无放回采样,确保训练集的权重分布与原数据一致:
import numpy as np import pandas as pd df = pd.read_csv('bavaria_stream_data.csv') split_ratio = 0.8 random_seed = 42 np.random.seed(random_seed) # 归一化权重,满足numpy.random.choice的概率参数要求 normalized_weights = df['weight'] / df['weight'].sum() # 无放回采样训练集 train_size = int(len(df) * split_ratio) train_idx = np.random.choice(df.index, size=train_size, replace=False, p=normalized_weights) test_idx = df.index.difference(train_idx) # 提取最终训练/测试集 X_train, y_train = df.loc[train_idx, 'features'], df.loc[train_idx, 'nitrate_concentration'] X_test, y_test = df.loc[test_idx, 'features'], df.loc[test_idx, 'nitrate_concentration']
验证拆分效果
拆分后可以通过直方图对比训练集、测试集与原数据的权重分布,确认空间分布一致性:
import matplotlib.pyplot as plt plt.hist(df['weight'], alpha=0.5, label='Original') plt.hist(df.loc[train_idx, 'weight'], alpha=0.5, label='Train') plt.hist(df.loc[test_idx, 'weight'], alpha=0.5, label='Test') plt.legend() plt.title('Spatial Density Weight Distribution Comparison') plt.show()
内容的提问来源于stack exchange,提问作者Karan Mahajan
相关产品推荐
相关产品推荐

