diff --git a/crates/ipu-codegen/src/planner/gemm_tests.rs b/crates/ipu-codegen/src/planner/gemm_tests.rs index b419bfc8..87c30cf0 100644 --- a/crates/ipu-codegen/src/planner/gemm_tests.rs +++ b/crates/ipu-codegen/src/planner/gemm_tests.rs @@ -4,6 +4,64 @@ use super::*; use crate::planner::tests::gelu; use ipu_target::Target; +#[test] +#[ignore] +fn inspect_sdk_fp16() { + let mut high = HighGraph::new(); + let x = high.host_input("x", [1, 729, 1152]).unwrap(); + let w1 = high.parameter("up", [1, 1152, 4304]).unwrap(); + let w2 = high.parameter("down", [1, 4304, 1152]).unwrap(); + let h = high.gemm(x, w1).unwrap(); + let a = high.gelu(h).unwrap(); + let y = high.gemm(a, w2).unwrap(); + high.set_outputs([y]).unwrap(); + let config = PipelineConfig::new(Target::Ipu21, 1472).with_input(x, TensorFormat { + precision: Precision::F16, layout: Layout::logical_linear(1472, 4), + }); + for (position, k, n, sdk, selected) in [(0,1152,4304,[15,12,8],[12,15,8]),(2,4304,1152,[6,16,15],[9,6,27])] { + let inputs = &high.operations()[position].inputs; + let origin = high.operations()[position].results[0]; + let shape = high.value_shape(origin).unwrap(); + let source = TensorType::new(vec![1,729,k], Precision::F16, Layout::logical_linear(1472,4)); + let weight = TensorType {shape: TensorShape(vec![1,k,n]),format: parameter_format(&TensorShape(vec![1,k,n]),false,false,Precision::F16,&config)}; + let assignments = choices([&source,&weight],shape,GemmOptions::default(),1472).unwrap(); + for (name, swapped, grid) in [("sdk",true,sdk),("selected",false,selected)] { + let mut reports = Vec::new(); + for c in assignments.iter().filter(|c| c.swapped==swapped && [c.operands[0].format.layout.tiling.axes[0].partitions,c.operands[1].format.layout.tiling.axes[0].partitions,c.operands[0].format.layout.tiling.axes[1].partitions]==grid) { + let live = LiveValues::from([(inputs[0],BoundaryValue {tensor:source.clone(),owners:OwnerMap::default()}),(inputs[1],BoundaryValue {tensor:weight.clone(),owners:OwnerMap::default()})]); + let mut candidate = Candidate::inputs(&high,&live,1472,position+1); + let output = append(&mut candidate.graph,[candidate.bindings[&inputs[0]],candidate.bindings[&inputs[1]]],c,GemmOptions::default(),shape,high.operations()[position].id,origin,Some(&Layout::logical_linear(1472,4))); + candidate.graph.outputs=vec![output]; + candidate.graph.compose_copies(); + let valid_tail = word_aligned_shards(&c.result) && candidate.graph.operations.last().is_none_or(|last| !matches!(last.kind, MidOperationKind::Copy {..}) || word_aligned_shards(&candidate.graph.values[last.inputs[0].index() as usize].tensor_type)); + let (cost,peak)=crate::estimate::coarse::analyze(Target::Ipu21,&candidate.graph).unwrap(); + reports.push((cost.total,valid_tail,peak.total,candidate)); + } + reports.sort_by_key(|r|r.0); + eprintln!("GRID {position} {name} count={} accepted={} costs={:?}",reports.len(),reports.iter().filter(|r|r.1).count(),reports.iter().map(|r|(r.0,r.1,r.2)).take(12).collect::>()); + if let Some((_,_,_,best))=reports.first() { + eprintln!("DETAIL {position} {name} {:?}",crate::estimate::analyze_mid(Target::Ipu21,&best.graph,&Default::default())); + for op in &best.graph.operations { + let cost=crate::estimate::operation_cost(Target::Ipu21,op,&best.graph.values,1472); + eprintln!("OP {:?} {:?}",op.kind,cost); + if cost.is_some_and(|v|v.0.total==u64::MAX) && matches!(op.kind,MidOperationKind::Copy {..}) { + let input=&best.graph.values[op.inputs[0].index() as usize].tensor_type; + let output=&best.graph.values[op.results[0].index() as usize].tensor_type; + let extents=output.format.layout.shard_extents(&output.shape).unwrap(); + let mut from=output.format.layout.clone(); from.order=input.format.layout.order; + eprintln!("SATURATED {:?}",crate::kernel::KernelCall::select(Target::Ipu21,&MidOperationKind::Rearrange{from,to:output.format.layout.clone()},&[crate::storage::TensorStorage{format:&output.format,extents:&extents[0].1}],&[crate::storage::TensorStorage{format:&output.format,extents:&extents[0].1}],None)); + } + } + } + } + } + let catalogue=super::super::candidates::catalogue(&high,&boundary_layouts(&high,&config),&config).unwrap(); + for (position,candidates) in catalogue.iter().enumerate() { + let swapped=candidates.iter().filter(|c|c.graph.operations.iter().any(|op|matches!(&op.kind,MidOperationKind::Gemm{axes,..} if axes.output_column==TensorAxis::FromEnd(2)))).count(); + eprintln!("CATALOGUE {position} total={} swapped={swapped}",candidates.len()); + } +} + #[test] #[ignore = "large catalogue measurement; run explicitly in release mode"] fn report_mlp_boundary_diversity() {