use furiosa_mapping::*;
use furiosa_opt_macro::primitive;
use crate::backend::Backend;
use crate::context::*;
use crate::engine::contraction::{ContractPacketTensor, ContractTimeTensor};
use crate::scalar::*;
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>
{
#[primitive(ContractPacketTensor::contract_time)]
pub fn contract_time<OutTime: M>(
self,
) -> ContractTimeTensor<'l, T, D, Chip, Cluster, Slice, Lane, OutTime, Packet, B> {
verify_contract_time(Time::to_value(), OutTime::to_value());
ContractTimeTensor::new(self.ctx, self.inner, Time::to_value())
}
}
pub(crate) fn verify_contract_time(time: Mapping, out_time: Mapping) {
furiosa_opt_lower::config_contract_time(&time, &out_time).unwrap_or_else(|message| panic!("{message}"));
}
#[cfg(test)]
mod tests {
use super::*;
axes![A = 4, B = 2, C = 4, D = 32, K = 64, M = 4, N = 8, O = 2, P = 8];
mod contract_time_subset {
use super::*;
use furiosa_mapping::M as _;
#[test]
fn valid_identity() {
verify_contract_time(<m![A, B]>::to_value(), <m![A, B]>::to_value());
}
#[test]
fn valid_reduce_inner() {
verify_contract_time(<m![A, B]>::to_value(), <m![A]>::to_value());
}
#[test]
fn valid_reduce_outer() {
verify_contract_time(<m![A, B]>::to_value(), <m![B]>::to_value());
}
}
}