use crate::{Elementary, Shape, Tensorial};
use super::Value;
impl<'network, Data: Elementary> Value<'network, Data> {
pub fn abs(self) -> Self {
self.maximum(-self)
}
}
impl<'network, Data: Tensorial> Value<'network, Data> {
pub fn softmax(self, axis: usize) -> Self {
self.log_softmax(axis).exp()
}
pub fn mean_along(self, axis: usize) -> Self {
let shape = self.shape();
assert!(axis < shape.rank(), "mean_along axis {axis} is out of rank");
let extent = shape.axes()[axis];
self.sum_along(axis) / Data::counted(shape.without_axis(axis), extent)
}
pub fn broadcast_to(self, shape: impl Into<Shape>) -> Self {
let target = shape.into();
let source = self.shape();
if source == target {
return self;
}
assert!(
target.rank() >= source.rank(),
"broadcast to {target} from {source} lowers the rank"
);
let offset = target.rank() - source.rank();
for (axis, &extent) in source.axes().iter().enumerate() {
let aligned = target.axes()[offset + axis];
assert!(
extent == aligned || extent == 1,
"broadcast to {target} from {source} cannot align source axis \
{axis} of extent {extent} to extent {aligned}"
);
}
if source.volume() == 1 {
let reference = self.literal(Data::counted(target, 0));
return self.broadcast_like(reference);
}
let mut current = if offset == 0 {
self
} else {
let mut axes = vec![1; offset];
axes.extend_from_slice(source.axes());
self.reshape(axes)
};
for axis in 0..target.rank() {
let aligned = target.axes()[axis];
if current.shape().axes()[axis] == aligned {
continue;
}
let mut axes = current.shape().axes().to_vec();
axes[axis] = aligned;
let reference = self.literal(Data::counted(Shape::new(axes), 0));
current = current.squeeze(axis).broadcast_along(axis, reference);
}
current
}
pub fn broadcast_pair(self, other: Self) -> (Self, Self) {
let common = broadcasted_shape(&self.shape(), &other.shape());
(
self.broadcast_to(common.clone()),
other.broadcast_to(common),
)
}
}
pub fn concat<'network, Data: Tensorial>(
values: &[Value<'network, Data>],
axis: usize,
) -> Value<'network, Data> {
let first = values.first().expect("concat requires at least one value");
let reference = first.shape();
assert!(
axis < reference.rank(),
"concat axis {axis} is out of rank for {reference}"
);
for value in &values[1..] {
let shape = value.shape();
assert_eq!(
shape.without_axis(axis),
reference.without_axis(axis),
"concat along axis {axis} requires equal shapes off the axis, \
got {shape} against {reference}"
);
}
if values.len() == 1 {
return *first;
}
let combined: usize = values.iter().map(|value| value.shape().axes()[axis]).sum();
let mut offset = 0;
let mut total: Option<Value<'network, Data>> = None;
for &value in values {
let padded = value.pad(axis, offset, combined);
offset += value.shape().axes()[axis];
total = Some(match total {
Some(sum) => sum + padded,
None => padded,
});
}
total.expect("concat combines at least one value")
}
pub fn stack<'network, Data: Tensorial>(
values: &[Value<'network, Data>],
axis: usize,
) -> Value<'network, Data> {
let lifted: Vec<Value<'network, Data>> =
values.iter().map(|&value| value.unsqueeze(axis)).collect();
concat(&lifted, axis)
}
fn broadcasted_shape(left: &Shape, right: &Shape) -> Shape {
let rank = left.rank().max(right.rank());
let mut axes = Vec::with_capacity(rank);
for offset in 0..rank {
let left_extent = extent_from_end(left, offset);
let right_extent = extent_from_end(right, offset);
assert!(
left_extent == right_extent || left_extent == 1 || right_extent == 1,
"broadcast of {left} and {right} cannot align extents \
{left_extent} and {right_extent}"
);
axes.push(left_extent.max(right_extent));
}
axes.reverse();
Shape::new(axes)
}
fn extent_from_end(shape: &Shape, offset: usize) -> usize {
let rank = shape.rank();
if offset < rank {
shape.axes()[rank - 1 - offset]
} else {
1
}
}
#[cfg(test)]
#[path = "tests/composite_tests.rs"]
mod tests;