use furiosa_mapping::*;
use furiosa_opt_macro::primitive;
use crate::backend::Backend;
use crate::cast::ContractionCast;
use crate::context::*;
use crate::engine::contraction::ContractPacketTensor;
use crate::engine::contraction::outer::ContractOuterTensor;
use crate::scalar::*;
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>
{
#[primitive(ContractOuterTensor::contract_packet)]
pub fn contract_packet<OutPacket: M>(
self,
) -> ContractPacketTensor<'l, T, D, Chip, Cluster, Slice, Lane, Time, OutPacket, B> {
verify_contract_packet::<Storage, Packet, OutPacket>();
ContractPacketTensor::new(self.ctx, self.inner)
}
}
pub(crate) fn verify_contract_packet<Storage: Scalar, Packet: M, OutPacket: M>() {
furiosa_opt_lower::config_contract_packet(&Packet::to_value(), &OutPacket::to_value(), Storage::BITS)
.unwrap_or_else(|message| panic!("{message}"));
}
#[cfg(test)]
mod tests {
use super::*;
use crate::scalar::bf16;
axes![A = 4, B = 2, C = 4, D = 32, K = 64, M = 4, N = 8, O = 2, P = 8];
#[test]
fn valid_full_reduction() {
verify_contract_packet::<i8, m![K], m![1]>();
}
#[test]
fn valid_partial_reduction() {
verify_contract_packet::<i8, m![K], m![K / 4]>();
}
#[test]
fn valid_partial_reduction_multi_axis() {
verify_contract_packet::<i8, m![A, D / 2], m![A, D / 8]>();
}
#[test]
fn valid_padded_packet_inner_reduction() {
verify_contract_packet::<i8, m![A # 16, C], m![A]>();
}
#[test]
fn valid_padded_packet_inner_reduction_with_padding() {
verify_contract_packet::<i8, m![A # 16, C], m![A # 16]>();
}
#[test]
fn valid_padded_packet_split() {
verify_contract_packet::<i8, m![B # 8, N], m![B]>();
}
#[test]
fn valid_no_spatial_reduction_bf16() {
verify_contract_packet::<bf16, m![D], m![D]>();
}
}