# Fast type-stable tensor products

**URL:** <https://discourse.julialang.org/t/fast-type-stable-tensor-products/35159>\
**Category:** Performance\
**Tags:** performance, linearalgebra, arrays, tensors\
**Created:** [February 26, 2020, 9:43am UTC](https://discourse.julialang.org/t/fast-type-stable-tensor-products/35159 "2020-02-26T09:43:32Z")\
**Posts on this page:** 19\
**Page:** 1

<div class="post-metadata">

**Author:** ![e3c6](https://avatars.discourse-cdn.com/v4/letter/e/e79b87/32.png) [@e3c6](https://discourse.julialang.org/u/e3c6)\
**Post date:** [February 26, 2020, 9:43am UTC](https://discourse.julialang.org/t/fast-type-stable-tensor-products/35159/1 "2020-02-26T09:43:32Z")

</div>

Let `A` and `B` be two multi-dimensional arrays. I want to form a product of them that contracts the last `n` dimensions of `A` with the first `n` dimensions of `B`. Calling this product `C`, one way to do it is with the following code:

```julia

function tensordot(A::AbstractArray, B::AbstractArray, n::Int)
    Amat = reshape(A, prod(size(A)[1:end-n]), prod(size(A)[end-n+1:end]))
    Bmat = reshape(B, prod(size(B)[1:n]), prod(size(B)[n+1:end]))
    Cmat = Amat * Bmat
    C = reshape(Cmat, size(A)[1:end-n]..., size(B)[n+1:end]...)
    return C
end

# example
A = randn(10,2,5,7);
B = randn(5,7,3,4,7); 
tensordot(A, B, 2)

```

Is this a fast way to do this kind of products in general? Or is there a better way?

---

<div class="post-metadata">

**Author:** ![mcabbott](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/mcabbott/32/6603_2.png) [@mcabbott](https://discourse.julialang.org/u/mcabbott)\
**Post date:** [February 26, 2020, 10:43am UTC](https://discourse.julialang.org/t/fast-type-stable-tensor-products/35159/2 "2020-02-26T10:43:30Z")

</div>

Not really, it’s a single call to `*`.

If the arrays are small, then you may see some advantage to calculating the sizes for `reshape` using `ntuple`isms, and not making `n` global.

---

<div class="post-metadata">

**Author:** ![e3c6](https://avatars.discourse-cdn.com/v4/letter/e/e79b87/32.png) [@e3c6](https://discourse.julialang.org/u/e3c6)\
**Post date:** [February 26, 2020, 10:46am UTC](https://discourse.julialang.org/t/fast-type-stable-tensor-products/35159/3 "2020-02-26T10:46:32Z")

</div>

> [@mcabbott](#):
>
> `ntuple` isms

Not sure what you mean here.

> [@mcabbott](#):
>
> Not really, it’s a single call to `*` .

So this is as fast as possible?

---

<div class="post-metadata">

**Author:** ![mcabbott](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/mcabbott/32/6603_2.png) [@mcabbott](https://discourse.julialang.org/u/mcabbott)\
**Post date:** [February 26, 2020, 11:02am UTC](https://discourse.julialang.org/t/fast-type-stable-tensor-products/35159/4 "2020-02-26T11:02:02Z")

</div>

I mean you can squeeze out a μs or so by being careful when calculating these sizes. Whether this matters at all will depend on how big the arrays are:

```julia
julia> @btime prod(size($A)[end-2:end])
  335.348 ns (2 allocations: 144 bytes)
70

julia> @btime prod(ntuple(d -> size($A,ndims($A)-d+1), 3))
  1.421 ns (0 allocations: 0 bytes)
70

```

---

<div class="post-metadata">

**Author:** ![e3c6](https://avatars.discourse-cdn.com/v4/letter/e/e79b87/32.png) [@e3c6](https://discourse.julialang.org/u/e3c6)\
**Post date:** [February 26, 2020, 11:09am UTC](https://discourse.julialang.org/t/fast-type-stable-tensor-products/35159/5 "2020-02-26T11:09:33Z")

</div>

Strange. Raised an issue here: [Unit range indexing of tuples · Issue #34884 · JuliaLang/julia · GitHub](https://github.com/JuliaLang/julia/issues/34884)

But back to the original topic, I think the running time will be dominated by the matrix multiply.

---

<div class="post-metadata">

**Author:** ![e3c6](https://avatars.discourse-cdn.com/v4/letter/e/e79b87/32.png) [@e3c6](https://discourse.julialang.org/u/e3c6)\
**Post date:** [February 26, 2020, 3:09pm UTC](https://discourse.julialang.org/t/fast-type-stable-tensor-products/35159/6 "2020-02-26T15:09:31Z")

</div>

> [@e3c6](#):
>
> ```julia
> function tensordot(A::AbstractArray, B::AbstractArray, n::Int)
> Amat = reshape(A, prod(size(A)[1:end-n]), prod(size(A)[end-n+1:end]))
> Bmat = reshape(B, prod(size(B)[1:n]), prod(size(B)[n+1:end]))
> Cmat = Amat * Bmat
> C = reshape(Cmat, size(A)[1:end-n]..., size(B)[n+1:end]...)
> return C
> end
> 
> ```

Unfortunately this is **not type-stable** , because the dimensions of the output array depend on `n`.  
Now in all the uses I will do of this function, `n` can actually be computed from type information.  
So any suggestions on how I can make a type-stable version of this?

---

<div class="post-metadata">

**Author:** ![raminammour](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/raminammour/32/13572_2.png) [@raminammour](https://discourse.julialang.org/u/raminammour)\
**Post date:** [February 26, 2020, 8:46pm UTC](https://discourse.julialang.org/t/fast-type-stable-tensor-products/35159/8 "2020-02-26T20:46:27Z")

</div>

Here is a “type stable” implementation, but if you benchmark, it doesn’t really matter for the runtime. The time is spent on matrix-matrix product. But if type stability is important for some other parts, then it may be useful:

```julia
 function tensordot(A::AbstractArray{T1,N1}, B::AbstractArray{T2,N2},::Val{n}) where {T1,T2,N1,N2,n}
           s1=ntuple(i->size(A,i),Val(N1-n))
           s2=ntuple(i->size(A,N1-i+1),Val(n))
           s3=ntuple(i->size(B,i),Val(n))
           s4=ntuple(i->size(B,i+n),Val(N2-n))
           s5=ntuple(i->( i<=N1-n ? s1[i] : s4[i-(N1-n)]),Val(N1+N2-2n))
           Amat = reshape(A, prod(s1), prod(s2))
           Bmat = reshape(B, prod(s3), prod(s4))
           Cmat = Amat * Bmat
           C = reshape(Cmat,s5)
           return C
           end

julia> A=rand(10,12,13,14);

julia> B=rand(13,14,15,16,17);

julia> tensordot(A, B, Val(2))==tensordot(A,B,2)
true

julia> @btime tensordot($A,$B,Val(2));
  958.343 μs (8 allocations: 3.74 MiB)

julia> @btime tensordot($A,$B,2);
  965.924 μs (23 allocations: 3.74 MiB)

julia> @code_warntype tensordot(A,B,Val(2))
Variables
  #self#::Core.Compiler.Const(tensordot, false)
  A::Array{Float64,4}
  B::Array{Float64,5}
  #unused#::Core.Compiler.Const(Val{2}(), false)
  #273::var"#273#278"{Array{Float64,4}}
  #274::var"#274#279"{4,Array{Float64,4}}
  #275::var"#275#280"{Array{Float64,5}}
  #276::var"#276#281"{2,Array{Float64,5}}
  #277::var"#277#282"{4,2,Tuple{Int64,Int64},Tuple{Int64,Int64,Int64}}
  s1::Tuple{Int64,Int64}
  s2::Tuple{Int64,Int64}
  s3::Tuple{Int64,Int64}
  s4::Tuple{Int64,Int64,Int64}
  s5::NTuple{5,Int64}
  Amat::Array{Float64,2}
  Bmat::Array{Float64,2}
  Cmat::Array{Float64,2}
  C::Array{Float64,5}

Body::Array{Float64,5}
1 ─ %1 = Main.:(var"#273#278")::Core.Compiler.Const(var"#273#278", false)
│ %2 = Core.typeof(A)::Core.Compiler.Const(Array{Float64,4}, false)
│ %3 = Core.apply_type(%1, %2)::Core.Compiler.Const(var"#273#278"{Array{Float64,4}}, false)
│ (#273 = %new(%3, A))
│ %5 = #273::var"#273#278"{Array{Float64,4}}
│ %6 = ($(Expr(:static_parameter, 3)) - $(Expr(:static_parameter, 5)))::Core.Compiler.Const(2, false)
│ %7 = Main.Val(%6)::Core.Compiler.Const(Val{2}(), true)
│ (s1 = Main.ntuple(%5, %7))
│ %9 = Main.:(var"#274#279")::Core.Compiler.Const(var"#274#279", false)
│ %10 = $(Expr(:static_parameter, 3))::Core.Compiler.Const(4, false)
│ %11 = Core.typeof(A)::Core.Compiler.Const(Array{Float64,4}, false)
│ %12 = Core.apply_type(%9, %10, %11)::Core.Compiler.Const(var"#274#279"{4,Array{Float64,4}}, false)
│ (#274 = %new(%12, A))
│ %14 = #274::var"#274#279"{4,Array{Float64,4}}
│ %15 = Main.Val($(Expr(:static_parameter, 5)))::Core.Compiler.Const(Val{2}(), true)
│ (s2 = Main.ntuple(%14, %15))
│ %17 = Main.:(var"#275#280")::Core.Compiler.Const(var"#275#280", false)
│ %18 = Core.typeof(B)::Core.Compiler.Const(Array{Float64,5}, false)
│ %19 = Core.apply_type(%17, %18)::Core.Compiler.Const(var"#275#280"{Array{Float64,5}}, false)
│ (#275 = %new(%19, B))
│ %21 = #275::var"#275#280"{Array{Float64,5}}
│ %22 = Main.Val($(Expr(:static_parameter, 5)))::Core.Compiler.Const(Val{2}(), true)
│ (s3 = Main.ntuple(%21, %22))
│ %24 = Main.:(var"#276#281")::Core.Compiler.Const(var"#276#281", false)
│ %25 = $(Expr(:static_parameter, 5))::Core.Compiler.Const(2, false)
│ %26 = Core.typeof(B)::Core.Compiler.Const(Array{Float64,5}, false)
│ %27 = Core.apply_type(%24, %25, %26)::Core.Compiler.Const(var"#276#281"{2,Array{Float64,5}}, false)
│ (#276 = %new(%27, B))
│ %29 = #276::var"#276#281"{2,Array{Float64,5}}
│ %30 = ($(Expr(:static_parameter, 4)) - $(Expr(:static_parameter, 5)))::Core.Compiler.Const(3, false)
│ %31 = Main.Val(%30)::Core.Compiler.Const(Val{3}(), true)
│ (s4 = Main.ntuple(%29, %31))
│ %33 = Main.:(var"#277#282")::Core.Compiler.Const(var"#277#282", false)
│ %34 = $(Expr(:static_parameter, 3))::Core.Compiler.Const(4, false)
│ %35 = $(Expr(:static_parameter, 5))::Core.Compiler.Const(2, false)
│ %36 = Core.typeof(s1)::Core.Compiler.Const(Tuple{Int64,Int64}, false)
│ %37 = Core.typeof(s4)::Core.Compiler.Const(Tuple{Int64,Int64,Int64}, false)
│ %38 = Core.apply_type(%33, %34, %35, %36, %37)::Core.Compiler.Const(var"#277#282"{4,2,Tuple{Int64,Int64},Tuple{Int64,Int64,Int64}}, false)
│ %39 = s1::Tuple{Int64,Int64}
│ (#277 = %new(%38, %39, s4))
│ %41 = #277::var"#277#282"{4,2,Tuple{Int64,Int64},Tuple{Int64,Int64,Int64}}
│ %42 = ($(Expr(:static_parameter, 3)) + $(Expr(:static_parameter, 4)))::Core.Compiler.Const(9, false)
│ %43 = (2 * $(Expr(:static_parameter, 5)))::Core.Compiler.Const(4, false)
│ %44 = (%42 - %43)::Core.Compiler.Const(5, false)
│ (s5 = Main.ntuple(%41, %44))
│ %46 = Main.prod(s1)::Int64
│ %47 = Main.prod(s2)::Int64
│ (Amat = Main.reshape(A, %46, %47))
│ %49 = Main.prod(s3)::Int64
│ %50 = Main.prod(s4)::Int64
│ (Bmat = Main.reshape(B, %49, %50))
│ (Cmat = Amat * Bmat)
│ (C = Main.reshape(Cmat, s5))
└── return C

```

The trick is to make the sizes (`N1,N2,n`) available at compile time, `Val` and `ntuple`…

Cheers!

---

<div class="post-metadata">

**Author:** ![simeonschaub](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/simeonschaub/32/216566_2.png) [@simeonschaub](https://discourse.julialang.org/u/simeonschaub)\
**Post date:** [February 26, 2020, 8:49pm UTC](https://discourse.julialang.org/t/fast-type-stable-tensor-products/35159/9 "2020-02-26T20:49:22Z")

</div>

@e3c6 Do you know about [https://github.com/Jutho/TensorOperations.jl](https://github.com/Jutho/TensorOperations.jl)? Seems like it already implements most of what you’re trying to do.

---

<div class="post-metadata">

**Author:** ![e3c6](https://avatars.discourse-cdn.com/v4/letter/e/e79b87/32.png) [@e3c6](https://discourse.julialang.org/u/e3c6)\
**Post date:** [February 26, 2020, 8:56pm UTC](https://discourse.julialang.org/t/fast-type-stable-tensor-products/35159/10 "2020-02-26T20:56:42Z")

</div>

Yes, but (at least from my knowledge of it) it doesn’t seem to support tensor products where the number of contracted dimensions is dynamic. Does it?

Also from the things I’ve tried, `TensorOperations.jl` is incompatible with Zygote because it modifies temporary arrays in-place.

These two things kept me away from it, but I could be mistaken.

---

<div class="post-metadata">

**Author:** ![e3c6](https://avatars.discourse-cdn.com/v4/letter/e/e79b87/32.png) [@e3c6](https://discourse.julialang.org/u/e3c6)\
**Post date:** [February 26, 2020, 9:04pm UTC](https://discourse.julialang.org/t/fast-type-stable-tensor-products/35159/11 "2020-02-26T21:04:15Z")

</div>

Thanks. It seems that you don’t need `Val` in the body of the function. In fact it can be simplified to:

```julia
function tensordot(A::AbstractArray, B::AbstractArray, ::Val{n}) where {n}
		Amat = reshape(A, prod(size(A,i) for i = 1:ndims(A)-n), :)
		Bmat = reshape(B, prod(size(B,i) for i=1:n), :)
		Cmat = Amat * Bmat
		C = reshape(Cmat, ntuple(i -> i ≤ ndims(A) - n ? size(A,i) : size(B, i - ndims(A) + 2n), ndims(A) + ndims(B) - 2n))
		return C
end

```

and it will also be type stable!

@raminammour See edit.

---

<div class="post-metadata">

**Author:** ![raminammour](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/raminammour/32/13572_2.png) [@raminammour](https://discourse.julialang.org/u/raminammour)\
**Post date:** [February 26, 2020, 9:06pm UTC](https://discourse.julialang.org/t/fast-type-stable-tensor-products/35159/12 "2020-02-26T21:06:16Z")

</div>

Yes, I can never figure out when it is needed and when the compiler infers without it, so I have taken to the habit of including it just in case 🙂

---

<div class="post-metadata">

**Author:** ![mcabbott](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/mcabbott/32/6603_2.png) [@mcabbott](https://discourse.julialang.org/u/mcabbott)\
**Post date:** [February 26, 2020, 11:52pm UTC](https://discourse.julialang.org/t/fast-type-stable-tensor-products/35159/13 "2020-02-26T23:52:55Z")

</div>

One other way you can do this is:

```julia
julia> using OMEinsum

julia> tensordot(A, B, 2) ≈ ein" abcd, cdefg -> abefg "(A, B)
true

julia> @macroexpand ein"abcd,cdefg -> abefg"(A, B)
:((EinCode{(('a', 'b', 'c', 'd'), ('c', 'd', 'e', 'f', 'g')),('a', 'b', 'e', 'f', 'g')}())(A, B))

```

You can write such codes at runtime. And the result should be Zygote-friendly.

However if you time it, it’s not as fast as your function here. I think it’s decomposing this into more operations than strictly necessary.

---

<div class="post-metadata">

**Author:** ![orialb](https://avatars.discourse-cdn.com/v4/letter/o/65b543/32.png) [@orialb](https://discourse.julialang.org/u/orialb)\
**Post date:** [February 27, 2020, 12:11am UTC](https://discourse.julialang.org/t/fast-type-stable-tensor-products/35159/14 "2020-02-27T00:11:57Z")

</div>

> it doesn’t seem to support tensor products where the number of contracted dimensions is dynamic

You can do that with `TensorOperations` using the `tensorcontract` function.

---

<div class="post-metadata">

**Author:** ![carstenbauer](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/carstenbauer/32/4981_2.png) [@carstenbauer](https://discourse.julialang.org/u/carstenbauer)\
**Post date:** [February 27, 2020, 6:56am UTC](https://discourse.julialang.org/t/fast-type-stable-tensor-products/35159/15 "2020-02-27T06:56:29Z")

</div>

Here a link to the function @orialb mentioned: [Functions · TensorOperations.jl](https://jutho.github.io/TensorOperations.jl/stable/functions/#TensorOperations.tensorcontract)

---

<div class="post-metadata">

**Author:** ![e3c6](https://avatars.discourse-cdn.com/v4/letter/e/e79b87/32.png) [@e3c6](https://discourse.julialang.org/u/e3c6)\
**Post date:** [February 27, 2020, 8:18am UTC](https://discourse.julialang.org/t/fast-type-stable-tensor-products/35159/16 "2020-02-27T08:18:06Z")

</div>

> [@mcabbott](#):
>
> `tensordot(A, B, 2) ≈ ein" abcd, cdefg -> abefg "(A, B)`

But that assumes you know _a priori_ what dimensions to contract. Note that in the function I wrote above, `n` is an argument that comes from outside and is unknown to the function (even though it can be inferred at compile time, hence the `Val`).

---

<div class="post-metadata">

**Author:** ![e3c6](https://avatars.discourse-cdn.com/v4/letter/e/e79b87/32.png) [@e3c6](https://discourse.julialang.org/u/e3c6)\
**Post date:** [February 27, 2020, 8:18am UTC](https://discourse.julialang.org/t/fast-type-stable-tensor-products/35159/17 "2020-02-27T08:18:30Z")

</div>

> [@carstenbauer](#):
>
> Here a link to the function @orialb mentioned: [https://jutho.github.io/TensorOperations.jl/stable/functions/#TensorOperations.tensorcontract](https://jutho.github.io/TensorOperations.jl/stable/functions/#TensorOperations.tensorcontract)

Thanks. But that is not Zygote friendly, correct?

---

<div class="post-metadata">

**Author:** ![mcabbott](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/mcabbott/32/6603_2.png) [@mcabbott](https://discourse.julialang.org/u/mcabbott)\
**Post date:** [February 27, 2020, 8:23am UTC](https://discourse.julialang.org/t/fast-type-stable-tensor-products/35159/18 "2020-02-27T08:23:19Z")

</div>

> [@e3c6](#):
>
> But that assumes you know _a priori_ what dimensions to contract

No:

> [@mcabbott](#):
>
> You can write such codes at runtime

You do not have to use the macro, and can make your function construct the EinCode object by itself.

---

<div class="post-metadata">

**Author:** ![carstenbauer](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/carstenbauer/32/4981_2.png) [@carstenbauer](https://discourse.julialang.org/u/carstenbauer)\
**Post date:** [February 27, 2020, 8:24am UTC](https://discourse.julialang.org/t/fast-type-stable-tensor-products/35159/19 "2020-02-27T08:24:19Z")

</div>

I don’t know but would expect compatibility at least with `disable_blas()`.

---

<div class="post-metadata">

**Author:** ![juthohaegeman](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/juthohaegeman/32/8620_2.png) [@juthohaegeman](https://discourse.julialang.org/u/juthohaegeman)\
**Post date:** [February 27, 2020, 9:56am UTC](https://discourse.julialang.org/t/fast-type-stable-tensor-products/35159/20 "2020-02-27T09:56:51Z")

</div>

There is now also an `ncon` (and `@ncon`) function and macro in the latest version of TensorOperations.jl.

But indeed, autodiff support is still on the todo list. I am a bit overwhelmed by the autodiff packages in Julia. Will Zygote.jl become the community standard (or is it already)? And as a third-party package I should only include ZygoteRules.jl and define the appropriate adjoints? And how does it indeed deal with in-place modification?
