use std::cell::{Cell, RefCell};
use crate::discretizedasset::DiscretizedAsset;
use crate::errors::QlResult;
use crate::math::array::Array;
use crate::math::comparison::close;
use crate::math::timegrid::TimeGrid;
use crate::methods::lattices::lattice::Lattice;
use crate::methods::lattices::tree::Tree;
use crate::require;
use crate::types::{Real, Size, Time};
pub trait TreeLatticeImpl {
type Tree: Tree;
fn tree(&self) -> &Self::Tree;
fn discount(&self, i: Size, index: Size) -> Real;
}
pub struct TreeLattice<I: TreeLatticeImpl> {
implementation: I,
time_grid: TimeGrid,
state_prices: RefCell<Vec<Array>>,
state_prices_limit: Cell<Size>,
}
impl<I: TreeLatticeImpl> TreeLattice<I> {
pub fn new(implementation: I, time_grid: TimeGrid) -> QlResult<Self> {
require!(
<I::Tree as Tree>::BRANCHES > 0,
"there is no zeronomial lattice!"
);
Ok(TreeLattice {
implementation,
time_grid,
state_prices: RefCell::new(vec![Array::filled(1, 1.0)]),
state_prices_limit: Cell::new(0),
})
}
pub fn implementation(&self) -> &I {
&self.implementation
}
pub fn time_grid(&self) -> &TimeGrid {
&self.time_grid
}
fn size(&self, i: Size) -> Size {
self.implementation.tree().size(i)
}
fn discount(&self, i: Size, index: Size) -> Real {
self.implementation.discount(i, index)
}
fn descendant(&self, i: Size, index: Size, branch: Size) -> Size {
self.implementation.tree().descendant(i, index, branch)
}
fn probability(&self, i: Size, index: Size, branch: Size) -> Real {
self.implementation.tree().probability(i, index, branch)
}
pub fn state_prices(&self, i: Size) -> Array {
if i > self.state_prices_limit.get() {
self.compute_state_prices(i);
}
self.state_prices.borrow()[i].clone()
}
fn compute_state_prices(&self, until: Size) {
let branches = <I::Tree as Tree>::BRANCHES;
let mut state_prices = self.state_prices.borrow_mut();
for i in self.state_prices_limit.get()..until {
state_prices.push(Array::filled(self.size(i + 1), 0.0));
for j in 0..self.size(i) {
let discount = self.discount(i, j);
let state_price = state_prices[i][j];
for branch in 0..branches {
let destination = self.descendant(i, j, branch);
let probability = self.probability(i, j, branch);
state_prices[i + 1][destination] += state_price * discount * probability;
}
}
}
self.state_prices_limit.set(until);
}
fn stepback(&self, i: Size, values: &Array, new_values: &mut Array) {
let branches = <I::Tree as Tree>::BRANCHES;
for j in 0..self.size(i) {
let mut value = 0.0;
for branch in 0..branches {
value += self.probability(i, j, branch) * values[self.descendant(i, j, branch)];
}
value *= self.discount(i, j);
new_values[j] = value;
}
}
pub fn initialize(&self, asset: &mut dyn DiscretizedAsset, t: Time) -> QlResult<()> {
let i = self.time_grid.index(t)?;
asset.set_time(t);
asset.reset(self.size(i))
}
pub fn rollback(&self, asset: &mut dyn DiscretizedAsset, to: Time) -> QlResult<()> {
self.partial_rollback(asset, to)?;
asset.adjust_values()
}
#[allow(clippy::neg_cmp_op_on_partial_ord)]
pub fn partial_rollback(&self, asset: &mut dyn DiscretizedAsset, to: Time) -> QlResult<()> {
let from = asset.time();
if close(from, to) {
return Ok(());
}
require!(
from > to,
"cannot roll the asset back to {to} (it is already at t = {from})"
);
let i_from = self.time_grid.index(from)?;
let i_to = self.time_grid.index(to)?;
for i in (i_to..i_from).rev() {
let mut new_values = Array::filled(self.size(i), 0.0);
self.stepback(i, asset.values(), &mut new_values);
asset.set_time(self.time_grid[i]);
*asset.values_mut() = new_values;
if i != i_to {
asset.adjust_values()?;
}
}
Ok(())
}
pub fn present_value(&self, asset: &mut dyn DiscretizedAsset) -> QlResult<Real> {
let i = self.time_grid.index(asset.time())?;
let state_prices = self.state_prices(i);
Ok(asset.values().dot(&state_prices))
}
}
pub struct TreeLattice1D<I: TreeLatticeImpl> {
base: TreeLattice<I>,
}
impl<I: TreeLatticeImpl> TreeLattice1D<I> {
pub fn new(implementation: I, time_grid: TimeGrid) -> QlResult<Self> {
Ok(TreeLattice1D {
base: TreeLattice::new(implementation, time_grid)?,
})
}
pub fn underlying(&self, i: Size, index: Size) -> Real {
self.base.implementation.tree().underlying(i, index)
}
}
impl<I: TreeLatticeImpl> std::ops::Deref for TreeLattice1D<I> {
type Target = TreeLattice<I>;
fn deref(&self) -> &TreeLattice<I> {
&self.base
}
}
impl<I: TreeLatticeImpl> Lattice for TreeLattice1D<I> {
fn time_grid(&self) -> &TimeGrid {
self.base.time_grid()
}
fn initialize(&self, asset: &mut dyn DiscretizedAsset, time: Time) -> QlResult<()> {
self.base.initialize(asset, time)
}
fn rollback(&self, asset: &mut dyn DiscretizedAsset, to: Time) -> QlResult<()> {
self.base.rollback(asset, to)
}
fn partial_rollback(&self, asset: &mut dyn DiscretizedAsset, to: Time) -> QlResult<()> {
self.base.partial_rollback(asset, to)
}
fn present_value(&self, asset: &mut dyn DiscretizedAsset) -> QlResult<Real> {
self.base.present_value(asset)
}
fn grid(&self, t: Time) -> QlResult<Array> {
let i = self.base.time_grid.index(t)?;
let size = self.base.size(i);
let mut grid = Array::filled(size, 0.0);
for (j, value) in grid.iter_mut().enumerate() {
*value = self.underlying(i, j);
}
Ok(grid)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::discretizedasset::{DiscretizedAssetBase, DiscretizedDiscountBond};
use crate::methods::lattices::trinomialtree::TrinomialTree;
use crate::processes::OrnsteinUhlenbeckProcess;
use crate::shared::{Shared, shared};
use crate::stochasticprocess::StochasticProcess1D;
use std::cell::Cell;
const SPEED: Real = 0.1;
const VOL: Real = 0.01;
const X0: Real = 0.10;
const LEVEL: Real = 0.05;
const R: Real = 0.05;
fn process() -> Shared<dyn StochasticProcess1D> {
shared(OrnsteinUhlenbeckProcess::new(SPEED, VOL, X0, LEVEL).unwrap())
}
fn trinomial(steps: Size, end: Time) -> (TrinomialTree, TimeGrid) {
let grid = TimeGrid::new(end, steps).unwrap();
let tree = TrinomialTree::new(process(), grid.clone(), false).unwrap();
(tree, grid)
}
struct FlatRate<T: Tree> {
tree: T,
grid: TimeGrid,
rate: Real,
}
impl<T: Tree> TreeLatticeImpl for FlatRate<T> {
type Tree = T;
fn tree(&self) -> &T {
&self.tree
}
fn discount(&self, i: Size, _index: Size) -> Real {
(-self.rate * self.grid.dt(i)).exp()
}
}
fn flat_lattice(steps: Size, end: Time) -> TreeLattice1D<FlatRate<TrinomialTree>> {
let (tree, grid) = trinomial(steps, end);
TreeLattice1D::new(
FlatRate {
tree,
grid: grid.clone(),
rate: R,
},
grid,
)
.unwrap()
}
struct NoDiscount<T: Tree> {
tree: T,
}
impl<T: Tree> TreeLatticeImpl for NoDiscount<T> {
type Tree = T;
fn tree(&self) -> &T {
&self.tree
}
fn discount(&self, _i: Size, _index: Size) -> Real {
1.0
}
}
struct SwapDescTree {
inner: TrinomialTree,
swap_slice: Size,
swap_node: Size,
}
impl Tree for SwapDescTree {
const BRANCHES: Size = 3;
fn columns(&self) -> Size {
self.inner.columns()
}
fn size(&self, i: Size) -> Size {
self.inner.size(i)
}
fn underlying(&self, i: Size, index: Size) -> Real {
self.inner.underlying(i, index)
}
fn probability(&self, i: Size, index: Size, branch: Size) -> Real {
self.inner.probability(i, index, branch)
}
fn descendant(&self, i: Size, index: Size, branch: Size) -> Size {
if i == self.swap_slice && index == self.swap_node {
let swapped = match branch {
0 => 2,
2 => 0,
b => b,
};
self.inner.descendant(i, index, swapped)
} else {
self.inner.descendant(i, index, branch)
}
}
}
#[derive(Default)]
struct AdjustCountingAsset {
base: DiscretizedAssetBase,
post_adjusts: Cell<u32>,
}
impl DiscretizedAsset for AdjustCountingAsset {
fn base(&self) -> &DiscretizedAssetBase {
&self.base
}
fn base_mut(&mut self) -> &mut DiscretizedAssetBase {
&mut self.base
}
fn as_asset_mut(&mut self) -> &mut dyn DiscretizedAsset {
self
}
fn reset(&mut self, size: Size) -> QlResult<()> {
*self.values_mut() = Array::filled(size, 0.0);
Ok(())
}
fn mandatory_times(&self) -> Vec<Time> {
Vec::new()
}
fn post_adjust_values_impl(&mut self) -> QlResult<()> {
self.post_adjusts.set(self.post_adjusts.get() + 1);
Ok(())
}
}
#[test]
fn zeronomial_lattice_is_rejected() {
struct NoBranch;
impl Tree for NoBranch {
const BRANCHES: Size = 0;
fn columns(&self) -> Size {
1
}
fn size(&self, _i: Size) -> Size {
1
}
fn underlying(&self, _i: Size, _index: Size) -> Real {
0.0
}
fn descendant(&self, _i: Size, _index: Size, _branch: Size) -> Size {
0
}
fn probability(&self, _i: Size, _index: Size, _branch: Size) -> Real {
0.0
}
}
let grid = TimeGrid::new(1.0, 2).unwrap();
let err = TreeLattice::new(NoDiscount { tree: NoBranch }, grid)
.err()
.expect("a zeronomial lattice must be rejected");
assert_eq!(err.message(), "there is no zeronomial lattice!");
}
#[test]
fn state_prices_discount_sum_to_the_zero_bond() {
let steps = 5;
let lattice = flat_lattice(steps, 1.0);
let grid = lattice.time_grid().clone();
for i in 0..=steps {
let sum: Real = lattice.state_prices(i).iter().sum();
let bond = (-R * grid[i]).exp();
assert!((sum - bond).abs() < 1e-13, "sum(sp[{i}]) = {sum} != {bond}");
}
}
#[test]
fn present_value_of_a_unit_bond_is_the_discount_sum() {
let steps = 5;
let i = 3;
let lattice: Shared<dyn Lattice> = shared(flat_lattice(steps, 1.0));
let grid = lattice.time_grid().clone();
let mut bond = DiscretizedDiscountBond::new();
bond.initialize(Shared::clone(&lattice), grid[i]).unwrap();
let pv = bond.present_value().unwrap();
assert!(
(pv - (-R * grid[i]).exp()).abs() < 1e-13,
"pv = {pv} != {}",
(-R * grid[i]).exp()
);
}
#[test]
#[allow(clippy::needless_range_loop)]
fn state_prices_match_an_independent_forward_recursion_elementwise() {
let steps = 4;
let (tree, grid) = trinomial(steps, 1.0);
let disc: Vec<Real> = (0..steps).map(|i| (-R * grid.dt(i)).exp()).collect();
let mut expected: Vec<Vec<Real>> = vec![vec![1.0]];
for i in 0..steps {
let mut next = vec![0.0; tree.size(i + 1)];
for j in 0..tree.size(i) {
let sp = expected[i][j];
for l in 0..3 {
let d = tree.descendant(i, j, l);
next[d] += sp * disc[i] * tree.probability(i, j, l);
}
}
expected.push(next);
}
let lattice = TreeLattice1D::new(
FlatRate {
tree,
grid: grid.clone(),
rate: R,
},
grid,
)
.unwrap();
for i in 0..=steps {
let sp = lattice.state_prices(i);
assert_eq!(sp.size(), expected[i].len(), "slice {i} size mismatch");
for j in 0..sp.size() {
assert!(
(sp[j] - expected[i][j]).abs() < 1e-14,
"sp[{i}][{j}] = {} != {}",
sp[j],
expected[i][j]
);
}
}
}
#[test]
fn descendant_perturbation_evades_sum_identity_but_the_elementwise_pin_catches_it() {
let steps = 3;
let end = 1.0;
let i = 2;
let (correct_tree, grid) = trinomial(steps, end);
let correct = TreeLattice1D::new(
FlatRate {
tree: correct_tree,
grid: grid.clone(),
rate: R,
},
grid.clone(),
)
.unwrap();
let (inner, _g) = trinomial(steps, end);
let perturbed = TreeLattice1D::new(
FlatRate {
tree: SwapDescTree {
inner,
swap_slice: 1,
swap_node: 0,
},
grid: grid.clone(),
rate: R,
},
grid.clone(),
)
.unwrap();
let sp_correct = correct.state_prices(i);
let sp_perturbed = perturbed.state_prices(i);
let bond = (-R * grid[i]).exp();
let sum_correct: Real = sp_correct.iter().sum();
let sum_perturbed: Real = sp_perturbed.iter().sum();
assert!(
(sum_correct - bond).abs() < 1e-13,
"correct sum {sum_correct} != {bond}"
);
assert!(
(sum_perturbed - bond).abs() < 1e-13,
"perturbed sum {sum_perturbed} != {bond}: the sum-identity must stay blind"
);
let max_elt_diff = (0..sp_correct.size())
.map(|j| (sp_correct[j] - sp_perturbed[j]).abs())
.fold(0.0_f64, Real::max);
assert!(
max_elt_diff > 1e-6,
"element-wise pin failed to distinguish the perturbation (max diff {max_elt_diff})"
);
let dt = grid.dt(i);
let fit_correct: Real = (0..sp_correct.size())
.map(|j| sp_correct[j] * (-correct.underlying(i, j) * dt).exp())
.sum();
let fit_perturbed: Real = (0..sp_perturbed.size())
.map(|j| sp_perturbed[j] * (-perturbed.underlying(i, j) * dt).exp())
.sum();
assert!(
(fit_correct - fit_perturbed).abs() > 1e-6,
"fit quantity failed to distinguish the perturbation: {fit_correct} vs {fit_perturbed}"
);
}
#[test]
fn dropping_the_discount_breaks_the_sum_identity() {
let steps = 3;
let (tree, grid) = trinomial(steps, 1.0);
let lattice = TreeLattice1D::new(NoDiscount { tree }, grid.clone()).unwrap();
for i in 1..=steps {
let sum: Real = lattice.state_prices(i).iter().sum();
assert!((sum - 1.0).abs() < 1e-13, "undiscounted sum {sum} != 1");
let bond = (-R * grid[i]).exp();
assert!(
(sum - bond).abs() > 1e-6,
"the sum-identity must FAIL without the discount (sum {sum} vs bond {bond})"
);
}
}
#[test]
fn constant_payoff_rolls_back_to_payoff_times_bond() {
let steps = 4;
let end = 1.0;
let k = 5.0;
let lattice: Shared<dyn Lattice> = shared(flat_lattice(steps, end));
let mut bond = DiscretizedDiscountBond::new();
bond.initialize(Shared::clone(&lattice), end).unwrap();
let size = bond.values().size();
*bond.values_mut() = Array::filled(size, k);
bond.rollback(0.0).unwrap();
let expected = k * (-R * end).exp();
assert!(
(bond.values()[0] - expected).abs() < 1e-13,
"rolled back to {} != {expected}",
bond.values()[0]
);
}
#[test]
fn partial_rollback_skips_the_destination_adjust_and_rollback_supplies_it() {
let steps = 4;
let end = 1.0;
let lattice: Shared<dyn Lattice> = shared(flat_lattice(steps, end));
let mut full_asset = AdjustCountingAsset::default();
full_asset.initialize(Shared::clone(&lattice), end).unwrap();
full_asset.rollback(0.0).unwrap();
let full = full_asset.post_adjusts.get();
let mut partial_asset = AdjustCountingAsset::default();
partial_asset
.initialize(Shared::clone(&lattice), end)
.unwrap();
partial_asset.partial_rollback(0.0).unwrap();
let partial = partial_asset.post_adjusts.get();
assert_eq!(
partial,
steps as u32 - 1,
"partial_rollback should adjust every intermediate node, skipping the destination"
);
assert_eq!(
full, steps as u32,
"rollback should add exactly the destination adjust"
);
assert_eq!(
full,
partial + 1,
"the skipped destination adjust is supplied once by rollback"
);
}
#[test]
fn stepback_gathers_each_branch_from_its_descendant() {
let steps = 4;
let end = 1.0;
let i = 3;
let (tree, grid) = trinomial(steps, end);
let lattice: Shared<dyn Lattice> = shared(flat_lattice(steps, end));
let mut bond = DiscretizedDiscountBond::new();
bond.initialize(Shared::clone(&lattice), grid[i]).unwrap();
let terminal: Vec<Real> = (0..tree.size(i)).map(|j| 1.0 + j as Real).collect();
*bond.values_mut() = Array::from(terminal.clone());
bond.rollback(grid[i - 1]).unwrap();
let disc = (-R * grid.dt(i - 1)).exp();
for j in 0..tree.size(i - 1) {
let mut expected = 0.0;
for l in 0..3 {
expected += tree.probability(i - 1, j, l) * terminal[tree.descendant(i - 1, j, l)];
}
expected *= disc;
assert!(
(bond.values()[j] - expected).abs() < 1e-13,
"stepback[{j}] = {} != {expected}",
bond.values()[j]
);
}
}
#[test]
fn grid_returns_the_underlying_state_nodes() {
let steps = 3;
let end = 1.0;
let i = 2;
let (tree, grid) = trinomial(steps, end);
let lattice = flat_lattice(steps, end);
let g = lattice.grid(grid[i]).unwrap();
assert_eq!(g.size(), tree.size(i));
for j in 0..tree.size(i) {
assert!(
(g[j] - tree.underlying(i, j)).abs() < 1e-15,
"grid[{j}] = {} != {}",
g[j],
tree.underlying(i, j)
);
}
}
}