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

cpp 复制代码
// crossval.cpp
// 编译: g++ -std=c++17 -O2 -o crossval crossval.cpp
// 运行: ./crossval_demo

#include <vector>
#include <functional>
#include <numeric>
#include <cmath>
#include <stdexcept>
#include <algorithm>
#include <random>
#include <iostream>

namespace ml {

// ==================== 配置与类型 ====================

struct CrossValOptions {
    int kfold = 10;                 // 折数,默认 10
    bool leaveout = false;          // Leave-one-out
    bool holdout = false;           // Holdout 模式
    double holdout_ratio = 0.2;     // Holdout 测试集比例
    unsigned int seed = 42;         // 随机种子
};

using Partition = std::vector<std::vector<size_t>>;

// ==================== 分区生成 ====================

inline Partition make_kfold_partition(size_t n, int k, unsigned int seed = 42) {
    if (k <= 1) throw std::invalid_argument("kfold must be > 1");
    if (k > static_cast<int>(n)) k = static_cast<int>(n);

    std::vector<size_t> idx(n);
    std::iota(idx.begin(), idx.end(), 0);

    std::mt19937 rng(seed);
    std::shuffle(idx.begin(), idx.end(), rng);

    Partition folds(k);
    for (size_t i = 0; i < n; ++i) {
        folds[i % k].push_back(idx[i]);
    }
    return folds;
}

inline Partition make_leaveout_partition(size_t n) {
    Partition folds(n);
    for (size_t i = 0; i < n; ++i) {
        folds[i] = { i };
    }
    return folds;
}

inline Partition make_holdout_partition(size_t n, double ratio, unsigned int seed = 42) {
    size_t ntest = std::max<size_t>(1, static_cast<size_t>(n * ratio));

    std::vector<size_t> idx(n);
    std::iota(idx.begin(), idx.end(), 0);

    std::mt19937 rng(seed);
    std::shuffle(idx.begin(), idx.end(), rng);

    Partition folds(1);
    folds[0].assign(idx.begin(), idx.begin() + ntest);
    return folds;
}

inline Partition make_partition(size_t n, const CrossValOptions& opts) {
    if (opts.leaveout) {
        return make_leaveout_partition(n);
    }
    if (opts.holdout) {
        return make_holdout_partition(n, opts.holdout_ratio, opts.seed);
    }
    return make_kfold_partition(n, opts.kfold, opts.seed);
}

inline std::vector<size_t> complement_indices(size_t n, const std::vector<size_t>& test_idx) {
    std::vector<bool> is_test(n, false);
    for (size_t i : test_idx) is_test[i] = true;

    std::vector<size_t> train_idx;
    train_idx.reserve(n - test_idx.size());
    for (size_t i = 0; i < n; ++i) {
        if (!is_test[i]) train_idx.push_back(i);
    }
    return train_idx;
}

// ==================== 核心 crossval 函数 ====================

// 通用交叉验证:对每个折调用 fun(train_idx, test_idx) 并收集结果
template <typename Func>
std::vector<double> crossval(Func&& fun,
                             size_t n_samples,
                             const CrossValOptions& opts = {}) {
    Partition folds = make_partition(n_samples, opts);
    std::vector<double> results;
    results.reserve(folds.size());

    for (const auto& test_idx : folds) {
        std::vector<size_t> train_idx = complement_indices(n_samples, test_idx);
        double val = fun(train_idx, test_idx);
        results.push_back(val);
    }
    return results;
}

// MSE 交叉验证
template <typename PredFun, typename YVec>
double crossval_mse(PredFun&& predfun,
                    const YVec& y,
                    const CrossValOptions& opts = {}) {
    size_t n = y.size();
    Partition folds = make_partition(n, opts);

    double total_se = 0.0;
    size_t total_count = 0;

    for (const auto& test_idx : folds) {
        std::vector<size_t> train_idx = complement_indices(n, test_idx);

        std::vector<double> yfit = predfun(train_idx, test_idx);
        if (yfit.size() != test_idx.size()) {
            throw std::runtime_error("predfun returned wrong number of predictions");
        }

        for (size_t j = 0; j < test_idx.size(); ++j) {
            double err = y[test_idx[j]] - yfit[j];
            total_se += err * err;
        }
        total_count += test_idx.size();
    }

    return total_se / static_cast<double>(total_count);
}

// MCR 交叉验证
template <typename PredFun, typename YVec>
double crossval_mcr(PredFun&& predfun,
                    const YVec& y,
                    const CrossValOptions& opts = {}) {
    size_t n = y.size();
    Partition folds = make_partition(n, opts);

    size_t misclassified = 0;
    size_t total_count = 0;

    for (const auto& test_idx : folds) {
        std::vector<size_t> train_idx = complement_indices(n, test_idx);

        auto yfit = predfun(train_idx, test_idx);
        if (yfit.size() != test_idx.size()) {
            throw std::runtime_error("predfun returned wrong number of predictions");
        }

        for (size_t j = 0; j < test_idx.size(); ++j) {
            if (yfit[j] != y[test_idx[j]]) {
                ++misclassified;
            }
        }
        total_count += test_idx.size();
    }

    return static_cast<double>(misclassified) / static_cast<double>(total_count);
}

// 值函数交叉验证
template <typename Func>
std::vector<double> crossval_values(Func&& fun,
                                    size_t n_samples,
                                    const CrossValOptions& opts = {}) {
    Partition folds = make_partition(n_samples, opts);
    std::vector<double> vals;
    vals.reserve(folds.size());

    for (const auto& test_idx : folds) {
        vals.push_back(fun(test_idx));
    }
    return vals;
}

} // namespace ml

// ==================== 示例程序 ====================

int main() {
    // 模拟数据:100 个样本,3 个特征,线性关系 y = 2*x1 + 3*x2 - x3 + noise
    const size_t n = 100;
    const size_t p = 3;
    std::vector<std::vector<double>> X(n, std::vector<double>(p));
    std::vector<double> y(n);

    std::mt19937 rng(123);
    std::normal_distribution<double> noise(0.0, 0.5);
    std::uniform_real_distribution<double> uni(-2.0, 2.0);

    for (size_t i = 0; i < n; ++i) {
        for (size_t j = 0; j < p; ++j) X[i][j] = uni(rng);
        y[i] = 2.0 * X[i][0] + 3.0 * X[i][1] - 1.0 * X[i][2] + noise(rng);
    }

    // 预测函数:最小二乘线性回归
    auto predfun = [&](const std::vector<size_t>& train_idx,
                       const std::vector<size_t>& test_idx) -> std::vector<double> {
        size_t ntrain = train_idx.size();
        std::vector<std::vector<double>> A(ntrain, std::vector<double>(p + 1));
        std::vector<double> b(ntrain);
        for (size_t i = 0; i < ntrain; ++i) {
            A[i][0] = 1.0;
            for (size_t j = 0; j < p; ++j) A[i][j + 1] = X[train_idx[i]][j];
            b[i] = y[train_idx[i]];
        }

        size_t d = p + 1;
        std::vector<std::vector<double>> AtA(d, std::vector<double>(d, 0.0));
        std::vector<double> Atb(d, 0.0);

        for (size_t i = 0; i < ntrain; ++i) {
            for (size_t r = 0; r < d; ++r) {
                for (size_t c = 0; c < d; ++c) {
                    AtA[r][c] += A[i][r] * A[i][c];
                }
                Atb[r] += A[i][r] * b[i];
            }
        }

        // 高斯消元求解正规方程
        for (size_t col = 0; col < d; ++col) {
            size_t pivot = col;
            for (size_t r = col + 1; r < d; ++r) {
                if (std::abs(AtA[r][col]) > std::abs(AtA[pivot][col])) pivot = r;
            }
            std::swap(AtA[col], AtA[pivot]);
            std::swap(Atb[col], Atb[pivot]);

            double div = AtA[col][col];
            for (size_t c = col; c < d; ++c) AtA[col][c] /= div;
            Atb[col] /= div;

            for (size_t r = 0; r < d; ++r) {
                if (r == col) continue;
                double factor = AtA[r][col];
                for (size_t c = col; c < d; ++c) AtA[r][c] -= factor * AtA[col][c];
                Atb[r] -= factor * Atb[col];
            }
        }

        std::vector<double> yfit(test_idx.size());
        for (size_t i = 0; i < test_idx.size(); ++i) {
            double pred = Atb[0];
            for (size_t j = 0; j < p; ++j) {
                pred += Atb[j + 1] * X[test_idx[i]][j];
            }
            yfit[i] = pred;
        }
        return yfit;
    };

    // ---- 5 折交叉验证 ----
    ml::CrossValOptions opts;
    opts.kfold = 5;
    opts.seed = 7;
    double mse5 = ml::crossval_mse(predfun, y, opts);
    std::cout << "5-fold CV MSE = " << mse5 << std::endl;

    // ---- 10 折交叉验证 ----
    opts.kfold = 10;
    double mse10 = ml::crossval_mse(predfun, y, opts);
    std::cout << "10-fold CV MSE = " << mse10 << std::endl;

    // ---- Leave-one-out ----
    opts.leaveout = true;
    double mse_loo = ml::crossval_mse(predfun, y, opts);
    std::cout << "Leave-one-out CV MSE = " << mse_loo << std::endl;

    // ---- 值函数形式:返回每折测试集样本数 ----
    opts.leaveout = false;
    opts.kfold = 10;
    auto fold_size_fun = [](const std::vector<size_t>& test_idx) -> double {
        return static_cast<double>(test_idx.size());
    };
    auto vals = ml::crossval_values(fold_size_fun, n, opts);
    std::cout << "Per-fold test sizes (10-fold): ";
    for (double v : vals) std::cout << v << " ";
    std::cout << std::endl;

    return 0;
}
相关推荐
ttwuai1 小时前
Go后台管理系统开源项目:4个官方仓库怎么追溯和核验
开发语言·golang·开源
谢亮_vipxieliang1 小时前
Go Worker Pool 设计——从原理到生产级实现
开发语言·后端·golang
青少儿编程课堂1 小时前
背包问题(0/1 背包与完全背包)解题精讲——动态规划入门
c++·python·算法·bfs·信息学竞赛
倒头就睡的小比特1 小时前
C++多态
c++
光影少年2 小时前
为什么 JavaScript 中 0.1 + 0.2 !== 0.3,如何让其相等?
前端·javascript·算法
Escalating_xu2 小时前
【C 语言】深入理解指针(5):sizeof、strlen、数组名语义与指针笔试题全解析
java·c语言·开发语言
HEJOO92 小时前
Linux 设备驱动系列——使用全局工作队列
java·linux·算法
Code_Solitude2 小时前
c语言:递归
java·c语言·算法
Nebula_g2 小时前
JavaSE加强:Stream流
java·开发语言·windows·stream·javase·流