use std::marker::PhantomData;
use crate::components::batch::partitioned_matmul::partition::{
GlobalPartitionMatmul, PartitionRangeDim, PartitionRanges,
};
use crate::components::batch::{BatchConfig as _, BatchMatmul, CubeCountInput};
use crate::components::global::{self, GlobalMatmul};
use crate::components::{AccG, batch::partitioned_matmul::config::PartitionedBatchConfig};
use crate::components::{LhsG, MatmulPrecision, RhsG};
use cubecl_core as cubecl;
use cubecl_core::prelude::*;
use cubecl_std::{
CubeOption,
tensor::{View, layout::Coords3d},
};
pub struct PartitionedBatchMatmul<
MP: MatmulPrecision,
GMM: global::GlobalMatmul<MP>,
S: GlobalPartitionMatmul,
> {
_mp: PhantomData<MP>,
_gmm: PhantomData<GMM>,
_s: PhantomData<S>,
}
#[cube]
impl<MP: MatmulPrecision, GMM: GlobalMatmul<MP>, GPMM: GlobalPartitionMatmul> BatchMatmul<MP>
for PartitionedBatchMatmul<MP, GMM, GPMM>
{
type Config = PartitionedBatchConfig<GMM::Config>;
fn execute(
a: View<Line<LhsG<MP>>, Coords3d>,
b: View<Line<RhsG<MP>>, Coords3d>,
c: CubeOption<View<Line<AccG<MP>>, Coords3d>>,
out: View<Line<AccG<MP>>, Coords3d, ReadWrite>,
cube_count_args: CubeCountInput,
#[comptime] config: Self::Config,
) {
let (_, _, problem_k) = a.shape();
let k_range = (0, problem_k);
let tiling_scheme = config.tiling_scheme();
let (m_index, n_index, batch_index) =
cube_count_args.cube_pos_to_tensor_pos(config.hypercube_config().global_order);
let ranges = PartitionRanges::new(
PartitionRangeDim::new(
m_index,
tiling_scheme.elements_in_stage_m(),
tiling_scheme.elements_in_global_partition_m(),
),
PartitionRangeDim::new(
n_index,
tiling_scheme.elements_in_stage_n(),
tiling_scheme.elements_in_global_partition_n(),
),
PartitionRangeDim::new(
batch_index,
1u32,
tiling_scheme.global_partition_size.batches,
),
);
let global_config = config.global_config();
let acc = GMM::init_accumulators(global_config);
GPMM::execute::<MP, GMM>(a, b, c, out, ranges, acc, k_range, global_config);
}
}