use crate::planning::plan::{
DiagPlan as InnerDiagPlan, DiagStage as InnerDiagStage, GemmPlan as InnerGemmPlan,
ReducePlan as InnerReducePlan, StepPlan as InnerStepPlan,
};
#[derive(Clone, Copy, Debug)]
pub struct PairwiseStepPlan<'a> {
inner: &'a InnerStepPlan,
}
impl<'a> PairwiseStepPlan<'a> {
pub(crate) fn new(inner: &'a InnerStepPlan) -> Self {
Self { inner }
}
#[must_use]
pub fn lhs_diag(&self) -> Option<DiagPlan<'a>> {
let inner: &'a InnerStepPlan = self.inner;
inner.diag_a.as_ref().map(DiagPlan::new)
}
#[must_use]
pub fn rhs_diag(&self) -> Option<DiagPlan<'a>> {
let inner: &'a InnerStepPlan = self.inner;
inner.diag_b.as_ref().map(DiagPlan::new)
}
#[must_use]
pub fn lhs_reduce(&self) -> Option<ReducePlan<'a>> {
let inner: &'a InnerStepPlan = self.inner;
inner.gemm.reduce_a.as_ref().map(ReducePlan::new)
}
#[must_use]
pub fn rhs_reduce(&self) -> Option<ReducePlan<'a>> {
let inner: &'a InnerStepPlan = self.inner;
inner.gemm.reduce_b.as_ref().map(ReducePlan::new)
}
#[must_use]
pub fn gemm(&self) -> GemmPlan<'a> {
let inner: &'a InnerStepPlan = self.inner;
GemmPlan::new(&inner.gemm)
}
}
#[derive(Clone, Copy, Debug)]
pub struct DiagPlan<'a> {
inner: &'a InnerDiagPlan,
}
impl<'a> DiagPlan<'a> {
fn new(inner: &'a InnerDiagPlan) -> Self {
Self { inner }
}
pub fn stages(self) -> impl ExactSizeIterator<Item = DiagStage<'a>> + 'a {
let inner: &'a InnerDiagPlan = self.inner;
inner.stages.iter().map(DiagStage::new)
}
#[must_use]
pub fn result_subs(&self) -> &'a [u32] {
let inner: &'a InnerDiagPlan = self.inner;
inner.result_subs.as_slice()
}
}
#[derive(Clone, Copy, Debug)]
pub struct DiagStage<'a> {
inner: &'a InnerDiagStage,
}
impl<'a> DiagStage<'a> {
fn new(inner: &'a InnerDiagStage) -> Self {
Self { inner }
}
#[must_use]
pub fn axis_pairs(&self) -> &'a [(usize, usize)] {
let inner: &'a InnerDiagStage = self.inner;
inner.axis_pairs.as_slice()
}
#[must_use]
pub fn result_subs(&self) -> &'a [u32] {
let inner: &'a InnerDiagStage = self.inner;
inner.result_subs.as_slice()
}
}
#[derive(Clone, Copy, Debug)]
pub struct ReducePlan<'a> {
inner: &'a InnerReducePlan,
}
impl<'a> ReducePlan<'a> {
fn new(inner: &'a InnerReducePlan) -> Self {
Self { inner }
}
#[must_use]
pub fn original_subs(&self) -> &'a [u32] {
let inner: &'a InnerReducePlan = self.inner;
inner.original_subs.as_slice()
}
#[must_use]
pub fn kept_subs(&self) -> &'a [u32] {
let inner: &'a InnerReducePlan = self.inner;
inner.kept_subs.as_slice()
}
#[must_use]
pub fn out_shape(&self) -> &'a [usize] {
let inner: &'a InnerReducePlan = self.inner;
inner.out_shape.as_slice()
}
}
#[derive(Clone, Copy, Debug)]
pub struct GemmPlan<'a> {
inner: &'a InnerGemmPlan,
}
impl<'a> GemmPlan<'a> {
fn new(inner: &'a InnerGemmPlan) -> Self {
Self { inner }
}
#[must_use]
pub fn left_only_modes(&self) -> &'a [u32] {
let inner: &'a InnerGemmPlan = self.inner;
inner.lo_modes.as_slice()
}
#[must_use]
pub fn left_only_shape(&self) -> &'a [usize] {
let inner: &'a InnerGemmPlan = self.inner;
inner.lo_sizes.as_slice()
}
#[must_use]
pub fn right_only_modes(&self) -> &'a [u32] {
let inner: &'a InnerGemmPlan = self.inner;
inner.ro_modes.as_slice()
}
#[must_use]
pub fn right_only_shape(&self) -> &'a [usize] {
let inner: &'a InnerGemmPlan = self.inner;
inner.ro_sizes.as_slice()
}
#[must_use]
pub fn contracted_modes(&self) -> &'a [u32] {
let inner: &'a InnerGemmPlan = self.inner;
inner.sum_modes.as_slice()
}
#[must_use]
pub fn contracted_shape(&self) -> &'a [usize] {
let inner: &'a InnerGemmPlan = self.inner;
inner.sum_sizes.as_slice()
}
#[must_use]
pub fn batch_modes(&self) -> &'a [u32] {
let inner: &'a InnerGemmPlan = self.inner;
let batch_start = inner.lo_modes.len() + inner.ro_modes.len();
&inner.canonical_modes[batch_start..]
}
#[must_use]
pub fn batch_shape(&self) -> &'a [usize] {
let inner: &'a InnerGemmPlan = self.inner;
inner.batch_sizes.as_slice()
}
#[must_use]
pub fn lhs_target_modes(&self) -> &'a [u32] {
let inner: &'a InnerGemmPlan = self.inner;
inner.target_a.as_slice()
}
#[must_use]
pub fn rhs_target_modes(&self) -> &'a [u32] {
let inner: &'a InnerGemmPlan = self.inner;
inner.target_b.as_slice()
}
#[must_use]
pub fn canonical_output_modes(&self) -> &'a [u32] {
let inner: &'a InnerGemmPlan = self.inner;
inner.canonical_modes.as_slice()
}
#[must_use]
pub fn m(&self) -> usize {
self.inner.m
}
#[must_use]
pub fn n(&self) -> usize {
self.inner.n
}
#[must_use]
pub fn k(&self) -> usize {
self.inner.k
}
#[must_use]
pub fn lhs_gemm_shape(&self) -> &'a [usize] {
let inner: &'a InnerGemmPlan = self.inner;
inner.a_gemm_shape.as_slice()
}
#[must_use]
pub fn rhs_gemm_shape(&self) -> &'a [usize] {
let inner: &'a InnerGemmPlan = self.inner;
inner.b_gemm_shape.as_slice()
}
#[must_use]
pub fn output_gemm_shape(&self) -> &'a [usize] {
let inner: &'a InnerGemmPlan = self.inner;
inner.c_gemm_shape.as_slice()
}
#[must_use]
pub fn expanded_output_shape(&self) -> &'a [usize] {
let inner: &'a InnerGemmPlan = self.inner;
inner.expanded_shape.as_slice()
}
#[must_use]
pub fn needs_final_permute(&self) -> bool {
self.inner.needs_final_permute
}
}