How to implement a 3D sparse_tensor_dense_matmul operation in pytorch (or tf)?

Viewed 456

If I have two tensor, a sparse tensor A and a dense tensor B, A.shape is [batch_size, m, n], B.shape is [batch_size, n, k], how can I implement a function f that can perform the following task efficiently:C = f(A, B), C.shape is [batch_size, m, k], and for any batch < batch_size, C[batch] = matmul(A[batch], B[batch]), this function should support backward method.

I try to use for loop and torch.sparse.mm, however, this method does not get the most of GPU. How can I parallelize these operations? (I don't mean torch.nn.DataParallel or something like that.) I am searching for a long time on net. But no use. Please help or try to give some ideas how to achieve this. Thanks in advance.

0 Answers
Related