已开启
[bug][auto-vectorize-v2][multi-consumer] 开启后性能劣化,待分析原因 #387
hujiajun创建于  28 天前
hujiajun成员
28 天前 创建

主干 dd3198ac66f277b9165f2e817d3c708b8615b42c 手动打开多 consumer 复现

bishengir-compile  --target=Ascend910_9589  --enable-auto-multi-buffer=True --enable-auto-bind-sub-block=True --disable-ffts --enable-hivm-graph-sync-solver=True --limit-auto-multi-buffer-of-local-buffer=no-limit --enable-mixed-cv=True --enable-flatten=False --enable-hfusion-compile=true --enable-triton-kernel-compile=true --hfusion-enable-multiple-consumer-fusion=true xx.mlir
#loc = loc("flash_attention_npu_v8_copy.py":821:0)
#loc71 = loc("flash_attention_npu_v8_copy.py":997:44)
#map = affine_map<(d0) -> (d0)>
module attributes {hacc.target = #hacc.target<"Ascend950PR_9589">} {
  func.func private @triton_indirect_load(memref<?xf32>, tensor<128xi64>, tensor<128xi1>, tensor<128xf32>) -> tensor<128xf32> loc(#loc100)
  func.func private @triton_indirect_load_0(memref<?xf32>, tensor<128xi64>, tensor<128xi1>, tensor<128xf32>) -> tensor<128xf32> loc(#loc101)
  func.func @bwd_qkv_kernel(%arg0: memref<?xi8> loc("flash_attention_npu_v8_copy.py":821:0), %arg1: memref<?xi8> loc("flash_attention_npu_v8_copy.py":821:0), %arg2: memref<?xbf16> {tt.divisibility = 16 : i32, tt.tensor_kind = 0 : i32} loc("flash_attention_npu_v8_copy.py":821:0), %arg3: memref<?xbf16> {tt.divisibility = 16 : i32, tt.tensor_kind = 0 : i32} loc("flash_attention_npu_v8_copy.py":821:0), %arg4: memref<?xbf16> {tt.divisibility = 16 : i32, tt.tensor_kind = 0 : i32} loc("flash_attention_npu_v8_copy.py":821:0), %arg5: memref<?xf32> {tt.divisibility = 16 : i32, tt.tensor_kind = 2 : i32} loc("flash_attention_npu_v8_copy.py":821:0), %arg6: memref<?xbf16> {tt.divisibility = 16 : i32, tt.tensor_kind = 1 : i32} loc("flash_attention_npu_v8_copy.py":821:0), %arg7: memref<?xbf16> {tt.divisibility = 16 : i32, tt.tensor_kind = 1 : i32} loc("flash_attention_npu_v8_copy.py":821:0), %arg8: memref<?xbf16> {tt.divisibility = 16 : i32, tt.tensor_kind = 0 : i32} loc("flash_attention_npu_v8_copy.py":821:0), %arg9: memref<?xf32> {tt.divisibility = 16 : i32, tt.tensor_kind = 0 : i32} loc("flash_attention_npu_v8_copy.py":821:0), %arg10: memref<?xf32> {tt.divisibility = 16 : i32, tt.tensor_kind = 0 : i32} loc("flash_attention_npu_v8_copy.py":821:0), %arg11: memref<?xi32> {tt.divisibility = 16 : i32} loc("flash_attention_npu_v8_copy.py":821:0), %arg12: memref<?xi32> {tt.divisibility = 16 : i32} loc("flash_attention_npu_v8_copy.py":821:0), %arg13: memref<?xi8> {tt.divisibility = 16 : i32, tt.tensor_kind = 0 : i32} loc("flash_attention_npu_v8_copy.py":821:0), %arg14: memref<?xi32> {tt.divisibility = 16 : i32, tt.tensor_kind = 0 : i32} loc("flash_attention_npu_v8_copy.py":821:0), %arg15: memref<?xi32> {tt.divisibility = 16 : i32, tt.tensor_kind = 0 : i32} loc("flash_attention_npu_v8_copy.py":821:0), %arg16: i32 loc("flash_attention_npu_v8_copy.py":821:0), %arg17: i32 {tt.divisibility = 16 : i32} loc("flash_attention_npu_v8_copy.py":821:0), %arg18: i32 {tt.divisibility = 16 : i32} loc("flash_attention_npu_v8_copy.py":821:0), %arg19: i32 loc("flash_attention_npu_v8_copy.py":821:0), %arg20: i32 loc("flash_attention_npu_v8_copy.py":821:0), %arg21: f32 loc("flash_attention_npu_v8_copy.py":821:0), %arg22: i32 loc("flash_attention_npu_v8_copy.py":821:0), %arg23: i32 loc("flash_attention_npu_v8_copy.py":821:0), %arg24: i32 loc("flash_attention_npu_v8_copy.py":821:0), %arg25: i32 loc("flash_attention_npu_v8_copy.py":821:0), %arg26: i32 loc("flash_attention_npu_v8_copy.py":821:0), %arg27: i32 loc("flash_attention_npu_v8_copy.py":821:0)) attributes {SyncBlockLockArgIdx = 0 : i64, WorkspaceArgIdx = 1 : i64, global_kernel = "local", mix_mode = "mix", parallel_mode = "mix_simd_simt"} {
    %c128 = arith.constant 128 : index loc(#loc4)
    %cst = arith.constant 0.000000e+00 : bf16 loc(#loc102)
    %c64 = arith.constant 64 : index loc(#loc103)
    %c1 = arith.constant 1 : index loc(#loc8)
    %c32_i32 = arith.constant 32 : i32 loc(#loc9)
    %c63_i32 = arith.constant 63 : i32 loc(#loc10)
    %c0_i8 = arith.constant 0 : i8 loc(#loc104)
    %c64_i64 = arith.constant 64 : i64 loc(#loc13)
    %c128_i32 = arith.constant 128 : i32 loc(#loc14)
    %c0_i32 = arith.constant 0 : i32 loc(#loc102)
    %c64_i32 = arith.constant 64 : i32 loc(#loc15)
    %c0_i64 = arith.constant 0 : i64 loc(#loc101)
    %cst_0 = arith.constant 0.000000e+00 : f32 loc(#loc16)
    %c0 = arith.constant 0 : index loc(#loc103)
    %cst_1 = arith.constant 1.44269502 : f32 loc(#loc17)
    %0 = tensor.empty() : tensor<1xf32> loc(#loc18)
    %1 = linalg.fill ins(%cst_1 : f32) outs(%0 : tensor<1xf32>) -> tensor<1xf32> loc(#loc18)
    %2 = tensor.empty() : tensor<128xf32> loc(#loc101)
    %3 = linalg.fill ins(%cst_0 : f32) outs(%2 : tensor<128xf32>) -> tensor<128xf32> loc(#loc101)
    %4 = tensor.empty() : tensor<128xi64> loc(#loc100)
    %5 = linalg.fill ins(%c0_i64 : i64) outs(%4 : tensor<128xi64>) -> tensor<128xi64> loc(#loc101)
    %6 = tensor.empty() : tensor<128x64xf32> loc(#loc19)
    %7 = linalg.fill ins(%cst_0 : f32) outs(%6 : tensor<128x64xf32>) -> tensor<128x64xf32> loc(#loc20)
    %8 = tensor.empty() : tensor<128x1xf32> loc(#loc17)
    %9 = linalg.fill ins(%cst_1 : f32) outs(%8 : tensor<128x1xf32>) -> tensor<128x1xf32> loc(#loc17)
    %10 = tensor.empty() : tensor<128x64xi8> loc(#loc21)
    %11 = linalg.fill ins(%c0_i8 : i8) outs(%10 : tensor<128x64xi8>) -> tensor<128x64xi8> loc(#loc21)
    %12 = tensor.empty() : tensor<64x64xf32> loc(#loc22)
    %13 = linalg.fill ins(%cst_0 : f32) outs(%12 : tensor<64x64xf32>) -> tensor<64x64xf32> loc(#loc16)
    %14 = arith.addi %arg18, %c63_i32 : i32 loc(#loc10)
    %15 = arith.divsi %14, %c64_i32 : i32 loc(#loc10)
    %16 = arith.muli %15, %arg19 : i32 loc(#loc23)
    %17 = arith.muli %16, %arg16 : i32 loc(#loc24)
    scf.for %arg28 = %arg25 to %17 step %c32_i32  : i32 {
      %18 = arith.divsi %arg28, %15 : i32 loc(#loc25)
      %19 = arith.remsi %arg28, %15 : i32 loc(#loc26)
      %20 = arith.divsi %18, %arg19 : i32 loc(#loc27)
      %21 = arith.remsi %18, %arg19 : i32 loc(#loc28)
      %22 = arith.divsi %arg19, %arg20 : i32 loc(#loc29)
      %23 = arith.divsi %21, %22 : i32 loc(#loc30)
      %24 = arith.index_cast %20 : i32 to index loc(#loc27)
      %reinterpret_cast = memref.reinterpret_cast %arg15 to offset: [%24], sizes: [1], strides: [1] : memref<?xi32> to memref<1xi32, strided<[1], offset: ?>> loc(#loc31)
      %25 = memref.load %reinterpret_cast[%c0] : memref<1xi32, strided<[1], offset: ?>> loc(#loc32)
      %26 = arith.addi %24, %c1 : index loc(#loc8)
      %reinterpret_cast_2 = memref.reinterpret_cast %arg15 to offset: [%26], sizes: [1], strides: [1] : memref<?xi32> to memref<1xi32, strided<[1], offset: ?>> loc(#loc8)
      %27 = memref.load %reinterpret_cast_2[%c0] : memref<1xi32, strided<[1], offset: ?>> loc(#loc33)
      %28 = arith.subi %27, %25 : i32 loc(#loc34)
      %29 = arith.muli %19, %c64_i32 : i32 loc(#loc35)
      %30 = arith.cmpi slt, %29, %28 : i32 loc(#loc36)
      scf.if %30 {
        %reinterpret_cast_3 = memref.reinterpret_cast %arg14 to offset: [%24], sizes: [1], strides: [1] : memref<?xi32> to memref<1xi32, strided<[1], offset: ?>> loc(#loc38)
        %31 = memref.load %reinterpret_cast_3[%c0] : memref<1xi32, strided<[1], offset: ?>> loc(#loc39)
        %reinterpret_cast_4 = memref.reinterpret_cast %arg14 to offset: [%26], sizes: [1], strides: [1] : memref<?xi32> to memref<1xi32, strided<[1], offset: ?>> loc(#loc40)
        %32 = memref.load %reinterpret_cast_4[%c0] : memref<1xi32, strided<[1], offset: ?>> loc(#loc41)
        %33 = arith.subi %32, %31 : i32 loc(#loc42)
        %34 = arith.divsi %29, %c128_i32 : i32 loc(#loc43)
        %35 = arith.muli %34, %c128_i32 : i32 loc(#loc44)
        %36 = arith.addi %35, %c128_i32 : i32 loc(#loc45)
        %37 = arith.minsi %36, %33 : i32 loc(#loc46)
        %inserted = tensor.insert %arg21 into %0[%c0] : tensor<1xf32> loc(#loc18)
        %38 = arith.mulf %inserted, %1 : tensor<1xf32> loc(#loc18)
        %extracted = tensor.extract %38[%c0] : tensor<1xf32> loc(#loc18)
        %39 = arith.extsi %31 : i32 to i64 loc(#loc47)
        %40 = arith.extsi %25 : i32 to i64 loc(#loc48)
        %41 = arith.extsi %arg19 : i32 to i64 loc(#loc49)
        %42 = arith.muli %39, %41 : i64 loc(#loc49)
        %43 = arith.muli %42, %c64_i64 : i64 loc(#loc50)
        %44 = arith.muli %21, %c64_i32 : i32 loc(#loc51)
        %45 = arith.index_cast %43 : i64 to index loc(#loc50)
        %46 = arith.index_cast %44 : i32 to index loc(#loc51)
        %47 = arith.addi %45, %46 : index loc(#loc52)
        %48 = arith.muli %arg19, %c64_i32 : i32 loc(#loc53)
        %49 = arith.extsi %33 : i32 to i64 loc(#loc54)
        %50 = arith.extsi %arg20 : i32 to i64 loc(#loc55)
        %51 = arith.muli %40, %50 : i64 loc(#loc55)
        %52 = arith.muli %51, %c64_i64 : i64 loc(#loc56)
        %53 = arith.muli %23, %c64_i32 : i32 loc(#loc57)
        %54 = arith.index_cast %52 : i64 to index loc(#loc56)
        %55 = arith.index_cast %53 : i32 to index loc(#loc57)
        %56 = arith.addi %54, %55 : index loc(#loc58)
        %57 = arith.muli %arg20, %c64_i32 : i32 loc(#loc15)
        %58 = arith.muli %40, %41 : i64 loc(#loc59)
        %59 = arith.muli %58, %c64_i64 : i64 loc(#loc60)
        %60 = arith.index_cast %59 : i64 to index loc(#loc60)
        %61 = arith.addi %60, %46 : index loc(#loc61)
        %62 = arith.maxsi %29, %c0_i32 : i32 loc(#loc62)
        %63 = arith.index_cast %62 : i32 to index loc(#loc62)
        %64 = arith.index_cast %48 : i32 to index loc(#loc54)
        %65 = arith.muli %63, %64 : index loc(#loc62)
        %66 = arith.addi %65, %61 : index loc(#loc62)
        %67 = arith.index_cast %28 : i32 to index loc(#loc63)
        %reinterpret_cast_5 = memref.reinterpret_cast %arg6 to offset: [%66], sizes: [64, 64], strides: [%64, 1] : memref<?xbf16> to memref<64x64xbf16, strided<[?, 1], offset: ?>> loc(#loc62)
        %reinterpret_cast_6 = memref.reinterpret_cast %arg7 to offset: [%66], sizes: [64, 64], strides: [%64, 1] : memref<?xbf16> to memref<64x64xbf16, strided<[?, 1], offset: ?>> loc(#loc64)
        %68 = arith.muli %20, %arg17 : i32 loc(#loc65)
        %69 = arith.muli %68, %arg18 : i32 loc(#loc66)
        %70 = arith.index_cast %69 : i32 to index loc(#loc66)
        %71 = arith.index_cast %57 : i32 to index loc(#loc63)
        %72 = arith.muli %63, %71 : index loc(#loc105)
        %73 = arith.addi %72, %56 : index loc(#loc105)
        %reinterpret_cast_7 = memref.reinterpret_cast %arg3 to offset: [%73], sizes: [64, 64], strides: [%71, 1] : memref<?xbf16> to memref<64x64xbf16, strided<[?, 1], offset: ?>> loc(#loc105)
        %alloc = memref.alloc() : memref<64x64xbf16> loc(#loc105)
        %74 = arith.divsi %72, %71 : index loc(#loc105)
        %75 = arith.subi %67, %74 : index loc(#loc105)
        %76 = arith.maxsi %75, %c0 : index loc(#loc105)
        %77 = arith.minsi %76, %c64 : index loc(#loc105)
        %78 = arith.subi %c0_i32, %29 : i32 loc(#loc105)
        %79 = arith.maxsi %78, %c0_i32 : i32 loc(#loc105)
        %80 = arith.index_cast %79 : i32 to index loc(#loc105)
        %81 = arith.minsi %80, %77 : index loc(#loc105)
        %82 = arith.subi %77, %81 : index loc(#loc105)
        %83 = arith.cmpi slt, %82, %c64 : index loc(#loc105)
        scf.if %83 {
          linalg.fill ins(%cst : bf16) outs(%alloc : memref<64x64xbf16>) loc(#loc105)
        } {hivm.unlikely_condition} loc(#loc105)
        %subview = memref.subview %reinterpret_cast_7[0, 0] [%82, 64] [1, 1] : memref<64x64xbf16, strided<[?, 1], offset: ?>> to memref<?x64xbf16, strided<[?, 1], offset: ?>> loc(#loc105)
        %subview_8 = memref.subview %alloc[%81, 0] [%82, 64] [1, 1] : memref<64x64xbf16> to memref<?x64xbf16, strided<[64, 1], offset: ?>> loc(#loc105)
        memref.copy %subview, %subview_8 : memref<?x64xbf16, strided<[?, 1], offset: ?>> to memref<?x64xbf16, strided<[64, 1], offset: ?>> loc(#loc105)
        %84 = bufferization.to_tensor %alloc restrict writable : memref<64x64xbf16> loc(#loc105)
        %85 = tensor.empty() : tensor<64x64xbf16> loc(#loc105)
        %transposed = linalg.transpose ins(%84 : tensor<64x64xbf16>) outs(%85 : tensor<64x64xbf16>) permutation = [1, 0]  loc(#loc105)
        %reinterpret_cast_9 = memref.reinterpret_cast %arg4 to offset: [%73], sizes: [64, 64], strides: [%71, 1] : memref<?xbf16> to memref<64x64xbf16, strided<[?, 1], offset: ?>> loc(#loc106)
        %alloc_10 = memref.alloc() : memref<64x64xbf16> loc(#loc106)
        scf.if %83 {
          linalg.fill ins(%cst : bf16) outs(%alloc_10 : memref<64x64xbf16>) loc(#loc106)
        } {hivm.unlikely_condition} loc(#loc106)
        %subview_11 = memref.subview %reinterpret_cast_9[0, 0] [%82, 64] [1, 1] : memref<64x64xbf16, strided<[?, 1], offset: ?>> to memref<?x64xbf16, strided<[?, 1], offset: ?>> loc(#loc106)
        %subview_12 = memref.subview %alloc_10[%81, 0] [%82, 64] [1, 1] : memref<64x64xbf16> to memref<?x64xbf16, strided<[64, 1], offset: ?>> loc(#loc106)
        memref.copy %subview_11, %subview_12 : memref<?x64xbf16, strided<[?, 1], offset: ?>> to memref<?x64xbf16, strided<[64, 1], offset: ?>> loc(#loc106)
        %86 = bufferization.to_tensor %alloc_10 restrict writable : memref<64x64xbf16> loc(#loc106)
        %transposed_13 = linalg.transpose ins(%86 : tensor<64x64xbf16>) outs(%85 : tensor<64x64xbf16>) permutation = [1, 0]  loc(#loc106)
        %87 = linalg.fill ins(%extracted : f32) outs(%6 : tensor<128x64xf32>) -> tensor<128x64xf32> loc(#loc19)
        %88 = arith.extsi %44 : i32 to i64 loc(#loc70)
        %89 = tensor.empty() : tensor<128xi32> loc(#loc71)
        %90 = linalg.generic {indexing_maps = [#map], iterator_types = ["parallel"]} outs(%89 : tensor<128xi32>) attrs =  {tt.from_make_range, tt.make_range_offset = 0 : index, tt.make_range_size = 128 : index} {
        ^bb0(%out: i32 loc("flash_attention_npu_v8_copy.py":997:44)):
          %103 = linalg.index 0 : index loc(#loc71)
          %104 = arith.index_cast %103 : index to i32 loc(#loc71)
          linalg.yield %104 : i32 loc(#loc71)
        } -> tensor<128xi32> loc(#loc71)
        %91 = linalg.fill ins(%arg21 : f32) outs(%6 : tensor<128x64xf32>) -> tensor<128x64xf32> loc(#loc72)
        %92:7 = scf.for %arg29 = %c0_i32 to %37 step %c128_i32 iter_args(%arg30 = %13, %arg31 = %13, %arg32 = %c0_i32, %arg33 = %c0_i32, %arg34 = %c0_i32, %arg35 = %c0_i32, %arg36 = %c0_i32) -> (tensor<64x64xf32>, tensor<64x64xf32>, i32, i32, i32, i32, i32)  : i32 {
          %103 = arith.maxsi %arg32, %c0_i32 : i32 loc(#loc16)
          %104 = arith.index_cast %103 : i32 to index loc(#loc16)
          %105 = arith.muli %104, %64 : index loc(#loc16)
          %106 = arith.addi %105, %47 : index loc(#loc16)
          %107 = arith.index_cast %33 : i32 to index loc(#loc54)
          %reinterpret_cast_17 = memref.reinterpret_cast %arg2 to offset: [%106], sizes: [128, 64], strides: [%64, 1] : memref<?xbf16> to memref<128x64xbf16, strided<[?, 1], offset: ?>> loc(#loc16)
          %108 = arith.maxsi %arg33, %c0_i32 : i32 loc(#loc16)
          %109 = arith.index_cast %108 : i32 to index loc(#loc16)
          %110 = arith.muli %109, %64 : index loc(#loc16)
          %111 = arith.addi %110, %47 : index loc(#loc16)
          %reinterpret_cast_18 = memref.reinterpret_cast %arg8 to offset: [%111], sizes: [128, 64], strides: [%64, 1] : memref<?xbf16> to memref<128x64xbf16, strided<[?, 1], offset: ?>> loc(#loc16)
          %112 = arith.maxsi %arg36, %c0_i32 : i32 loc(#loc16)
          %113 = arith.index_cast %112 : i32 to index loc(#loc16)
          %114 = arith.index_cast %arg18 : i32 to index loc(#loc10)
          %115 = arith.muli %113, %114 : index loc(#loc16)
          %116 = arith.addi %115, %70 : index loc(#loc16)
          %117 = arith.addi %116, %63 : index loc(#loc16)
          %reinterpret_cast_19 = memref.reinterpret_cast %arg13 to offset: [%117], sizes: [128, 64], strides: [%114, 1] : memref<?xi8> to memref<128x64xi8, strided<[?, 1], offset: ?>> loc(#loc16)
          %alloc_20 = memref.alloc() : memref<128x64xi8> loc(#loc104)
          %118 = arith.subi %117, %70 : index loc(#loc104)
          %119 = arith.divsi %118, %114 : index loc(#loc104)
          %120 = arith.subi %107, %119 : index loc(#loc104)
          %121 = arith.maxsi %120, %c0 : index loc(#loc104)
          %122 = arith.minsi %121, %c128 : index loc(#loc104)
          %123 = arith.remsi %118, %114 : index loc(#loc104)
          %124 = arith.subi %67, %123 : index loc(#loc104)
          %125 = arith.maxsi %124, %c0 : index loc(#loc104)
          %126 = arith.minsi %125, %c64 : index loc(#loc104)
          %127 = arith.subi %c0_i32, %arg36 : i32 loc(#loc104)
          %128 = arith.maxsi %127, %c0_i32 : i32 loc(#loc104)
          %129 = arith.index_cast %128 : i32 to index loc(#loc104)
          %130 = arith.minsi %129, %122 : index loc(#loc104)
          %131 = arith.subi %122, %130 : index loc(#loc104)
          %132 = arith.minsi %80, %126 : index loc(#loc104)
          %133 = arith.subi %126, %132 : index loc(#loc104)
          %134 = arith.cmpi slt, %131, %c128 : index loc(#loc104)
          %135 = arith.cmpi slt, %133, %c64 : index loc(#loc104)
          %136 = arith.ori %134, %135 : i1 loc(#loc104)
          scf.if %136 {
            linalg.fill ins(%c0_i8 : i8) outs(%alloc_20 : memref<128x64xi8>) loc(#loc104)
          } {hivm.unlikely_condition} loc(#loc104)
          %subview_21 = memref.subview %reinterpret_cast_19[0, 0] [%131, %133] [1, 1] : memref<128x64xi8, strided<[?, 1], offset: ?>> to memref<?x?xi8, strided<[?, 1], offset: ?>> loc(#loc104)
          %subview_22 = memref.subview %alloc_20[%130, %132] [%131, %133] [1, 1] : memref<128x64xi8> to memref<?x?xi8, strided<[64, 1], offset: ?>> loc(#loc104)
          memref.copy %subview_21, %subview_22 : memref<?x?xi8, strided<[?, 1], offset: ?>> to memref<?x?xi8, strided<[64, 1], offset: ?>> loc(#loc104)
          %137 = bufferization.to_tensor %alloc_20 restrict writable : memref<128x64xi8> loc(#loc104)
          %alloc_23 = memref.alloc() : memref<128x64xbf16> loc(#loc107)
          %138 = arith.divsi %105, %64 : index loc(#loc107)
          %139 = arith.subi %107, %138 : index loc(#loc107)
          %140 = arith.maxsi %139, %c0 : index loc(#loc107)
          %141 = arith.minsi %140, %c128 : index loc(#loc107)
          %142 = arith.subi %c0_i32, %arg32 : i32 loc(#loc107)
          %143 = arith.maxsi %142, %c0_i32 : i32 loc(#loc107)
          %144 = arith.index_cast %143 : i32 to index loc(#loc107)
          %145 = arith.minsi %144, %141 : index loc(#loc107)
          %146 = arith.subi %141, %145 : index loc(#loc107)
          %147 = arith.cmpi slt, %146, %c128 : index loc(#loc107)
          scf.if %147 {
            linalg.fill ins(%cst : bf16) outs(%alloc_23 : memref<128x64xbf16>) loc(#loc107)
          } {hivm.unlikely_condition} loc(#loc107)
          %subview_24 = memref.subview %reinterpret_cast_17[0, 0] [%146, 64] [1, 1] : memref<128x64xbf16, strided<[?, 1], offset: ?>> to memref<?x64xbf16, strided<[?, 1], offset: ?>> loc(#loc107)
          %subview_25 = memref.subview %alloc_23[%145, 0] [%146, 64] [1, 1] : memref<128x64xbf16> to memref<?x64xbf16, strided<[64, 1], offset: ?>> loc(#loc107)
          memref.copy %subview_24, %subview_25 : memref<?x64xbf16, strided<[?, 1], offset: ?>> to memref<?x64xbf16, strided<[64, 1], offset: ?>> loc(#loc107)
          %148 = bufferization.to_tensor %alloc_23 restrict writable : memref<128x64xbf16> loc(#loc107)
          %149 = linalg.matmul {input_precision = "ieee"} ins(%148, %transposed : tensor<128x64xbf16>, tensor<64x64xbf16>) outs(%7 : tensor<128x64xf32>) -> tensor<128x64xf32> loc(#loc74)
          %150 = arith.extsi %21 : i32 to i64 loc(#loc100)
          %151 = arith.addi %150, %42 : i64 loc(#loc100)
          %152 = arith.extsi %arg34 : i32 to i64 loc(#loc100)
          %153 = arith.muli %152, %41 : i64 loc(#loc100)
          %154 = arith.addi %151, %153 : i64 loc(#loc100)
          %155 = linalg.fill ins(%154 : i64) outs(%4 : tensor<128xi64>) -> tensor<128xi64> loc(#loc100)
          %156 = arith.extsi %90 : tensor<128xi32> to tensor<128xi64> loc(#loc100)
          %157 = linalg.fill ins(%41 : i64) outs(%4 : tensor<128xi64>) -> tensor<128xi64> loc(#loc100)
          %158 = arith.muli %156, %157 : tensor<128xi64> loc(#loc100)
          %159 = arith.addi %155, %158 : tensor<128xi64> loc(#loc100)
          %160 = linalg.fill ins(%arg34 : i32) outs(%89 : tensor<128xi32>) -> tensor<128xi32> loc(#loc100)
          %161 = arith.addi %90, %160 : tensor<128xi32> loc(#loc100)
          %162 = arith.extsi %161 : tensor<128xi32> to tensor<128xi64> loc(#loc100)
          %163 = linalg.fill ins(%49 : i64) outs(%4 : tensor<128xi64>) -> tensor<128xi64> loc(#loc100)
          %164 = arith.cmpi sge, %162, %5 : tensor<128xi64> loc(#loc100)
          %165 = arith.cmpi slt, %162, %163 : tensor<128xi64> loc(#loc100)
          %166 = arith.andi %164, %165 : tensor<128xi1> loc(#loc100)
          %167 = func.call @triton_indirect_load(%arg9, %159, %166, %3) : (memref<?xf32>, tensor<128xi64>, tensor<128xi1>, tensor<128xf32>) -> tensor<128xf32> loc(#loc100)
          %168 = arith.mulf %149, %87 : tensor<128x64xf32> loc(#loc19)
          %expanded = tensor.expand_shape %167 [[0, 1]] output_shape [128, 1] : tensor<128xf32> into tensor<128x1xf32> loc(#loc75)
          %169 = arith.mulf %expanded, %9 : tensor<128x1xf32> loc(#loc17)
          %collapsed = tensor.collapse_shape %169 [[0, 1]] : tensor<128x1xf32> into tensor<128xf32> loc(#loc76)
          %broadcasted = linalg.broadcast ins(%collapsed : tensor<128xf32>) outs(%6 : tensor<128x64xf32>) dimensions = [1]  loc(#loc76)
          %170 = arith.subf %168, %broadcasted : tensor<128x64xf32> loc(#loc76)
          %171 = math.exp2 %170 : tensor<128x64xf32> loc(#loc77)
          %172 = arith.cmpi ne, %137, %11 : tensor<128x64xi8> loc(#loc21)
          %173 = arith.select %172, %171, %7 : tensor<128x64xi1>, tensor<128x64xf32> loc(#loc21)
          %alloc_26 = memref.alloc() : memref<128x64xbf16> loc(#loc102)
          %174 = arith.divsi %110, %64 : index loc(#loc102)
          %175 = arith.subi %107, %174 : index loc(#loc102)
          %176 = arith.maxsi %175, %c0 : index loc(#loc102)
          %177 = arith.minsi %176, %c128 : index loc(#loc102)
          %178 = arith.subi %c0_i32, %arg33 : i32 loc(#loc102)
          %179 = arith.maxsi %178, %c0_i32 : i32 loc(#loc102)
          %180 = arith.index_cast %179 : i32 to index loc(#loc102)
          %181 = arith.minsi %180, %177 : index loc(#loc102)
          %182 = arith.subi %177, %181 : index loc(#loc102)
          %183 = arith.cmpi slt, %182, %c128 : index loc(#loc102)
          scf.if %183 {
            linalg.fill ins(%cst : bf16) outs(%alloc_26 : memref<128x64xbf16>) loc(#loc102)
          } {hivm.unlikely_condition} loc(#loc102)
          %subview_27 = memref.subview %reinterpret_cast_18[0, 0] [%182, 64] [1, 1] : memref<128x64xbf16, strided<[?, 1], offset: ?>> to memref<?x64xbf16, strided<[?, 1], offset: ?>> loc(#loc102)
          %subview_28 = memref.subview %alloc_26[%181, 0] [%182, 64] [1, 1] : memref<128x64xbf16> to memref<?x64xbf16, strided<[64, 1], offset: ?>> loc(#loc102)
          memref.copy %subview_27, %subview_28 : memref<?x64xbf16, strided<[?, 1], offset: ?>> to memref<?x64xbf16, strided<[64, 1], offset: ?>> loc(#loc102)
          %184 = bufferization.to_tensor %alloc_26 restrict writable : memref<128x64xbf16> loc(#loc102)
          %185 = arith.truncf %173 : tensor<128x64xf32> to tensor<128x64xbf16> loc(#loc78)
          %186 = tensor.empty() : tensor<64x128xbf16> loc(#loc79)
          %transposed_29 = linalg.transpose ins(%185 : tensor<128x64xbf16>) outs(%186 : tensor<64x128xbf16>) permutation = [1, 0]  loc(#loc79)
          %187 = linalg.matmul {input_precision = "ieee"} ins(%transposed_29, %184 : tensor<64x128xbf16>, tensor<128x64xbf16>) outs(%arg31 : tensor<64x64xf32>) -> tensor<64x64xf32> loc(#loc80)
          %188 = arith.extsi %arg35 : i32 to i64 loc(#loc101)
          %189 = arith.muli %188, %41 : i64 loc(#loc101)
          %190 = arith.addi %151, %189 : i64 loc(#loc101)
          %191 = linalg.fill ins(%190 : i64) outs(%4 : tensor<128xi64>) -> tensor<128xi64> loc(#loc101)
          %192 = arith.addi %191, %158 : tensor<128xi64> loc(#loc101)
          %193 = linalg.fill ins(%arg35 : i32) outs(%89 : tensor<128xi32>) -> tensor<128xi32> loc(#loc101)
          %194 = arith.addi %90, %193 : tensor<128xi32> loc(#loc101)
          %195 = arith.extsi %194 : tensor<128xi32> to tensor<128xi64> loc(#loc101)
          %196 = arith.cmpi sge, %195, %5 : tensor<128xi64> loc(#loc101)
          %197 = arith.cmpi slt, %195, %163 : tensor<128xi64> loc(#loc101)
          %198 = arith.andi %196, %197 : tensor<128xi1> loc(#loc101)
          %199 = func.call @triton_indirect_load_0(%arg10, %192, %198, %3) : (memref<?xf32>, tensor<128xi64>, tensor<128xi1>, tensor<128xf32>) -> tensor<128xf32> loc(#loc101)
          %200 = linalg.matmul {input_precision = "ieee"} ins(%184, %transposed_13 : tensor<128x64xbf16>, tensor<64x64xbf16>) outs(%7 : tensor<128x64xf32>) -> tensor<128x64xf32> loc(#loc81)
          %broadcasted_30 = linalg.broadcast ins(%199 : tensor<128xf32>) outs(%6 : tensor<128x64xf32>) dimensions = [1]  loc(#loc82)
          %201 = arith.subf %200, %broadcasted_30 : tensor<128x64xf32> loc(#loc82)
          %202 = arith.mulf %173, %201 : tensor<128x64xf32> loc(#loc83)
          %203 = arith.select %172, %202, %7 : tensor<128x64xi1>, tensor<128x64xf32> loc(#loc84)
          %204 = arith.truncf %203 : tensor<128x64xf32> to tensor<128x64xbf16> loc(#loc85)
          %205 = arith.extsi %arg29 : i32 to i64 loc(#loc86)
          %206 = arith.addi %39, %205 : i64 loc(#loc86)
          %207 = arith.muli %206, %41 : i64 loc(#loc87)
          %208 = arith.muli %207, %c64_i64 : i64 loc(#loc13)
          %209 = arith.addi %208, %88 : i64 loc(#loc70)
          annotation.mark %204 {break_vf} : tensor<128x64xbf16> loc(#loc88)
          %210 = linalg.matmul {input_precision = "ieee"} ins(%204, %84 : tensor<128x64xbf16>, tensor<64x64xbf16>) outs(%7 : tensor<128x64xf32>) -> tensor<128x64xf32> loc(#loc20)
          %211 = arith.mulf %210, %91 : tensor<128x64xf32> loc(#loc72)
          %212 = arith.index_cast %209 : i64 to index loc(#loc70)
          %213 = arith.index_cast %arg19 : i32 to index loc(#loc)
          %214 = arith.muli %213, %c64 : index loc(#loc89)
          %reinterpret_cast_31 = memref.reinterpret_cast %arg5 to offset: [%212], sizes: [128, 64], strides: [%214, 1] : memref<?xf32> to memref<128x64xf32, strided<[?, 1], offset: ?>> loc(#loc89)
          %215 = arith.index_cast %arg29 : i32 to index loc(#loc16)
          %216 = arith.addi %215, %c128 : index loc(#loc4)
          %217 = arith.index_cast %37 : i32 to index loc(#loc46)
          %218 = arith.maxsi %215, %217 : index loc(#loc4)
          %219 = arith.minsi %216, %218 : index loc(#loc4)
          %220 = arith.subi %219, %215 : index loc(#loc4)
          %subview_32 = memref.subview %reinterpret_cast_31[0, 0] [%220, 64] [1, 1] : memref<128x64xf32, strided<[?, 1], offset: ?>> to memref<?x64xf32, strided<[?, 1], offset: ?>> loc(#loc4)
          %extracted_slice_33 = tensor.extract_slice %211[0, 0] [%220, 64] [1, 1] : tensor<128x64xf32> to tensor<?x64xf32> loc(#loc4)
          hivm.hir.store ins(%extracted_slice_33 : tensor<?x64xf32>) outs(%subview_32 : memref<?x64xf32, strided<[?, 1], offset: ?>>) atomic = <add> loc(#loc4)
          %transposed_34 = linalg.transpose ins(%204 : tensor<128x64xbf16>) outs(%186 : tensor<64x128xbf16>) permutation = [1, 0]  loc(#loc90)
          %221 = linalg.matmul {input_precision = "ieee"} ins(%transposed_34, %148 : tensor<64x128xbf16>, tensor<128x64xbf16>) outs(%arg30 : tensor<64x64xf32>) -> tensor<64x64xf32> loc(#loc91)
          %222 = arith.addi %arg32, %c128_i32 : i32 loc(#loc92)
          %223 = arith.addi %arg33, %c128_i32 : i32 loc(#loc93)
          %224 = arith.addi %arg34, %c128_i32 : i32 loc(#loc94)
          %225 = arith.addi %arg35, %c128_i32 : i32 loc(#loc95)
          %226 = arith.addi %arg36, %c128_i32 : i32 loc(#loc14)
          scf.yield %221, %187, %222, %223, %224, %225, %226 : tensor<64x64xf32>, tensor<64x64xf32>, i32, i32, i32, i32, i32 loc(#loc96)
        } loc(#loc16)
        %93 = linalg.fill ins(%arg21 : f32) outs(%12 : tensor<64x64xf32>) -> tensor<64x64xf32> loc(#loc22)
        %94 = arith.mulf %92#0, %93 : tensor<64x64xf32> loc(#loc22)
        %95 = arith.truncf %94 : tensor<64x64xf32> to tensor<64x64xbf16> loc(#loc97)
        %96 = arith.divsi %65, %64 : index loc(#loc103)
        %97 = arith.subi %67, %96 : index loc(#loc103)
        %98 = arith.maxsi %97, %c0 : index loc(#loc103)
        %99 = arith.minsi %98, %c64 : index loc(#loc103)
        %100 = arith.minsi %80, %99 : index loc(#loc103)
        %101 = arith.subi %99, %100 : index loc(#loc103)
        %extracted_slice = tensor.extract_slice %95[%100, 0] [%101, 64] [1, 1] : tensor<64x64xbf16> to tensor<?x64xbf16> loc(#loc103)
        %subview_14 = memref.subview %reinterpret_cast_5[0, 0] [%101, 64] [1, 1] : memref<64x64xbf16, strided<[?, 1], offset: ?>> to memref<?x64xbf16, strided<[?, 1], offset: ?>> loc(#loc103)
        bufferization.materialize_in_destination %extracted_slice in writable %subview_14 : (tensor<?x64xbf16>, memref<?x64xbf16, strided<[?, 1], offset: ?>>) -> () loc(#loc103)
        %102 = arith.truncf %92#1 : tensor<64x64xf32> to tensor<64x64xbf16> loc(#loc98)
        %extracted_slice_15 = tensor.extract_slice %102[%100, 0] [%101, 64] [1, 1] : tensor<64x64xbf16> to tensor<?x64xbf16> loc(#loc108)
        %subview_16 = memref.subview %reinterpret_cast_6[0, 0] [%101, 64] [1, 1] : memref<64x64xbf16, strided<[?, 1], offset: ?>> to memref<?x64xbf16, strided<[?, 1], offset: ?>> loc(#loc108)
        bufferization.materialize_in_destination %extracted_slice_15 in writable %subview_16 : (tensor<?x64xbf16>, memref<?x64xbf16, strided<[?, 1], offset: ?>>) -> () loc(#loc108)
      } loc(#loc37)
    } loc(#loc9)
    return loc(#loc)
  } loc(#loc)
} loc(#loc)
#loc1 = loc("flash_attention_npu_v8_copy.py":120:23)
#loc2 = loc("flash_attention_npu_v8_copy.py":983:52)
#loc3 = loc("flash_attention_npu_v8_copy.py":990:52)
#loc4 = loc("flash_attention_npu_v8_copy.py":1005:24)
#loc5 = loc("flash_attention_npu_v8_copy.py":986:54)
#loc6 = loc("flash_attention_npu_v8_copy.py":130:28)
#loc7 = loc("flash_attention_npu_v8_copy.py":1018:56)
#loc8 = loc("flash_attention_npu_v8_copy.py":853:49)
#loc9 = loc("flash_attention_npu_v8_copy.py":846:51)
#loc10 = loc("flash_attention_npu_v8_copy.py":843:37)
#loc11 = loc("flash_attention_npu_v8_copy.py":122:23)
#loc12 = loc("flash_attention_npu_v8_copy.py":979:54)
#loc13 = loc("flash_attention_npu_v8_copy.py":996:66)
#loc14 = loc("flash_attention_npu_v8_copy.py":1015:60)
#loc15 = loc("flash_attention_npu_v8_copy.py":888:38)
#loc16 = loc("flash_attention_npu_v8_copy.py":974:45)
#loc17 = loc("flash_attention_npu_v8_copy.py":984:65)
#loc18 = loc("flash_attention_npu_v8_copy.py":872:31)
#loc19 = loc("flash_attention_npu_v8_copy.py":984:41)
#loc20 = loc("flash_attention_npu_v8_copy.py":1001:36)
#loc21 = loc("flash_attention_npu_v8_copy.py":985:42)
#loc22 = loc("flash_attention_npu_v8_copy.py":1017:18)
#loc23 = loc("flash_attention_npu_v8_copy.py":844:32)
#loc24 = loc("flash_attention_npu_v8_copy.py":844:41)
#loc25 = loc("flash_attention_npu_v8_copy.py":847:36)
#loc26 = loc("flash_attention_npu_v8_copy.py":848:30)
#loc27 = loc("flash_attention_npu_v8_copy.py":849:34)
#loc28 = loc("flash_attention_npu_v8_copy.py":850:34)
#loc29 = loc("flash_attention_npu_v8_copy.py":851:43)
#loc30 = loc("flash_attention_npu_v8_copy.py":851:33)
#loc31 = loc("flash_attention_npu_v8_copy.py":852:42)
#loc32 = loc("flash_attention_npu_v8_copy.py":852:27)
#loc33 = loc("flash_attention_npu_v8_copy.py":853:24)
#loc34 = loc("flash_attention_npu_v8_copy.py":854:24)
#loc35 = loc("flash_attention_npu_v8_copy.py":855:21)
#loc36 = loc("flash_attention_npu_v8_copy.py":855:31)
#loc37 = loc("flash_attention_npu_v8_copy.py":855:11)
#loc38 = loc("flash_attention_npu_v8_copy.py":856:46)
#loc39 = loc("flash_attention_npu_v8_copy.py":856:31)
#loc40 = loc("flash_attention_npu_v8_copy.py":857:53)
#loc41 = loc("flash_attention_npu_v8_copy.py":857:28)
#loc42 = loc("flash_attention_npu_v8_copy.py":858:28)
#loc43 = loc("flash_attention_npu_v8_copy.py":866:58)
#loc44 = loc("flash_attention_npu_v8_copy.py":866:68)
#loc45 = loc("flash_attention_npu_v8_copy.py":866:78)
#loc46 = loc("flash_attention_npu_v8_copy.py":866:87)
#loc47 = loc("flash_attention_npu_v8_copy.py":875:34)
#loc48 = loc("flash_attention_npu_v8_copy.py":876:34)
#loc49 = loc("flash_attention_npu_v8_copy.py":878:39)
#loc50 = loc("flash_attention_npu_v8_copy.py":878:48)
#loc51 = loc("flash_attention_npu_v8_copy.py":878:68)
#loc52 = loc("flash_attention_npu_v8_copy.py":878:57)
#loc53 = loc("flash_attention_npu_v8_copy.py":880:34)
#loc54 = loc("flash_attention_npu_v8_copy.py":883:16)
#loc55 = loc("flash_attention_npu_v8_copy.py":886:39)
#loc56 = loc("flash_attention_npu_v8_copy.py":886:49)
#loc57 = loc("flash_attention_npu_v8_copy.py":886:70)
#loc58 = loc("flash_attention_npu_v8_copy.py":886:58)
#loc59 = loc("flash_attention_npu_v8_copy.py":902:40)
#loc60 = loc("flash_attention_npu_v8_copy.py":902:49)
#loc61 = loc("flash_attention_npu_v8_copy.py":902:58)
#loc62 = loc("flash_attention_npu_v8_copy.py":907:16)
#loc63 = loc("flash_attention_npu_v8_copy.py":891:16)
#loc64 = loc("flash_attention_npu_v8_copy.py":915:16)
#loc65 = loc("flash_attention_npu_v8_copy.py":959:49)
#loc66 = loc("flash_attention_npu_v8_copy.py":959:61)
#loc67 = loc("flash_attention_npu_v8_copy.py":118:23)
#loc68 = loc("flash_attention_npu_v8_copy.py":970:43)
#loc69 = loc("flash_attention_npu_v8_copy.py":971:43)
#loc70 = loc("flash_attention_npu_v8_copy.py":996:75)
#loc72 = loc("flash_attention_npu_v8_copy.py":1002:26)
#loc73 = loc("flash_attention_npu_v8_copy.py":981:52)
#loc74 = loc("flash_attention_npu_v8_copy.py":982:34)
#loc75 = loc("flash_attention_npu_v8_copy.py":984:54)
#loc76 = loc("flash_attention_npu_v8_copy.py":984:52)
#loc77 = loc("flash_attention_npu_v8_copy.py":984:37)
#loc78 = loc("flash_attention_npu_v8_copy.py":987:34)
#loc79 = loc("flash_attention_npu_v8_copy.py":989:42)
#loc80 = loc("flash_attention_npu_v8_copy.py":989:51)
#loc81 = loc("flash_attention_npu_v8_copy.py":991:36)
#loc82 = loc("flash_attention_npu_v8_copy.py":992:35)
#loc83 = loc("flash_attention_npu_v8_copy.py":992:30)
#loc84 = loc("flash_attention_npu_v8_copy.py":993:44)
#loc85 = loc("flash_attention_npu_v8_copy.py":994:31)
#loc86 = loc("flash_attention_npu_v8_copy.py":996:46)
#loc87 = loc("flash_attention_npu_v8_copy.py":996:57)
#loc88 = loc("flash_attention_npu_v8_copy.py":1000:47)
#loc89 = loc("flash_attention_npu_v8_copy.py":1004:34)
#loc90 = loc("flash_attention_npu_v8_copy.py":1009:42)
#loc91 = loc("flash_attention_npu_v8_copy.py":1009:47)
#loc92 = loc("flash_attention_npu_v8_copy.py":1010:54)
#loc93 = loc("flash_attention_npu_v8_copy.py":1011:56)
#loc94 = loc("flash_attention_npu_v8_copy.py":1012:54)
#loc95 = loc("flash_attention_npu_v8_copy.py":1013:54)
#loc96 = loc("flash_attention_npu_v8_copy.py":1015:16)
#loc97 = loc("flash_attention_npu_v8_copy.py":1018:41)
#loc98 = loc("flash_attention_npu_v8_copy.py":1019:41)
#loc99 = loc("flash_attention_npu_v8_copy.py":1019:56)
#loc100 = loc(callsite(#loc1 at #loc2))
#loc101 = loc(callsite(#loc1 at #loc3))
#loc102 = loc(callsite(#loc1 at #loc5))
#loc103 = loc(callsite(#loc6 at #loc7))
#loc104 = loc(callsite(#loc11 at #loc12))
#loc105 = loc(callsite(#loc67 at #loc68))
#loc106 = loc(callsite(#loc67 at #loc69))
#loc107 = loc(callsite(#loc1 at #loc73))
#loc108 = loc(callsite(#loc6 at #loc99))
likedislike
ascend-robotascend-robot成员
28 天前 添加了label:bug
SSL25成员
27 天前 关联了看板:AscendNPU IR项目
SL25成员
27 天前 评论:

/label add triaged

likedislike
ascend-robotascend-robot成员
27 天前 添加了label:triaged
Hhujiajun成员
24 天前 关联了pull request:feat: promote Hu_JJN to be reviewr of ascendnpu-ir repo