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

如何在C++中通过条件逻辑初始化std::random中的分布

问题描述

我尝试根据条件逻辑从不同分布中采样随机数,但找不到合适的实现方式。我定义了如下表示分布的结构体:

struct Distribution {
    std::string name;
    double args[2];
};

标准正态分布可表示为:

Distribution normal = {"normal", {0, 1}};

我的目标是:给定一个Distribution数组,为每个分布生成数千个样本。但std::random中的不同分布类型不同,导致我无法在采样前初始化分布。

我希望能在采样循环外完成分布的初始化(类似如下伪代码逻辑),但不同分布类型无法存入同一个变量:

struct Distribution {
    std::string name;
    double args[2];
};

int main(void) {
    Distribution distrs[2] = {
        {"uniform", {0, 1}},
        {"normal",  {0, 1}}
    };

    int n_samples = 100;
    double samples[200];

    std::random_device rd;
    std::mt19937 mt(rd());

    for (int ix = 0; ix < 2; ix++) {
        // 这里需要一个通用类型存储不同分布,但直接写会编译失败
        some_common_type_for_distr sampler; 
        std::string distr_name = distrs[ix].name;
        double args[2] = distrs[ix].args;
        if (distr_name == "uniform") {
            sampler = std::uniform_real_distribution<double>(args[0], args[1]);
        } else if (distr_name == "normal") {
            sampler = std::normal_distribution<double>(args[0], args[1]);
        }

        // 采样循环,希望只初始化一次分布,多次采样
        for (int jx = 0; jx < n_samples; jx++) { 
            samples[jx + ix*n_samples] = sampler(mt);
        }
    }

    // 处理样本
}

目前只能在采样循环内重复初始化分布,每次采样都要重新创建分布对象,效率较低:

struct Distribution {
    std::string name;
    double args[2];
};

int main(void) {
    Distribution distrs[2] = {
        {"uniform", {0, 1}},
        {"normal",  {0, 1}}
    };

    int n_samples = 100;
    double samples[200];

    std::random_device rd;
    std::mt19937 mt(rd());

    for (int ix = 0; ix < 2; ix++) {
        for (int jx = 0; jx < n_samples; jx++) { 
            std::string distr_name = distrs[ix].name;
            double args[2] = distrs[ix].args;
            if (distr_name == "uniform") {
                std::uniform_real_distribution<double> sampler(args[0], args[1]);
                samples[jx + ix*n_samples] = sampler(mt);
            } else if (distr_name == "normal") {
                std::normal_distribution<double> sampler(args[0], args[1]);
                samples[jx + ix*n_samples] = sampler(mt);
            }
        }    
    }

    // 处理样本
}

我的问题是:是否可以像第一个main函数那样,在采样循环外完成分布的初始化?


解决方案

当然可以实现循环外初始化分布,下面提供几种可行的方法:

方法1:使用多态(基类 + 派生类)

定义一个抽象基类封装采样操作,为每种分布实现派生类,通过基类指针/引用统一管理不同分布对象。

#include <random>
#include <string>
#include <memory>

// 抽象基类
class Sampler {
public:
    virtual double sample(std::mt19937& mt) = 0;
    virtual ~Sampler() = default;
};

// 均匀分布采样器
class UniformSampler : public Sampler {
private:
    std::uniform_real_distribution<double> dist;
public:
    UniformSampler(double min, double max) : dist(min, max) {}
    double sample(std::mt19937& mt) override {
        return dist(mt);
    }
};

// 正态分布采样器
class NormalSampler : public Sampler {
private:
    std::normal_distribution<double> dist;
public:
    NormalSampler(double mean, double stddev) : dist(mean, stddev) {}
    double sample(std::mt19937& mt) override {
        return dist(mt);
    }
};

// 根据配置创建采样器
std::unique_ptr<Sampler> create_sampler(const std::string& name, double args[2]) {
    if (name == "uniform") {
        return std::make_unique<UniformSampler>(args[0], args[1]);
    } else if (name == "normal") {
        return std::make_unique<NormalSampler>(args[0], args[1]);
    }
    return nullptr;
}

struct Distribution {
    std::string name;
    double args[2];
};

int main() {
    Distribution distrs[2] = {
        {"uniform", {0, 1}},
        {"normal",  {0, 1}}
    };

    const int n_samples = 100;
    double samples[200];

    std::random_device rd;
    std::mt19937 mt(rd());

    for (int ix = 0; ix < 2; ix++) {
        // 循环外初始化采样器
        auto sampler = create_sampler(distrs[ix].name, distrs[ix].args);
        if (!sampler) continue;

        // 多次采样,无需重复初始化分布
        for (int jx = 0; jx < n_samples; jx++) {
            samples[jx + ix * n_samples] = sampler->sample(mt);
        }
    }

    // 处理样本
    return 0;
}

方法2:使用std::variant(C++17及以上)

std::variant可存储不同类型的对象,通过std::visit调用对应采样操作,无需多态,代码更简洁。

#include <random>
#include <string>
#include <variant>

struct Distribution {
    std::string name;
    double args[2];
};

int main() {
    Distribution distrs[2] = {
        {"uniform", {0, 1}},
        {"normal",  {0, 1}}
    };

    const int n_samples = 100;
    double samples[200];

    std::random_device rd;
    std::mt19937 mt(rd());

    // 定义可存储两种分布的variant类型
    using DistVariant = std::variant<std::uniform_real_distribution<double>,
                                    std::normal_distribution<double>>;

    for (int ix = 0; ix < 2; ix++) {
        DistVariant sampler;
        const auto& distr = distrs[ix];
        if (distr.name == "uniform") {
            sampler = std::uniform_real_distribution<double>(distr.args[0], distr.args[1]);
        } else if (distr.name == "normal") {
            sampler = std::normal_distribution<double>(distr.args[0], distr.args[1]);
        }

        // 使用std::visit调用采样逻辑
        for (int jx = 0; jx < n_samples; jx++) {
            samples[jx + ix * n_samples] = std::visit([&mt](auto& dist) {
                return dist(mt);
            }, sampler);
        }
    }

    // 处理样本
    return 0;
}

方法3:使用函数对象包装(std::function)

把分布的采样操作包装成std::function<double(std::mt19937&)>,直接存储不同分布的采样逻辑。

#include <random>
#include <string>
#include <functional>

struct Distribution {
    std::string name;
    double args[2];
};

int main() {
    Distribution distrs[2] = {
        {"uniform", {0, 1}},
        {"normal",  {0, 1}}
    };

    const int n_samples = 100;
    double samples[200];

    std::random_device rd;
    std::mt19937 mt(rd());

    for (int ix = 0; ix < 2; ix++) {
        std::function<double(std::mt19937&)> sampler;
        const auto& distr = distrs[ix];
        if (distr.name == "uniform") {
            std::uniform_real_distribution<double> dist(distr.args[0], distr.args[1]);
            sampler = [dist](std::mt19937& mt) mutable { return dist(mt); };
        } else if (distr.name == "normal") {
            std::normal_distribution<double> dist(distr.args[0], distr.args[1]);
            sampler = [dist](std::mt19937& mt) mutable { return dist(mt); };
        }

        for (int jx = 0; jx < n_samples; jx++) {
            samples[jx + ix * n_samples] = sampler(mt);
        }
    }

    // 处理样本
    return 0;
}

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.05 16:55:25