use crate::OptimizerOpts;
use crate::BurnBackend::NdArray32;
use crate::BurnBackend::Wgpu32;
use crate::optimize_flat;
use crate::optimizer::Algorithm;
use burn::backend::ndarray::NdArray;
use burn_tensor::{Shape, Tensor};
use cxx::CxxVector;
#[cxx::bridge(namespace = "mixt")]
mod ffi {
extern "Rust" {
fn rcg_optl_cpu(
logl: &CxxVector<f32>,
log_times_observed: &CxxVector<f32>,
alpha0: &CxxVector<f32>,
tol: f64,
max_iters: usize,
) -> Vec<f32>;
fn rcg_optl_gpu(
logl: &CxxVector<f32>,
log_times_observed: &CxxVector<f32>,
alpha0: &CxxVector<f32>,
tol: f64,
max_iters: usize,
) -> Vec<f32>;
fn em_cpu(
logl: &CxxVector<f32>,
log_times_observed: &CxxVector<f32>,
alpha0: &CxxVector<f32>,
tol: f64,
max_iters: usize,
) -> Vec<f32>;
fn em_gpu(
logl: &CxxVector<f32>,
log_times_observed: &CxxVector<f32>,
alpha0: &CxxVector<f32>,
tol: f64,
max_iters: usize,
) -> Vec<f32>;
fn mixture_components(
probs: &CxxVector<f32>,
log_times_observed: &CxxVector<f32>,
) -> Vec<f32>;
}
}
fn run_optimizer(
logl: &CxxVector<f32>,
log_times_observed: &CxxVector<f32>,
alpha0: &CxxVector<f32>,
options: OptimizerOpts,
) -> Vec<f32> {
let (_, probs) = optimize_flat(logl.as_slice(), log_times_observed.as_slice(), alpha0.as_slice(), Some(options)).unwrap();
probs
}
pub fn rcg_optl_cpu(
logl: &CxxVector<f32>,
log_times_observed: &CxxVector<f32>,
alpha0: &CxxVector<f32>,
tolerance: f64,
max_iters: usize,
) -> Vec<f32> {
let options = OptimizerOpts { tolerance, max_iters, device: NdArray32, algorithm: Algorithm::RCG };
run_optimizer(logl, log_times_observed, alpha0, options)
}
#[allow(unused_variables)]
pub fn rcg_optl_gpu(
logl: &CxxVector<f32>,
log_times_observed: &CxxVector<f32>,
alpha0: &CxxVector<f32>,
tolerance: f64,
max_iters: usize,
) -> Vec<f32> {
let options = OptimizerOpts { tolerance, max_iters, device: Wgpu32, algorithm: Algorithm::RCG };
#[cfg(any(feature = "wgpu", feature = "webgpu", feature = "vulkan"))]
return run_optimizer(logl, log_times_observed, alpha0, options);
#[cfg(not(any(feature = "wgpu", feature = "webgpu", feature = "vulkan")))]
panic!("mixt: mixt was not compiled with GPU support.")
}
pub fn em_cpu(
logl: &CxxVector<f32>,
log_times_observed: &CxxVector<f32>,
alpha0: &CxxVector<f32>,
tolerance: f64,
max_iters: usize,
) -> Vec<f32> {
let options = OptimizerOpts { tolerance, max_iters, device: NdArray32, algorithm: Algorithm::EM };
run_optimizer(logl, log_times_observed, alpha0, options)
}
#[allow(unused_variables)]
pub fn em_gpu(
logl: &CxxVector<f32>,
log_times_observed: &CxxVector<f32>,
alpha0: &CxxVector<f32>,
tolerance: f64,
max_iters: usize,
) -> Vec<f32> {
let options = OptimizerOpts { tolerance, max_iters, device: Wgpu32, algorithm: Algorithm::EM };
#[cfg(any(feature = "wgpu", feature = "webgpu", feature = "vulkan"))]
return run_optimizer(logl, log_times_observed, alpha0, options);
#[cfg(not(any(feature = "wgpu", feature = "webgpu", feature = "vulkan")))]
panic!("mixt: mixt was not compiled with GPU support.")
}
pub fn mixture_components(
probs: &cxx::CxxVector<f32>,
log_times_observed: &cxx::CxxVector<f32>,
) -> Vec<f32> {
let n_obs = log_times_observed.len();
let n_targets = probs.len()/n_obs;
let device = Default::default();
type Backend = NdArray<f32>;
let probs_t: Tensor::<Backend, 2> = Tensor::<Backend, 1>::from_data(probs.as_slice(), &device).reshape(Shape::new([n_targets, n_obs]));
let log_counts_t = Tensor::<Backend, 1>::from_data(log_times_observed.as_slice(), &device);
let thetas_t = crate::optimizer::mixture_components(probs_t, log_counts_t);
thetas_t.into_data().to_vec().unwrap()
}