Line data Source code
1 : /**
2 : * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3 : * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4 : * CANN Open Software License Agreement Version 2.0 (the "License").
5 : * Please refer to the License for details. You may not use this file except in compliance with the License.
6 : * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7 : * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8 : * See LICENSE in the root of the software repository for the full text of the License.
9 : */
10 :
11 : #ifndef __KFC_DUMP_INF_H__
12 : #define __KFC_DUMP_INF_H__
13 :
14 : #include "kfc_dump_base.h"
15 :
16 : namespace KfcDumpStat {
17 :
18 : template <typename T, bool isStatPosInf>
19 344 : __aicore__ inline void ComputeInfSingleType(
20 : TQue<QuePosition::VECOUT, BUFFER_NUM>& workQueue, TBuf<TPosition::VECCALC>& maskBuf,
21 : TBuf<TPosition::VECCALC>& cacheBuf1, LocalTensor<float>& procesLocal, LocalTensor<uint8_t>& calMidAddr,
22 : int64_t curLoop, int64_t curProcessCount)
23 : {
24 344 : LocalTensor<uint8_t> curInfAddr = calMidAddr[sizeof(float)]; // 四个字节偏移作为当前最大值存放的地址
25 344 : LocalTensor<uint8_t> compareResult = maskBuf.Get<uint8_t>();
26 344 : LocalTensor<uint8_t> workLocal = workQueue.AllocTensor<uint8_t>();
27 344 : LocalTensor<int16_t> cacheTensor = cacheBuf1.Get<int16_t>();
28 :
29 344 : Compare(compareResult, procesLocal, procesLocal, CMPMODE::EQ, curProcessCount);
30 :
31 344 : float inputVal(0);
32 344 : LocalTensor<float> cacluateLocal = workLocal.ReinterpretCast<float>();
33 344 : LocalTensor<float> selectLocal = cacheTensor.ReinterpretCast<float>();
34 344 : Duplicate<float>(cacluateLocal, inputVal, curProcessCount);
35 344 : PipeBarrier<PIPE_ALL>();
36 :
37 344 : Select(
38 : selectLocal, compareResult, cacluateLocal, static_cast<float>(1), SELMODE::VSEL_TENSOR_SCALAR_MODE,
39 : curProcessCount);
40 344 : PipeBarrier<PIPE_ALL>();
41 :
42 344 : ReduceSum<float>(selectLocal, selectLocal, selectLocal, curProcessCount);
43 344 : PipeBarrier<PIPE_ALL>();
44 :
45 344 : int32_t nanNum = static_cast<int32_t>(selectLocal.GetValue(0));
46 : if (isStatPosInf) {
47 172 : CompareScalar(compareResult, procesLocal, FP_INF, CMPMODE::LT, curProcessCount);
48 : } else {
49 172 : CompareScalar(compareResult, procesLocal, -FP_INF, CMPMODE::GT, curProcessCount);
50 : }
51 344 : Duplicate<float>(cacluateLocal, inputVal, curProcessCount);
52 344 : PipeBarrier<PIPE_ALL>();
53 :
54 344 : Select(
55 : selectLocal, compareResult, cacluateLocal, static_cast<float>(1), SELMODE::VSEL_TENSOR_SCALAR_MODE,
56 : curProcessCount);
57 344 : PipeBarrier<PIPE_ALL>();
58 :
59 344 : ReduceSum<float>(selectLocal, selectLocal, selectLocal, curProcessCount);
60 344 : PipeBarrier<PIPE_ALL>();
61 :
62 344 : int32_t infNum = static_cast<int32_t>(selectLocal.GetValue(0) - nanNum);
63 :
64 : // 更新最终结果
65 344 : LocalTensor<int32_t> curInf = curInfAddr.ReinterpretCast<int32_t>();
66 344 : if (curLoop == 0) {
67 340 : curInf.SetValue(0, infNum);
68 : } else {
69 4 : curInf.SetValue(0, curInf.GetValue(0) + infNum);
70 : }
71 344 : PipeBarrier<PIPE_ALL>();
72 :
73 344 : workQueue.FreeTensor(workLocal);
74 344 : }
75 :
76 : // inf 统计计算回调:供 ComputeFloatStatSkeleton 在 cast 完成后调用
77 : template <typename T, bool isStatPosInf>
78 : struct InfComputeFunc {
79 344 : __aicore__ inline void operator()(
80 : TQue<QuePosition::VECOUT, BUFFER_NUM>& workQue, TBuf<TPosition::VECCALC>& mask,
81 : TBuf<TPosition::VECCALC>& cache1, LocalTensor<float>& inputLocal, LocalTensor<uint8_t>& midAddr, int64_t loop,
82 : int64_t processCount) const
83 : {
84 344 : ComputeInfSingleType<T, isStatPosInf>(workQue, mask, cache1, inputLocal, midAddr, loop, processCount);
85 344 : }
86 : };
87 :
88 : template <typename T, bool isStatPosInf>
89 398 : __aicore__ inline void ComputeInfMultiDataType(
90 : TQue<QuePosition::VECIN, BUFFER_NUM>& xQue, TBuf<TPosition::VECCALC>& castXBuf,
91 : TQue<QuePosition::VECOUT, BUFFER_NUM>& workQueue, TBuf<TPosition::VECCALC>& maskBuf,
92 : TBuf<TPosition::VECCALC>& cacheBuf1, LocalTensor<uint8_t>& calMidAddr, int64_t curLoop, int64_t curProcessCount,
93 : int64_t appendNum)
94 : {
95 398 : ComputeFloatStatSkeleton<T>(
96 : xQue, castXBuf, workQueue, maskBuf, cacheBuf1, calMidAddr, curLoop, curProcessCount, appendNum,
97 : InfComputeFunc<T, isStatPosInf>());
98 398 : }
99 :
100 : template <typename T, bool isStatPosInf>
101 376 : __aicore__ inline void ProcessInf(
102 : TQue<QuePosition::VECIN, BUFFER_NUM>& xQue, TBuf<TPosition::VECCALC>& calMiddleBuf,
103 : TBuf<TPosition::VECCALC>& castXBuf, TQue<QuePosition::VECOUT, BUFFER_NUM>& workQueue, GlobalTensor<T>& xGm,
104 : TBuf<TPosition::VECCALC>& maskBuf, TBuf<TPosition::VECCALC>& cacheBuf1, int64_t innerLoopTime,
105 : int64_t tileLengthMean, int64_t tileLengthEnd, int64_t perBlockCount, int64_t blockOffset, int64_t tileNumEnd,
106 : uint64_t xDtypeSize)
107 : {
108 376 : StatClass statClass = isStatPosInf ? StatClass::STAT_POS_INF : StatClass::STAT_NEG_INF;
109 376 : LocalTensor<uint8_t> calMidAddr = calMiddleBuf.Get<uint8_t>()[static_cast<int64_t>(statClass) * BLOCK_SIZE];
110 376 : LocalTensor<uint8_t> curInfAddr = calMidAddr[sizeof(float)];
111 :
112 400 : for (int64_t curLoopTime = 0; curLoopTime < innerLoopTime; ++curLoopTime) {
113 24 : CopyInX(xQue, xGm, blockOffset + curLoopTime * tileLengthMean, tileLengthMean, perBlockCount);
114 24 : int64_t appendNum = 0;
115 24 : ComputeInfMultiDataType<T, isStatPosInf>(
116 : xQue, castXBuf, workQueue, maskBuf, cacheBuf1, calMidAddr, curLoopTime, tileLengthMean, appendNum);
117 : }
118 :
119 376 : if (tileNumEnd) {
120 374 : CopyInX(xQue, xGm, blockOffset + innerLoopTime * tileLengthMean, tileLengthEnd, perBlockCount);
121 374 : auto appendNum = CeilAlign(tileLengthEnd * xDtypeSize, COMMAND_SIZE) / xDtypeSize - tileLengthEnd;
122 374 : ComputeInfMultiDataType<T, isStatPosInf>(
123 374 : xQue, castXBuf, workQueue, maskBuf, cacheBuf1, calMidAddr, innerLoopTime, tileLengthEnd + appendNum,
124 : appendNum);
125 : }
126 :
127 376 : WriteStatOutput<int32_t>(calMidAddr, curInfAddr);
128 376 : }
129 :
130 : } // namespace KfcDumpStat
131 :
132 : #endif // __KFC_DUMP_INF_H__
|