diff --git a/openequivariance/extension/libtorch_tp_jit.cpp b/openequivariance/extension/libtorch_tp_jit.cpp index 156e462f..1d52a994 100644 --- a/openequivariance/extension/libtorch_tp_jit.cpp +++ b/openequivariance/extension/libtorch_tp_jit.cpp @@ -167,8 +167,8 @@ tuple jit_tp_backward( const torch::Tensor &L3_grad) { int64_t num_batch = L1_in.sizes()[0]; - torch::Tensor L1_grad = torch::empty(L1_in.sizes(), L1_in.options()); - torch::Tensor L2_grad = torch::empty(L2_in.sizes(), L2_in.options()); + torch::Tensor L1_grad = torch::zeros(L1_in.sizes(), L1_in.options()); + torch::Tensor L2_grad = torch::zeros(L2_in.sizes(), L2_in.options()); torch::Tensor W_grad = torch::empty(W.sizes(), W.options()); if(jit_instance->shared_weights == 1) { @@ -207,8 +207,8 @@ tuple jit_tp_double_ Stream stream = get_current_stream(); int64_t num_batch = L1_in.sizes()[0]; // Declaring outputs - torch::Tensor L1_grad = torch::empty(L1_in.sizes(), L1_in.options()); - torch::Tensor L2_grad = torch::empty(L2_in.sizes(), L2_in.options()); + torch::Tensor L1_grad = torch::zeros(L1_in.sizes(), L1_in.options()); + torch::Tensor L2_grad = torch::zeros(L2_in.sizes(), L2_in.options()); torch::Tensor W_grad = torch::empty(W.sizes(), W.options()); torch::Tensor L3_dgrad = torch::empty(L3_grad.sizes(), L3_grad.options()); @@ -402,7 +402,7 @@ tuple jit_conv_backward( int64_t nnz = rows.sizes()[0]; int64_t node_count = L1_in.sizes()[0]; torch::Tensor L1_grad = torch::zeros(L1_in.sizes(), L1_in.options()); - torch::Tensor L2_grad = torch::empty(L2_in.sizes(), L2_in.options()); + torch::Tensor L2_grad = torch::zeros(L2_in.sizes(), L2_in.options()); torch::Tensor W_grad = torch::empty(W.sizes(), W.options()); torch::Tensor L1_in_contig = L1_in.contiguous(); @@ -452,7 +452,7 @@ tuple jit_conv_doubl int64_t nnz = rows.sizes()[0]; int64_t node_count = L1_in.sizes()[0]; torch::Tensor L1_grad = torch::zeros(L1_in.sizes(), L1_in.options()); - torch::Tensor L2_grad = torch::empty(L2_in.sizes(), L2_in.options()); + torch::Tensor L2_grad = torch::zeros(L2_in.sizes(), L2_in.options()); torch::Tensor W_grad = torch::empty(W.sizes(), W.options()); torch::Tensor L3_dgrad = torch::zeros(L3_grad.sizes(), L3_grad.options()); diff --git a/openequivariance/implementations/TensorProduct.py b/openequivariance/implementations/TensorProduct.py index 88bb94db..7a740f19 100644 --- a/openequivariance/implementations/TensorProduct.py +++ b/openequivariance/implementations/TensorProduct.py @@ -125,8 +125,8 @@ def backward_helper( weights: torch.Tensor, L3_grad: torch.Tensor, ) -> typing.List[torch.Tensor]: - L1_grad = torch.empty_like(L1_in) - L2_grad = torch.empty_like(L2_in) + L1_grad = torch.zeros_like(L1_in) + L2_grad = torch.zeros_like(L2_in) weights_grad = torch.empty_like(weights) if self.config.shared_weights: diff --git a/openequivariance/implementations/convolution/TensorProductConv.py b/openequivariance/implementations/convolution/TensorProductConv.py index bfc4d25b..27ece018 100644 --- a/openequivariance/implementations/convolution/TensorProductConv.py +++ b/openequivariance/implementations/convolution/TensorProductConv.py @@ -198,7 +198,7 @@ def backward_helper( transpose_perm: Optional[torch.Tensor] = None, ) -> List[torch.Tensor]: L1_grad = torch.zeros_like(L1_in) - L2_grad = torch.empty_like(L2_in) + L2_grad = torch.zeros_like(L2_in) weights_grad = torch.empty_like(weights) if self.config.shared_weights: @@ -287,7 +287,7 @@ def double_backward_helper( transpose_perm: Optional[torch.Tensor] = None, ) -> List[torch.Tensor]: L1_grad = torch.zeros_like(L1_in) - L2_grad = torch.empty_like(L2_in) + L2_grad = torch.zeros_like(L2_in) W_grad = torch.empty_like(W) L3_dgrad = torch.zeros_like(L3_grad)