pub mod lane;
pub mod outer;
pub mod packet;
pub mod time;
use std::marker::PhantomData;
pub use lane::LaneMode;
pub use outer::ContractOuterTensor;
use furiosa_mapping::Mapping as MappingValue;
use furiosa_mapping::*;
use crate::backend::Backend;
use crate::constraints;
use crate::context::*;
use crate::runtime::CurrentBackend;
use crate::scalar::*;
use crate::tensor::Tensor;
use crate::tensor::tu::{Position, TuTensor};
#[derive(Debug, Clone)]
pub(crate) struct LazyContraction<D: Scalar, B: Backend> {
pub(crate) lhs: B::Storage<D>,
pub(crate) rhs: B::Storage<D>,
pub(crate) lhs_map: MappingValue,
pub(crate) rhs_map: MappingValue,
pub(crate) pre_reduce: MappingValue,
}
pub(crate) use furiosa_opt_lower::TEMPORAL_ACCUMULATOR_COLS;
fn assert_packet_pow2_within_accumulator_cols<Packet: M>() {
const {
let size = Packet::SIZE;
assert!(
size != 0 && size & (size - 1) == 0 && size <= TEMPORAL_ACCUMULATOR_COLS,
"Packet element count must be a power of two and at most the accumulator column count (32)"
);
};
}
#[derive(Debug)]
pub struct PositionContraction;
impl Position for PositionContraction {}
#[derive(Debug)]
pub struct ContractPacketTensor<
'l,
const T: Tu,
D: Scalar,
Chip: M,
Cluster: M,
Slice: M,
Lane: M,
Time: M,
Packet: M,
B: Backend = CurrentBackend,
> {
pub(crate) ctx: &'l mut TuContext<{ T }>,
pub(crate) inner: LazyContraction<D, B>,
pub(crate) _axes: PhantomData<(Chip, Cluster, Slice, Lane, Time, Packet)>,
}
#[derive(Debug)]
pub struct ContractTimeTensor<
'l,
const T: Tu,
D: Scalar,
Chip: M,
Cluster: M,
Slice: M,
Lane: M,
Time: M,
Packet: M,
B: Backend = CurrentBackend,
> {
pub(crate) ctx: &'l mut TuContext<{ T }>,
pub(crate) inner: LazyContraction<D, B>,
pub(crate) pre_reduce_time: Mapping,
pub(crate) _axes: PhantomData<(Chip, Cluster, Slice, Lane, Time, Packet)>,
}
impl<'l, const T: Tu, D: Scalar, Chip: M, Cluster: M, Slice: M, Lane: M, Time: M, Packet: M, B: Backend>
ContractPacketTensor<'l, T, D, Chip, Cluster, Slice, Lane, Time, Packet, B>
{
fn check_constraints() {
constraints::assert_cluster_size::<Cluster>();
constraints::assert_slice_size::<Slice>();
assert_packet_pow2_within_accumulator_cols::<Packet>();
}
#[doc(hidden)]
pub(crate) fn new(ctx: &'l mut TuContext<{ T }>, inner: LazyContraction<D, B>) -> Self {
Self::check_constraints();
Self {
ctx,
inner,
_axes: PhantomData,
}
}
}
impl<'l, const T: Tu, D: Scalar, Chip: M, Cluster: M, Slice: M, Lane: M, Time: M, Packet: M, B: Backend>
ContractTimeTensor<'l, T, D, Chip, Cluster, Slice, Lane, Time, Packet, B>
{
fn check_constraints() {
constraints::assert_cluster_size::<Cluster>();
constraints::assert_slice_size::<Slice>();
assert_packet_pow2_within_accumulator_cols::<Packet>();
}
#[doc(hidden)]
pub(crate) fn new(ctx: &'l mut TuContext<{ T }>, inner: LazyContraction<D, B>, pre_reduce_time: Mapping) -> Self {
Self::check_constraints();
Self {
ctx,
inner,
pre_reduce_time,
_axes: PhantomData,
}
}
}
pub type ContractTensor<'l, const T: Tu, D, Chip, Cluster, Slice, Time, Packet, B = CurrentBackend> =
TuTensor<'l, { T }, PositionContraction, D, Chip, Cluster, Slice, Time, Packet, B>;
impl<'l, const T: Tu, D: Scalar, Chip: M, Cluster: M, Slice: M, Time: M, Packet: M, B: Backend>
ContractTensor<'l, T, D, Chip, Cluster, Slice, Time, Packet, B>
{
fn check_constraints() {
constraints::assert_cluster_size::<Cluster>();
constraints::assert_slice_size::<Slice>();
constraints::assert_packet_one_flit::<D, Packet>();
}
#[doc(hidden)]
pub fn new(ctx: &'l mut TuContext<{ T }>, inner: Tensor<D, Self::Mapping, B>) -> Self {
Self::check_constraints();
Self {
ctx,
inner,
_position: PhantomData,
}
}
}