acyclib 0.3.0

ML library for directed acyclic tensor graphs.
Documentation
use crate::{
    device::operation::DiffableFromOutput,
    graph::ir::{
        BackendMarker,
        operation::{GraphIROperationCompilable, sparse::SparseAffineActivate},
    },
};

use super::{GraphBuilderNode, InitSettings};

#[derive(Clone, Copy)]
pub struct Affine<'a, B: BackendMarker> {
    pub weights: GraphBuilderNode<'a, B>,
    pub bias: GraphBuilderNode<'a, B>,
}

impl<'a, B: BackendMarker> Affine<'a, B>
where
    SparseAffineActivate: GraphIROperationCompilable<B>,
{
    pub fn forward(self, input: GraphBuilderNode<'a, B>) -> GraphBuilderNode<'a, B> {
        self.weights.matmul(input) + self.bias
    }

    pub fn init_with_effective_input_size(&self, size: usize) {
        let builder = self.weights.builder.ir();
        let id = builder.get_id(self.weights.node.idx).unwrap();
        *self.weights.builder.init().get_mut(&id).unwrap() =
            InitSettings::Normal { mean: 0.0, stdev: (2.0 / size as f32).sqrt() };
    }

    pub fn forward_sparse_with_values(
        self,
        stm: GraphBuilderNode<'a, B>,
        vals: GraphBuilderNode<'a, B>,
    ) -> GraphBuilderNode<'a, B> {
        stm.builder.apply(SparseAffineActivate {
            weights: self.weights.node,
            indices: stm.node,
            values: Some(vals.node),
            biases: Some(self.bias.node),
            activation: DiffableFromOutput::Identity,
        })
    }
}