pub fn outer_product(a: &Tensor, b: &Tensor) -> Result<Tensor>
Flat outer product of two [B, N, d] tensors.
[B, N, d]
Each (i,j) entry is the flattened outer product of row i from a and row j from b.
a
b
[B, N, N, da * db]