use crate::{ModelError, ParameterSlice};
#[derive(Debug)]
pub struct BlockObjective<'a, O> {
full_objective: &'a mut O,
full_beta: Vec<f64>,
working_beta: Vec<f64>,
full_grad: Vec<f64>,
block: ParameterSlice,
}
impl<'a, O> BlockObjective<'a, O>
where
O: Objective,
{
pub fn try_new(
full_objective: &'a mut O,
full_beta: Vec<f64>,
block: ParameterSlice,
) -> Result<Self, ModelError> {
let full_grad = vec![0.0; full_objective.dim()];
validate_block_range(&block, full_beta.len())?;
if full_grad.len() != full_beta.len() {
return Err(ModelError::BetaLength {
expected: full_grad.len(),
actual: full_beta.len(),
});
}
let working_beta = full_beta[block.range.clone()].to_vec();
debug_assert_eq!(working_beta.len(), block.range.len());
Ok(Self {
full_objective,
full_beta,
working_beta,
full_grad,
block,
})
}
fn update_block_beta(&mut self, block_beta: &[f64]) {
self.working_beta.copy_from_slice(block_beta);
self.full_beta[self.block.range.clone()].copy_from_slice(&self.working_beta);
}
}
impl<O> Objective for BlockObjective<'_, O>
where
O: Objective,
O::Error: From<ModelError>,
{
type Error = O::Error;
fn dim(&self) -> usize {
self.block.range.len()
}
fn value(&mut self, block_beta: &[f64]) -> Result<f64, Self::Error> {
validate_block_len("parameters", block_beta.len(), self.block.range.len())?;
self.update_block_beta(block_beta);
self.full_objective.value(&self.full_beta)
}
fn gradient(&mut self, block_beta: &[f64], grad: &mut [f64]) -> Result<(), Self::Error> {
self.value_gradient(block_beta, grad).map(|_| ())
}
fn value_gradient(&mut self, block_beta: &[f64], grad: &mut [f64]) -> Result<f64, Self::Error> {
validate_block_len("parameters", block_beta.len(), self.block.range.len())?;
validate_block_len("gradient", grad.len(), self.block.range.len())?;
self.update_block_beta(block_beta);
let value = self
.full_objective
.value_gradient(&self.full_beta, &mut self.full_grad)?;
grad.copy_from_slice(&self.full_grad[self.block.range.clone()]);
Ok(value)
}
}
pub trait Objective {
type Error;
fn dim(&self) -> usize;
fn value(&mut self, parameters: &[f64]) -> Result<f64, Self::Error>;
fn value_gradient(&mut self, parameters: &[f64], grad: &mut [f64]) -> Result<f64, Self::Error>;
fn gradient(&mut self, parameters: &[f64], grad: &mut [f64]) -> Result<(), Self::Error> {
self.value_gradient(parameters, grad).map(|_| ())
}
}
fn validate_block_len(
name: &'static str,
actual: usize,
expected: usize,
) -> Result<(), ModelError> {
if actual == expected {
Ok(())
} else if name == "gradient" {
Err(ModelError::GradientLength { expected, actual })
} else {
Err(ModelError::BetaLength { expected, actual })
}
}
const fn validate_block_range(block: &ParameterSlice, dim: usize) -> Result<(), ModelError> {
if block.range.start <= block.range.end && block.range.end <= dim {
Ok(())
} else {
Err(ModelError::BlockRangeOutOfBounds {
parameter: block.name,
start: block.range.start,
end: block.range.end,
dim,
})
}
}
#[cfg(test)]
mod tests {
use super::{BlockObjective, Objective};
use crate::{ModelError, ParameterSlice};
#[derive(Debug)]
struct QuadraticObjective {
dim: usize,
}
impl Objective for QuadraticObjective {
type Error = ModelError;
fn dim(&self) -> usize {
self.dim
}
fn value(&mut self, parameters: &[f64]) -> Result<f64, Self::Error> {
Ok(0.5 * parameters.iter().map(|value| value * value).sum::<f64>())
}
fn value_gradient(
&mut self,
parameters: &[f64],
grad: &mut [f64],
) -> Result<f64, Self::Error> {
grad.copy_from_slice(parameters);
self.value(parameters)
}
}
#[test]
#[allow(clippy::float_cmp)]
fn block_objective_reuses_working_buffers_on_repeated_calls() {
let mut full = QuadraticObjective { dim: 3 };
let mut objective = BlockObjective::try_new(
&mut full,
vec![1.0, 2.0, 3.0],
ParameterSlice {
name: "sigma",
range: 1..3,
},
)
.unwrap();
let full_beta_capacity = objective.full_beta.capacity();
let working_beta_capacity = objective.working_beta.capacity();
let grad_capacity = objective.full_grad.capacity();
let mut grad = vec![0.0; objective.dim()];
assert_eq!(objective.dim(), 2);
assert_eq!(objective.value(&[4.0, 5.0]).unwrap(), 21.0);
assert_eq!(
objective.value_gradient(&[6.0, 7.0], &mut grad).unwrap(),
43.0
);
assert_eq!(grad, vec![6.0, 7.0]);
assert_eq!(objective.full_beta, vec![1.0, 6.0, 7.0]);
assert_eq!(objective.working_beta, vec![6.0, 7.0]);
assert_eq!(objective.full_beta.capacity(), full_beta_capacity);
assert_eq!(objective.working_beta.capacity(), working_beta_capacity);
assert_eq!(objective.full_grad.capacity(), grad_capacity);
assert_eq!(objective.value(&[8.0, 9.0]).unwrap(), 73.0);
assert_eq!(objective.full_beta.capacity(), full_beta_capacity);
assert_eq!(objective.working_beta.capacity(), working_beta_capacity);
assert_eq!(objective.full_grad.capacity(), grad_capacity);
}
#[test]
fn try_new_rejects_wrong_full_beta_length() {
let mut full = QuadraticObjective { dim: 3 };
assert_eq!(
BlockObjective::try_new(
&mut full,
vec![1.0, 2.0],
ParameterSlice {
name: "mu",
range: 0..2,
},
)
.unwrap_err(),
ModelError::BetaLength {
expected: 3,
actual: 2,
}
);
}
#[test]
fn try_new_rejects_block_outside_full_beta() {
let mut full = QuadraticObjective { dim: 3 };
assert_eq!(
BlockObjective::try_new(
&mut full,
vec![1.0, 2.0, 3.0],
ParameterSlice {
name: "mu",
range: 2..4,
},
)
.unwrap_err(),
ModelError::BlockRangeOutOfBounds {
parameter: "mu",
start: 2,
end: 4,
dim: 3,
}
);
}
}