use laddu_compile::CompiledModel;
use laddu_runtime::PreparedModel;
use numpy::{PyArray1, PyArray2};
use pyo3::{
exceptions::PyTypeError,
prelude::*,
types::{PyAny, PyDict},
};
use super::{data::PyDataset, error::to_py_err, expr::PyExpr, runtime::PyExecution};
pub fn model_free_values(model: &CompiledModel, values: &Bound<'_, PyAny>) -> PyResult<Vec<f64>> {
if let Ok(values) = values.extract::<Vec<f64>>() {
return Ok(values);
}
if let Ok(mapping) = values.cast::<PyDict>() {
let mut out = model.params().initial_free_values();
for (index, id) in model.params().free_params().iter().enumerate() {
let name = model.params().name(*id).map_err(to_py_err)?;
if let Some(value) = mapping.get_item(name)? {
out[index] = value.extract()?;
}
}
return Ok(out);
}
Err(PyTypeError::new_err(
"parameters must be a numeric sequence or dict keyed by parameter name",
))
}
#[pyclass(name = "Model", module = "laddu", frozen, skip_from_py_object)]
#[derive(Clone)]
pub struct PyModel {
pub(crate) inner: CompiledModel,
}
#[pymethods]
impl PyModel {
#[new]
fn new(expr: &PyExpr) -> PyResult<Self> {
Ok(Self {
inner: CompiledModel::from_expr(&expr.inner).map_err(to_py_err)?,
})
}
fn __repr__(&self) -> String {
format!("Model(parameters={:?})", self.parameter_names())
}
#[getter]
fn parameter_names(&self) -> Vec<String> {
self.inner
.params()
.free_params()
.iter()
.map(|id| {
self.inner
.params()
.name(*id)
.unwrap_or("<invalid>")
.to_owned()
})
.collect()
}
#[getter]
fn default_parameters(&self) -> Vec<f64> {
self.inner.params().initial_free_values()
}
#[pyo3(signature = (*, seed=0))]
fn sample_parameters(&self, seed: u64) -> Vec<f64> {
self.inner.params().sample_initial(seed)
}
fn projection(&self, tags: Vec<String>) -> PyResult<Self> {
Ok(Self {
inner: self
.inner
.project_tags(tags.iter().map(String::as_str))
.map_err(to_py_err)?,
})
}
fn fix(&self, name: &str, value: f64) -> PyResult<Self> {
Ok(Self {
inner: self.inner.fix_parameter(name, value).map_err(to_py_err)?,
})
}
fn free(&self, name: &str) -> PyResult<Self> {
Ok(Self {
inner: self.inner.free_parameter(name).map_err(to_py_err)?,
})
}
#[pyo3(signature = (
dataset,
*,
parameters: "Sequence[float] | dict[str, float] | None" = None,
execution=None,
real=false
) -> "Sequence[float]")]
fn evaluate<'py>(
&self,
py: Python<'py>,
dataset: &PyDataset,
parameters: Option<&Bound<'_, PyAny>>,
execution: Option<&PyExecution>,
real: bool,
) -> PyResult<Bound<'py, PyAny>> {
let execution = execution
.cloned()
.map(Ok)
.unwrap_or_else(PyExecution::default_inner)?;
let free = match parameters {
Some(values) => model_free_values(&self.inner, values)?,
None => self.inner.params().initial_free_values(),
};
let params = self.inner.params().values(&free).map_err(to_py_err)?;
let plan = PreparedModel::prepare(&self.inner, &execution.inner).map_err(to_py_err)?;
let dataset = dataset.inner.clone();
let values = py
.detach(move || {
let mut values = Vec::new();
for batch in dataset
.batches()
.map_err(|error| laddu_runtime::RuntimeError::Data(error.to_string()))?
{
values.extend(plan.evaluate_batch(
¶ms,
&batch.map_err(|error| {
laddu_runtime::RuntimeError::Data(error.to_string())
})?,
)?);
}
Ok::<_, laddu_runtime::RuntimeError>(values)
})
.map_err(to_py_err)?;
if real {
Ok(
PyArray1::from_vec(py, values.into_iter().map(|value| value.re).collect())
.into_any(),
)
} else {
Ok(PyArray1::from_vec(py, values).into_any())
}
}
#[pyo3(signature = (
dataset,
*,
parameters: "Sequence[float] | dict[str, float] | None" = None,
execution=None,
real=false
) -> "tuple[Sequence[float], Sequence[Sequence[float]]]")]
fn value_and_gradient<'py>(
&self,
py: Python<'py>,
dataset: &PyDataset,
parameters: Option<&Bound<'_, PyAny>>,
execution: Option<&PyExecution>,
real: bool,
) -> PyResult<(Bound<'py, PyAny>, Bound<'py, PyAny>)> {
let execution = execution
.cloned()
.map(Ok)
.unwrap_or_else(PyExecution::default_inner)?;
let free = match parameters {
Some(values) => model_free_values(&self.inner, values)?,
None => self.inner.params().initial_free_values(),
};
let params = self.inner.params().values(&free).map_err(to_py_err)?;
let plan = PreparedModel::prepare(&self.inner, &execution.inner).map_err(to_py_err)?;
let dataset = dataset.inner.clone();
let evaluations = py
.detach(move || {
let mut evaluations = Vec::new();
for batch in dataset
.batches()
.map_err(|error| laddu_runtime::RuntimeError::Data(error.to_string()))?
{
evaluations.extend(plan.evaluate_batch_with_gradient(
¶ms,
&batch.map_err(|error| {
laddu_runtime::RuntimeError::Data(error.to_string())
})?,
)?);
}
Ok::<_, laddu_runtime::RuntimeError>(evaluations)
})
.map_err(to_py_err)?;
if real {
let values = evaluations.iter().map(|value| value.value().re).collect();
let gradients = evaluations
.iter()
.map(|value| value.gradient().iter().map(|entry| entry.re).collect())
.collect::<Vec<Vec<f64>>>();
Ok((
PyArray1::from_vec(py, values).into_any(),
PyArray2::from_vec2(py, &gradients)?.into_any(),
))
} else {
let values = evaluations.iter().map(|value| value.value()).collect();
let gradients = evaluations
.iter()
.map(|value| value.gradient().to_vec())
.collect::<Vec<_>>();
Ok((
PyArray1::from_vec(py, values).into_any(),
PyArray2::from_vec2(py, &gradients)?.into_any(),
))
}
}
}