Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 6 additions & 6 deletions openequivariance/extension/libtorch_tp_jit.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -167,8 +167,8 @@ tuple<torch::Tensor, torch::Tensor, torch::Tensor> 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) {
Expand Down Expand Up @@ -207,8 +207,8 @@ tuple<torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor> 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());

Expand Down Expand Up @@ -402,7 +402,7 @@ tuple<torch::Tensor, torch::Tensor, torch::Tensor> 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();
Expand Down Expand Up @@ -452,7 +452,7 @@ tuple<torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor> 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());

Expand Down
4 changes: 2 additions & 2 deletions openequivariance/implementations/TensorProduct.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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)

Expand Down