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

PyTorch中train_test_split实现与字符串特征编码最佳方法

PyTorch 训练测试集划分与字符串特征编码方案

1. 训练集与测试集划分实现

PyTorch 官方torch.utils.data 模块内置了数据集划分工具,可以和后续数据加载、训练流水线无缝打通,不需要手动写切片逻辑。

  • 如果你手上是未拆分的全量数据集,直接用random_split 方法即可,固定随机种子就能保证结果可复现:
import torch
import numpy as np
import pandas as pd
from torch.utils.data import Dataset, DataLoader, random_split

# 自定义数据集封装
class FootballDataset(Dataset):
    def __init__(self, df):
        # 分离特征与标签,测试集标签为NaN时可以先做填充占位
        self.features = df.drop(columns=['home_score', 'away_score']).values.astype(np.float32)
        self.labels = df[['home_score', 'away_score']].fillna(-1).values.astype(np.float32)
    
    def __len__(self):
        return len(self.features)
    
    def __getitem__(self, idx):
        return self.features[idx], self.labels[idx]

# 加载全量数据后封装
full_data = FootballDataset(pd.read_csv('football_matches.csv'))
# 按8:2比例拆分训练、测试集
train_len = int(0.8 * len(full_data))
test_len = len(full_data) - train_len
train_set, test_set = random_split(
    full_data,
    [train_len, test_len],
    generator=torch.Generator().manual_seed(42) # 固定随机种子保证可复现
)
# 直接生成可输入模型的DataLoader
train_loader = DataLoader(train_set, batch_size=64, shuffle=True)
test_loader = DataLoader(test_set, batch_size=64, shuffle=False)
  • 如果你已经提前拿到了拆分好的训练、测试DataFrame(比如你当前已经有独立的df_train、df_test),不需要额外做拆分,分别封装成Dataset 实例即可。

注意:足球赛事数据有明确的时间顺序,不要用纯随机拆分。建议按比赛时间排序后,用时间点做切分(比如用2023年之前的数据做训练,2023年之后的数据做测试),避免用未来数据训练导致验证结果失真。

2. 字符串类型分类特征编码实现

你当前用sklearn LabelEncoder的逻辑可以跑通,但和PyTorch训练流水线割裂,遇到未知类别值会直接报错,也不方便后续接Embedding层,更推荐用PyTorch生态原生的torchtext.vocab 实现编码,鲁棒性更强。
针对你数据集中country、league、home_team、away_team这4个分类特征,实现代码如下:

from collections import Counter
from torchtext.vocab import vocab

cat_cols = ['country', 'league', 'home_team', 'away_team']
col_vocab = {}

# 仅用训练集数据构建词表,禁止把测试集数据加入构建过程,避免数据泄露
for col in cat_cols:
    count = Counter(df_train[col].values)
    # 加入<unk>标记处理未见过的类别,设置默认索引为<unk>
    cur_vocab = vocab(count, specials=['<unk>'], min_freq=1)
    cur_vocab.set_default_index(cur_vocab['<unk>'])
    col_vocab[col] = cur_vocab

# 批量编码训练集、测试集
for col in cat_cols:
    df_train[col] = df_train[col].apply(lambda x: col_vocab[col][x])
    df_test[col] = df_test[col].apply(lambda x: col_vocab[col][x])
  • 对比sklearn的LabelEncoder,这个方案的优势:
    • 原生支持未知类别容错:测试集或者线上推理时遇到训练集没出现过的球队、联赛,不会抛出ValueError,会自动映射到<unk>索引,适配生产环境的不确定输入
    • 与PyTorch组件无缝兼容:后续给分类特征加Embedding层时,可以直接用len(col_vocab[col]) 获取对应特征的词表大小作为Embedding层的输入维度,不需要额外统计类别数
    • 支持灵活的频次过滤:可以通过调整min_freq参数,把出现次数过低的小众球队、联赛自动归为<unk>,减少特征冗余,降低过拟合风险

提示:home_team、away_team属于高基数分类特征,不建议做OneHot编码,编码后接对应维度的Embedding层做低维稠密映射,模型效果会明显优于直接把标签编码输入全连接层。

针对当前数据集的额外提醒

  • 测试集中的home_score、away_score为待预测标签,构建数据集时不要把这两个字段纳入特征列
  • 赔率、月份、星期这类数值特征建议做标准化处理,标准化的均值、标准差只能在训练集上计算,再应用到测试集,避免数据泄露
  • 如果需要合并训练集和测试集的类别构建词表,一定要确认后续不会有超出词表的输入,否则线上推理很容易出问题

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.29 05:24:23