use nalgebra::*;
use nalgebra::storage::*;
pub struct PatchIterator<'a, N, S, W, V>
where
N : Scalar,
S : Storage<N, Dynamic, Dynamic>,
W : Dim,
V : Dim
{
source : &'a Matrix<N, Dynamic, Dynamic, S>,
size : (usize, usize),
curr_pos : (usize, usize),
_c_stride : usize,
step_v : V,
step_h : W,
_row_wise : bool,
pool_dims : (usize, usize)
}
impl<'a, N, S, W, V> PatchIterator<'a, N, S, W, V>
where
N : Scalar,
S : Storage<N, Dynamic, Dynamic>,
W : Dim,
V : Dim
{
pub fn pool<F>(
mut self,
mut f : F
) -> Matrix<N, Dynamic, Dynamic, VecStorage<N, Dynamic, Dynamic>>
where
F : FnMut(Matrix<N, Dynamic, Dynamic, SliceStorage<'a, N, Dynamic, Dynamic, S::RStride, S::CStride>>)->N
{
let mut data : Vec<N> = Vec::with_capacity(self.pool_dims.1 * self.pool_dims.0);
while let Some(w) = self.next() {
let s = f(w);
data.push(s);
}
let mut ans = DMatrix::<N>::from_vec(self.pool_dims.1, self.pool_dims.0, data);
ans.transpose_mut();
ans
}
}
pub trait WindowIterate<N, S>
where
N : Scalar,
S : Storage<N, Dynamic, Dynamic>,
{
fn windows(&self, win_sz : (usize, usize)) -> PatchIterator<N, S, U1, U1>;
}
pub trait ChunkIterate<N, S>
where
N : Scalar,
S : Storage<N, Dynamic, Dynamic>,
{
fn chunks(&self, sz : (usize, usize)) -> PatchIterator<N, S, Dynamic, Dynamic>;
}
impl<'a, N, S, W, V> Iterator for PatchIterator<'a, N, S, W, V>
where
N : Scalar,
S : Storage<N, Dynamic, Dynamic>,
W : Dim,
V : Dim
{
type Item = Matrix<N, Dynamic, Dynamic, SliceStorage<'a, N, Dynamic, Dynamic, S::RStride, S::CStride>>;
fn next(&mut self) -> Option<Self::Item> {
let win = if self.curr_pos.0 + self.size.0 <= self.source.nrows() && self.curr_pos.1 + self.size.1 <= self.source.ncols() {
Some(self.source.slice(self.curr_pos, self.size))
} else {
None
};
self.curr_pos.1 += self.step_h.value(); if self.curr_pos.1 + self.size.1 > self.source.ncols() { self.curr_pos.1 = 0;
self.curr_pos.0 += self.step_v.value();
}
win
}
}
impl<N, S> WindowIterate<N, S> for Matrix<N, Dynamic, Dynamic, S>
where
N : Scalar,
S : Storage<N, Dynamic, Dynamic>,
{
fn windows(
&self,
sz : (usize, usize)
) -> PatchIterator<N, S, U1, U1> {
if self.nrows() % sz.0 != 0 || self.ncols() % sz.1 != 0 {
panic!("Matrix size should be a multiple of window size");
}
let pool_dims = (self.nrows() - sz.0 + 1, self.ncols() - sz.1 + 1);
PatchIterator::<N, S, U1, U1> {
source : &self,
size : sz,
curr_pos : (0, 0),
_c_stride : self.nrows(),
step_h : U1{},
step_v : U1{},
_row_wise : false,
pool_dims
}
}
}
impl<N, S> ChunkIterate<N, S> for Matrix<N, Dynamic, Dynamic, S>
where
N : Scalar,
S : Storage<N, Dynamic, Dynamic>,
{
fn chunks(
&self,
sz : (usize, usize)
) -> PatchIterator<N, S, Dynamic, Dynamic> {
let step_v = Dim::from_usize(sz.0);
let step_h = Dim::from_usize(sz.1);
if self.nrows() % sz.0 != 0 || self.ncols() % sz.1 != 0 {
panic!("Matrix size should be a multiple of window size");
}
let pool_dims = (self.nrows() / sz.0, self.ncols() / sz.1);
PatchIterator::<N, S, Dynamic, Dynamic> {
source : &self,
size : sz,
curr_pos : (0, 0),
_c_stride : self.nrows(),
step_v,
step_h,
_row_wise : false,
pool_dims
}
}
}