已合并
fix: 修复 FFT 算子 host 侧缓冲区尺寸计算的整数溢出 #33
fix: 修复 FFT 算子 host 侧缓冲区尺寸计算的整数溢出 #33
已合并
Tian_1122创建于 14 天前
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 & copy71 // 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 & copy174 // 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 & copy147 // 7. Allocate & copy
148 // R2C: input is n real floats, output is (n/2+1) complex = (n/2+1)*2 floats148 // 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));