// The outer pipeline queues complete guarded row reductions. Each caller
// chooses its outer and inner schedule independently.
template.decl @guide.sum_guarded_rows(%values: view<63x32x4xf32>, %count: index, %columns: index, %lane: index, %active_lanes: index, %depth: index, %factor: index, %inner_depth: index) -> (f32)

template.def<@guide.sum_guarded_rows> @sum_guarded_rows_impl(%values: view<63x32x4xf32>, %count: index, %columns: index, %lane: index, %active_lanes: index, %depth: index, %factor: index, %inner_depth: index) -> (f32) {
  %begin = index.constant 0 : index
  %step = index.constant 1 : index
  %identity = scalar.constant 0.0 : f32
  %tile_width = index.constant 4 : index
  %tile_bias = index.sub %tile_width, %step : index
  %padded_count = index.add %count, %tile_bias : index
  %tile_count = index.div %padded_count, %tile_width : index
  %end = index.mul %tile_count, %tile_width : index
  %active = index.cmp slt, %lane, %active_lanes : index
  %total = scf.for %row = [%begin to %end step %step](%sum = %identity : f32) -> (f32) pipeline(%depth) unroll(%factor) schedule(recurrence) {
    %valid = index.cmp slt, %row, %count : index
    %partial = scf.if %valid -> (f32) {
      %row_sum = scf.for %column = [%begin to %columns step %step](%component_sum = %identity : f32) -> (f32) pipeline(%inner_depth) {
        %value = scf.if %active -> (f32) {
          %loaded = view.load %values[%row, %lane, %column] : view<63x32x4xf32> -> f32
          scf.yield %loaded : f32
        } else {
          scf.yield %identity : f32
        }
        %next_component = scalar.addf %component_sum, %value : f32
        scf.yield %next_component : f32
      }
      scf.yield %row_sum : f32
    } else {
      scf.yield %identity : f32
    }
    %next = scalar.addf %sum, %partial : f32
    scf.yield %next : f32
  }
  template.return %total : f32
}

// This caller pipelines rows at depth three and unrolls pairs of iterations.
kernel.def export("sum_guarded_rows") @sum_guarded_rows() {
  %unit = index.constant 1 : index
  %width = index.constant 32 : index
  kernel.launch.config workgroups(%unit, %unit, %unit) workgroup_size(%width, %unit, %unit) : index
} launch(%row_count: index, %column_count: index, %lane_count: index, %input: buffer, %output: buffer) {
  %count = index.assume %row_count [range(%row_count, 0, 63)] : index
  %columns = index.assume %column_count [range(%column_count, 0, 4)] : index
  %active_lanes = index.assume %lane_count [range(%lane_count, 0, 32)] : index
  %lane = kernel.workitem.id<x> : index
  %base = index.constant 0 : offset
  %aligned_input = buffer.assume.alignment %input {minimum_alignment = 4} : buffer
  %aligned_output = buffer.assume.alignment %output {minimum_alignment = 4} : buffer
  %values = buffer.view %aligned_input[%base] : buffer -> view<63x32x4xf32>
  %sums = buffer.view %aligned_output[%base] : buffer -> view<32xf32>
  %serial_policy = index.constant 1 : index
  %pair_policy = index.constant 2 : index
  %triple_policy = index.constant 3 : index
  %result = template.apply<@guide.sum_guarded_rows>(%values, %count, %columns, %lane, %active_lanes, %triple_policy, %pair_policy, %serial_policy) : (view<63x32x4xf32>, index, index, index, index, index, index, index) -> (f32)
  view.store %result, %sums[%lane] : f32, view<32xf32>
  kernel.return
}
