已合并
[matmulbackward] 修改MatmulBackwardKernelNpuOpApi.cpp原始shape为1维场景bug #1638
AtomGit-Bot创建于 2024年6月22日
[matmulbackward] 修改MatmulBackwardKernelNpuOpApi.cpp原始shape为1维场景bug #1638
已合并
从refs/pull/1638/head合入到master
共 2 个文件变更+87-82
| @@ -87,101 +87,101 @@ static c10::SmallVector<int64_t, op_infer::SIZE> get_output_size(const at::Tenso | |||
| 87 | at::Tensor matmul_mat1_backward(const at::Tensor self, | 87 | at::Tensor matmul_mat1_backward(const at::Tensor self, |
| 88 | const at::Tensor other, | 88 | const at::Tensor other, |
| 89 | const at::Tensor grad_output) { | 89 | const at::Tensor grad_output) { |
| 90 | /*mat1_grad = grad * mat2^T*/ | 90 | /*mat1_grad = grad * mat2^T*/ |
| 91 | at::Tensor mat1 = self; | 91 | at::Tensor mat1 = self; |
| 92 | at::Tensor mat2 = other; | 92 | at::Tensor mat2 = other; |
| 93 | at::Tensor grad = grad_output; | 93 | at::Tensor grad = grad_output; |
| 94 | 94 | ||
| 95 | // strip mat: (1, 1, m, n)-> (m, n) | 95 | // strip mat: (1, 1, m, n)-> (m, n) |
| 96 | while (mat1.dim() > 2 && mat1.size(0) == 1) { | 96 | while (mat1.dim() > 2 && mat1.size(0) == 1) { |
| 97 | mat1 = mat1.squeeze(0); | 97 | mat1 = mat1.squeeze(0); |
| 98 | } | 98 | } |
| 99 | // unsqueese: (5)*(5)^ -> (1*5)*(1,5)^ | 99 | // unsqueese: (5)*(5)^ -> (1*5)*(1,5)^ |
| 100 | if (mat2.dim() == 1) { | 100 | if (mat2.dim() == 1) { |
| 101 | mat2 = mat2.unsqueeze(-1); | 101 | mat2 = mat2.unsqueeze(-1); |
| 102 | grad = grad.unsqueeze(-1); | 102 | grad = grad.unsqueeze(-1); |
| 103 | } | 103 | } |
| 104 | if (mat1.dim() == 1) { | 104 | if (mat1.dim() == 1) { |
| 105 | mat1 = mat1.unsqueeze(0); | 105 | mat1 = mat1.unsqueeze(0); |
| 106 | grad = grad.unsqueeze(-2); | 106 | grad = grad.unsqueeze(-2); |
| 107 | } | 107 | } |
| 108 | at::Tensor output; | 108 | at::Tensor output; |
| 109 | if (mat1.dim() == 2 && mat2.dim() > 2) { // mm | 109 | if (mat1.dim() == 2 && mat2.dim() > 2) { // mm |
| 110 | output = at_npu::native::OpPreparation::apply_tensor_without_format(mat1.sizes(), grad.options()); | 110 | output = at_npu::native::OpPreparation::apply_tensor_without_format(mat1.sizes(), grad.options()); |
| 111 | mat2 = mat2.transpose(-2, -1); | 111 | mat2 = mat2.transpose(-2, -1); |
| 112 | mat2 = mat2.reshape({-1, mat2.size(-1)}); | 112 | mat2 = mat2.reshape({-1, mat2.size(-1)}); |
| 113 | grad = grad.view({grad.size(-2), -1}); | 113 | grad = grad.view({grad.size(-2), -1}); |
| 114 | matmul_implement_npu(output, grad, mat2); | 114 | matmul_implement_npu(output, grad, mat2); |
| 115 | output = output.reshape(self.sizes()); | 115 | output = output.reshape(self.sizes()); |
| 116 | } else { // bmm | 116 | } else { // bmm |
| 117 | mat2 = mat2.transpose(-2, -1); | 117 | mat2 = mat2.transpose(-2, -1); |
| 118 | auto expend_sizes = get_output_size(grad, mat2); | 118 | auto expend_sizes = get_output_size(grad, mat2); |
| 119 | output = at_npu::native::OpPreparation::apply_tensor_without_format(expend_sizes, grad.options()); | 119 | output = at_npu::native::OpPreparation::apply_tensor_without_format(expend_sizes, grad.options()); |
| 120 | matmul_implement_npu(output, grad, mat2); | 120 | matmul_implement_npu(output, grad, mat2); |
| 121 | } | 121 | } |
| 122 | return output; | 122 | return output; |
| 123 | } | 123 | } |
| 124 | 124 | ||
| 125 | at::Tensor matmul_mat2_backward(const at::Tensor self, | 125 | at::Tensor matmul_mat2_backward(const at::Tensor self, |
| 126 | const at::Tensor other, | 126 | const at::Tensor other, |
| 127 | const at::Tensor grad_output) { | 127 | const at::Tensor grad_output) { |
| 128 | /*mat2_grad = mat1^T * grad*/ | 128 | /*mat2_grad = mat1^T * grad*/ |
| 129 | at::Tensor mat1 = self; | 129 | at::Tensor mat1 = self; |
| 130 | at::Tensor mat2 = other; | 130 | at::Tensor mat2 = other; |
| 131 | at::Tensor grad = grad_output; | 131 | at::Tensor grad = grad_output; |
| 132 | // strip mat: (1, 1, m, n)-> (m, n) | 132 | // strip mat: (1, 1, m, n)-> (m, n) |
| 133 | while (mat2.dim() > 2 && mat2.size(0) == 1) { | 133 | while (mat2.dim() > 2 && mat2.size(0) == 1) { |
| 134 | mat2 = mat2.squeeze(0); | 134 | mat2 = mat2.squeeze(0); |
| 135 | } | 135 | } |
| 136 | // unsqueese: (5)*(5)^ -> (1*5)*(1,5)^ | 136 | // unsqueese: (5)*(5)^ -> (1*5)*(1,5)^ |
| 137 | if (mat2.dim() == 1) { | 137 | if (mat2.dim() == 1) { |
| 138 | mat2 = mat2.unsqueeze(-1); | 138 | mat2 = mat2.unsqueeze(-1); |
| 139 | grad = grad.unsqueeze(-1); | 139 | grad = grad.unsqueeze(-1); |
| 140 | } | 140 | } |
| 141 | if (mat1.dim() == 1) { | 141 | if (mat1.dim() == 1) { |
| 142 | mat1 = mat1.unsqueeze(0); | 142 | mat1 = mat1.unsqueeze(0); |
| 143 | grad = grad.unsqueeze(-2); | 143 | grad = grad.unsqueeze(-2); |
| 144 | } | 144 | } |
| 145 | at::Tensor output; | 145 | at::Tensor output; |
| 146 | if (mat2.dim() == 2 && mat1.dim() > 2) { // mm | 146 | if (mat2.dim() == 2 && mat1.dim() > 2) { // mm |
| 147 | output = at_npu::native::OpPreparation::apply_tensor_without_format(mat2.sizes(), mat1.options()); | 147 | output = at_npu::native::OpPreparation::apply_tensor_without_format(mat2.sizes(), mat1.options()); |
| 148 | mat1 = mat1.reshape({-1, mat1.size(-1)}); | 148 | mat1 = mat1.reshape({-1, mat1.size(-1)}); |
| 149 | grad = grad.reshape({-1, grad.size(-1)}); | 149 | grad = grad.reshape({-1, grad.size(-1)}); |
| 150 | mat1 = mat1.transpose(-2, -1); | 150 | mat1 = mat1.transpose(-2, -1); |
| 151 | matmul_implement_npu(output, mat1, grad); | 151 | matmul_implement_npu(output, mat1, grad); |
| 152 | output = output.reshape(other.sizes()); | 152 | output = output.reshape(other.sizes()); |
| 153 | } else { // bmm | 153 | } else { // bmm |
| 154 | mat1 = mat1.transpose(-2, -1); | 154 | mat1 = mat1.transpose(-2, -1); |
| 155 | auto expend_sizes = get_output_size(mat1, grad); | 155 | auto expend_sizes = get_output_size(mat1, grad); |
| 156 | output = at_npu::native::OpPreparation::apply_tensor_without_format(expend_sizes, mat1.options()); | 156 | output = at_npu::native::OpPreparation::apply_tensor_without_format(expend_sizes, mat1.options()); |
| 157 | matmul_implement_npu(output, mat1, grad); | 157 | matmul_implement_npu(output, mat1, grad); |
| 158 | } | 158 | } |
| 159 | return output; | 159 | return output; |
| 160 | } | 160 | } |
| 161 | 161 | ||
| 162 | std::tuple<at::Tensor, at::Tensor> matmul_backward(const at::Tensor &grad, | 162 | std::tuple<at::Tensor, at::Tensor> matmul_backward(const at::Tensor &grad, |
| 163 | const at::Tensor &self, | 163 | const at::Tensor &self, |
| 164 | const at::Tensor &other, | 164 | const at::Tensor &other, |
| 165 | std::array<bool, 2> grad_input_mask) { | 165 | std::array<bool, 2> grad_input_mask) { |
| 166 | if (!grad.defined()) { | 166 | if (!grad.defined()) { |
| 167 | return std::make_tuple(at::Tensor(), at::Tensor()); | 167 | return std::make_tuple(at::Tensor(), at::Tensor()); |
| 168 | } | 168 | } |
| 169 | // backward mat1 and mat2 separately | 169 | // backward mat1 and mat2 separately |
| 170 | at::Tensor self_grad; | 170 | at::Tensor self_grad; |
| 171 | at::Tensor other_grad; | 171 | at::Tensor other_grad; |
| 172 | if (grad_input_mask[1]) { | 172 | if (grad_input_mask[1]) { |
| 173 | other_grad = matmul_mat2_backward(self, other, grad); | 173 | other_grad = matmul_mat2_backward(self, other, grad); |
| 174 | } | 174 | } |
| 175 | 175 | ||
| 176 | if (grad_input_mask[0]) { | 176 | if (grad_input_mask[0]) { |
| 177 | self_grad = matmul_mat1_backward(self, other, grad); | 177 | self_grad = matmul_mat1_backward(self, other, grad); |
| 178 | } | 178 | } |
| 179 | 179 | ||
| 180 | // strip added dim: (5,1)->(5) | 180 | // strip added dim: (5,1)->(5) |
| 181 | if (other.dim() == 1 && other_grad.size(-1) == 1) { | 181 | if (other.dim() == 1 && other_grad.size(-1) == 1 && other_grad.dim() != 1) { |
| 182 | other_grad = other_grad.squeeze(-1); | 182 | other_grad = other_grad.squeeze(-1); |
| 183 | } | 183 | } |
| 184 | return std::make_tuple(self_grad, other_grad); | 184 | return std::make_tuple(self_grad, other_grad); |
| 185 | } | 185 | } |
| 186 | 186 | ||
| 187 | } // namespace op_api | 187 | } // namespace op_api |
| @@ -120,6 +120,12 @@ class TestMatMul(TestCase): | |||
| 120 | ] | 120 | ] |
| 121 | self.matmul_backward_result(shape_format) | 121 | self.matmul_backward_result(shape_format) |
| 122 | 122 | ||
| 123 | def test_matmul_backward_shape_format_fp16_case10(self): | ||
| 124 | shape_format = [ | ||
| 125 | [[np.float16, 2, [9, 1]], [np.float16, 2, [1]]], | ||
| 126 | ] | ||
| 127 | self.matmul_backward_result(shape_format) | ||
| 128 | |||
| 123 | def test_matmul_allow_hf32(self): | 129 | def test_matmul_allow_hf32(self): |
| 124 | torch.npu.matmul.allow_hf32 = True | 130 | torch.npu.matmul.allow_hf32 = True |
| 125 | shape_format = [ | 131 | shape_format = [ |
| @@ -130,7 +136,6 @@ class TestMatMul(TestCase): | |||
| 130 | self.matmul_backward_result(shape_format) | 136 | self.matmul_backward_result(shape_format) |
| 131 | torch.npu.matmul.allow_hf32 = False | 137 | torch.npu.matmul.allow_hf32 = False |
| 132 | 138 | ||
| 133 | |||
| 134 | if __name__ == "__main__": | 139 | if __name__ == "__main__": |
| 135 | np.random.seed(1234) | 140 | np.random.seed(1234) |
| 136 | run_tests() | 141 | run_tests() |