use smallvec::smallvec;
use crate::{Element, Recordable, Shape, Tensor};
use super::{Cotangents, Operation, Reads, unary};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) struct Narrow {
pub(crate) axis: usize,
pub(crate) start: usize,
pub(crate) len: usize,
}
impl Narrow {
pub(crate) fn arity(&self) -> usize {
1
}
pub(crate) fn reads(&self) -> Reads {
Reads::NOTHING
}
pub(crate) fn infer_shape(&self, operands: &[Shape]) -> Shape {
let operand = unary(operands);
assert!(
self.axis < operand.rank(),
"narrow axis {} is out of rank for {operand}",
self.axis
);
assert!(self.len > 0, "narrow window must hold at least one element");
let extent = operand.axes()[self.axis];
let end = self
.start
.checked_add(self.len)
.expect("narrow window end overflows `usize`");
assert!(
end <= extent,
"narrow window {}..{end} exceeds axis {} extent {extent}",
self.start,
self.axis
);
Shape::new(
operand
.axes()
.iter()
.enumerate()
.map(|(index, &e)| if index == self.axis { self.len } else { e }),
)
}
}
impl Narrow {
pub(crate) fn forward<E: Element>(&self, operands: &[&Tensor<E>]) -> Tensor<E> {
unary(operands).narrow(self.axis, self.start, self.len)
}
}
impl<Rule: Recordable> Operation<Rule> for Narrow {
fn backward(&self, operands: &[&Rule], _output: &Rule, gradient: &Rule) -> Cotangents<Rule> {
let &operand = unary(operands);
let full_extent = operand.shape().axes()[self.axis];
smallvec![Some(gradient.pad(self.axis, self.start, full_extent))]
}
}