已合并
修复conv2d dx PreDilation PostDilation问题 #3953
jiangqi创建于 4月17日
修复conv2d dx PreDilation PostDilation问题 #3953
已合并
jiangqi创建于 4月17日
2 个文件变更+62-46
@@ -636,17 +636,6 @@ static const aclTensor *PreDilation(ConvolutionBackwardInputTensor &inputTensor,
636 aclOpExecutor *executor) {636 aclOpExecutor *executor) {
637 const aclTensor *preDilationGradOutputNC1HWC0 = nullptr;637 const aclTensor *preDilationGradOutputNC1HWC0 = nullptr;
638 int64_t preDilationDilationsVector[] = {1, 1, (*params.stride)[kSTRIDEHIdx], (*params.stride)[kSTRIDEWIdx], 1};638 int64_t preDilationDilationsVector[] = {1, 1, (*params.stride)[kSTRIDEHIdx], (*params.stride)[kSTRIDEWIdx], 1};
639- /* 当输出的h/w为1时,把传给Dilation算子的dilation值修正为fmap_h/w,避免在dilation值超大时,Dilation算子超时 */
640- if (inputTensor.gradOutput->GetViewShape().GetDim(kHDimNCHWIdx) == 1 &&
641- (*params.stride)[kSTRIDEHIdx] > inputTensor.input->GetViewShape().GetDim(kHDimNCHWIdx)) {
642- preDilationDilationsVector[kHDimNC1HWC0Idx] = inputTensor.input->GetViewShape().GetDim(kHDimNCHWIdx);
643- }
644- if (inputTensor.gradOutput->GetViewShape().GetDim(kWDimNCHWIdx) == 1 &&
645- (*params.stride)[kSTRIDEWIdx] > inputTensor.input->GetViewShape().GetDim(kWDimNCHWIdx)) {
646- preDilationDilationsVector[kWDimNC1HWC0Idx] = inputTensor.input->GetViewShape().GetDim(kWDimNCHWIdx);
647- }
648- aclIntArray *preDilationDilations = executor->AllocIntArray(preDilationDilationsVector, 5);
649- CHECK_RET(preDilationDilations != nullptr, nullptr);
650 int64_t preDilationPadUp = 0;639 int64_t preDilationPadUp = 0;
651 int64_t preDilationPadLeft = 0;640 int64_t preDilationPadLeft = 0;
652 int64_t preDilationPadH = 0;641 int64_t preDilationPadH = 0;
@@ -658,6 +647,19 @@ static const aclTensor *PreDilation(ConvolutionBackwardInputTensor &inputTensor,
658 preDilationPadH = 2 * (*params.padding)[kPADDINGUPIdx];647 preDilationPadH = 2 * (*params.padding)[kPADDINGUPIdx];
659 preDilationPadW = 2 * (*params.padding)[kPADDINGLEFTIdx];648 preDilationPadW = 2 * (*params.padding)[kPADDINGLEFTIdx];
660 }649 }
650+ 
651+ /* 当输出的h/w为1时,把传给Dilation算子的dilation值修正为fmap_h/w + pad,避免在dilation值超大时,Dilation算子超时 */
652+ if (inputTensor.gradOutput->GetViewShape().GetDim(kHDimNCHWIdx) == 1 &&
653+ (*params.stride)[kSTRIDEHIdx] > inputTensor.input->GetViewShape().GetDim(kHDimNCHWIdx) + preDilationPadH) {
654+ preDilationDilationsVector[kHDimNC1HWC0Idx] = inputTensor.input->GetViewShape().GetDim(kHDimNCHWIdx) + preDilationPadH;
655+ }
656+ if (inputTensor.gradOutput->GetViewShape().GetDim(kWDimNCHWIdx) == 1 &&
657+ (*params.stride)[kSTRIDEWIdx] > inputTensor.input->GetViewShape().GetDim(kWDimNCHWIdx) + preDilationPadW) {
658+ preDilationDilationsVector[kWDimNC1HWC0Idx] = inputTensor.input->GetViewShape().GetDim(kWDimNCHWIdx) + preDilationPadW;
659+ }
660+ aclIntArray *preDilationDilations = executor->AllocIntArray(preDilationDilationsVector, 5);
661+ CHECK_RET(preDilationDilations != nullptr, nullptr);
662+ 
661 int64_t preDilationPadDown =663 int64_t preDilationPadDown =
662 inputTensor.input->GetViewShape().GetDim(kHDimNCHWIdx) -664 inputTensor.input->GetViewShape().GetDim(kHDimNCHWIdx) -
663 (inputTensor.weight->GetViewShape().GetDim(kHDimNCHWIdx) - 1) * (*params.dilation)[kDILATIONHIdx] +665 (inputTensor.weight->GetViewShape().GetDim(kHDimNCHWIdx) - 1) * (*params.dilation)[kDILATIONHIdx] +
@@ -684,33 +686,39 @@ static const aclTensor *PostDilation(const aclTensor *dxGradInputNC1HWC0, Convol
684 ConvolutionBackwardParams &params, aclOpExecutor *executor) {686 ConvolutionBackwardParams &params, aclOpExecutor *executor) {
685 const aclTensor *postDilationGradInputNC1HWC0 = nullptr;687 const aclTensor *postDilationGradInputNC1HWC0 = nullptr;
686 int64_t postDilationDilationsVector[] = {1, 1, (*params.stride)[kSTRIDEHIdx], (*params.stride)[kSTRIDEWIdx], 1};688 int64_t postDilationDilationsVector[] = {1, 1, (*params.stride)[kSTRIDEHIdx], (*params.stride)[kSTRIDEWIdx], 1};
687- /* 当输出的h/w为1时,把传给Dilation算子的dilation值修正为fmap_h/w,避免在dilation值超大时,Dilation算子超时 */
688- if (inputTensor.gradOutput->GetViewShape().GetDim(kHDimNCHWIdx) == 1 &&
689- (*params.stride)[kSTRIDEHIdx] > inputTensor.input->GetViewShape().GetDim(kHDimNCHWIdx)) {
690- postDilationDilationsVector[kHDimNC1HWC0Idx] = inputTensor.input->GetViewShape().GetDim(kHDimNCHWIdx);
691- }
692- if (inputTensor.gradOutput->GetViewShape().GetDim(kWDimNCHWIdx) == 1 &&
693- (*params.stride)[kSTRIDEWIdx] > inputTensor.input->GetViewShape().GetDim(kWDimNCHWIdx)) {
694- postDilationDilationsVector[kWDimNC1HWC0Idx] = inputTensor.input->GetViewShape().GetDim(kWDimNCHWIdx);
695- }
696- aclIntArray *post_dilation_dilations = executor->AllocIntArray(postDilationDilationsVector, 5);
697- CHECK_RET(post_dilation_dilations != nullptr, nullptr);
698- 
699 int64_t postDilationPadUp = 0;689 int64_t postDilationPadUp = 0;
700 int64_t postDilationPadLeft = 0;690 int64_t postDilationPadLeft = 0;
701 int64_t padUp = 0;691 int64_t padUp = 0;
702 int64_t padLeft = 0;692 int64_t padLeft = 0;
693+ int64_t padH = 0;
694+ int64_t padW = 0;
703 if (params.padding->Size() == 4) {695 if (params.padding->Size() == 4) {
704 postDilationPadUp = -(*params.padding)[kPadding4UpIdx];696 postDilationPadUp = -(*params.padding)[kPadding4UpIdx];
705 postDilationPadLeft = -(*params.padding)[kPadding4LeftIdx];697 postDilationPadLeft = -(*params.padding)[kPadding4LeftIdx];
706 padUp = (*params.padding)[kPadding4UpIdx];698 padUp = (*params.padding)[kPadding4UpIdx];
707 padLeft = (*params.padding)[kPadding4LeftIdx];699 padLeft = (*params.padding)[kPadding4LeftIdx];
700+ padH = (*params.padding)[kPadding4UpIdx] + (*params.padding)[kPadding4DownIdx];
701+ padW = (*params.padding)[kPadding4LeftIdx] + (*params.padding)[kPadding4RightIdx];
708 } else {702 } else {
709 postDilationPadUp = -(*params.padding)[kPADDINGUPIdx];703 postDilationPadUp = -(*params.padding)[kPADDINGUPIdx];
710 postDilationPadLeft = -(*params.padding)[kPADDINGLEFTIdx];704 postDilationPadLeft = -(*params.padding)[kPADDINGLEFTIdx];
711 padUp = (*params.padding)[kPADDINGUPIdx];705 padUp = (*params.padding)[kPADDINGUPIdx];
712 padLeft = (*params.padding)[kPADDINGLEFTIdx];706 padLeft = (*params.padding)[kPADDINGLEFTIdx];
707+ padH = 2 * (*params.padding)[kPADDINGUPIdx];
708+ padW = 2 * (*params.padding)[kPADDINGLEFTIdx];
713 }709 }
710+ /* 当输出的h/w为1时,把传给Dilation算子的dilation值修正为fmap_h/w + pad,避免在dilation值超大时,Dilation算子超时 */
711+ if (inputTensor.gradOutput->GetViewShape().GetDim(kHDimNCHWIdx) == 1 &&
712+ (*params.stride)[kSTRIDEHIdx] > inputTensor.input->GetViewShape().GetDim(kHDimNCHWIdx) + padH) {
713+ postDilationDilationsVector[kHDimNC1HWC0Idx] = inputTensor.input->GetViewShape().GetDim(kHDimNCHWIdx) + padH;
714+ }
715+ if (inputTensor.gradOutput->GetViewShape().GetDim(kWDimNCHWIdx) == 1 &&
716+ (*params.stride)[kSTRIDEWIdx] > inputTensor.input->GetViewShape().GetDim(kWDimNCHWIdx) + padW) {
717+ postDilationDilationsVector[kWDimNC1HWC0Idx] = inputTensor.input->GetViewShape().GetDim(kWDimNCHWIdx) + padW;
718+ }
719+ aclIntArray *post_dilation_dilations = executor->AllocIntArray(postDilationDilationsVector, 5);
720+ CHECK_RET(post_dilation_dilations != nullptr, nullptr);
721+ 
714 int64_t postDilationPadDown =722 int64_t postDilationPadDown =
715 inputTensor.input->GetViewShape().GetDim(kHDimNCHWIdx) + padUp -723 inputTensor.input->GetViewShape().GetDim(kHDimNCHWIdx) + padUp -
716 (inputTensor.gradOutput->GetViewShape().GetDim(kHDimNCHWIdx) - 1) * (*params.stride)[kSTRIDEHIdx] - 1;724 (inputTensor.gradOutput->GetViewShape().GetDim(kHDimNCHWIdx) - 1) * (*params.stride)[kSTRIDEHIdx] - 1;
@@ -105,17 +105,6 @@ static const aclTensor *PreDilation(ConvolutionBackwardInputTensorForAvgPool2d &
105 aclOpExecutor *executor) {105 aclOpExecutor *executor) {
106 const aclTensor *preDilationGradOutputNC1HWC0 = nullptr;106 const aclTensor *preDilationGradOutputNC1HWC0 = nullptr;
107 int64_t preDilationDilationsVector[] = {1, 1, (*params.stride)[kSTRIDEHIdx], (*params.stride)[kSTRIDEWIdx], 1};107 int64_t preDilationDilationsVector[] = {1, 1, (*params.stride)[kSTRIDEHIdx], (*params.stride)[kSTRIDEWIdx], 1};
108- /* 当输出的h/w为1时,把传给Dilation算子的dilation值修正为fmap_h/w,避免在dilation值超大时,Dilation算子超时 */
109- if (inputTensor.gradOutput->GetViewShape().GetDim(kHDimNCHWIdx) == 1 &&
110- (*params.stride)[kSTRIDEHIdx] > inputTensor.input->GetViewShape().GetDim(kHDimNCHWIdx)) {
111- preDilationDilationsVector[kHDimNC1HWC0Idx] = inputTensor.input->GetViewShape().GetDim(kHDimNCHWIdx);
112- }
113- if (inputTensor.gradOutput->GetViewShape().GetDim(kWDimNCHWIdx) == 1 &&
114- (*params.stride)[kSTRIDEWIdx] > inputTensor.input->GetViewShape().GetDim(kWDimNCHWIdx)) {
115- preDilationDilationsVector[kWDimNC1HWC0Idx] = inputTensor.input->GetViewShape().GetDim(kWDimNCHWIdx);
116- }
117- aclIntArray *preDilationDilations = executor->AllocIntArray(preDilationDilationsVector, 5);
118- CHECK_RET(preDilationDilations != nullptr, nullptr);
119 int64_t preDilationPadUp = 0;108 int64_t preDilationPadUp = 0;
120 int64_t preDilationPadLeft = 0;109 int64_t preDilationPadLeft = 0;
121 int64_t preDilationPadH = 0;110 int64_t preDilationPadH = 0;
@@ -127,6 +116,19 @@ static const aclTensor *PreDilation(ConvolutionBackwardInputTensorForAvgPool2d &
127 preDilationPadH = NUM2 * (*params.padding)[kPADDINGUPIdx];116 preDilationPadH = NUM2 * (*params.padding)[kPADDINGUPIdx];
128 preDilationPadW = NUM2 * (*params.padding)[kPADDINGLEFTIdx];117 preDilationPadW = NUM2 * (*params.padding)[kPADDINGLEFTIdx];
129 }118 }
119+ 
120+ /* 当输出的h/w为1时,把传给Dilation算子的dilation值修正为fmap_h/w + pad,避免在dilation值超大时,Dilation算子超时 */
121+ if (inputTensor.gradOutput->GetViewShape().GetDim(kHDimNCHWIdx) == 1 &&
122+ (*params.stride)[kSTRIDEHIdx] > inputTensor.input->GetViewShape().GetDim(kHDimNCHWIdx) + preDilationPadH) {
123+ preDilationDilationsVector[kHDimNC1HWC0Idx] = inputTensor.input->GetViewShape().GetDim(kHDimNCHWIdx) + preDilationPadH;
124+ }
125+ if (inputTensor.gradOutput->GetViewShape().GetDim(kWDimNCHWIdx) == 1 &&
126+ (*params.stride)[kSTRIDEWIdx] > inputTensor.input->GetViewShape().GetDim(kWDimNCHWIdx) + preDilationPadW) {
127+ preDilationDilationsVector[kWDimNC1HWC0Idx] = inputTensor.input->GetViewShape().GetDim(kWDimNCHWIdx) + preDilationPadW;
128+ }
129+ aclIntArray *preDilationDilations = executor->AllocIntArray(preDilationDilationsVector, 5);
130+ CHECK_RET(preDilationDilations != nullptr, nullptr);
131+ 
130 int64_t preDilationPadDown =132 int64_t preDilationPadDown =
131 inputTensor.input->GetViewShape().GetDim(kHDimNCHWIdx) -133 inputTensor.input->GetViewShape().GetDim(kHDimNCHWIdx) -
132 (inputTensor.weight->GetViewShape().GetDim(kHDimNCHWIdx) - 1) * (*params.dilation)[kDILATIONHIdx] +134 (inputTensor.weight->GetViewShape().GetDim(kHDimNCHWIdx) - 1) * (*params.dilation)[kDILATIONHIdx] +
@@ -223,33 +225,39 @@ static const aclTensor *PostDilation(const aclTensor *dxGradInputNC1HWC0, Convol
223 ConvolutionBackwardParamsForAvgPool2d &params, aclOpExecutor *executor) {225 ConvolutionBackwardParamsForAvgPool2d &params, aclOpExecutor *executor) {
224 const aclTensor *postDilationGradInputNC1HWC0 = nullptr;226 const aclTensor *postDilationGradInputNC1HWC0 = nullptr;
225 int64_t postDilationDilationsVector[] = {1, 1, (*params.stride)[kSTRIDEHIdx], (*params.stride)[kSTRIDEWIdx], 1};227 int64_t postDilationDilationsVector[] = {1, 1, (*params.stride)[kSTRIDEHIdx], (*params.stride)[kSTRIDEWIdx], 1};
226- /* 当输出的h/w为1时,把传给Dilation算子的dilation值修正为fmap_h/w,避免在dilation值超大时,Dilation算子超时 */
227- if (inputTensor.gradOutput->GetViewShape().GetDim(kHDimNCHWIdx) == 1 &&
228- (*params.stride)[kSTRIDEHIdx] > inputTensor.input->GetViewShape().GetDim(kHDimNCHWIdx)) {
229- postDilationDilationsVector[kHDimNC1HWC0Idx] = inputTensor.input->GetViewShape().GetDim(kHDimNCHWIdx);
230- }
231- if (inputTensor.gradOutput->GetViewShape().GetDim(kWDimNCHWIdx) == 1 &&
232- (*params.stride)[kSTRIDEWIdx] > inputTensor.input->GetViewShape().GetDim(kWDimNCHWIdx)) {
233- postDilationDilationsVector[kWDimNC1HWC0Idx] = inputTensor.input->GetViewShape().GetDim(kWDimNCHWIdx);
234- }
235- aclIntArray *post_dilation_dilations = executor->AllocIntArray(postDilationDilationsVector, 5);
236- CHECK_RET(post_dilation_dilations != nullptr, nullptr);
237- 
238 int64_t postDilationPadUp = 0;228 int64_t postDilationPadUp = 0;
239 int64_t postDilationPadLeft = 0;229 int64_t postDilationPadLeft = 0;
240 int64_t padUp = 0;230 int64_t padUp = 0;
241 int64_t padLeft = 0;231 int64_t padLeft = 0;
232+ int64_t padH = 0;
233+ int64_t padW = 0;
242 if (params.padding->Size() == NUM4) {234 if (params.padding->Size() == NUM4) {
243 postDilationPadUp = -(*params.padding)[kPadding4UpIdx];235 postDilationPadUp = -(*params.padding)[kPadding4UpIdx];
244 postDilationPadLeft = -(*params.padding)[kPadding4LeftIdx];236 postDilationPadLeft = -(*params.padding)[kPadding4LeftIdx];
245 padUp = (*params.padding)[kPadding4UpIdx];237 padUp = (*params.padding)[kPadding4UpIdx];
246 padLeft = (*params.padding)[kPadding4LeftIdx];238 padLeft = (*params.padding)[kPadding4LeftIdx];
239+ padH = (*params.padding)[kPadding4UpIdx] + (*params.padding)[kPadding4DownIdx];
240+ padW = (*params.padding)[kPadding4LeftIdx] + (*params.padding)[kPadding4RightIdx];
247 } else {241 } else {
248 postDilationPadUp = -(*params.padding)[kPADDINGUPIdx];242 postDilationPadUp = -(*params.padding)[kPADDINGUPIdx];
249 postDilationPadLeft = -(*params.padding)[kPADDINGLEFTIdx];243 postDilationPadLeft = -(*params.padding)[kPADDINGLEFTIdx];
250 padUp = (*params.padding)[kPADDINGUPIdx];244 padUp = (*params.padding)[kPADDINGUPIdx];
251 padLeft = (*params.padding)[kPADDINGLEFTIdx];245 padLeft = (*params.padding)[kPADDINGLEFTIdx];
246+ padH = NUM2 * (*params.padding)[kPADDINGUPIdx];
247+ padW = NUM2 * (*params.padding)[kPADDINGLEFTIdx];
252 }248 }
249+ /* 当输出的h/w为1时,把传给Dilation算子的dilation值修正为fmap_h/w + pad,避免在dilation值超大时,Dilation算子超时 */
250+ if (inputTensor.gradOutput->GetViewShape().GetDim(kHDimNCHWIdx) == 1 &&
251+ (*params.stride)[kSTRIDEHIdx] > inputTensor.input->GetViewShape().GetDim(kHDimNCHWIdx) + padH) {
252+ postDilationDilationsVector[kHDimNC1HWC0Idx] = inputTensor.input->GetViewShape().GetDim(kHDimNCHWIdx) + padH;
253+ }
254+ if (inputTensor.gradOutput->GetViewShape().GetDim(kWDimNCHWIdx) == 1 &&
255+ (*params.stride)[kSTRIDEWIdx] > inputTensor.input->GetViewShape().GetDim(kWDimNCHWIdx) + padW) {
256+ postDilationDilationsVector[kWDimNC1HWC0Idx] = inputTensor.input->GetViewShape().GetDim(kWDimNCHWIdx) + padW;
257+ }
258+ aclIntArray *post_dilation_dilations = executor->AllocIntArray(postDilationDilationsVector, 5);
259+ CHECK_RET(post_dilation_dilations != nullptr, nullptr);
260+ 
253 int64_t postDilationPadDown =261 int64_t postDilationPadDown =
254 inputTensor.input->GetViewShape().GetDim(kHDimNCHWIdx) + padUp -262 inputTensor.input->GetViewShape().GetDim(kHDimNCHWIdx) + padUp -
255 (inputTensor.gradOutput->GetViewShape().GetDim(kHDimNCHWIdx) - 1) * (*params.stride)[kSTRIDEHIdx] - 1;263 (inputTensor.gradOutput->GetViewShape().GetDim(kHDimNCHWIdx) - 1) * (*params.stride)[kSTRIDEHIdx] - 1;