PyTable表按8:2拆分为机器学习训练集测试集的工具函数问询
PyTables数据集按比例拆分训练/测试集的实现方法
你可以直接使用成熟的工具函数实现需求,无需手动重复造轮子,最常用的方案是调用scikit-learn的train_test_split函数,同时支持随机拆分、按标签分层拆分等常用能力。
场景1:数据集可全量装入内存
如果你的内存足够容纳全部8909条256*256的浮点数组,可以直接读取全部数据后拆分:
from sklearn.model_selection import train_test_split import tables # 读取PyTables数据 h5file = tables.open_file('your_file_path.h5', mode='r') data_node = h5file.get_node('/neg/data') X = data_node.cols.data[:] y = data_node.cols.label[:] # 按8:2比例拆分,stratify参数保证训练/测试集标签分布一致 X_train, X_test, y_train, y_test = train_test_split( X, y, test_size=0.2, random_state=42, stratify=y )
参数说明:
test_size=0.2:指定测试集占总样本的20%random_state=42:固定随机种子,保证每次运行拆分结果可复现stratify=y:按标签列分层拆分,避免分类任务出现样本倾斜问题
场景2:数据集过大无法全量装入内存
如果全量读取数据会触发内存溢出,可以仅读取标签列和生成索引完成拆分,后续按需按索引读取对应样本即可:
from sklearn.model_selection import train_test_split import numpy as np import tables h5file = tables.open_file('your_file_path.h5', mode='r') data_node = h5file.get_node('/neg/data') total_samples = data_node.nrows # 仅读取标签列(内存占用极低)和生成全量索引 y = data_node.cols.label[:] idx = np.arange(total_samples) # 拆分索引而非原始数据 train_idx, test_idx, y_train, y_test = train_test_split( idx, y, test_size=0.2, random_state=42, stratify=y ) # 后续需要读取数据时,按索引批量/逐行读取即可,示例: # train_sample = data_node.cols.data[train_idx[0]] # test_sample = data_node.cols.data[test_idx[0]]
无依赖纯numpy实现方案
如果不想引入scikit-learn依赖,也可以用numpy原生方法完成随机拆分:
import numpy as np total_samples = 8909 idx = np.arange(total_samples) # 固定随机种子保证可复现 np.random.seed(42) np.random.shuffle(idx) # 按8:2切分索引 train_size = int(total_samples * 0.8) train_idx = idx[:train_size] test_idx = idx[train_size:]
注意该方案默认不支持分层拆分,需要保证原数据集标签分布均匀,或手动实现分层逻辑。
内容的提问来源于stack exchange,提问作者Sadman Sakib
相关产品推荐
相关产品推荐

