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

如何扩展代码用sklearn.train_test_split训练双目标X_S与X_S2?

双目标分类任务的数据集划分代码修正

错误原因分析

你遇到的_append_dispatcher() missing 1 required positional argument: 'values'错误,核心问题有两个:

  • np.append()仅支持两个数组+axis参数的输入,你一次性传入三个数组(X_S, X_S2, X_BG),违反函数参数规则。
  • 修改后的函数丢失了test_sample_size参数,导致train_test_split的test_size无有效值可用。

修正后的完整代码

数据集划分函数

import numpy as np
from sklearn.model_selection import train_test_split
from tensorflow.keras.utils import to_categorical

def split(X_BG, X_S, X_S2, test_sample_size, do_shuffle=True):
    np.random.shuffle(X_BG)

    # 标签赋值:X_S和X_S2归为正类(标签1),X_BG归为负类(标签0)
    y_S = np.ones(len(X_S))
    y_S2 = np.ones(len(X_S2))
    y_BG = np.zeros(len(X_BG)) 

    # 合并多组数据:用np.concatenate更高效(替代嵌套np.append)
    X = np.concatenate([X_S, X_S2, X_BG], axis=0)
    y = np.concatenate([y_S, y_S2, y_BG], axis=0)

    # 调整输入形状适配CNN的单通道要求
    X = X.reshape((X.shape[0], X.shape[1], X.shape[2], 1))

    # 划分训练集与测试集
    X_train, X_test, y_train, y_test = train_test_split(
        X, y, test_size=test_sample_size, shuffle=do_shuffle
    )

    # 转换标签为独热编码,适配分类模型输出
    y = to_categorical(y, 2)
    y_train = to_categorical(y_train, 2)
    y_test = to_categorical(y_test, 2)

    input_shape = X[0].shape

    return X, X_train, X_test, y_train, y_test, input_shape

函数调用示例

# 传入测试集比例(如0.25表示25%数据作为测试集)
X_2500, X_train_2500, X_test_2500, y_train_2500, y_test_2500, input_shape_2500 = split(X_BG_2500, X_S_2500, X_S2_2500, 0.25)
X_1500, X_train_1500, X_test_1500, y_train_1500, y_test_1500, input_shape_1500 = split(X_BG_1500, X_S_1500, X_S2_1500, 0.25)
X_600, X_train_600, X_test_600, y_train_600, y_test_600, input_shape_600 = split(X_BG_600, X_S_600, X_S2_600, 0.25)
X_300, X_train_300, X_test_300, y_train_300, y_test_300, input_shape_300 = split(X_BG_300, X_S_300, X_S2_300, 0.25)

print("Input vector shape:", input_shape_2500)
print("Number of input samples:", len(X_2500))

额外优化建议

  • 优先使用np.concatenate合并多数组:相比嵌套np.append,concatenate直接接收数组列表,避免生成中间临时数组,效率更高。
  • 确保X_S、X_S2、X_BG的维度完全一致,否则数组合并时会抛出维度不匹配错误。
  • do_shuffle参数设为默认值True,如果不需要打乱数据,调用时显式传入do_shuffle=False即可。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.27 15:12:04