extern crate itertools;
pub mod eval;
pub mod workstack;
mod dot;
mod mul;
mod add;
mod index;
use std::fmt::{Debug, Write};
use std::ops::Range;
use crate::matrix::Matrix;
use itertools::izip;
use workstack::WorkStack;
use super::matrix;
use crate::utils::{iter::*, ApplyPermutationEx, ApplyPermutationMutEx, Permutation, ShapeToStridesEx};
use std::iter::{Peekable,Zip};
use std::slice::Iter;
pub use dot::{RightDottable,ExprDot};
pub use mul::*;
pub use add::*;
pub use super::domain;
pub use index::{ModelExprIndexElement,ModelExprIndex};
pub struct ExprEvalError {
file : &'static str,
line : u32,
msg : String
}
impl ExprEvalError {
fn new<S>(file : &'static str, line : u32, msg : S) -> ExprEvalError where S : Into<String> { ExprEvalError{ file,line,msg:msg.into() } }
}
impl ToString for ExprEvalError {
fn to_string(&self) -> String {
format!("{}:{}: {}",self.file,self.line,self.msg)
}
}
impl Debug for ExprEvalError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(self.file)?;
f.write_char(':')?;
self.line.fmt(f)?;
f.write_char(':')?;
f.write_str(self.msg.as_str())
}
}
pub trait ExprTrait<const N : usize> {
fn eval(&self,rs : & mut WorkStack, ws : & mut WorkStack, xs : & mut WorkStack) -> Result<(),ExprEvalError>;
fn eval_finalize(&self,rs : & mut WorkStack, ws : & mut WorkStack, xs : & mut WorkStack) -> Result<(),ExprEvalError> {
self.eval(ws,rs,xs)?;
eval::eval_finalize(rs,ws,xs)
}
fn dynamic<'a>(self) -> ExprDynamic<'a,N> where Self : Sized+'a { ExprDynamic::new(self) }
fn axispermute(self,perm : &[usize; N]) -> ExprPermuteAxes<N,Self> where Self:Sized { ExprPermuteAxes{item : self, perm: *perm } }
fn sum(self) -> ExprSum<N,Self> where Self:Sized { ExprSum{item:self} }
fn neg(self) -> ExprMulScalar<N,Self> where Self:Sized {
self.mul(-1.0)
}
fn sum_on<const K : usize>(self, axes : &[usize; K]) -> ExprReduceShape<N,K,ExprSumLastDims<N,ExprPermuteAxes<N,Self>>> where Self:Sized {
if K > N {
panic!("Invalid axis specification")
}
else if axes.iter().zip(axes[1..].iter()).any(|(a,b)| a >= b) {
panic!("Axis specification is unsorted or contains duplicates: {:?}",axes)
}
else if let Some(&last) = axes.last() {
if last >= N {
panic!("Axis specification is unsorted or contains duplicates")
}
}
let mut perm = [0usize; N];
perm[0..K].clone_from_slice(axes);
{
let (_,perm1) = perm.split_at_mut(K);
let mut i = 0;
let mut j = 0;
for &a in axes {
for ii in i..a {
unsafe { *perm1.get_unchecked_mut(j) = ii };
j += 1;
}
i = a+1;
}
for ii in i..N {
unsafe { *perm1.get_unchecked_mut(j) = ii };
j += 1;
}
}
ExprReduceShape{
item : ExprSumLastDims{
num : N-K,
item : ExprPermuteAxes{
item : self,
perm
}
}
}
}
fn add<RHS>(self, rhs : RHS) -> ExprAdd<N,Self,RHS::Result>
where
RHS : IntoExpr<N>,
Self : Sized
{
ExprAdd::new(self,rhs.into(),1.0,1.0)
}
fn sub<RHS>(self, rhs : RHS) -> ExprAdd<N,Self,RHS::Result>
where
RHS : IntoExpr<N>,
Self : Sized
{
ExprAdd::new(self,rhs.into(),1.0,-1.0)
}
fn dot<RHS>(self,rhs: RHS) -> RHS::Result where RHS: RightDottable<N,Self>, Self : Sized { rhs.dot(self) }
fn mul_elem<RHS>(self, other : RHS) -> RHS::Result where Self : Sized, RHS : ExprRightElmMultipliable<N,Self> { other.mul_elem(self) }
fn dot_rows<M>(self, other : M) -> ExprDotRows<Self>
where
Self : ExprTrait<2>+Sized,
M : Matrix
{
let (mshape,msp,mdata) = other.dissolve();
ExprDotRows{
item : self,
mshape,
msp,
mdata
}
}
fn vstack<E>(self,other : E) -> ExprStack<N,Self,E::Result> where Self:Sized, E:IntoExpr<N> { ExprStack::new(self,other.into(),0) }
fn hstack<E>(self,other : E) -> ExprStack<N,Self,E::Result> where Self:Sized,E:IntoExpr<N> { ExprStack::new(self,other.into(),1) }
fn stack<E>(self,dim : usize, other : E) -> ExprStack<N,Self,E::Result> where Self:Sized, E:IntoExpr<N>{ ExprStack::new(self,other.into(),dim) }
fn repeat(self,dim : usize, num : usize) -> ExprRepeat<N,Self> where Self:Sized { ExprRepeat{ expr : self, dim, num } }
fn index<I>(self, idx : I) -> I::Output where I : ModelExprIndex<Self>, Self:Sized {
idx.index(self)
}
fn reshape<const M : usize>(self,shape : &[usize; M]) -> ExprReshape<N,M,Self> where Self:Sized { ExprReshape{item:self,shape:*shape} }
fn into_symmetric(self, dim : usize) -> ExprIntoSymmetric<N,Self> where Self:Sized {
if dim > N-2 {
panic!("Invalid symmetrization dimension");
}
ExprIntoSymmetric{
dim,
expr : self
}
}
fn flatten(self) -> ExprReshapeOneRow<N,1,Self> where Self:Sized { ExprReshapeOneRow { item:self, dim : 0 } }
fn into_column(self) -> ExprReshapeOneRow<N,2,Self> where Self:Sized { ExprReshapeOneRow { item:self, dim : 0 } }
fn into_vec<const M : usize>(self, i : usize) -> ExprReshapeOneRow<N,M,Self> where Self:Sized+ExprTrait<1> {
if i >= M {
panic!("Invalid dimension index")
}
ExprReshapeOneRow{item:self, dim : i }
}
fn gather(self) -> ExprGatherToVec<N,Self> where Self:Sized { ExprGatherToVec{item:self} }
fn map<const M : usize,F>(self, shape : &[usize;M], f : F) -> ExprMap<N,M,F,Self>
where
F : Clone+FnMut(&[usize;N]) -> Option<[usize;M]>,
Self : Sized
{
ExprMap{ item : self, shape : *shape, f}
}
fn flip(self, dims : &[bool;N]) -> ExprFlip<N,Self> where Self:Sized
{
ExprFlip{
dims : *dims,
expr : self
}
}
fn mul<RHS>(self,other : RHS) -> RHS::Result where Self: Sized, RHS : ExprRightMultipliable<N,Self> { other.mul_right(self) }
fn rev_mul<LHS>(self, lhs: LHS) -> LHS::Result where Self: Sized, LHS : ExprLeftMultipliable<N,Self> { lhs.mul(self) }
fn transpose(self) -> ExprPermuteAxes<2,Self> where Self:Sized+ExprTrait<2> { ExprPermuteAxes{ item : self, perm : [1,0]} }
fn tril(self,with_diag:bool) -> ExprTriangularPart<Self> where Self:Sized+ExprTrait<2> { ExprTriangularPart{item:self,upper:false,with_diag} }
fn triu(self,with_diag:bool) -> ExprTriangularPart<Self> where Self:Sized+ExprTrait<2> { ExprTriangularPart{item:self,upper:true,with_diag} }
fn trilvec(self,with_diag:bool) -> ExprGatherToVec<2,ExprTriangularPart<Self>> where Self:Sized+ExprTrait<2> { ExprGatherToVec{ item:ExprTriangularPart{item:self,upper:false,with_diag} } }
fn triuvec(self,with_diag:bool) -> ExprGatherToVec<2,ExprTriangularPart<Self>> where Self:Sized+ExprTrait<2> { ExprGatherToVec{ item:ExprTriangularPart{item:self,upper:true,with_diag} } }
fn diag(self) -> ExprDiag<Self> where Self:Sized+ExprTrait<2> { ExprDiag{ item : self, anti : false, index : 0 } }
fn square_diag(self) -> ExprSquareDiag<Self> where Self:Sized+ExprTrait<1> { ExprSquareDiag{ item : self }}
fn mul_any_scalar(self, c : f64) -> ExprMulScalar<N,Self> where Self : Sized { ExprMulScalar{ item : self, lhs : c } }
fn mul_matrix_const_matrix<M>(self, m : &M) -> ExprMulRight<Self> where Self : Sized+ExprTrait<2>, M : Matrix {
ExprMulRight{
item : self,
shape : m.shape(),
data : m.data().to_vec(),
sp : m.sparsity().map(|v| v.to_vec())
}
}
fn mul_rev_matrix_const_matrix<M>(self, m : &M) -> ExprMulLeft<Self> where Self:Sized+ExprTrait<2>, M : Matrix {
ExprMulLeft{
item : self,
shape : m.shape(),
data : m.data().to_vec(),
sp : m.sparsity().map(|v| v.to_vec())
}
}
fn mul_matrix_vec(self,v : Vec<f64>) -> ExprReshapeOneRow<2,1,ExprMulRight<Self>> where Self:Sized+ExprTrait<2> {
ExprReshapeOneRow{
item : ExprMulRight{
item : self,
shape : [ v.len(),1],
data : v,
sp : None },
dim : 0
}
}
fn mul_rev_matrix_vec(self, v : Vec<f64>) -> ExprReshapeOneRow<2,1,ExprMulLeft<Self>> where Self:Sized+ExprTrait<2> {
ExprReshapeOneRow{
item : ExprMulLeft{
item : self,
shape : [1,v.len()],
data : v,
sp : None },
dim : 0
}
}
fn mul_vec_matrix<M>(self, m : &M) -> ExprReshapeOneRow<2,1,ExprMulRight<ExprReshapeOneRow<1,2,Self>>> where Self:Sized+ExprTrait<1>, M : Matrix {
ExprReshapeOneRow{
item : ExprMulRight{
item : ExprReshapeOneRow{ item : self, dim : 1 },
shape : m.shape(),
data : m.data().to_vec(),
sp : m.sparsity().map(|v| v.to_vec()) },
dim : 0
}
}
fn mul_rev_vec_matrix<M>(self, m : &M) -> ExprReshapeOneRow<2,1,ExprMulRight<ExprReshapeOneRow<1,2,Self>>> where Self:Sized+ExprTrait<1>, M : Matrix {
ExprReshapeOneRow{
item : ExprMulRight{
item : ExprReshapeOneRow{ item : self, dim : 0 },
shape : m.shape(),
data : m.data().to_vec(),
sp : m.sparsity().map(|v| v.to_vec()) },
dim : 0
}
}
fn mul_scalar_matrix<M>(self, m : &M) -> ExprReshape<1, 2, ExprMulElm<1, ExprRepeat<1, ExprReshape<0, 1, Self>>>> where Self : Sized+ExprTrait<0>, M : Matrix {
ExprReshape{
item : ExprMulElm{
expr : ExprRepeat {
expr : ExprReshape{ item : self, shape : [1] },
dim : 0,
num : m.height()*m.width()
},
data : m.data().to_vec(),
datasparsity : m.sparsity().map(|s| s.to_vec()),
datashape : [m.height()*m.width()]
},
shape : m.shape()
}
}
}
pub trait IntoExpr<const N : usize> {
type Result : ExprTrait<N>;
fn into(self) -> Self::Result;
fn into_expr(self) -> Self::Result where Self : Sized { self.into() }
}
#[derive(Clone)]
pub struct Expr<const N : usize> {
shape : [usize; N],
aptr : Vec<usize>,
asubj : Vec<usize>,
acof : Vec<f64>,
sparsity : Option<Vec<usize>>
}
impl<const N : usize> Expr<N> {
pub fn new(shape : &[usize;N],
sparsity : Option<Vec<usize>>,
aptr : Vec<usize>,
asubj : Vec<usize>,
acof : Vec<f64>) -> Expr<N> {
let fullsize = shape.iter().product();
if aptr.is_empty() { panic!("Invalid aptr"); }
if ! aptr[0..aptr.len()-1].iter().zip(aptr[1..].iter()).all(|(a,b)| a <= b) {
panic!("Invalid aptr: Not sorted");
}
let & sz = aptr.last().unwrap();
if sz != asubj.len() || sz != acof.len() {
panic!("Mismatching aptr ({}) and lengths of asubj (= {}) and acof (= {})",sz,asubj.len(),acof.len());
}
if let Some(ref sp) = sparsity {
if sp.len() != aptr.len()-1 {
panic!("Sparsity pattern length (= {})does not match length of aptr (={})",sp.len(),aptr.len());
}
if sp.iter().max().map(|&i| i >= fullsize).unwrap_or(false) {
panic!("Sparsity pattern out of bounds");
}
if ! sp.is_empty() && ! sp.iter().zip(sp[1..].iter()).all(|(&i0,&i1)| i0 < i1) {
panic!("Sparsity is not sorted or contains duplicates");
}
}
else if fullsize != aptr.len()-1 {
panic!("Shape does not match number of elements");
}
Expr{
aptr,
asubj,
acof,
shape:*shape,
sparsity
}
}
pub fn reshape<const M : usize>(self,shape:&[usize;M]) -> Expr<M> {
if self.shape.iter().product::<usize>() != shape.iter().product::<usize>() {
panic!("Invalid shape {:?} for expression with shape {:?}",shape,self.shape);
}
Expr{
aptr : self.aptr,
asubj : self.asubj,
acof : self.acof,
shape : *shape,
sparsity : self.sparsity
}
}
}
pub struct ExprDotRows<E> where E : ExprTrait<2> {
mshape : [usize;2],
msp : Option<Vec<usize>>,
mdata : Vec<f64>,
item : E
}
impl<E> ExprTrait<1> for ExprDotRows<E> where E : ExprTrait<2> {
fn eval(&self, rs : & mut WorkStack, ws : & mut WorkStack, xs : & mut WorkStack) -> Result<(),ExprEvalError> {
self.item.eval(ws,rs,xs)?;
eval::dot_rows(self.mshape,
if let Some(ref msp) = self.msp { Some(msp.as_slice()) } else { None },
self.mdata.as_slice(),
rs,ws,xs)
}
}
pub struct ExprScalarList<E> where E : ExprTrait<0> {
exprs : Vec<E>
}
impl<E> ExprTrait<1> for ExprScalarList<E> where E : ExprTrait<0> {
fn eval(&self, rs : & mut WorkStack, ws : & mut WorkStack, xs : & mut WorkStack) -> Result<(),ExprEvalError> {
let n = self.exprs.len();
for e in self.exprs.iter() { e.eval(ws,rs,xs)?; }
let es = ws.pop_exprs(n);
let rnnz = es.iter().map(|(_,_,_,subj,_)| subj.len() ).sum::<usize>();
let rnelm = n;
let (rptr,_rsp,rsubj,rcof) = rs.alloc_expr(&[n], rnnz, rnelm);
rptr[0] = 0;
rptr[1..].iter_mut().zip(es.iter()).for_each(|(rp,(_,_,_,subj,_))| *rp = subj.len());
_ = rptr.iter_mut().fold(0,|c,rp| { *rp += c; *rp });
izip!(rptr.iter(),rptr[1..].iter(),es.iter())
.for_each(|(&pb,&pe, (_,_,_,subj,cof))| {
rsubj[pb..pe].copy_from_slice(subj);
rcof[pb..pe].copy_from_slice(cof);
});
Ok(())
}
}
pub fn from_sparse_iter<const N : usize,I,E>(shape : [usize;N], it : I) -> ExprScatter<N, ExprScalarList<E>>
where
I : Iterator<Item = (usize, E)>,
E : ExprTrait<0>
{
let mut es : Vec<(usize,E)> = it.take(shape.iter().product()).collect();
es.sort_by_key(|(idx,_)| *idx);
if es.iter().zip(es[1..].iter()).any(|(a,b)| a.0 == b.0) {
panic!("Sparsity indexes contains duplicates");
}
let mut sparsity = Vec::with_capacity(es.len());
let mut exprs = Vec::with_capacity(es.len());
for (i,e) in es {
sparsity.push(i);
exprs.push(e);
}
ExprScatter {
shape,
item : ExprScalarList { exprs },
sparsity
}
}
pub fn from_dense_iter<const N : usize, I,E>(shape : [usize; N], it : I) -> ExprReshape<1, N, ExprScalarList<E>>
where
I : Iterator<Item = E>,
E : ExprTrait<0>
{
let nelm = shape.iter().product();
let exprs : Vec<E> = it.take(nelm).collect();
if exprs.len() != nelm {
panic!("Insufficient expressions for shape");
}
ExprReshape{
shape,
item : ExprScalarList {
exprs
}
}
}
pub fn from_iter<const N : usize, I,E>(shape : [usize; N], it : I) -> ExprScatter<N, ExprScalarList<E>>
where
I : Iterator<Item = Option<E>>,
E : ExprTrait<0>
{
let nelm = shape.iter().product();
from_sparse_iter(shape,it.take(nelm).enumerate().filter(|v| v.1.is_some()).map(|v| (v.0,v.1.unwrap())))
}
impl<const N : usize> ExprTrait<N> for Expr<N> {
fn eval(&self,rs : & mut WorkStack, _ws : & mut WorkStack, _xs : & mut WorkStack) -> Result<(),ExprEvalError> {
let nnz = self.asubj.len();
let nelm = self.aptr.len()-1;
let (aptr,sp,asubj,acof) = rs.alloc_expr(self.shape.as_slice(),nnz,nelm);
if let (Some(ref ssp),Some(dsp)) = (&self.sparsity,sp) {
dsp.clone_from_slice(ssp.as_slice())
}
aptr.clone_from_slice(self.aptr.as_slice());
asubj.clone_from_slice(self.asubj.as_slice());
acof.clone_from_slice(self.acof.as_slice());
Ok(())
}
}
#[derive(Clone)]
pub struct ExprNil<const N : usize> { shape : [usize; N] }
impl<const N : usize> ExprTrait<N> for ExprNil<N> {
fn eval(&self,rs : & mut WorkStack, _ws : & mut WorkStack, _xs : & mut WorkStack) -> Result<(),ExprEvalError> {
let (rptr,_,_,_) = rs.alloc_expr(self.shape.as_slice(),0,0);
rptr[0] = 0;
Ok(())
}
}
pub fn zeros<const N : usize>(shape : &[usize;N]) -> Expr<N> {
Expr{
shape : *shape,
aptr : vec![0],
asubj : vec![],
acof : vec![],
sparsity : Some(vec![]),
}
}
pub fn const_expr<const N : usize>(shape : &[usize;N], value : f64) -> Expr<N> {
let nelm : usize = shape.iter().product();
Expr{
shape : *shape,
aptr : (0..(nelm+1)).collect(),
asubj : vec![0usize; nelm],
acof : vec![value; nelm],
sparsity : None
}
}
pub fn ones<const N : usize>(shape : &[usize;N]) -> Expr<N> {
const_expr(shape,1.0)
}
pub fn const_diag(n : usize,value:f64) -> Expr<2> {
Expr{
shape : [n,n],
aptr : (0..n+1).collect(),
asubj : vec![0usize; n],
acof : vec![value; n],
sparsity : Some((0..n*n).step_by(n+1).collect())
}
}
pub fn eye(n : usize) -> Expr<2> {
const_diag(n,1.0)
}
pub fn constants<const N : usize>(shape : &[usize;N], values : &[f64]) -> Expr<N> {
if shape.iter().product::<usize>() != values.len() {
panic!("Data and shape do not match");
}
Expr{
shape : *shape,
aptr : (0..values.len()+1).collect(),
asubj : vec![0; values.len()],
acof : values.to_vec(),
sparsity : None
}
}
pub fn nil<const N : usize>(shape : &[usize; N]) -> ExprNil<N> {
if shape.len() > 0 && shape.iter().product::<usize>() != 0 {
panic!("Shape must have at least one zero-dimension");
}
ExprNil{shape:*shape}
}
pub struct ExprReduceShape<const N : usize, const M : usize, E> where E : ExprTrait<N>+Sized { item : E }
impl<const N : usize, const M : usize, E> ExprTrait<M> for ExprReduceShape<N,M,E>
where E : ExprTrait<N>
{
fn eval(&self, rs : & mut WorkStack, ws : & mut WorkStack, xs : & mut WorkStack) -> Result<(),ExprEvalError> {
self.item.eval(rs,ws,xs)?;
eval::inplace_reduce_shape(M, rs, xs)
}
}
pub struct ExprReshapeOneRow<const N : usize, const M : usize, E:ExprTrait<N>> { item : E, dim : usize }
impl<const N : usize,const M : usize,E> ExprReshapeOneRow<N,M,E>
where
E:ExprTrait<N>
{
pub fn new(dim : usize, item : E) -> ExprReshapeOneRow<N,M,E> {
ExprReshapeOneRow{item,dim}
}
}
impl<const N : usize, const M : usize, E:ExprTrait<N>> ExprTrait<M> for ExprReshapeOneRow<N,M,E> {
fn eval(&self, rs : & mut WorkStack, ws : & mut WorkStack, xs : & mut WorkStack) -> Result<(),ExprEvalError> {
if self.dim >= M { panic!("Invalid dimension given"); }
self.item.eval(rs,ws,xs)?;
eval::inplace_reshape_one_row(M, self.dim, rs, xs)
}
}
pub struct ExprReshape<const N : usize, const M : usize, E:ExprTrait<N>> { item : E, shape : [usize; M] }
impl<const N : usize, const M : usize, E:ExprTrait<N>> ExprTrait<M> for ExprReshape<N,M,E> {
fn eval(&self, rs : & mut WorkStack, ws : & mut WorkStack, xs : & mut WorkStack) -> Result<(),ExprEvalError> {
self.item.eval(rs,ws,xs)?;
eval::inplace_reshape(self.shape.as_slice(), rs, xs)
}
}
pub struct ExprScatter<const M : usize, E:ExprTrait<1>> { item : E, shape : [usize; M], sparsity : Vec<usize> }
impl<const M : usize, E:ExprTrait<1>> ExprScatter<M,E> {
pub fn new(item : E,
shape : &[usize; M],
sparsity : Vec<usize>) -> ExprScatter<M,E> {
if sparsity.iter().max().map(|&v| v >= shape.iter().product()).unwrap_or(false) {
panic!("Sparsity pattern element out of bounds");
}
if sparsity.iter().zip(sparsity[1..].iter()).any(|(&i0,&i1)| i1 <= i0) {
let mut perm : Vec<usize> = (0..sparsity.len()).collect();
perm.sort_by_key(|&p| unsafe{ *sparsity.get_unchecked(p)});
if perm.iter().zip(perm[1..].iter()).any(|(&p0,&p1)| unsafe{ *sparsity.get_unchecked(p0) >= *sparsity.get_unchecked(p1) }) {
panic!("Sparsity pattern contains duplicates");
}
ExprScatter{ item,
shape:*shape,
sparsity : perm.iter().map(|&p| unsafe{ *sparsity.get_unchecked(p)}).collect() }
}
else {
ExprScatter{ item, shape: *shape, sparsity }
}
}
}
impl<const M : usize, E:ExprTrait<1>> ExprTrait<M> for ExprScatter<M,E> {
fn eval(&self, rs : & mut WorkStack, ws : & mut WorkStack, xs : & mut WorkStack) -> Result<(),ExprEvalError> {
self.item.eval(ws,rs,xs)?;
eval::scatter(self.shape.as_slice(),self.sparsity.as_slice(), rs, ws, xs)
}
}
pub struct ExprMap<const N : usize, const M : usize, F, E> where
F : Clone+FnMut(&[usize;N]) -> Option<[usize;M]>,
E : ExprTrait<N>
{
item : E,
shape : [usize;M],
f : F
}
impl<const N : usize, const M : usize, F, E> ExprTrait<M> for ExprMap<N,M,F,E>
where
F : Clone+FnMut(&[usize;N]) -> Option<[usize;M]>,
E : ExprTrait<N>
{
fn eval(&self,rs : & mut WorkStack, ws : & mut WorkStack, xs : & mut WorkStack) -> Result<(),ExprEvalError> {
self.item.eval(ws,rs,xs)?;
let (shape,ptr,sp,subj,cof) = ws.pop_expr();
let nelm : usize = sp.map(|v| v.len()).unwrap_or(ptr.len()-1);
let (irest,_) = xs.alloc(nelm*7, 0);
let (xptrb,xptre,xsp) = {
let (xptrb,irest) = irest.split_at_mut(nelm);
let (xptre,irest) = irest.split_at_mut(nelm);
let (xsp,irest) = irest.split_at_mut(nelm);
let mut src_shape = [0usize; N]; src_shape.copy_from_slice(shape);
let src_st = src_shape.to_strides();
let tgt_st = self.shape.to_strides();
let mut f = self.f.clone();
let xnelm = {
if let Some(sp) = sp {
izip!(ptr.iter(),
ptr[1..].iter(),
sp.iter().map(|&i| src_st.to_index(i)))
.filter_map(|(p0,p1,i)| {
let r = f(&i)
.and_then(|v| tgt_st.from_coords_checked(&v))
.and_then(|i| Some((i,*p0,*p1)));
if let Some(r) = r { println!("{:?} -> {:?}",i,r); }
r
})
.zip(izip!(xptrb.iter_mut(),xptre.iter_mut(),xsp.iter_mut()))
.fold(0,|n,((j,p0,p1),(xpb,xpe,xspi))| { println!("j = {}, range = {}..{}",j,p0,p1); *xspi = j; *xpb = p0; *xpe = p1; n+1 })
}
else {
izip!(ptr.iter(),
ptr[1..].iter(),
src_shape.index_iterator())
.filter_map(|(p0,p1,i)| {
f(&i)
.and_then(|v| tgt_st.from_coords_checked(&v))
.and_then(|i| Some((i,*p0,*p1)))
})
.zip(izip!(xptrb.iter_mut(),xptre.iter_mut(),xsp.iter_mut()))
.fold(0,|n,((j,p0,p1),(xpb,xpe,xspi))| { *xspi = j; *xpb = p0; *xpe = p1; n+1 })
}
};
let xptrb = &mut xptrb[..xnelm];
let xptre = &mut xptre[..xnelm];
let xsp = &mut xsp[..xnelm];
if xsp.is_empty() || xsp.iter().zip(xsp[1..].iter()).all(|(&i,&j)| i < j) {
(xptrb,xptre,xsp)
}
else {
let (xperm, irest) = irest.split_at_mut(xnelm);
let (xosp, irest) = irest.split_at_mut(xnelm);
let (xoptrb, irest) = irest.split_at_mut(xnelm);
let (xoptre,_irest) = irest.split_at_mut(xnelm);
xperm.copy_from_iter(0..nelm);
xperm.sort_by_key(|&i| unsafe{ *xsp.get_unchecked(i) });
let xperm = Permutation::from(xperm);
xoptrb.copy_from_iter(xperm.apply(xptrb).unwrap().iter().cloned());
xoptre.copy_from_iter(xperm.apply(xptre).unwrap().iter().cloned());
xosp.copy_from_iter(xperm.apply(xsp).unwrap().iter().cloned());
(xoptrb,xoptre,xosp)
}
};
let rnnz = xptrb.iter().zip(xptre.iter()).map(|(a,b)| b-a).sum();
let rnelm = if xsp.is_empty() { 0 } else { 1+xsp.iter().zip(xsp[1..].iter()).filter(|(a,b)| a!=b).count() };
let (rptr,rsp,rsubj,rcof) = rs.alloc_expr(&self.shape, rnnz, rnelm);
rptr[0] = 0;
if rnelm == xsp.len() {
if let Some(rsp) = rsp { rsp.copy_from_slice(xsp); }
izip!(rptr[1..].iter_mut(),xptrb.iter(),xptre.iter())
.fold(0,|n,(rp,&pb,&pe)| { *rp = n+pe-pb; *rp });
}
else {
let it =
izip!(xsp.iter(),
xsp[1..].iter().chain(std::iter::once(&usize::MAX)),
xptrb.iter(),
xptre.iter())
.scan(0usize,|n,(&spi0,&spi1,&pb,&pe)| if spi0 == spi1 { *n += pe-pb; Some(None) } else { let oldn = *n; *n = 0; Some(Some((oldn,spi0))) } )
.filter_map(|v| v)
;
if let Some(rsp) = rsp {
it
.zip(rsp.iter_mut().zip(rptr[1..].iter_mut()))
.fold(0,|n,((nz,i),(ri,rp))| { *ri = i; *rp = n+nz; *rp });
}
else {
it
.zip(rptr[1..].iter_mut())
.fold(0,|n,((nz,_),rp)| { *rp = n+nz; *rp });
}
}
izip!(cof.chunks_ptr2(xptrb, xptre),
rcof.chunks_ptr_mut(rptr,&rptr[1..]))
.for_each(|(cof,rcof)| rcof.clone_from_slice(cof));
izip!(subj.chunks_ptr2(xptrb, xptre),
rsubj.chunks_ptr_mut(rptr, &rptr[1..]))
.for_each(|(subj,rsubj)| rsubj.clone_from_slice(subj));
Ok(())
}
}
pub struct ExprGatherToVec<const N : usize, E:ExprTrait<N>> { item : E }
impl<const N : usize, E:ExprTrait<N>> ExprTrait<1> for ExprGatherToVec<N,E> {
fn eval(&self, rs : & mut WorkStack, ws : & mut WorkStack, xs : & mut WorkStack) -> Result<(),ExprEvalError> {
self.item.eval(ws,rs,xs)?;
eval::gather_to_vec(rs, ws, xs)
}
}
#[macro_export]
macro_rules! hstack {
[ $x0:expr ] => { $x0 . into_expr() };
[ $x0:expr , $( $x:expr ),* ] => {
{
$x0 . into_expr() $( .hstack( $x . into_expr() ) )*
}
}
}
#[macro_export]
macro_rules! vstack {
[ $x0:expr ] => { into_expr() $x0 };
[ $x0:expr , $( $x:expr ),* ] => {
{
$x0 . into_expr() $( .vstack( $x . into_expr() ))*
}
}
}
#[macro_export]
macro_rules! stack {
[ $n:expr ; $x0:expr ] => { $x0 };
[ $n:expr ; $x0:expr , $( $x:expr ),* ] => {
{
let n = $n;
$x0 . into_expr() $( .stack( n , $x . into_expr() ))*
}
}
}
#[macro_export]
macro_rules! exprcat {
[ $e0:expr ] => { $e0 };
[ $e0:expr , $( $es:expr ),+ ] => { hstack![ $e0 $( , $es )* ] };
[ $e0:expr ; $( $rest:tt )+ ] => { $e0 . vstack( exprcat![ $( $rest )* ]) };
[ $e0:expr , $( $es:expr ),+ ; $( $rest:tt )+ ] => { hstack![ $e0 $( , $es )* ].vstack( exprcat![ $( $rest )*] ) };
}
pub struct ExprStackVec<const N : usize, E : ExprTrait<N>> {
items : Vec<E>,
dim : usize
}
pub struct ExprStack<const N : usize,E1:ExprTrait<N>,E2:ExprTrait<N>> {
item1 : E1,
item2 : E2,
dim : usize
}
pub struct ExprStackRec<const N : usize,E1,E2>
where
E1:ExprStackRecTrait<N>,
E2:ExprTrait<N>
{
item1 : E1,
item2 : E2,
dim : usize
}
impl<const N : usize, E> ExprTrait<N> for ExprStackVec<N,E> where E : ExprTrait<N> {
fn eval(&self, rs : & mut WorkStack, ws : & mut WorkStack, xs : & mut WorkStack) -> Result<(),ExprEvalError> {
for item in self.items.iter().rev() {
item.eval(ws,rs,xs)?;
}
eval::stack(self.dim,self.items.len(),rs,ws,xs)
}
}
pub trait ExprStackRecTrait<const N : usize> : ExprTrait<N> {
fn stack_dim(&self) -> usize;
fn eval_rec(&self, rs : & mut WorkStack, ws : & mut WorkStack, xs : & mut WorkStack) -> Result<usize,ExprEvalError>;
}
impl<const N : usize, E1:ExprTrait<N>,E2:ExprTrait<N>> ExprStack<N,E1,E2> {
pub fn new(item1 : E1, item2 : E2, dim : usize) -> Self {
if dim > N {
panic!("Stacking dimension out of bounds");
}
ExprStack{item1,item2,dim}
}
pub fn stack<T:IntoExpr<N>>(self, dim : usize, other : T) -> ExprStackRec<N,Self,T::Result> { ExprStackRec{item1:self,item2:other.into(),dim} }
pub fn vstack<T:IntoExpr<N>>(self, other : T) -> ExprStackRec<N,Self,T::Result> { ExprStackRec{item1:self,item2:other.into(),dim:0} }
pub fn hstack<T:IntoExpr<N>>(self, other : T) -> ExprStackRec<N,Self,T::Result> { ExprStackRec{item1:self,item2:other.into(),dim:1} }
}
impl<const N : usize, E1:ExprStackRecTrait<N>,E2:ExprTrait<N>> ExprStackRec<N,E1,E2> {
pub fn stack<T:IntoExpr<N>>(self, dim : usize, other : T) -> ExprStackRec<N,Self,T::Result> { ExprStackRec{item1:self,item2:other.into(),dim} }
pub fn vstack<T:IntoExpr<N>>(self, other : T) -> ExprStackRec<N,Self,T::Result> { ExprStackRec{item1:self,item2:other.into(),dim:0} }
pub fn hstack<T:IntoExpr<N>>(self, other : T) -> ExprStackRec<N,Self,T::Result> { ExprStackRec{item1:self,item2:other.into(),dim:1} }
}
impl<const N : usize,E1:ExprTrait<N>,E2:ExprTrait<N>> ExprTrait<N> for ExprStack<N,E1,E2> {
fn eval(&self, rs : & mut WorkStack, ws : & mut WorkStack, xs : & mut WorkStack) -> Result<(),ExprEvalError> {
let n = self.eval_rec(ws,rs,xs)?;
eval::stack(self.dim,n,rs,ws,xs)
}
}
impl<const N : usize, E1:ExprTrait<N>,E2:ExprTrait<N>> ExprStackRecTrait<N> for ExprStack<N,E1,E2> {
fn stack_dim(&self) -> usize { self.dim }
fn eval_rec(&self, rs : & mut WorkStack, ws : & mut WorkStack, xs : & mut WorkStack) -> Result<usize,ExprEvalError> {
self.item2.eval(rs,ws,xs)?;
self.item1.eval(rs,ws,xs)?;
Ok(2)
}
}
impl<const N : usize, E1:ExprStackRecTrait<N>,E2:ExprTrait<N>> ExprTrait<N> for ExprStackRec<N,E1,E2> {
fn eval(&self, rs : & mut WorkStack, ws : & mut WorkStack, xs : & mut WorkStack) -> Result<(),ExprEvalError> {
let n = self.eval_rec(ws,rs,xs)?;
eval::stack(self.dim,n,rs,ws,xs)
}
}
impl<const N : usize, E1:ExprStackRecTrait<N>,E2:ExprTrait<N>> ExprStackRecTrait<N> for ExprStackRec<N,E1,E2> {
fn stack_dim(&self) -> usize { self.dim }
fn eval_rec(&self, rs : & mut WorkStack, ws : & mut WorkStack, xs : & mut WorkStack) -> Result<usize,ExprEvalError> {
self.item2.eval(rs,ws,xs)?;
if self.dim == self.item1.stack_dim() {
Ok(1+self.item1.eval_rec(rs,ws,xs)?)
}
else {
self.item1.eval(rs,ws,xs)?;
Ok(2)
}
}
}
pub struct ExprRepeat<const N : usize, E : ExprTrait<N>> {
expr : E,
dim : usize,
num : usize
}
impl<const N : usize, E : ExprTrait<N>> ExprTrait<N> for ExprRepeat<N,E> {
fn eval(&self, rs : & mut WorkStack, ws : & mut WorkStack, xs : & mut WorkStack) -> Result<(),ExprEvalError> {
self.expr.eval(ws,rs,xs)?;
eval::repeat(self.dim,self.num,rs,ws,xs)
}
}
pub struct ExprDynamic<'a,const N : usize> {
expr : Box<dyn ExprTrait<N>+'a>
}
impl<'a,const N : usize> ExprDynamic<'a,N> {
fn new<E>(e : E) -> ExprDynamic<'a,N> where E : ExprTrait<N>+'a {
ExprDynamic{
expr : Box::new(e)
}
}
}
impl<'a,const N : usize> ExprTrait<N> for ExprDynamic<'a,N> {
fn eval(&self, rs : & mut WorkStack, ws : & mut WorkStack, xs : & mut WorkStack) -> Result<(),ExprEvalError> {
self.expr.eval(rs,ws,xs)
}
}
pub struct ExprDynStack<const N : usize> {
exprs : Vec<ExprDynamic<'static,N>>,
dim : usize
}
impl<const N : usize> ExprTrait<N> for ExprDynStack<N> {
fn eval(&self, rs : & mut WorkStack, ws : & mut WorkStack, xs : & mut WorkStack) -> Result<(),ExprEvalError> {
let n = self.exprs.len();
for e in self.exprs.iter() {
e.eval(ws,rs,xs)?;
}
eval::stack(self.dim,n,rs,ws,xs)
}
}
pub fn stack<const N : usize>(dim : usize, exprs : Vec<ExprDynamic<'static,N>>) -> ExprDynStack<N> {
ExprDynStack{exprs,dim}
}
pub fn vstack<const N : usize>(exprs : Vec<ExprDynamic<'static, N>>) -> ExprDynStack<N> {
ExprDynStack{exprs,dim:0}
}
pub fn hstack<const N : usize>(exprs : Vec<ExprDynamic<'static,N>>) -> ExprDynStack<N> {
ExprDynStack{exprs,dim:1}
}
pub fn stackvec<const N : usize,E>(dim : usize, exprs : Vec<E>) -> ExprStackVec<N,E::Result>
where
E : IntoExpr<N>
{
ExprStackVec{
dim,
items:exprs.into_iter().map(|v| v.into()).collect() }
}
#[allow(unused)]
pub struct ExprSumVec<const N : usize,E> where E : ExprTrait<N>
{
exprs : Vec<E>
}
impl<const N : usize, E> ExprTrait<N> for ExprSumVec<N,E> where E : ExprTrait<N> {
fn eval(&self, rs : & mut WorkStack,ws : & mut WorkStack, xs : & mut WorkStack) -> Result<(),ExprEvalError> {
let n = self.exprs.len();
if n == 0 {
panic!("Cannot sum 0 expressions");
}
else if n == 1 {
self.exprs[0].eval(rs,ws,xs)
}
else {
for e in self.exprs.iter() {
e.eval(ws,rs,xs)?
}
let vals = ws.pop_exprs(n);
if let Some(((s0,_,_,_,_),(s1,_,_,_,_))) = vals.iter().zip(vals[1..].iter()).find(|((s0,_,_,_,_),(s1,_,_,_,_))| *s0 != *s1) {
panic!("Mismarching operand shapes {:?} vs. {:?}", s0,s1);
}
let is_dense = vals.iter().any(|vv| vv.2.is_none() );
let rnnz = vals.iter().map(|vv| *(vv.1.last().unwrap())).sum::<usize>();
let mut rshape = [0usize;N]; rshape.copy_from_slice(vals[0].0);
if is_dense {
let rnelm = rshape.iter().product();
let (rptr,_,rsubj,rcof) = rs.alloc_expr(&rshape, rnnz, rnelm);
rptr.iter_mut().for_each(|p| *p = 0);
for (_,ptr,sp,_,_) in vals.iter() {
if let Some(sp) = sp {
for (&pb,&pe,&i) in izip!(ptr.iter(),ptr[1..].iter(),sp.iter()) { rptr[i] += pe-pb; }
}
else {
for (&pb,&pe,rp) in izip!(ptr.iter(),ptr[1..].iter(),rptr.iter_mut()) { *rp += pe-pb; }
}
}
rptr.iter_mut().fold(0usize, |c,rp| { let tmp = *rp; *rp = c; tmp + c });
for (_,ptr,sp,subj,cof) in vals.iter() {
if let Some(sp) = sp {
for (&pb,&pe,&i) in izip!(ptr.iter(),ptr[1..].iter(),sp.iter()) {
let rp = rptr[i];
rsubj[rp..rp+pe-pb].copy_from_slice(&subj[pb..pe]);
rcof[rp..rp+pe-pb].copy_from_slice(&cof[pb..pe]);
rptr[i] += pe-pb;
}
}
else {
for (&pb,&pe,rp) in izip!(ptr.iter(),ptr[1..].iter(),rptr.iter_mut()) {
rsubj[*rp..*rp+pe-pb].copy_from_slice(&subj[pb..pe]);
rcof[*rp..*rp+pe-pb].copy_from_slice(&cof[pb..pe]);
*rp += pe-pb;
}
}
}
rptr.iter_mut().fold(0usize, |c,rp| { let tmp = *rp; *rp = c; tmp });
}
else {
let mut rnelm = 0usize;
{
let mut spit = vals.iter()
.map(| (_,_,sp,_,_) | sp.unwrap().iter().peekable())
.collect::<Vec<Peekable<Iter<usize>>>>();
while let Some(&i) = spit.iter_mut().filter_map(|it| it.peek()).min() {
rnelm += 1;
spit.iter_mut().for_each(|it| if let Some(&ii) = it.peek() { if ii == i { it.next(); } } );
}
}
let (rptr,rsp,rsubj,rcof) = rs.alloc_expr(&rshape,rnnz,rnelm);
rptr[0] = 0;
rptr[0] = 0;
if let Some(rsp) = rsp {
let mut nzi = 0usize;
let mut nelmi = 0usize;
let mut spit = vals.iter()
.map(| (_,_,sp,_,_) | sp.unwrap().iter().peekable())
.collect::<Vec<Peekable<Iter<usize>>>>();
let mut datait = vals.iter()
.map(| (_,ptr,_,subj,cof) |
subj.chunks_ptr(ptr)
.zip(cof.chunks_ptr(ptr)))
.collect::<Vec<Zip<ChunksByIter<usize,Zip<Iter<usize>,Iter<usize>>>,
ChunksByIter<f64,Zip<Iter<usize>,Iter<usize>>>>>>();
while let Some(&&i) = spit.iter_mut().filter_map(|it| it.peek()).min() {
rsp[nelmi] = i;
spit.iter_mut().zip(datait.iter_mut())
.filter_map(| (spit,data) | if let Some(&&ii) = spit.peek() { if ii == i { _ = spit.next(); Some(data) } else { None } } else { None } )
.filter_map(| data | data.next())
.for_each(|(subj,cof)| {
rsubj[nzi..nzi+n].copy_from_slice(subj);
rcof[nzi..nzi+n].copy_from_slice(cof);
nzi += n;
});
nelmi += 1;
rptr[nelmi] = nzi;
}
}
else {
rptr.iter_mut().for_each(|p| *p = 0) ;
for (_,ptr,sp,_,_) in vals.iter() {
let sp = sp.unwrap();
izip!(rptr[1..].permute_by_mut(sp),
ptr.iter(),
ptr[1..].iter())
.for_each(|(rp,&pb,&pe)| *rp += pe-pb );
}
rptr.iter_mut().fold(0usize, |c,p| { *p += c; *p });
for (_,ptr,sp,subj,cof) in vals.iter() {
let sp = sp.unwrap();
izip!(rptr.permute_by_mut(sp),
ptr.iter(),
ptr[1..].iter())
.for_each(|(rp,&pb,&pe)| {
let n = pe-pb;
rsubj[*rp..*rp+n].copy_from_slice(&subj[pb..pe]);
rcof[*rp..*rp+n].copy_from_slice(&cof[pb..pe]);
*rp += n;
});
}
rptr.iter_mut().fold(0usize,|c,p| { let tmp = *p; *p = c; tmp });
}
}
rs.check();
Ok(())
}
}
}
pub fn sumvec<const N : usize,E>(exprs : Vec<E>) -> ExprSumVec<N,E::Result> where E : IntoExpr<N> {
if exprs.is_empty() {
panic!("Empty operand list");
}
ExprSumVec{
exprs : exprs.into_iter().map(|e| e.into()).collect()
}
}
pub struct ExprSlice<const N : usize, E : ExprTrait<N>> {
expr : E,
begin : [usize; N],
end : [usize; N]
}
impl<const N : usize, E> ExprTrait<N> for ExprSlice<N,E> where E : ExprTrait<N> {
fn eval(&self, rs : & mut WorkStack, ws : & mut WorkStack, xs : & mut WorkStack) -> Result<(),ExprEvalError> {
self.expr.eval(ws,rs,xs)?;
eval::slice(&self.begin,&self.end,rs,ws,xs)
}
}
pub struct ExprSlice2<const N : usize, E : ExprTrait<N>> {
expr : E,
ranges : [Range<Option<usize>>;N]
}
impl<const N : usize,E> ExprTrait<N> for ExprSlice2<N,E> where E : ExprTrait<N> {
fn eval(&self,rs : & mut WorkStack, ws : & mut WorkStack, xs : & mut WorkStack) -> Result<(),ExprEvalError> {
self.expr.eval(ws,rs,xs)?;
let mut begin = vec![0usize; N];
let mut end = vec![0usize; N];
{
let (shape,_,_,_,_) = ws.peek_expr();
for (b,e,d,r) in izip!(begin.iter_mut(),end.iter_mut(),shape.iter(),self.ranges.iter()) {
*b = r.start.unwrap_or(0);
*e = r.end.unwrap_or(*d);
}
}
eval::slice(begin.as_slice(),end.as_slice(),rs,ws,xs)
}
}
pub struct ExprSum<const N : usize, T:ExprTrait<N>> {
item : T
}
pub struct ExprSumLastDims<const N : usize, T : ExprTrait<N>> {
item : T,
num : usize
}
impl<const N : usize, T:ExprTrait<N>> ExprTrait<0> for ExprSum<N,T> {
fn eval(&self, rs : & mut WorkStack, ws : & mut WorkStack, xs : & mut WorkStack) -> Result<(),ExprEvalError> {
self.item.eval(ws,rs,xs)?;
eval::sum(rs,ws,xs)
}
}
impl<const N : usize, E:ExprTrait<N>> ExprTrait<N> for ExprSumLastDims<N,E> {
fn eval(&self, rs : & mut WorkStack, ws : & mut WorkStack, xs : & mut WorkStack) -> Result<(),ExprEvalError> {
self.item.eval(ws,rs,xs)?;
eval::sum_last(self.num,rs,ws,xs)
}
}
pub struct ExprTriangularPart<T:ExprTrait<2>> {
item : T,
upper : bool,
with_diag : bool
}
impl<T:ExprTrait<2>> ExprTrait<2> for ExprTriangularPart<T> {
fn eval(&self, rs : & mut WorkStack, ws : & mut WorkStack, xs : & mut WorkStack) -> Result<(),ExprEvalError> {
self.item.eval(ws,rs,xs)?;
eval::triangular_part(self.upper, self.with_diag, rs, ws, xs)
}
}
pub struct ExprDiag<E:ExprTrait<2>> {
item : E,
anti : bool,
index : i64
}
impl<E:ExprTrait<2>> ExprTrait<1> for ExprDiag<E> {
fn eval(&self, rs : & mut WorkStack, ws : & mut WorkStack, xs : & mut WorkStack) -> Result<(),ExprEvalError> {
self.item.eval(ws,rs,xs)?;
eval::diag(self.anti, self.index, rs, ws, xs)
}
}
pub struct ExprSquareDiag<E : ExprTrait<1>> {
item : E
}
impl<E:ExprTrait<1>> ExprTrait<2> for ExprSquareDiag<E> {
fn eval(&self, rs : & mut WorkStack, ws : & mut WorkStack, xs : & mut WorkStack) -> Result<(),ExprEvalError> {
self.item.eval(ws,rs,xs).unwrap();
let (shape,ptr,sp,subj,cof) = ws.pop_expr();
if shape.len() != 1 { panic!("Operand has invalid shape {:?}, expected a vector", shape); }
let n = shape[0];
let rshape = [n,n];
let rnnz = *ptr.last().unwrap();
let rnelm = n;
let (rptr,rsp,rsubj,rcof) = rs.alloc_expr(&rshape, rnnz, rnelm);
rptr.copy_from_slice(ptr);
rsubj.copy_from_slice(subj);
rcof.copy_from_slice(cof);
if let Some(sp) = sp {
rsp.unwrap().iter_mut().zip(sp.iter()).for_each(|(ri,&i)| *ri = i * (n+1));
}
else {
rsp.unwrap().iter_mut().enumerate().for_each(|(i,ri)| *ri = i * (n+1));
}
rs.check();
Ok(())
}
}
#[allow(unused)]
pub struct ExprIntoSymmetric<const N : usize, E : ExprTrait<N>> {
dim : usize,
expr : E
}
#[allow(unused)]
impl<const N : usize, E : ExprTrait<N>> ExprIntoSymmetric<N,E> {
fn eval(&self, rs : & mut WorkStack, ws : & mut WorkStack, xs : & mut WorkStack) -> Result<(),ExprEvalError> {
self.expr.eval(ws,rs,xs)?;
eval::into_symmetric(self.dim,rs,ws,xs)
}
}
pub struct ExprFlip<const N : usize, E : ExprTrait<N>> {
dims : [bool; N],
expr : E
}
impl<const N : usize, E : ExprTrait<N>> ExprTrait<N> for ExprFlip<N,E> {
fn eval(&self,rs : & mut WorkStack, ws : & mut WorkStack, xs : & mut WorkStack) -> Result<(),ExprEvalError> {
self.expr.eval(ws,rs,xs)?;
let (sshape,ptr,sp,subj,cof) = ws.pop_expr();
let nelm = ptr.len()-1;
let nnz = subj.len();
let (rptr,rsp,rsubj,rcof) = rs.alloc_expr(sshape, nnz, nelm);
let mut shape = [0usize;N]; shape.copy_from_slice(sshape);
let (irest,_) = xs.alloc(nelm*4, 0);
let (perm,irest) = irest.split_at_mut(nelm);
let (ridxs,irest) = irest.split_at_mut(nelm);
let (xptrb,irest) = irest.split_at_mut(nelm);
let (xptre,_irest) = irest.split_at_mut(nelm);
perm.copy_from_iter(0..nelm);
let st = shape.to_strides();
if let Some(sp) = sp {
ridxs.copy_from_iter(sp.iter().map(|&i| {
let mut idx = st.to_index(i);
izip!(idx.iter_mut(),
shape.iter(),
self.dims.iter())
.for_each(|(j,&d,&f)| if f { *j = d-*j-1 } );
st.to_linear(&idx) }));
}
else {
ridxs.copy_from_iter((0..nelm).scan([0usize; N], |idx,_| {
let mut ii = *idx;
izip!(ii.iter_mut(),
shape.iter(),
self.dims.iter())
.for_each(|(j,&d,&f)| if f { *j = d-*j-1 } );
idx.iter_mut().zip(shape.iter()).rev().fold(1,|carry,(i,&d)| { *i += carry; if *i >= d { *i = 0; 1 } else { 0 } } );
Some(st.to_linear(&ii))
}));
}
perm.sort_by_key(|&i| unsafe{ *ridxs.get_unchecked_mut(i) });
let perm = Permutation::from(perm);
xptrb.copy_from_iter(ptr.permute_by(&perm).cloned());
xptre.copy_from_iter(ptr[1..].permute_by(&perm).cloned());
rptr[0] = 0;
izip!(ptr.permute_by(&perm),
ptr[1..].permute_by(&perm),
rptr[1..].iter_mut())
.fold(0,|p,(&p0,&p1,rp)| { *rp = p+p1-p0; *rp });
if let Some(rsp) = rsp {
rsp.copy_from_iter(ridxs.permute_by(&perm).cloned());
println!("flip : sp = {:?}",rsp);
}
izip!(rsubj.chunks_ptr_mut(rptr,&rptr[1..]),
rcof.chunks_ptr_mut(rptr,&rptr[1..]),
subj.chunks_ptr2(xptrb,xptre),
cof.chunks_ptr2(xptrb,xptre))
.for_each(|(rj,rc,j,c)| { rj.copy_from_slice(j); rc.copy_from_slice(c); });
Ok(())
}
}
pub struct ExprPermuteAxes<const N : usize, E:ExprTrait<N>> {
item : E,
perm : [usize; N]
}
impl<const N : usize, E:ExprTrait<N>> ExprTrait<N> for ExprPermuteAxes<N,E> {
fn eval(&self, rs : & mut WorkStack, ws : & mut WorkStack, xs : & mut WorkStack) -> Result<(),ExprEvalError> {
self.item.eval(ws,rs,xs)?;
eval::permute_axes(&self.perm,rs,ws,xs)
}
}
pub enum ExprEither<const N : usize,EL,ER> where EL : ExprTrait<N>, ER : ExprTrait<N> {
Left(EL),
Right(ER)
}
impl<const N : usize, EL,ER> ExprTrait<N> for ExprEither<N,EL,ER> where EL : ExprTrait<N>, ER : ExprTrait<N> {
fn eval(&self, rs : & mut WorkStack, ws : & mut WorkStack, xs : & mut WorkStack) -> Result<(),ExprEvalError> {
match self {
ExprEither::Left(e) => e.eval(rs,ws,xs),
ExprEither::Right(e) => e.eval(rs,ws,xs)
}
}
}
pub use ExprEither::Left as ExprLeft;
pub use ExprEither::Right as ExprRight;
impl From<f64> for Expr<0> {
fn from(v : f64) -> Expr<0> { Expr::new(&[], None, vec![0,1], vec![0], vec![v]) }
}
impl IntoExpr<0> for f64 {
type Result = Expr<0>;
fn into(self) -> Self::Result { Expr::new(&[], None, vec![0,1], vec![0], vec![self]) }
}
impl From<&[f64]> for Expr<1> {
fn from(v : &[f64]) -> Expr<1> { Expr::new(&[v.len()], None, (0..v.len()+1).collect(), vec![0; v.len()], v.to_vec()) }
}
impl From<Vec<f64>> for Expr<1> {
fn from(v : Vec<f64>) -> Expr<1> { Expr::new(&[v.len()], None, (0..v.len()+1).collect(), vec![0; v.len()], v) }
}
impl<const N : usize, E> IntoExpr<N> for E where E : ExprTrait<N>+Sized {
type Result = E;
fn into(self) -> Self::Result { self }
}
impl IntoExpr<1> for &[f64] {
type Result = Expr<1>;
fn into(self) -> Self::Result { Expr::from(self) }
}
impl IntoExpr<1> for Vec<f64> {
type Result = Expr<1>;
fn into(self) -> Self::Result { Expr::from(self) }
}
#[allow(unused)]
#[cfg(test)]
mod test {
use crate::*;
use crate::matrix::*;
use crate::expr::*;
use crate::variable::*;
use crate::dummy::Model;
fn eq<T:std::cmp::Eq>(a : &[T], b : &[T]) -> bool {
a.len() == b.len() && a.iter().zip(b.iter()).all(|(a,b)| *a == *b )
}
fn dense_expr() -> Expr<2> {
super::Expr::new(&[3,3],
None,
vec![0,1,2,3,4,5,6,7,8,9],
vec![0,1,2,0,1,2,0,1,2],
vec![1.1,1.2,1.3,2.1,2.2,2.3,3.1,3.2,3.3])
}
fn sparse_expr() -> Expr<2> {
super::Expr::new(&[3,3],
Some(vec![0,4,5,6,7]),
vec![0,1,2,3,4,5],
vec![0,1,2,3,4],
vec![1.1,2.2,3.3,4.4,5.5])
}
#[allow(non_snake_case)]
#[test]
fn slice() {
let mut m = Model::new(None);
let t = m.variable(Some("t"),unbounded().with_shape(&[2])); let X = m.variable(Some("X"), in_psd_cone().with_dim(4)); let Y = m.variable(Some("Y"), in_psd_cone().with_dim(2)); let mx = dense([2,2], vec![1.1,2.2,3.3,4.4]);
println!("X = {:?}, Y = {:?}",&X,&Y);
m.constraint(Some("X-Y"), X.index([0..2,0..2]).sub(Y.sub((&mx).mul_right(t.index(0)))), domain::zeros(&[2,2]));
let mut rs = WorkStack::new(512);
let mut ws = WorkStack::new(512);
let mut xs = WorkStack::new(512);
{
rs.clear(); ws.clear(); xs.clear();
(&X).into_expr().eval(&mut rs,&mut ws,&mut xs);
let (shape,ptr,sp,subj,cof) = rs.pop_expr();
assert_eq!(shape,&[4,4]);
assert_eq!(ptr,&[0,1,2,3,4,5,6,7,8,9,10,11,12,13,14,15,16]);
assert_eq!(subj,&[2,3,5,8, 3,4,6,9, 5,6,7,10, 8,9,10,11]);
}
{
rs.clear(); ws.clear(); xs.clear();
X.index([0..2,0..2]).into_expr().eval(&mut rs,&mut ws,&mut xs);
let (shape,ptr,sp,subj,cof) = rs.pop_expr();
assert_eq!(shape,&[2,2]);
assert_eq!(ptr,&[0,1,2,3,4]);
assert_eq!(subj,&[2,3,3,4]);
println!("subj = {:?}",subj);
}
{
rs.clear(); ws.clear(); xs.clear();
X.index([0..2,0..2]).into_expr().sub(Y.sub((&mx).mul_right(t.index(0)))).eval(&mut rs,&mut ws,&mut xs);
let (shape,ptr,sp,subj,cof) = rs.pop_expr();
assert_eq!(shape,&[2,2]);
assert_eq!(ptr,&[0,3,6,9,12]);
assert_eq!(subj,&[0,12,2, 0,13,3, 0,13,3, 0,14,4]);
println!("subj = {:?}",subj);
}
}
#[allow(non_snake_case)]
#[test]
fn permute_axes() {
let mut m = Model::new(None);
let u = m.variable(None,&[2,3,4,5,6,7]);
let v = m.variable(None,&[2,3,4,5,6,7]);
let w = m.variable(None,&[2,3,4,5,6,7]);
m.constraint(None, u.add(v).add(w).axispermute(&[3,4,5,0,1,2]).axispermute(&[4,3,2,0,1,5]).axispermute(&[5,4,3,2,1,0]), unbounded().with_shape(&[4,6,5,7,2,3]));
}
#[test]
fn into_symmetric() {
{ let mut rs = WorkStack::new(512);
let mut ws = WorkStack::new(512);
let mut xs = WorkStack::new(512);
let e = Expr::new(&[6,1],
None,
vec![0,1,2,3,4,5,6],
vec![0,1,2,3,4,5],
vec![1.1,2.1,2.2,3.1,3.2,3.3]);
let es = e.into_symmetric(0);
es.eval(& mut rs,& mut ws,& mut xs);
let (shape,ptr,sp,subj,cof) = rs.pop_expr();
assert!(sp.is_none());
assert!(shape.len() == 2);
assert!(shape[0] == 3 && shape[1] == 3);
assert_eq!(ptr, &[0,1,2,3,4,5,6,7,8,9]);
assert_eq!(subj, &[0,1,3,1,2,4,3,4,5]);
assert!(ws.is_empty());
assert!(rs.is_empty());
}
{ let mut rs = WorkStack::new(512);
let mut ws = WorkStack::new(512);
let mut xs = WorkStack::new(512);
let e = Expr::new(&[6,1],
Some(vec![0,2,3,5]),
vec![0,1,2,3,4],
vec![0,2,3,5],
vec![1.1,2.2,3.1,3.3]);
let es = e.into_symmetric(0);
es.eval(& mut rs,& mut ws,& mut xs);
let (shape,ptr,sp,subj,cof) = rs.pop_expr();
assert_eq!(sp.unwrap(), &[0,2,4,6,8]);
assert!(shape.len() == 2);
assert!(shape[0] == 3 && shape[1] == 3);
assert_eq!(ptr, &[0,1,2,3,4,5]);
assert_eq!(subj, &[0,3,2,3,5]);
assert!(ws.is_empty());
assert!(rs.is_empty());
}
}
#[test]
fn mul_left() {
let mut rs = WorkStack::new(512);
let mut ws = WorkStack::new(512);
let mut xs = WorkStack::new(512);
let e0 = dense_expr();
let e1 = sparse_expr();
let m1 = matrix::dense([3,2],vec![1.0,2.0,3.0,4.0,5.0,6.0]);
let m2 = matrix::dense([2,3],vec![1.0,2.0,3.0,4.0,5.0,6.0]);
let e0_1 = m2.clone().mul(e0.clone());
let e0_2 = e0.clone().mul(2.0);
let e1_1 = m2.clone().mul(e1.clone());
let e1_2 = e1.clone().mul(2.0);
e0.eval(& mut rs,& mut ws,& mut xs); assert!(ws.is_empty()); rs.clear();
e1.eval(& mut rs,& mut ws,& mut xs); assert!(ws.is_empty()); rs.clear();
e0_1.eval(& mut rs,& mut ws,& mut xs); assert!(ws.is_empty()); rs.clear();
e0_2.eval(& mut rs,& mut ws,& mut xs); assert!(ws.is_empty()); rs.clear();
e1_1.eval(& mut rs,& mut ws,& mut xs); assert!(ws.is_empty()); rs.clear();
e1_2.eval(& mut rs,& mut ws,& mut xs); assert!(ws.is_empty()); rs.clear();
}
#[test]
fn mul_right() {
let mut rs = WorkStack::new(512);
let mut ws = WorkStack::new(512);
let mut xs = WorkStack::new(512);
let m1 = matrix::dense([3,2],vec![1.0,2.0,3.0,4.0,5.0,6.0]);
let m2 = matrix::dense([2,3],vec![1.0,2.0,3.0,4.0,5.0,6.0]);
let e0 = dense_expr();
let e1 = sparse_expr();
let e0_1 = e0.clone().mul(m1.clone());
let e0_2 = e0.clone().mul(2.0);
let e1_1 = e1.clone().mul(m1.clone());
let e1_2 = e1.clone().mul(2.0);
e0_1.eval(& mut rs,& mut ws,& mut xs); assert!(ws.is_empty()); rs.clear();
e0_2.eval(& mut rs,& mut ws,& mut xs); assert!(ws.is_empty()); rs.clear();
e1_1.eval(& mut rs,& mut ws,& mut xs); assert!(ws.is_empty()); rs.clear();
e1_2.eval(& mut rs,& mut ws,& mut xs); assert!(ws.is_empty()); rs.clear();
}
#[test]
fn add() {
let mut rs = WorkStack::new(512);
let mut ws = WorkStack::new(512);
let mut xs = WorkStack::new(512);
let m1 = matrix::dense([3,3],vec![1.0,2.0,3.0,4.0,5.0,6.0,7.0,8.0,9.0]);
let e0 = dense_expr().add(sparse_expr()).add(dense_expr().mul(m1));
e0.eval(& mut rs,& mut ws,& mut xs); assert!(ws.is_empty()); rs.clear();
}
#[test]
fn repeat() {
let ed = super::Expr::new(&[3,2,1],
None,
(0..7).collect(),
(0..6).collect(),
(0..6).map(|v| v as f64 * 1.1).collect());
let es = super::Expr::new(&[3,2,1],
Some(vec![0,2,3,5]),
(0..5).collect(),
vec![6,8,9,11],
(0..4).map(|v| v as f64 * 1.1).collect());
let mut rs = WorkStack::new(512);
let mut ws = WorkStack::new(512);
let mut xs = WorkStack::new(512);
ed.clone().repeat(0,2).eval(& mut rs,& mut ws,& mut xs);
let (shape,ptr,sp,subj,_cof) = rs.pop_expr();
assert_eq!(shape.len(),3);
assert_eq!(shape[0],6);
assert_eq!(shape[1],2);
assert_eq!(shape[2],1);
assert_eq!(*ptr.last().unwrap(), 12);
assert!(sp.is_none());
assert_eq!(ptr,&[0,1,2,3,4,5,6,7,8,9,10,11,12]);
assert_eq!(subj,&[0,1,2,3,4,5,0,1,2,3,4,5]);
rs.clear();
ws.clear();
xs.clear();
ed.clone().repeat(1,2).eval(& mut rs,& mut ws,& mut xs);
let (shape,ptr,sp,subj,_cof) = rs.pop_expr();
assert_eq!(shape.len(),3);
assert_eq!(shape[0],3);
assert_eq!(shape[1],4);
assert_eq!(shape[2],1);
assert_eq!(*ptr.last().unwrap(), 12);
assert!(sp.is_none());
assert_eq!(ptr,&[0,1,2,3,4,5,6,7,8,9,10,11,12]);
assert_eq!(subj,&[0,1,0,1,2,3,2,3,4,5,4,5]);
rs.clear();
ws.clear();
xs.clear();
ed.clone().repeat(2,2).eval(& mut rs,& mut ws,& mut xs);
let (shape,ptr,sp,subj,_cof) = rs.pop_expr();
assert_eq!(shape.len(),3);
assert_eq!(shape[0],3);
assert_eq!(shape[1],2);
assert_eq!(shape[2],2);
assert_eq!(*ptr.last().unwrap(), 12);
assert!(sp.is_none());
assert_eq!(ptr,&[0,1,2,3,4,5,6,7,8,9,10,11,12]);
assert_eq!(subj,&[0,0,1,1,2,2,3,3,4,4,5,5]);
rs.clear();
ws.clear();
xs.clear();
es.clone().repeat(0,2).eval(&mut rs, &mut ws, &mut xs);
let (shape,ptr,sp,subj,_cof) = rs.pop_expr();
assert_eq!(shape.len(),3);
assert_eq!(shape,&[6,2,1]);
assert_eq!(*ptr.last().unwrap(), 8);
assert!(sp.is_some());
assert_eq!(ptr,&[0,1,2,3,4,5,6,7,8]);
assert_eq!(subj,&[6,8,9,11,6,8,9,11]);
rs.clear();
ws.clear();
xs.clear();
es.clone().repeat(1,2).eval(&mut rs, &mut ws, &mut xs);
let (shape,ptr,sp,subj,_cof) = rs.pop_expr();
assert_eq!(shape.len(),3);
assert_eq!(shape,&[3,4,1]);
assert_eq!(*ptr.last().unwrap(), 8);
assert!(sp.is_some());
assert_eq!(sp.unwrap(),&[0,2,4,5,6,7,9,11]);
assert_eq!(ptr,&[0,1,2,3,4,5,6,7,8]);
assert_eq!(subj,&[6,6,8,9,8,9,11,11]);
rs.clear();
ws.clear();
xs.clear();
es.clone().repeat(2,2).eval(&mut rs, &mut ws, &mut xs);
let (shape,ptr,sp,subj,_cof) = rs.pop_expr();
assert_eq!(shape.len(),3);
assert_eq!(shape,&[3,2,2]);
assert_eq!(*ptr.last().unwrap(), 8);
assert!(sp.is_some());
assert_eq!(sp.unwrap(),&[0,1,4,5,6,7,10,11]);
assert_eq!(ptr,&[0,1,2,3,4,5,6,7,8]);
assert_eq!(subj,&[6, 6, 8, 8,9,9,11,11]);
}
#[test]
fn stack() {
let e0 = super::Expr::new(&[3,2,1],
None,
(0..7).collect(),
(0..6).collect(),
(0..6).map(|v| v as f64 * 1.1).collect());
let e1 = super::Expr::new(&[3,2,1],
Some(vec![0,2,3,5]),
(0..5).collect(),
vec![6,8,9,11],
(0..4).map(|v| v as f64 * 1.1).collect());
let s1_0 = e0.clone().stack(0,e0.clone());
let s1_1 = e0.clone().stack(1,e0.clone());
let s1_2 = e0.clone().stack(2,e0.clone());
let s2_0 = e0.clone().stack(0,e1.clone());
let s2_1 = e0.clone().stack(1,e1.clone());
let s2_2 = e0.clone().stack(2,e1.clone());
let mut rs = WorkStack::new(512);
let mut ws = WorkStack::new(512);
let mut xs = WorkStack::new(512);
s1_0.eval(& mut rs,& mut ws,& mut xs);
let (shape,ptr,sp,subj,cof) = rs.pop_expr();
assert!(eq(shape,&[6,2,1]));
assert!(eq(ptr,&[0,1,2,3,4,5,6,7,8,9,10,11,12]));
assert!(eq(subj,&[0,1,2,3,4,5,0,1,2,3,4,5]));
assert!(rs.is_empty());
assert!(ws.is_empty());
s1_1.eval(& mut rs,& mut ws,& mut xs);
let (shape,ptr,sp,subj,cof) = rs.pop_expr();
assert!(eq(shape,&[3,4,1]));
assert!(eq(ptr,&[0,1,2,3,4,5,6,7,8,9,10,11,12]));
assert!(eq(subj,&[0,1,0,1,2,3,2,3,4,5,4,5]));
assert!(rs.is_empty());
assert!(ws.is_empty());
s1_2.eval(& mut rs,& mut ws,& mut xs);
let (shape,ptr,sp,subj,cof) = rs.pop_expr();
assert!(eq(shape,&[3,2,2]));
assert!(eq(ptr,&[0,1,2,3,4,5,6,7,8,9,10,11,12]));
assert!(eq(subj,&[0,0,1,1,2,2,3,3,4,4,5,5]));
assert!(rs.is_empty());
assert!(ws.is_empty());
s2_0.eval(& mut rs,& mut ws,& mut xs);
let (shape,ptr,sp,subj,cof) = rs.pop_expr();
assert!(eq(shape,&[6,2,1]));
assert!(eq(sp.unwrap(),&[0,1,2,3,4,5,6,8,9,11]));
assert!(eq(ptr,&[0,1,2,3,4,5,6,7,8,9,10]));
assert!(eq(subj,&[0,1,2,3,4,5,6,8,9,11]));
assert!(rs.is_empty());
assert!(ws.is_empty());
s2_1.eval(& mut rs,& mut ws,& mut xs);
let (shape,ptr,sp,subj,cof) = rs.pop_expr();
assert!(eq(shape,&[3,4,1]));
assert!(eq(sp.unwrap(),&[0,1,2,4,5,6,7,8,9,11]));
assert!(eq(ptr,&[0,1,2,3,4,5,6,7,8,9,10]));
assert!(eq(subj,&[0,1,6,2,3,8,9,4,5,11]));
assert!(rs.is_empty());
assert!(ws.is_empty());
s2_2.eval(& mut rs,& mut ws,& mut xs);
let (shape,ptr,sp,subj,_cof) = rs.pop_expr();
assert!(eq(shape,&[3,2,2]));
assert!(eq(sp.unwrap(),&[0,1,2,4,5,6,7,8,10,11]));
assert!(eq(ptr,&[0,1,2,3,4,5,6,7,8,9,10]));
assert!(eq(subj,&[0,6,1,2,8,3,9,4,5,11]));
assert!(rs.is_empty());
assert!(ws.is_empty());
let s3_0 = e1.clone().stack(0,e1.clone());
s3_0.eval(& mut rs,& mut ws,& mut xs);
let (shape,ptr,sp,subj,_cof) = rs.pop_expr();
assert!(eq(shape,&[6,2,1]));
assert!(eq(sp.unwrap(),&[0,2,3,5,6,8,9,11]));
assert!(eq(ptr,&[0,1,2,3,4,5,6,7,8]));
assert!(eq(subj,&[6,8,9,11,6,8,9,11]));
assert!(rs.is_empty());
assert!(ws.is_empty());
let s3_1 = e1.clone().stack(1,e1.clone());
s3_1.eval(& mut rs,& mut ws,& mut xs);
let (shape,ptr,sp,subj,_cof) = rs.pop_expr();
assert!(eq(shape,&[3,4,1]));
assert!(eq(sp.unwrap(),&[0,2,4,5,6,7,9,11]));
assert!(eq(ptr,&[0,1,2,3,4,5,6,7,8]));
assert!(eq(subj,&[6,6,8,9,8,9,11,11]));
assert!(rs.is_empty());
assert!(ws.is_empty());
let s3_2 = e1.clone().stack(2,e1.clone());
s3_2.eval(& mut rs,& mut ws,& mut xs);
let (shape,ptr,sp,subj,_cof) = rs.pop_expr();
assert!(eq(shape,&[3,2,2]));
assert!(eq(sp.unwrap(),&[0,1,4,5,6,7,10,11]));
assert!(eq(ptr,&[0,1,2,3,4,5,6,7,8]));
assert!(eq(subj,&[6,6,8,8,9,9,11,11]));
assert!(rs.is_empty());
assert!(ws.is_empty());
e0.clone().stack(0,e1.clone()).stack(0,e0.clone()).eval(& mut rs,& mut ws,& mut xs);
let (shape,ptr,sp,subj,_cof) = rs.pop_expr();
assert!(eq(shape,&[9,2,1]));
assert!(eq(sp.unwrap(),&[0,1,2,3,4,5,
6,8,9,11,
12,13,14,15,16,17]));
assert!(eq(ptr,&[0,1,2,3,4,5,6,7,8,9,10,11,12,13,14,15,16]));
assert!(eq(subj,&[0,1,2,3,4,5,
6,8,9,11,
0,1,2,3,4,5]));
assert!(rs.is_empty());
assert!(ws.is_empty());
e0.clone().stack(1,e1.clone()).stack(1,e0.clone()).eval(& mut rs,& mut ws,& mut xs);
let (shape,ptr,sp,subj,_cof) = rs.pop_expr();
assert!(eq(shape,&[3,6,1]));
assert!(eq(sp.unwrap(),&[0,1,2,4,5,
6,7,8,9,10,11,
12,13,15,16,17]));
assert!(eq(ptr,&[0,1,2,3,4,5,6,7,8,9,10,11,12,13,14,15,16]));
assert!(eq(subj,&[0,1,6,0,1,2,3,8,9,2,3,4,5,11,4,5]));
assert!(rs.is_empty());
assert!(ws.is_empty());
e0.clone().stack(2,e1.clone()).stack(2,e0.clone()).eval(& mut rs,& mut ws,& mut xs);
let (shape,ptr,sp,subj,_cof) = rs.pop_expr();
assert!(eq(shape,&[3,2,3]));
assert!(eq(sp.unwrap(),&[0,1,2,3,5,6,7,8,9,10,11,12,14,15,16,17]));
assert!(eq(ptr,&[0,1,2,3,4,5,6,7,8,9,10,11,12,13,14,15,16]));
assert!(eq(subj,&[0,6,0,
1,1,
2,8,2,
3,9,3,
4,4,
5,11,5]));
assert!(rs.is_empty());
assert!(ws.is_empty());
{
let mut rs = WorkStack::new(512);
let mut ws = WorkStack::new(512);
let mut xs = WorkStack::new(512);
let ed = Expr::new(&[2,2],None,vec![0,1,2,3,4],vec![1,2,3,4],vec![1.1,1.2,2.1,2.2]);
let es = Expr::new(&[2,2],Some(vec![]), vec![0], vec![], vec![]);
vstack![hstack![ed.clone(),es.clone()], hstack![es,ed]].eval(& mut rs,& mut ws,& mut xs);
let (shape,ptr,sp,subj,cof) = rs.pop_expr();
assert_eq!(shape,&[4,4]);
assert_eq!(ptr,&[0,1,2,3,4,5,6,7,8]);
assert!(sp.is_some());
assert_eq!(sp.unwrap(),&[0,1,4,5,10,11,14,15]);
assert_eq!(subj,&[1,2,3,4,1,2,3,4]);
}
{
let mut rs = WorkStack::new(512);
let mut ws = WorkStack::new(512);
let mut xs = WorkStack::new(512);
let ed = Expr::new(&[2,2],None,vec![0,1,2,3,4],vec![1,2,3,4],vec![1.1,1.2,2.1,2.2]);
let es = Expr::new(&[2,2],Some(vec![]), vec![0], vec![], vec![]);
exprcat![
ed.clone(), es.clone() ;
es.clone(), ed.clone() ].eval(& mut rs,& mut ws,& mut xs);
let (shape,ptr,sp,subj,cof) = rs.pop_expr();
assert_eq!(shape,&[4,4]);
assert_eq!(ptr,&[0,1,2,3,4,5,6,7,8]);
assert!(sp.is_some());
assert_eq!(sp.unwrap(),&[0,1,4,5,10,11,14,15]);
assert_eq!(subj,&[1,2,3,4,1,2,3,4]);
}
}
#[allow(non_snake_case)]
#[test]
fn sum_on() {
let mut m = Model::new(None);
let u = m.variable(None,&[3,4,3,4]);
let v = m.variable(None,&[4,3,4,3]);
let w = m.variable(None,&[4,3,3,4]);
{
let u = u.clone();
let v = v.clone();
let w = w.clone();
m.constraint(
None,
u.add(v.axispermute(&[1,0,1,0])).add(w.axispermute(&[1,0,2,3])).sum_on(&[0,3]),
unbounded().with_shape(&[3,4]));
}
{
let u = u.clone();
let v = v.clone();
let w = w.clone();
m.constraint(None, u.add(v.axispermute(&[1,0,1,0])).add(w.axispermute(&[1,0,2,3])).sum_on(&[1,3]), unbounded().with_shape(&[4,4]));
}
{
let u = u.clone();
let v = v.clone();
let w = w.clone();
m.constraint(None, u.add(v.axispermute(&[1,0,1,0])).add(w.axispermute(&[1,0,2,3])).sum_on(&[2]), unbounded().with_shape(&[3]));
}
}
#[allow(non_snake_case)]
#[test]
fn dot_rows_x() {
let dmx = NDArray::new([512,512], None, vec![1.0; 512*512]).unwrap();
let mut M = Model::new(None);
let dv = M.variable(None, unbounded().with_shape(&[512,512])); let dw = M.variable(None, unbounded().with_shape(&[512,512]));
let mut rs = WorkStack::new(1024);
let mut ws = WorkStack::new(1024);
let mut xs = WorkStack::new(1024);
dv.clone().add(dw.clone()).dot_rows(dmx.clone()).eval(&mut rs, &mut ws, &mut xs).unwrap();
}
#[allow(non_snake_case)]
#[test]
fn dot_rows() {
let dmx = NDArray::new([3,3], None, vec![1.1,1.2,1.3,2.1,2.2,2.3,3.1,3.2,3.3]).unwrap();
let smx = NDArray::new([3,3], Some(vec![0,2,7]), vec![1.1,1.3,3.2]).unwrap();
let mut M = Model::new(None);
let dv = M.variable(None, unbounded().with_shape(&[3,3])); let dw = M.variable(None, unbounded().with_shape(&[3,3])); let sv = M.variable(None, unbounded().with_shape_and_sparsity(&[3,3],&[[0,0],[0,1],[0,2],[2,0],[2,1],[2,2]])); let sw = M.variable(None, unbounded().with_shape_and_sparsity(&[3,3],&[[0,0],[0,2],[2,1]]));
let mut rs = WorkStack::new(1024);
let mut ws = WorkStack::new(1024);
let mut xs = WorkStack::new(1024);
{
dw.clone().add(dv.clone()).dot_rows(dmx.clone()).eval(&mut rs,&mut ws,&mut xs);
let (shape,ptr,sp,subj,cof) = rs.pop_expr();
assert_eq!(shape,[3]);
assert!(sp.is_none());
assert_eq!(ptr,[0,6,12,18]);
assert_eq!(subj,[0,9,1,10,2,11,3,12,4,13,5,14,6,15,7,16,8,17]);
assert_eq!(cof,[1.1,1.1,1.2,1.2,1.3,1.3,2.1,2.1,2.2,2.2,2.3,2.3,3.1,3.1,3.2,3.2,3.3,3.3]);
}
{
rs.clear(); ws.clear(); xs.clear();
dw.clone().add(dv.clone()).dot_rows(smx.clone()).eval(&mut rs,&mut ws,&mut xs);
let (shape,ptr,sp,subj,cof) = rs.pop_expr();
assert_eq!(shape,[3]);
assert!(sp.is_some());
assert_eq!(sp.unwrap(),[0,2]);
assert_eq!(ptr,[0,4,6]);
assert_eq!(subj,[0,9,2,11,7,16]);
assert_eq!(cof,[1.1,1.1,1.3,1.3,3.2,3.2]);
}
{
rs.clear(); ws.clear(); xs.clear();
sw.clone().add(sv.clone()).dot_rows(dmx.clone()).eval(&mut rs,&mut ws,&mut xs);
let (shape,ptr,sp,subj,cof) = rs.pop_expr();
assert_eq!(shape,[3]);
assert!(sp.is_some());
assert_eq!(sp.unwrap(),[0,2]);
assert_eq!(ptr,[0,5,9]);
assert_eq!(subj,[18,24,19,20,25,21,22,26,23]);
assert_eq!(cof,[1.1,1.1,1.2,1.3,1.3,3.1,3.2,3.2,3.3]);
}
{
rs.clear(); ws.clear(); xs.clear();
sw.clone().add(sv.clone()).dot_rows(smx.clone()).eval(&mut rs,&mut ws,&mut xs);
let (shape,ptr,sp,subj,cof) = rs.pop_expr();
println!("ptr = {:?}",ptr);
println!("sp = {:?}",sp);
println!("subj = {:?}",subj);
println!("cof = {:?}",cof);
assert_eq!(shape,[3]);
assert!(sp.is_some());
assert_eq!(sp.unwrap(),[0,2]);
assert_eq!(ptr,[0,4,6]);
assert_eq!(subj,[18,24,20,25,22,26]);
assert_eq!(cof,[1.1,1.1,1.3,1.3,3.2,3.2]);
}
}
#[allow(non_snake_case)]
#[test]
fn mul_elem() {
let dmx = NDArray::new([3,3], None, vec![1.1,1.2,1.3,2.1,2.2,2.3,3.1,3.2,3.3]).unwrap();
let smx = NDArray::new([3,3], Some(vec![0,2,7]), vec![1.1,1.3,3.2]).unwrap();
let mut M = Model::new(None);
let dv = M.variable(None, unbounded().with_shape(&[3,3])); let dw = M.variable(None, unbounded().with_shape(&[3,3])); let sv = M.variable(None, unbounded().with_shape_and_sparsity(&[3,3],&[[0,0],[0,1],[0,2],[2,0],[2,1],[2,2]])); let sw = M.variable(None, unbounded().with_shape_and_sparsity(&[3,3],&[[0,0],[0,2],[2,1]]));
let mut rs = WorkStack::new(1024);
let mut ws = WorkStack::new(1024);
let mut xs = WorkStack::new(1024);
{
dw.clone().add(dv.clone()).mul_elem(dmx.clone()).eval(&mut rs,&mut ws,&mut xs);
let (shape,ptr,sp,subj,cof) = rs.pop_expr();
assert_eq!(shape,[3,3]);
assert!(sp.is_none());
assert_eq!(ptr,[0,2,4,6,8,10,12,14,16,18]);
assert_eq!(subj,[0,9,1,10,2,11,3,12,4,13,5,14,6,15,7,16,8,17]);
assert_eq!(cof,[1.1,1.1,1.2,1.2,1.3,1.3,2.1,2.1,2.2,2.2,2.3,2.3,3.1,3.1,3.2,3.2,3.3,3.3]);
}
{
rs.clear(); ws.clear(); xs.clear();
dw.clone().add(dv.clone()).mul_elem(smx.clone()).eval(&mut rs,&mut ws,&mut xs);
let (shape,ptr,sp,subj,cof) = rs.pop_expr();
assert_eq!(shape,[3,3]);
assert!(sp.is_some());
assert_eq!(sp.unwrap(),[0,2,7]);
assert_eq!(ptr,[0,2,4,6]);
assert_eq!(subj,[0,9,2,11,7,16]);
assert_eq!(cof,[1.1,1.1,1.3,1.3,3.2,3.2]);
}
{
rs.clear(); ws.clear(); xs.clear();
sw.clone().add(sv.clone()).mul_elem(dmx.clone()).eval(&mut rs,&mut ws,&mut xs);
let (shape,ptr,sp,subj,cof) = rs.pop_expr();
assert_eq!(shape,[3,3]);
assert!(sp.is_some());
assert_eq!(sp.unwrap(),[0,1,2,6,7,8]);
assert_eq!(ptr,[0,2,3,5,6,8,9]);
assert_eq!(subj,[18,24,19,20,25,21,22,26,23]);
assert_eq!(cof,[1.1,1.1,1.2,1.3,1.3,3.1,3.2,3.2,3.3]);
}
{
rs.clear(); ws.clear(); xs.clear();
sw.clone().add(sv.clone()).mul_elem(smx.clone()).eval(&mut rs,&mut ws,&mut xs);
let (shape,ptr,sp,subj,cof) = rs.pop_expr();
assert_eq!(shape,[3,3]);
assert!(sp.is_some());
assert_eq!(sp.unwrap(),[0,2,7]);
assert_eq!(ptr,[0,2,4,6]);
assert_eq!(subj,[18,24,20,25,22,26]);
assert_eq!(cof,[1.1,1.1,1.3,1.3,3.2,3.2]);
}
}
#[test]
fn map_expr() {
let mut model = Model::new(None);
let x = model.variable(None,&[5,5,5]); let s = model.variable(None,unbounded().with_shape(&[5,5,5]).with_sparsity(&[[0,0,0],[1,0,1],[1,2,1],[2,1,2],[3,2,2]]));
let mut rs = WorkStack::new(1024);
let mut ws = WorkStack::new(1024);
let mut xs = WorkStack::new(1024);
println!("x = {:?}",x);
println!("s = {:?}",s);
{
rs.clear(); ws.clear(); xs.clear();
(&s).into_expr().map(&[5,5],|i| if i[0] == i[2] && i[0] > i[1] { Some([i[0],i[1]]) } else { None }).eval(&mut rs,&mut ws,&mut xs).unwrap();
let (shape,ptr,sp,subj,cof) = rs.pop_expr();
assert_eq!(shape,&[5,5]);
assert_eq!(sp.unwrap(),&[5,11]);
assert_eq!(ptr,&[0,1,2]);
assert_eq!(subj,&[126,128]);
}
println!("---------------------------------");
{
rs.clear(); ws.clear(); xs.clear();
(&x).into_expr().map(&[5,5],|i| if i[0] == i[2] && i[0] > i[1] { Some([i[0],i[1]]) } else { None }).eval(&mut rs,&mut ws,&mut xs).unwrap();
let (shape,ptr,sp,subj,cof) = rs.pop_expr();
assert_eq!(shape,&[5,5]);
assert_eq!(sp.unwrap(),&[5usize,10,11,15,16,17,20,21,22,23]);
assert_eq!(ptr,&[0,1,2,3,4,5,6,7,8,9,10]);
assert_eq!(subj,&[26,52,57,78,83,88,104,109,114,119]);
}
println!("---------------------------------");
{
rs.clear(); ws.clear(); xs.clear();
(&x).into_expr().map(&[5,5],|i| if i[0] == i[2] && i[0] > i[1] { Some([4-i[0],i[1]]) } else { None }).eval(&mut rs,&mut ws,&mut xs).unwrap();
let (shape,ptr,sp,subj,cof) = rs.pop_expr();
assert_eq!(shape,&[5,5]);
assert_eq!(sp.unwrap(),&[0,1,2,3,5,6,7,10,11,15]);
assert_eq!(ptr,&[0,1,2,3,4,5,6,7,8,9,10]);
assert_eq!(subj,&[104,109,114,119, 78,83,88, 52,57, 26]);
}
}
#[test]
fn expr_flip() {
let mut model = Model::new(None);
let x = model.variable(None,&[3,2,3]); let s = model.variable(None,unbounded().with_shape(&[3,2,3]).with_sparsity(&[[0,0,0],[1,0,1],[1,1,1],[1,1,2],[2,1,2]])); let mut rs = WorkStack::new(1024);
let mut ws = WorkStack::new(1024);
let mut xs = WorkStack::new(1024);
{
(&x).into_expr().flip(&[true,false,true]).eval(&mut rs,&mut ws,&mut xs).unwrap();
let (shape,ptr,sp,subj,_cof) = rs.pop_expr();
assert_eq!(shape,&[3,2,3]);
assert_eq!(ptr,&[0, 1,2,3, 4,5,6,
7,8,9, 10,11,12,
13,14,15, 16,17,18]);
assert_eq!(subj,&[14,13,12, 17,16,15, 8,7,6, 11,10,9, 2,1,0, 5,4,3 ]);
}
{
(&s).into_expr().flip(&[true,false,true]).eval(&mut rs,&mut ws,&mut xs).unwrap();
let (shape,ptr,sp,subj,_cof) = rs.pop_expr();
assert_eq!(shape,&[3,2,3]);
assert_eq!(sp.unwrap(),&[3,7,9,10,14]);
assert_eq!(ptr,&[0,1,2,3,4,5]);
assert_eq!(subj,&[22,19,21,20,18]);
}
{
s.add(&x).flip(&[true,false,true]).eval(&mut rs,&mut ws,&mut xs).unwrap();
let (shape,ptr,sp,subj,_cof) = rs.pop_expr();
assert_eq!(shape,&[3,2,3]);
assert_eq!(ptr,&[0,1,2,3,5,6,7, 8,10,11,13,15,16, 17,18,20,21,22,23 ]);
assert_eq!(subj,&[14,13,12, 17,22,16,15,
8,7,19,6,11,21,10,20,9,
2,1,0,18,5,4,3 ]);
}
}
}