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

是否有与sklearn的train_test_split等效的Keras数据划分方法?

Great question! Keras doesn't have a direct drop-in replacement for sklearn's train_test_split() with the stratify parameter, but there are a couple of straightforward ways to get the exact same stratified train-test split behavior. Let's break them down:

1. Use sklearn's train_test_split() directly (simplest approach)

Since Keras works seamlessly with numpy arrays and TensorFlow tensors, you can just use sklearn's well-tested splitting function first, then feed the split data into your Keras model. This is the most low-effort way to replicate the behavior you're used to.

Here's a quick example:

from sklearn.model_selection import train_test_split
from tensorflow.keras.models import Sequential

# Assume X (features) and Y (labels) are your dataset (numpy arrays or tf tensors)
X_train, X_test, Y_train, Y_test = train_test_split(X, Y, test_size=0.3, stratify=Y)

# Proceed with Keras model training as usual
model = Sequential([
    # Your model layers here
])
model.compile(optimizer='adam', loss='sparse_categorical_crossentropy', metrics=['accuracy'])
model.fit(X_train, Y_train, validation_data=(X_test, Y_test), epochs=10)

This preserves the stratified label distribution exactly like the sklearn method you're familiar with, and you don't have to reinvent any logic.

2. Implement stratified splitting with TensorFlow/Keras tools (no sklearn dependency)

If you want to stay entirely within the TensorFlow/Keras ecosystem, you can manually split your data by grouping samples by their labels, splitting each group proportionally, then combining the results.

Here's a reusable function to do this:

import tensorflow as tf

def stratified_train_test_split(X, Y, test_size=0.3):
    # Pack features and labels into a TensorFlow Dataset
    full_dataset = tf.data.Dataset.from_tensor_slices((X, Y))
    
    # Get all unique labels in the dataset
    unique_labels = tf.unique(Y)[0].numpy()
    
    train_ds = None
    test_ds = None
    
    for label in unique_labels:
        # Filter samples for the current label
        label_subset = full_dataset.filter(lambda x, y: tf.equal(y, label))
        # Calculate how many samples go to test vs train for this label
        total_samples = tf.data.experimental.cardinality(label_subset).numpy()
        test_samples = int(total_samples * test_size)
        train_samples = total_samples - test_samples
        
        # Split the subset into train and test
        label_train = label_subset.take(train_samples)
        label_test = label_subset.skip(train_samples)
        
        # Merge into the full train/test datasets
        if train_ds is None:
            train_ds = label_train
            test_ds = label_test
        else:
            train_ds = train_ds.concatenate(label_train)
            test_ds = test_ds.concatenate(label_test)
    
    # Shuffle the final datasets (optional but recommended for training)
    train_ds = train_ds.shuffle(buffer_size=tf.data.experimental.cardinality(train_ds).numpy())
    test_ds = test_ds.shuffle(buffer_size=tf.data.experimental.cardinality(test_ds).numpy())
    
    # Convert back to numpy arrays if needed (skip if using Dataset directly in model.fit)
    X_train = tf.concat(list(train_ds.map(lambda x, y: x)), axis=0).numpy()
    Y_train = tf.concat(list(train_ds.map(lambda x, y: y)), axis=0).numpy()
    X_test = tf.concat(list(test_ds.map(lambda x, y: x)), axis=0).numpy()
    Y_test = tf.concat(list(test_ds.map(lambda x, y: y)), axis=0).numpy()
    
    return X_train, X_test, Y_train, Y_test

# Usage example
X_train, X_test, Y_train, Y_test = stratified_train_test_split(X, Y, test_size=0.3)

This function ensures each label's proportion is maintained in both train and test sets, just like stratify=Y in sklearn. You can also skip converting back to numpy arrays and use the train_ds/test_ds directly with model.fit() if you prefer working with TensorFlow Datasets.


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.19 03:08:02