use std::ops::Add;
use static_assertions::assert_impl_all;
use crate::{Element, Tensor};
use super::Field;
use super::{Kinship, Network, Origin, SlotStore, Symbol, ValueId};
assert_impl_all!(Parameters<f64>: Send, Sync);
#[derive(Debug, Clone)]
pub struct Parameters<E> {
origin: Origin,
store: SlotStore<Tensor<E>>,
}
impl<E: Element> Parameters<E> {
pub(crate) fn new(origin: Origin, store: SlotStore<Tensor<E>>) -> Self {
Self { origin, store }
}
pub(crate) fn from_rows(
origin: Origin,
rows: impl IntoIterator<Item = (ValueId, Tensor<E>)>,
) -> Self {
Self {
origin,
store: SlotStore::from_rows(rows),
}
}
pub(crate) fn origin(&self) -> Origin {
self.origin
}
fn kinship(&self) -> Kinship {
Kinship::over(self.origin, self.len())
}
pub fn len(&self) -> usize {
self.store.len()
}
pub fn is_empty(&self) -> bool {
self.store.len() == 0
}
pub fn payloads(&self) -> &[Tensor<E>] {
self.store.payloads()
}
pub fn of(&self, symbol: Symbol) -> &Tensor<E> {
self.kinship().family(symbol);
let Some(slot) = self.store.slot_of(symbol.id) else {
panic!("symbol does not name a parameter these parameters carry");
};
&self.store.payloads()[slot.index()]
}
pub fn map(&self, transform: impl Fn(&Tensor<E>) -> Tensor<E>) -> Self {
Self {
origin: self.origin,
store: self
.store
.with_payloads(self.store.payloads().iter().map(transform).collect()),
}
}
pub fn zip(&self, other: &Self, combine: impl Fn(&Tensor<E>, &Tensor<E>) -> Tensor<E>) -> Self {
self.assert_compatible(other);
Self {
origin: self.origin,
store: self.store.with_payloads(
self.store
.payloads()
.iter()
.zip(other.store.payloads())
.map(|(left, right)| combine(left, right))
.collect(),
),
}
}
pub(super) fn filled_from(&self, field: &Field<E>) -> Self {
assert!(
field.origin() == self.origin,
"field belongs to a different network"
);
if let Some(last) = self.store.last_node() {
assert!(
last.index() < field.len(),
"field is stale: it does not cover every parameter"
);
}
let payloads = self
.store
.iter()
.map(|(node, _)| field.payloads()[node.index()].clone())
.collect();
Self {
origin: self.origin,
store: self.store.with_payloads(payloads),
}
}
fn assert_compatible(&self, other: &Self) {
assert!(
self.origin == other.origin,
"parameter tables belong to different networks"
);
assert_eq!(
self.store.len(),
other.store.len(),
"parameter tables cover different slots"
);
}
pub fn step(
&self,
direction: &Parameters<E>,
mut rule: impl FnMut(&Tensor<E>, &Tensor<E>) -> Tensor<E>,
) -> Self {
self.step_each(direction, move |_, current, direction| {
rule(current, direction)
})
}
pub fn step_each(
&self,
direction: &Parameters<E>,
mut rule: impl FnMut(Symbol, &Tensor<E>, &Tensor<E>) -> Tensor<E>,
) -> Self {
assert!(
direction.origin == self.origin,
"direction belongs to a different network"
);
assert_eq!(
direction.store.len(),
self.store.len(),
"direction covers different parameter slots"
);
let mut payloads = Vec::with_capacity(self.store.len());
for ((node, current), direction) in self.store.iter().zip(direction.store.payloads()) {
let symbol = Symbol {
origin: self.origin,
id: node,
};
let next = rule(symbol, current, direction);
assert_eq!(
next.shape(),
current.shape(),
"step must preserve the parameter's shape"
);
payloads.push(next);
}
Self {
origin: self.origin,
store: self.store.with_payloads(payloads),
}
}
pub fn carried(&self, network: &Network<E>) -> Self {
assert!(
network.origin() == self.origin,
"parameters belong to a different network"
);
let fresh = network.parameters();
assert!(
self.len() <= fresh.len(),
"parameters cover more slots than the network records"
);
let mut payloads: Vec<Tensor<E>> = Vec::with_capacity(fresh.len());
payloads.extend(self.store.payloads().iter().cloned());
payloads.extend(fresh.store.payloads()[self.len()..].iter().cloned());
Self {
origin: self.origin,
store: fresh.store.with_payloads(payloads),
}
}
pub fn with_payloads(
&self,
replacements: impl IntoIterator<Item = (Symbol, Tensor<E>)>,
) -> Self {
let mut payloads = self.store.payloads().to_vec();
for (symbol, payload) in replacements {
self.kinship().family(symbol);
let Some(slot) = self.store.slot_of(symbol.id) else {
panic!("symbol does not name a parameter these parameters carry");
};
assert_eq!(
payload.shape(),
payloads[slot.index()].shape(),
"a replacement must preserve the parameter's shape"
);
payloads[slot.index()] = payload;
}
Self {
origin: self.origin,
store: self.store.with_payloads(payloads),
}
}
}
impl<E: Element> Parameters<E> {
pub fn scale(&self, factor: &Tensor<E>) -> Self {
self.map(|value| value.clone() * factor.broadcast_like(value))
}
}
impl<E: Element> Add for &Parameters<E> {
type Output = Parameters<E>;
fn add(self, rhs: Self) -> Parameters<E> {
self.zip(rhs, |left, right| left.clone() + right.clone())
}
}
impl<E: Element> Add for Parameters<E> {
type Output = Parameters<E>;
fn add(self, rhs: Self) -> Parameters<E> {
&self + &rhs
}
}
#[cfg(test)]
#[path = "tests/parameters_tests.rs"]
mod tests;