C++代码实现MATLAB中的cvpartition函数功能

cpp 复制代码
#include <iostream>
#include <vector>
#include <string>
#include <random>
#include <algorithm>
#include <numeric>
#include <unordered_map>
#include <stdexcept>
#include <cmath>
// ============================================================
// CVPartition: 模拟 MATLAB cvpartition 的核心行为
// ============================================================
class CVPartition {
public:
enum class Type { KFold, HoldOut, LeaveOut, Resubstitution };
// 构造函数:基于观测数量 n(非分层)
CVPartition(size_t n, Type type = Type::KFold, double k = 10.0)
: n_(n), type_(type), numTestSets_(1), isStratified_(false),
rng_(std::random_device{}()) {
validateAndSetup(n, type, k);
generateNonStratifiedPartition(k);
}
// 构造函数:基于分组标签 y(默认分层)
template<typename T>
CVPartition(const std::vector<T>& y, Type type = Type::KFold, double k
= 10.0)
: n_(y.size()), type_(type), numTestSets_(1), isStratified_(true),
rng_(std::random_device{}()) {
validateAndSetup(n_, type, k);
generateStratifiedPartition(y, k);
}
// 获取训练集索引(第 i 折,从 1 开始)
std::vector<bool> training(size_t i = 1) const {
if (i < 1 || i > numTestSets_)
throw std::out_of_range("Fold index out of range");
std::vector<bool> idx(n_, true);
auto testIdx = test(i);
for (size_t j = 0; j < n_; ++j) idx[j] = !testIdx[j];
return idx;
}
// 获取测试集索引(第 i 折,从 1 开始)
std::vector<bool> test(size_t i = 1) const {
if (i < 1 || i > numTestSets_)
throw std::out_of_range("Fold index out of range");
return testSets_[i - 1];
}
// 属性访问
size_t NumObservations() const { return n_; }
size_t NumTestSets() const { return numTestSets_; }
Type PartitionType() const { return type_; }
bool IsStratified() const { return isStratified_; }
// 获取每个折的训练集大小和测试集大小
std::vector<size_t> TrainSize() const {
std::vector<size_t> sizes;
for (size_t i = 0; i < numTestSets_; ++i) {
size_t testCount = std::count(testSets_[i].begin(),
testSets_[i].end(), true);
sizes.push_back(n_ - testCount);
}
return sizes;
}
std::vector<size_t> TestSize() const {
std::vector<size_t> sizes;
for (size_t i = 0; i < numTestSets_; ++i) {
sizes.push_back(std::count(testSets_[i].begin(),
testSets_[i].end(), true));
}
return sizes;
}
// 打印摘要信息(模仿 MATLAB 的显示格式)
void disp() const {
std::cout << typeName() << " cross validation partition\n";
std::cout << " NumObservations: " << n_ << "\n";
std::cout << " NumTestSets: " << numTestSets_ << "\n";
auto trSizes = TrainSize();
auto teSizes = TestSize();
std::cout << " TrainSize: ";
for (size_t i = 0; i < trSizes.size(); ++i) {
if (i > 0) std::cout << " ";
std::cout << trSizes[i];
}
std::cout << "\n";
std::cout << " TestSize: ";
for (size_t i = 0; i < teSizes.size(); ++i) {
if (i > 0) std::cout << " ";
std::cout << teSizes[i];
}
std::cout << "\n";
}
private:
size_t n_;
Type type_;
size_t numTestSets_;
bool isStratified_;
std::mt19937 rng_;
std::vector<std::vector<bool>> testSets_;
std::string typeName() const {
switch (type_) {
case Type::KFold: return "K-fold";
case Type::HoldOut: return "Hold-out";
case Type::LeaveOut: return "Leave-one-out";
case Type::Resubstitution: return "Resubstitution";
}
return "";
}
void validateAndSetup(size_t n, Type type, double k) {
if (n == 0) throw std::invalid_argument("Number of observations must be positive.");
switch (type) {
case Type::KFold:
if (k < 2 || k > static_cast<double>(n))
throw std::invalid_argument("k must be in [2, n].");
numTestSets_ = static_cast<size_t>(k);
break;
case Type::HoldOut:
numTestSets_ = 1;
break;
case Type::LeaveOut:
numTestSets_ = n;
break;
case Type::Resubstitution:
numTestSets_ = 1;
break;
}
}
// ======================== 非分层划分 ========================
void generateNonStratifiedPartition(double k) {
testSets_.clear();
std::vector<size_t> indices(n_);
std::iota(indices.begin(), indices.end(), 0);
std::shuffle(indices.begin(), indices.end(), rng_);
switch (type_) {
case Type::KFold: {
size_t baseSize = n_ / numTestSets_;
size_t remainder = n_ % numTestSets_;
size_t offset = 0;
for (size_t f = 0; f < numTestSets_; ++f) {
size_t foldSize = baseSize + (f < remainder ? 1 : 0);
std::vector<bool> testIdx(n_, false);
for (size_t j = 0; j < foldSize; ++j)
testIdx[indices[offset + j]] = true;
testSets_.push_back(testIdx);
offset += foldSize;
}
break;
}
case Type::HoldOut: {
double p = 0.1; // MATLAB 默认 HoldOut 比例为 0.1
if (k > 0 && k < 1.0) p = k;
size_t testSize = static_cast<size_t>(std::round(p * n_));
if (testSize == 0) testSize = 1;
if (testSize >= n_) testSize = n_ - 1;
std::vector<bool> testIdx(n_, false);
for (size_t j = 0; j < testSize; ++j)
testIdx[indices[j]] = true;
testSets_.push_back(testIdx);
break;
}
case Type::LeaveOut: {
for (size_t j = 0; j < n_; ++j) {
std::vector<bool> testIdx(n_, false);
testIdx[indices[j]] = true;
testSets_.push_back(testIdx);
}
break;
}
case Type::Resubstitution: {
std::vector<bool> testIdx(n_, true);
testSets_.push_back(testIdx);
break;
}
}
}
// ======================== 分层划分 ========================
template<typename T>
void generateStratifiedPartition(const std::vector<T>& y, double k) {
testSets_.clear();
// 按类别分组索引
std::unordered_map<T, std::vector<size_t>> classIndices;
for (size_t i = 0; i < n_; ++i)
classIndices[y[i]].push_back(i);
// 每类内部打乱
for (auto& kv : classIndices)
std::shuffle(kv.second.begin(), kv.second.end(), rng_);
std::vector<std::vector<bool>> testSets(numTestSets_,
std::vector<bool>(n_, false));
switch (type_) {
case Type::KFold: {
// 每类轮流分配到各折,保证比例近似
for (auto& kv : classIndices) {
const auto& vec = kv.second;
size_t fold = 0;
for (size_t idx : vec) {
testSets[fold][idx] = true;
fold = (fold + 1) % numTestSets_;
}
}
break;
}
case Type::HoldOut: {
double p = 0.1;
if (k > 0 && k < 1.0) p = k;
for (auto& kv : classIndices) {
const auto& vec = kv.second;
size_t testSize = static_cast<size_t>(std::round(p * vec.size()));
if (testSize == 0 && !vec.empty()) testSize = 1;
for (size_t j = 0; j < testSize; ++j)
testSets[0][vec[j]] = true;
}
break;
}
case Type::LeaveOut: {
for (size_t i = 0; i < n_; ++i)
testSets[i][i] = true;
break;
}
case Type::Resubstitution: {
std::fill(testSets[0].begin(), testSets[0].end(), true);
break;
}
}
testSets_ = std::move(testSets);
}
};
// ============================================================
// 主函数:演示与 MATLAB 一致的用法
// ============================================================
int main() {
// 示例 1:非分层 K 折(n=10, k=3)
std::cout << "===== 非分层 KFold (n=10, k=3) =====\n";
CVPartition cvp1(10, CVPartition::Type::KFold, 3);
cvp1.disp();
std::cout << "\n";
// 示例 2:分层 K 折(基于标签)
std::cout << "===== 分层 KFold (9 样本, 3 类, k=3) =====\n";
std::vector<int> labels = {0, 0, 0, 1, 1, 1, 2, 2, 2};
CVPartition cvp2(labels, CVPartition::Type::KFold, 3);
cvp2.disp();
std::cout << "\n";
// 示例 3:HoldOut(默认比例 0.1)
std::cout << "===== 非分层 HoldOut (n=20, 默认 p=0.1) =====\n";
CVPartition cvp3(20, CVPartition::Type::HoldOut);
cvp3.disp();
std::cout << "\n";
// 示例 4:分层 HoldOut(指定比例 0.3)
std::cout << "===== 分层 HoldOut (labels, p=0.3) =====\n";
std::vector<std::string> groups = {
"A","A","A","A","A",
"B","B","B","B","B",
"C","C","C","C","C"
};
CVPartition cvp4(groups, CVPartition::Type::HoldOut, 0.3);
cvp4.disp();
std::cout << "\n";
// 示例 5:LeaveOut
std::cout << "===== LeaveOut (n=5) =====\n";
CVPartition cvp5(5, CVPartition::Type::LeaveOut);
cvp5.disp();
std::cout << "\n";
return 0;
}
相关推荐
黑妹天下第一乖1 小时前
小智改造实战解读-首 token 延迟去哪了:云端与端侧大模型的分段对照
开发语言·人工智能·python·嵌入式硬件·自然语言处理·iot
。小二1 小时前
Go 泛型并发利刃:async 库深度评测——370K ops/s、零依赖、生产就绪
开发语言·javascript·golang
(Charon)1 小时前
【C++面试】线程同步机制:互斥锁、条件变量、原子操作与读写锁
开发语言·c++·算法·面试
七牛云行业应用2 小时前
刚刚!GPT-6.1 Sol Ultrafast 上线:速度、价格、API 接入全指南
算法
Omics Pro2 小时前
研究证实AI虚拟细胞可用于药物靶点发现
数据库·人工智能·算法·机器学习·自然语言处理
在所不辞兄2 小时前
常微分方程详解与应用
人工智能·神经网络·算法·机器学习
-dzk-2 小时前
【动态规划】LC 416.分割等和子集
算法·动态规划·代理模式
CV工程师丁Sir2 小时前
ArkWeb 手记 06|window.open 与页面跳转拦截
开发语言·php·harmonyos
夜不会漫长2 小时前
C++:内存管理
java·开发语言·c++