主干 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))
/label add triaged
主干 dd3198ac66f277b9165f2e817d3c708b8615b42c 手动打开多 consumer 复现