如何在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
相关产品推荐
相关产品推荐

