// Copyright 2026 The IREE Authors
//
// Licensed under the Apache License v2.0 WITH LLVM Exceptions.
// See https://llvm.org/LICENSE.txt for license information.
// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception

// Two waves repeatedly exchange records through one LDS tile. GFX12 and
// GFX12.5 providers release the tile before private work and wait before the
// next overwrite; other targets use the complete-barrier fallback behind the
// same template contract.
amdgpu.target<gfx12-generic> @gfx12

amdgpu.target<gfx12-5-generic> @gfx12_5

func.def pure @long_private_work(%value: i32, %sum: i32) -> (i32) {
  %c5 = scalar.constant 5 : i32
  %c7 = scalar.constant 7 : i32
  %c11 = scalar.constant 11 : i32
  %multiplier0 = scalar.constant 1664525 : i32
  %multiplier1 = scalar.constant 1103515245 : i32
  %sum0 = scalar.addi %sum, %value : i32
  %shift0 = scalar.shrui %sum0, %c11 : i32
  %mix0 = scalar.xori %sum0, %shift0 : i32
  %product0 = scalar.muli %mix0, %multiplier0 : i32
  %shift1 = scalar.shli %product0, %c5 : i32
  %mix1 = scalar.xori %product0, %shift1 : i32
  %product1 = scalar.muli %mix1, %multiplier1 : i32
  %shift2 = scalar.shrui %product1, %c7 : i32
  %mix2 = scalar.xori %product1, %shift2 : i32
  func.return %mix2 : i32
}

func.def pure @short_private_work(%value: i32, %sum: i32) -> (i32) {
  %sum0 = scalar.addi %sum, %value : i32
  %mix = scalar.xori %sum0, %value : i32
  func.return %mix : i32
}

func.def pure @private_work(%takes_long_path: i1, %value: i32, %sum: i32) -> (i32) {
  %updated = scf.if %takes_long_path -> (i32) {
    %long = func.call pure @long_private_work(%value, %sum) : (i32, i32) -> (i32)
    scf.yield %long : i32
  } else {
    %short = func.call pure @short_private_work(%value, %sum) : (i32, i32) -> (i32)
    scf.yield %short : i32
  }
  func.return %updated : i32
}

template.decl @split_barrier.finish_shared_read(%takes_long_path: i1, %value: i32, %sum: i32) -> (i32)

template.def<@split_barrier.finish_shared_read> target(@gfx12) priority(20) @finish_shared_read_gfx12(%takes_long_path: i1, %value: i32, %sum: i32) -> (i32) {
  %phase = kernel.barrier.arrive<workgroup> scope(workgroup) ordering(acq_rel) -> kernel.barrier.phase
  %updated = func.call pure @private_work(%takes_long_path, %value, %sum) : (i1, i32, i32) -> (i32)
  kernel.barrier.wait %phase : kernel.barrier.phase
  template.return %updated : i32
}

template.def<@split_barrier.finish_shared_read> target(@gfx12_5) priority(20) @finish_shared_read_gfx12_5(%takes_long_path: i1, %value: i32, %sum: i32) -> (i32) {
  %phase = kernel.barrier.arrive<workgroup> scope(workgroup) ordering(acq_rel) -> kernel.barrier.phase
  %updated = func.call pure @private_work(%takes_long_path, %value, %sum) : (i1, i32, i32) -> (i32)
  kernel.barrier.wait %phase : kernel.barrier.phase
  template.return %updated : i32
}

template.def<@split_barrier.finish_shared_read> priority(1) @finish_shared_read_fallback(%takes_long_path: i1, %value: i32, %sum: i32) -> (i32) {
  %updated = func.call pure @private_work(%takes_long_path, %value, %sum) : (i1, i32, i32) -> (i32)
  kernel.barrier<workgroup> scope(workgroup) ordering(acq_rel)
  template.return %updated : i32
}

kernel.def @selected_barrier_reuse() {
  %unit = index.constant 1 : index
  %width = index.constant 64 : index
  kernel.launch.config workgroups(%unit, %unit, %unit) workgroup_size(%width, %unit, %unit) : index
} launch(%input: buffer, %output: buffer) {
  %c0 = index.constant 0 : index
  %c1 = index.constant 1 : index
  %c4 = index.constant 4 : index
  %c32 = index.constant 32 : index
  %last = index.constant 63 : index
  %lane = kernel.workitem.id<x> : index
  // GFX12 and GFX12.5 use wave32, making this branch subgroup-uniform.
  %takes_long_path = index.cmp ult, %lane, %c32 : index
  %peer = index.sub %last, %lane : index
  %base = index.constant 0 : offset
  %scratch_bytes = index.constant 256 : offset
  %source = buffer.view %input[%base] : buffer -> view<4x64xi32>
  %destination = buffer.view %output[%base] : buffer -> view<64xi32>
  %scratch = buffer.alloca<workgroup> align(16) %scratch_bytes : buffer
  %tile = buffer.view %scratch[%base] : buffer -> view<64xi32>
  %initial = scalar.constant 0 : i32
  %total = scf.for %record = [%c0 to %c4 step %c1](%sum = %initial : i32) -> (i32) {
    %current = view.load %source[%record, %lane] : view<4x64xi32> -> i32
    view.store %current, %tile[%lane] : i32, view<64xi32>
    kernel.barrier<workgroup> scope(workgroup) ordering(acq_rel)
    %observed = view.load %tile[%peer] : view<64xi32> -> i32
    %updated = template.apply<@split_barrier.finish_shared_read>(%takes_long_path, %observed, %sum) : (i1, i32, i32) -> (i32)
    scf.yield %updated : i32
  }
  view.store %total, %destination[%lane] : i32, view<64xi32>
  kernel.return
}

kernel.def @full_barrier_reuse() {
  %unit = index.constant 1 : index
  %width = index.constant 64 : index
  kernel.launch.config workgroups(%unit, %unit, %unit) workgroup_size(%width, %unit, %unit) : index
} launch(%input: buffer, %output: buffer) {
  %c0 = index.constant 0 : index
  %c1 = index.constant 1 : index
  %c4 = index.constant 4 : index
  %c32 = index.constant 32 : index
  %last = index.constant 63 : index
  %lane = kernel.workitem.id<x> : index
  // Match the provider kernel's private work in the full-barrier reference.
  %takes_long_path = index.cmp ult, %lane, %c32 : index
  %peer = index.sub %last, %lane : index
  %base = index.constant 0 : offset
  %scratch_bytes = index.constant 256 : offset
  %source = buffer.view %input[%base] : buffer -> view<4x64xi32>
  %destination = buffer.view %output[%base] : buffer -> view<64xi32>
  %scratch = buffer.alloca<workgroup> align(16) %scratch_bytes : buffer
  %tile = buffer.view %scratch[%base] : buffer -> view<64xi32>
  %initial = scalar.constant 0 : i32
  %total = scf.for %record = [%c0 to %c4 step %c1](%sum = %initial : i32) -> (i32) {
    %current = view.load %source[%record, %lane] : view<4x64xi32> -> i32
    view.store %current, %tile[%lane] : i32, view<64xi32>
    kernel.barrier<workgroup> scope(workgroup) ordering(acq_rel)
    %observed = view.load %tile[%peer] : view<64xi32> -> i32
    %updated = func.call pure @private_work(%takes_long_path, %observed, %sum) : (i1, i32, i32) -> (i32)
    kernel.barrier<workgroup> scope(workgroup) ordering(acq_rel)
    scf.yield %updated : i32
  }
  view.store %total, %destination[%lane] : i32, view<64xi32>
  kernel.return
}

check.case public @split_barrier_reuse_case {
  %input = check.generate.iota offset(1) step(1) : tensor<4x64xi32>
  %split_output = check.generate.fill value(-1) : tensor<64xi32>
  %full_output = check.generate.fill value(-2) : tensor<64xi32>
  kernel.launch @selected_barrier_reuse(%input, %split_output) : (tensor<4x64xi32>, tensor<64xi32>)
  kernel.launch @full_barrier_reuse(%input, %full_output) : (tensor<4x64xi32>, tensor<64xi32>)
  check.expect.bitwise actual(%split_output) expected(%full_output) : tensor<64xi32>
  check.return
}
