京公网安备 11010802034615号
经营许可证编号:京B2-20210330
PyTorch是一种流行的深度学习框架,它提供了许多方便的工具来处理数据集并构建模型。在深度学习中,我们通常需要对训练数据进行交叉验证,以评估模型的性能和确定超参数的最佳值。本文将介绍如何使用PyTorch实现10折交叉验证。
首先,我们需要加载数据集。假设我们有一个包含1000个样本的训练集,每个样本有10个特征和一个标签。我们可以使用PyTorch的Dataset和DataLoader类来加载和处理数据集。下面是一个示例代码片段:
import torch
from torch.utils.data import Dataset, DataLoader
class MyDataset(Dataset):
def __init__(self, data):
self.data = data
def __len__(self):
return len(self.data)
def __getitem__(self, idx):
x = torch.tensor(self.data[idx][:10], dtype=torch.float32)
y = torch.tensor(self.data[idx][10], dtype=torch.long)
return x, y
data = [[1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 0],
[2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 1],
...
[1000, 999, 998, 997, 996, 995, 994, 993, 992, 991, 9]]
dataset = MyDataset(data)
dataloader = DataLoader(dataset, batch_size=32, shuffle=True)
在这里,我们定义了一个名为MyDataset的自定义数据集类,它从数据列表中返回一个样本。每个样本分别由10个特征和1个标签组成。然后,我们使用Dataset和DataLoader类将数据集加载到内存中,并将其分成大小为32的批次。我们也可以选择在每个时期迭代时随机打乱数据集(shuffle=True)。
接下来,我们需要将训练集划分为10个不同的子集。我们可以使用Scikit-learn的StratifiedKFold类来将数据集划分为k个连续的折叠,并确保每个折叠中的类别比例与整个数据集相同。下面是一个示例代码片段:
from sklearn.model_selection import StratifiedKFold
kfold = StratifiedKFold(n_splits=10)
X = torch.stack([x for x, y in dataset])
y = torch.tensor([y for x, y in dataset])
for fold, (train_index, val_index) in enumerate(kfold.split(X, y)):
train_dataset = torch.utils.data.Subset(dataset, train_index)
val_dataset = torch.utils.data.Subset(dataset, val_index)
train_dataloader = DataLoader(train_dataset, batch_size=32, shuffle=True)
val_dataloader = DataLoader(val_dataset, batch_size=32, shuffle=False)
# Train and evaluate model on this fold
# ...
在这里,我们使用StratifiedKFold类将数据集划分为10个连续的折叠。然后,我们使用Subset类从原始数据集中选择训练集和验证集。最后,我们使用DataLoader类将每个子集分成批次,并分别对其进行训练和评估。
在每个折叠上训练和评估模型时,我们需要编写适当的代码。以下是一个简单的示例模型和训练代码:
import torch.nn as nn
import torch.optim as optim
class MyModel(nn.Module):
def __init__(self):
super(MyModel, self).__init__()
self.fc1 = nn.Linear(10, 64)
self.fc2 = nn.Linear(64, 2)
def forward(self, x):
x = self.fc1(x)
x = nn.functional.relu(x) x = self.fc2(x) return x
model = MyModel() criterion = nn.CrossEntropyLoss() optimizer = optim.Adam(model.parameters(), lr=0.001)
for epoch in range(10): for i, (inputs, labels) in enumerate(train_dataloader): optimizer.zero_grad()
outputs = model(inputs)
loss = criterion(outputs, labels)
loss.backward()
optimizer.step()
# Evaluate on validation set
with torch.no_grad():
total_correct = 0
total_samples = 0
for inputs, labels in val_dataloader:
outputs = model(inputs)
_, predicted = torch.max(outputs, 1)
total_correct += (predicted == labels).sum().item()
total_samples += labels.size(0)
accuracy = total_correct / total_samples
print(f"Fold {fold + 1}, Epoch {epoch + 1}: Validation accuracy={accuracy}")
在这里,我们定义了一个名为MyModel的简单模型,并使用Adam优化器和交叉熵损失函数进行训练。对于每个时期和每个批次,我们计算输出、损失和梯度,并更新模型参数。然后,我们使用no_grad()上下文管理器在验证集上进行评估,并计算准确性。
4. 汇总结果
最后,我们需要将10个折叠的结果合并以获得最终结果。可以使用numpy来跟踪每个折叠的测试损失和准确性,并计算平均值和标准差。以下是一个示例代码片段:
```python
import numpy as np
test_losses = []
test_accuracies = []
for fold, (train_index, test_index) in enumerate(kfold.split(X, y)):
test_dataset = torch.utils.data.Subset(dataset, test_index)
test_dataloader = DataLoader(test_dataset, batch_size=32, shuffle=False)
# Evaluate on test set
with torch.no_grad():
total_correct = 0
total_loss = 0
total_samples = 0
for inputs, labels in test_dataloader:
outputs = model(inputs)
loss = criterion(outputs, labels)
_, predicted = torch.max(outputs, 1)
total_correct += (predicted == labels).sum().item()
total_loss += loss.item() * labels.size(0)
total_samples += labels.size(0)
loss = total_loss / total_samples
accuracy = total_correct / total_samples
test_losses.append(loss)
test_accuracies.append(accuracy)
mean_test_loss = np.mean(test_losses)
std_test_loss = np.std(test_losses)
mean_test_accuracy = np.mean(test_accuracies)
std_test_accuracy = np.std(test_accuracies)
print(f"Final results: Test loss={mean_test_loss} ± {std_test_loss}, Test accuracy={mean_test_accuracy} ± {std_test_accuracy}")
在这里,我们使用Subset类创建测试集,并在每个折叠上评估模型。然后,我们使用numpy计算测试损失和准确性的平均值和标准差,并将它们打印出来。
总之,使用PyTorch实现10折交叉验证相对简单,只需使用Dataset、DataLoader、StratifiedKFold和Subset类即可。重点是编写适当的模型和训练代码,并汇总所有10个折叠的结果。这种方法可以帮助我们更好地评估模型的性能并确定超参数的最佳值。
数据分析咨询请扫描二维码
若不方便扫码,搜微信号:CDAshujufenxi
AB实验是互联网产品迭代、营销优化、功能升级的核心科学验证手段,通过流量随机分组、对照组与实验组对比,科学验证策略、功能、 ...
2026-08-10在MySQL数据库优化中,索引是提升查询效率、降低数据库IO开销、优化系统性能的核心手段。普通单列索引仅适配简单查询场景,面对 ...
2026-08-10 很多数据分析师每天盯着几十个指标,但当被问到“这套指标要支撑什么业务目标”“指标之间是什么逻辑关系”“业务变化时如何 ...
2026-08-10在数字化市场调研体系中,大数据与小数据是两类核心调研数据形态,分别对应海量行为统计与精准样本深度调研。行业普遍存在认知误 ...
2026-08-07数据透视表是Excel、WPS中最核心的数据分析工具,凭借快速汇总、分组统计、动态筛选的优势,被广泛应用于销量统计、业绩复盘、数 ...
2026-08-07 很多数据分析师每天盯着GMV、DAU、转化率,但当被问到“哪些指标在所有行业都适用”“哪些指标只对电商有意义”“二者如何搭 ...
2026-08-07在商品销量、市场需求、营收规模等业务数据中,季节性波动是最普遍、最核心的数据特征。零售快消、食品餐饮、家电服饰、电商行业 ...
2026-08-06在流量红利消退、市场竞争白热化的商业环境中,传统依托经验、跟风投放、广撒网式的营销模式,逐渐暴露出成本高、精准度低、转化 ...
2026-08-06 很多数据分析师每天盯着GMV、DAU、转化率,但当被问到“什么是指标”“指标和维度有什么区别”“如何定义指标值的计算规则和 ...
2026-08-06【核心关键词】知识、设备、工程师、数字化、建模、算法、工业大数据、数据分析、控制算法、机器学习、深度学习、业务目标、工 ...
2026-08-05增长率是数据分析、业务报表、可视化看板中最核心的指标之一,常用于衡量营收、销量、用户量、流量等业务数据的涨跌幅度与发展趋 ...
2026-08-05小陈是某电商平台的数据分析师。老板交给他一个任务:“我们平台的注册用户已经突破1000万了,想了解一下用户的平均月消费金额。 ...
2026-08-05在机器学习建模、数据预处理、抽样划分、模型初始化训练的全流程中,随机种子是控制实验随机性、保证结果可复现、提升模型稳定性 ...
2026-08-04Python作为面向对象的主流编程语言,核心编程体系由变量、方法、类三大基础要素构成。三者层层递进、相互协作,支撑起所有基础代 ...
2026-08-04 很多数据分析师沉迷于复杂的机器学习算法,却忽略了数据分析最基础也最核心的能力——描述性统计。事实上,80%的商业分析问 ...
2026-08-04为什么学习数据分析? 当下,我们已然步入数据要素价值全面释放的智能时代。数据不再只是零散的数字记录,更是驱动新质生产力运 ...
2026-08-03CDA等级认证证书有效期为三年,三年进行一次年审,主要考察CDA持证人的职业发展情况;学历、工作单位及职业变更情况;继续教育 ...
2026-08-03【核心关键词】金融、岗位、企业、算法、知识、专业、理论、课程、软件、数据分析、业务类型、应用场景、经营管理、大型企业、 ...
2026-08-03在日常数据分析、业务报表复盘、业绩预测工作中,平均增长率是衡量数据长期变化趋势的核心指标,广泛用于营收增长、用户增长、销 ...
2026-08-03很多人把统计学理解为“一堆公式和计算”,却忽略了它的本质——一门让数据“开口说话”的科学。真正的数据分析高手,不是会算平 ...
2026-08-03