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

Python中利用空间密度权重实现训练测试集空间一致拆分的方法

基于空间密度权重的训练/测试集拆分方案

可用工具

直接用numpy和pandas就能实现,不需要额外复杂库;如果偏好现成工具组合,scikit-learn的GroupShuffleSplit可以结合权重分箱后的组来完成拆分,但自定义实现会更灵活适配你的空间密度权重场景。

自定义实现思路

方法1:权重分箱分层拆分(稳定优先)

通过将权重分组,保证每组内按比例拆分,最大化保留原数据的空间密度分布,适合权重差异较大的场景:

  1. 对空间密度权重做等频分箱(建议10-20个箱),把样本按权重区间归类
  2. 每个箱内独立按预设比例拆分训练/测试样本
  3. 合并所有箱的拆分结果
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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.25 05:40:33