已合并
fix: 修复 FFT 算子 host 侧缓冲区尺寸计算的整数溢出 #33
Tian_1122创建于 14 天前
fix: 修复 FFT 算子 host 侧缓冲区尺寸计算的整数溢出 #33
已合并
共 8 个文件变更+19-19
| @@ -142,8 +142,8 @@ extern "C" aclError aclfftFft1DB(float *x, float *y, uint32_t n, | |||
| 142 | std::vector<uint32_t> radixVec; | 142 | std::vector<uint32_t> radixVec; |
| 143 | InitRadixB(n, radixVec); | 143 | InitRadixB(n, radixVec); |
| 144 | 144 | ||
| 145 | - uint32_t inputSize = n * batches * sizeof(float) * 2; | 145 | + size_t inputSize = static_cast<size_t>(n) * batches * sizeof(float) * 2; |
| 146 | - uint32_t outputSize = inputSize; | 146 | + size_t outputSize = inputSize; |
| 147 | 147 | ||
| 148 | std::vector<float> wMatrixHost; | 148 | std::vector<float> wMatrixHost; |
| 149 | for (size_t i = 0; i < radixVec.size(); i++) { | 149 | for (size_t i = 0; i < radixVec.size(); i++) { |
| @@ -69,7 +69,7 @@ extern "C" aclError aclfftFft1DMix(float *x, float *y, uint32_t n, | |||
| 69 | int64_t totalWs = wsIn + wsOut + wsSync + wsC2c + wsAux; | 69 | int64_t totalWs = wsIn + wsOut + wsSync + wsC2c + wsAux; |
| 70 | 70 | ||
| 71 | // 6. Allocate & copy | 71 | // 6. Allocate & copy |
| 72 | - uint32_t inputSize = n * batches * sizeof(float) * 2; | 72 | + size_t inputSize = static_cast<size_t>(n) * batches * sizeof(float) * 2; |
| 73 | void *dIn=nullptr,*dOut=nullptr,*dDft=nullptr,*dTw=nullptr,*dRadix=nullptr,*dWs=nullptr,*dTil=nullptr; | 73 | void *dIn=nullptr,*dOut=nullptr,*dDft=nullptr,*dTw=nullptr,*dRadix=nullptr,*dWs=nullptr,*dTil=nullptr; |
| 74 | CHECK_ACL(aclrtMalloc(&dIn, inputSize, ACL_MEM_MALLOC_HUGE_FIRST)); | 74 | CHECK_ACL(aclrtMalloc(&dIn, inputSize, ACL_MEM_MALLOC_HUGE_FIRST)); |
| 75 | CHECK_ACL(aclrtMalloc(&dOut, inputSize, ACL_MEM_MALLOC_HUGE_FIRST)); | 75 | CHECK_ACL(aclrtMalloc(&dOut, inputSize, ACL_MEM_MALLOC_HUGE_FIRST)); |
| @@ -265,8 +265,8 @@ extern "C" aclError aclfftFft1DN(float *x, float *y, uint32_t n, int32_t norm, | |||
| 265 | } | 265 | } |
| 266 | int32_t repeatBatchSize = 0; | 266 | int32_t repeatBatchSize = 0; |
| 267 | ComputeRepeatBatchSize(n, batches, repeatBatchSize); | 267 | ComputeRepeatBatchSize(n, batches, repeatBatchSize); |
| 268 | - const uint32_t inputSize = n * batches * sizeof(float) * 2; | 268 | + const size_t inputSize = static_cast<size_t>(n) * batches * sizeof(float) * 2; |
| 269 | - const uint32_t outputSize = inputSize; | 269 | + const size_t outputSize = inputSize; |
| 270 | std::vector<float> wMatrixHost; | 270 | std::vector<float> wMatrixHost; |
| 271 | for (size_t i = 0; i < radixVec.size(); i++) { | 271 | for (size_t i = 0; i < radixVec.size(); i++) { |
| 272 | uint32_t radix = radixVec[i]; | 272 | uint32_t radix = radixVec[i]; |
| @@ -131,13 +131,13 @@ aclError aclfftFft1DStride(float *x, float *y, uint32_t n, uint32_t stride, | |||
| 131 | 131 | ||
| 132 | uint32_t s0 = ComputeS0(n, stride); | 132 | uint32_t s0 = ComputeS0(n, stride); |
| 133 | 133 | ||
| 134 | - uint32_t inputSize = n * stride * sizeof(float) * 2; | 134 | + size_t inputSize = static_cast<size_t>(n) * stride * sizeof(float) * 2; |
| 135 | uint32_t sMatrixSize = sMatrixHost.size() * sizeof(float); | 135 | uint32_t sMatrixSize = sMatrixHost.size() * sizeof(float); |
| 136 | - uint32_t outputSize = inputSize; | 136 | + size_t outputSize = inputSize; |
| 137 | - uint32_t kernelWorkspaceSize = n * s0 * sizeof(float) * 4; | 137 | + size_t kernelWorkspaceSize = static_cast<size_t>(n) * s0 * sizeof(float) * 4; |
| 138 | uint32_t tilingSize = sizeof(Fft1DStrideTilingData); | 138 | uint32_t tilingSize = sizeof(Fft1DStrideTilingData); |
| 139 | uint32_t sysWorkspaceSize = ascendcPlatform->GetLibApiWorkSpaceSize(); | 139 | uint32_t sysWorkspaceSize = ascendcPlatform->GetLibApiWorkSpaceSize(); |
| 140 | - uint32_t totalWorkspaceSize = kernelWorkspaceSize + sysWorkspaceSize; | 140 | + size_t totalWorkspaceSize = kernelWorkspaceSize + sysWorkspaceSize; |
| 141 | 141 | ||
| 142 | void *dev_input = nullptr; | 142 | void *dev_input = nullptr; |
| 143 | void *dev_s_matrix = nullptr; | 143 | void *dev_s_matrix = nullptr; |
| @@ -76,15 +76,15 @@ extern "C" aclError aclfftFft2DDd(float *x, float *y, uint32_t fftX, uint32_t ff | |||
| 76 | return ACL_ERROR_INVALID_PARAM; | 76 | return ACL_ERROR_INVALID_PARAM; |
| 77 | } | 77 | } |
| 78 | 78 | ||
| 79 | - uint32_t inputSize = batches * fftX * fftY * sizeof(float) * 2; | 79 | + size_t inputSize = static_cast<size_t>(batches) * fftX * fftY * sizeof(float) * 2; |
| 80 | - uint32_t outputSize = inputSize; | 80 | + size_t outputSize = inputSize; |
| 81 | 81 | ||
| 82 | std::vector<float> pMatrixHost = InitPQMatrix(fftX, isForward != 0, true); | 82 | std::vector<float> pMatrixHost = InitPQMatrix(fftX, isForward != 0, true); |
| 83 | std::vector<float> qMatrixHost = InitPQMatrix(fftY, isForward != 0, false); | 83 | std::vector<float> qMatrixHost = InitPQMatrix(fftY, isForward != 0, false); |
| 84 | uint32_t pMatrixSize = pMatrixHost.size() * sizeof(float); | 84 | uint32_t pMatrixSize = pMatrixHost.size() * sizeof(float); |
| 85 | uint32_t qMatrixSize = qMatrixHost.size() * sizeof(float); | 85 | uint32_t qMatrixSize = qMatrixHost.size() * sizeof(float); |
| 86 | 86 | ||
| 87 | - uint64_t workspaceSize = batches * fftX * fftY * sizeof(float) * 2; | 87 | + uint64_t workspaceSize = static_cast<uint64_t>(batches) * fftX * fftY * sizeof(float) * 2; |
| 88 | 88 | ||
| 89 | void *dev_input = nullptr; | 89 | void *dev_input = nullptr; |
| 90 | void *dev_output = nullptr; | 90 | void *dev_output = nullptr; |
| @@ -172,8 +172,8 @@ extern "C" aclError aclfftIrfft1DC2RFft(float *x, float *y, uint32_t n, | |||
| 172 | int64_t totalWs = wsIn + wsOut + wsSync + wsC2c + wsAux; | 172 | int64_t totalWs = wsIn + wsOut + wsSync + wsC2c + wsAux; |
| 173 | 173 | ||
| 174 | // 7. Allocate & copy | 174 | // 7. Allocate & copy |
| 175 | - uint32_t inputSize = (n / 2 + 1) * batches * sizeof(float) * 2; | 175 | + size_t inputSize = static_cast<size_t>(n / 2 + 1) * batches * sizeof(float) * 2; |
| 176 | - uint32_t outputSize = n * batches * sizeof(float); | 176 | + size_t outputSize = static_cast<size_t>(n) * batches * sizeof(float); |
| 177 | void *dIn=nullptr,*dOut=nullptr,*dDft=nullptr,*dTw=nullptr,*dRadix=nullptr,*dWs=nullptr,*dTil=nullptr; | 177 | void *dIn=nullptr,*dOut=nullptr,*dDft=nullptr,*dTw=nullptr,*dRadix=nullptr,*dWs=nullptr,*dTil=nullptr; |
| 178 | void *dInIdx=nullptr,*dA=nullptr,*dB=nullptr,*dOutIdx=nullptr; | 178 | void *dInIdx=nullptr,*dA=nullptr,*dB=nullptr,*dOutIdx=nullptr; |
| 179 | CHECK_ACL(aclrtMalloc(&dIn, inputSize, ACL_MEM_MALLOC_HUGE_FIRST)); | 179 | CHECK_ACL(aclrtMalloc(&dIn, inputSize, ACL_MEM_MALLOC_HUGE_FIRST)); |
| @@ -141,12 +141,12 @@ extern "C" aclError aclfftIrfft1DFft(float *x, float *y, uint32_t n, int32_t nor | |||
| 141 | tempN = M; | 141 | tempN = M; |
| 142 | } | 142 | } |
| 143 | 143 | ||
| 144 | - const uint32_t inputSize = batches * (n / 2 + 1) * sizeof(float) * 2; | 144 | + const size_t inputSize = static_cast<size_t>(batches) * (n / 2 + 1) * sizeof(float) * 2; |
| 145 | - const uint32_t outputSize = batches * n * sizeof(float); | 145 | + const size_t outputSize = static_cast<size_t>(batches) * n * sizeof(float); |
| 146 | const uint32_t fftMatrixSize = allfftMatrices.size() * sizeof(float); | 146 | const uint32_t fftMatrixSize = allfftMatrices.size() * sizeof(float); |
| 147 | const uint32_t twSize = allTwiddleFactors.size() * sizeof(float); | 147 | const uint32_t twSize = allTwiddleFactors.size() * sizeof(float); |
| 148 | const uint32_t radixListSize = tilingData.radixListLen * sizeof(float); | 148 | const uint32_t radixListSize = tilingData.radixListLen * sizeof(float); |
| 149 | - const uint32_t workspaceSize = 2 * batches * n * sizeof(float) * 2; | 149 | + const size_t workspaceSize = 2 * static_cast<size_t>(batches) * n * sizeof(float) * 2; |
| 150 | const uint32_t tilingSize = sizeof(Irfft1DfftTilingData); | 150 | const uint32_t tilingSize = sizeof(Irfft1DfftTilingData); |
| 151 | 151 | ||
| 152 | void *dev_input = nullptr; | 152 | void *dev_input = nullptr; |
| @@ -146,8 +146,8 @@ extern "C" aclError aclfftRfft1DR2CFft(float *x, float *y, uint32_t n, | |||
| 146 | 146 | ||
| 147 | // 7. Allocate & copy | 147 | // 7. Allocate & copy |
| 148 | // R2C: input is n real floats, output is (n/2+1) complex = (n/2+1)*2 floats | 148 | // R2C: input is n real floats, output is (n/2+1) complex = (n/2+1)*2 floats |
| 149 | - uint32_t inputSize = n * batches * sizeof(float); | 149 | + size_t inputSize = static_cast<size_t>(n) * batches * sizeof(float); |
| 150 | - uint32_t outputSize = (n / 2 + 1) * batches * sizeof(float) * 2; | 150 | + size_t outputSize = static_cast<size_t>(n / 2 + 1) * batches * sizeof(float) * 2; |
| 151 | void *dIn=nullptr,*dOut=nullptr,*dDft=nullptr,*dTw=nullptr,*dRadix=nullptr,*dWs=nullptr,*dTil=nullptr; | 151 | void *dIn=nullptr,*dOut=nullptr,*dDft=nullptr,*dTw=nullptr,*dRadix=nullptr,*dWs=nullptr,*dTil=nullptr; |
| 152 | void *dInIdx=nullptr,*dA=nullptr,*dB=nullptr,*dOutIdx=nullptr; | 152 | void *dInIdx=nullptr,*dA=nullptr,*dB=nullptr,*dOutIdx=nullptr; |
| 153 | CHECK_ACL(aclrtMalloc(&dIn, inputSize, ACL_MEM_MALLOC_HUGE_FIRST)); | 153 | CHECK_ACL(aclrtMalloc(&dIn, inputSize, ACL_MEM_MALLOC_HUGE_FIRST)); |