pub(super) mod stream_adapter;
pub(super) mod trf_sequencer;
use std::marker::PhantomData;
use furiosa_mapping::*;
use furiosa_opt_macro::primitive;
use crate::backend::Backend;
use crate::cast::{Cast, ContractionCast, ContractionWeight};
use crate::constraints;
use crate::context::*;
use crate::engine::CanApplyContractOuter;
use crate::engine::contraction::LazyContraction;
use crate::runtime::CurrentBackend;
use crate::scalar::*;
use crate::tensor::memory::TrfTensor;
use crate::tensor::tu::TuTensor;
#[derive(Debug)]
pub struct ContractOuterTensor<
'l,
const T: Tu,
D: Scalar,
Storage: ContractionCast<Output = D>,
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) _storage: PhantomData<Storage>,
pub(crate) _axes: PhantomData<(Chip, Cluster, Slice, Lane, Time, Packet)>,
}
impl<
'l,
const T: Tu,
D: Scalar,
Storage: ContractionCast<Output = D>,
Chip: M,
Cluster: M,
Slice: M,
Lane: M,
Time: M,
Packet: M,
B: Backend,
> ContractOuterTensor<'l, T, D, Storage, Chip, Cluster, Slice, Lane, Time, Packet, B>
{
fn check_constraints() {
constraints::assert_cluster_size::<Cluster>();
constraints::assert_slice_size::<Slice>();
constraints::assert_packet_one_or_two_flit::<Storage, Packet>();
}
#[doc(hidden)]
pub(crate) fn new(ctx: &'l mut TuContext<{ T }>, inner: LazyContraction<D, B>) -> Self {
Self::check_constraints();
Self {
ctx,
inner,
_storage: PhantomData,
_axes: PhantomData,
}
}
}
impl<
'l,
const T: Tu,
P: CanApplyContractOuter,
D: Scalar + ContractionCast,
Chip: M,
Cluster: M,
Slice: M,
Time: M,
Packet: M,
B: Backend,
> TuTensor<'l, T, P, D, Chip, Cluster, Slice, Time, Packet, B>
{
#[primitive(TuTensor::contract_outer)]
pub fn contract_outer<OutTime: M, OutPacket: M, Lane: M, TrfElement: M, TrfD>(
self,
trf_tensor: &TrfTensor<TrfD, Chip, Cluster, Slice, Lane, TrfElement, B>,
) -> ContractOuterTensor<'l, T, <D as ContractionCast>::Output, D, Chip, Cluster, Slice, Lane, OutTime, OutPacket, B>
where
D: Cast<<D as ContractionCast>::Output>,
TrfD: Scalar + ContractionWeight<D> + Cast<<D as ContractionCast>::Output>,
{
type Out<D> = <D as ContractionCast>::Output;
stream_adapter::verify_stream_adapter::<D, Lane, Time, Packet, OutTime, OutPacket>();
trf_sequencer::verify_trf_sequencer::<TrfD, Lane, TrfElement, OutTime, OutPacket>();
let lhs_map = <m![{ Chip }, { Cluster }, { Slice }, { Time }, { Packet }]>::to_value();
let rhs_map = <m![{ Chip }, { Cluster }, { Slice }, { Lane }, { TrfElement }]>::to_value();
let pre_reduce = <m![{ Chip }, { Cluster }, { Slice }, { Lane }, { OutTime }, { OutPacket }]>::to_value();
let contraction = LazyContraction {
lhs: self.inner.map_bounded(|v| -> Out<D> { v.cast() }).inner,
rhs: trf_tensor.inner.map_bounded(|v| -> Out<D> { v.cast() }).inner,
lhs_map,
rhs_map,
pre_reduce,
};
ContractOuterTensor::new(self.ctx, contraction)
}
}