# Efficient approach to generate optimised methods for special cases of a general function

**URL:** <https://discourse.julialang.org/t/efficient-approach-to-generate-optimised-methods-for-special-cases-of-a-general-function/110892>\
**Category:** Performance\
**Tags:** question, performance\
**Created:** [February 28, 2024, 11:55am UTC](https://discourse.julialang.org/t/efficient-approach-to-generate-optimised-methods-for-special-cases-of-a-general-function/110892 "2024-02-28T11:55:30Z")\
**Posts on this page:** 6\
**Page:** 1

<div class="post-metadata">

**Author:** ![TimHargreaves](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/timhargreaves/32/207284_2.png) [@TimHargreaves](https://discourse.julialang.org/u/TimHargreaves)\
**Post date:** [February 28, 2024, 11:55am UTC](https://discourse.julialang.org/t/efficient-approach-to-generate-optimised-methods-for-special-cases-of-a-general-function/110892/1 "2024-02-28T11:55:30Z")

</div>

_I am a fairly new user of Julia so may have misconceptions about the language’s capabilities and behaviour._

**Problem Statement**

I have implemented the most general form of an algorithm I am interested in, which is parameterised by multiple vectors and matrices. Although this general form is sometimes used, most of the time I will only be interested in the cases where some of the vectors/matrices are zero, or the matrices are the identity/multiples of the identity.

As a minimal example, let’s consider the function

```julia
function linear_transformation(x::Vector{Float64}, A::Matrix{Float64}, b::Vector{Float64})
    return A * x + b
end

```

I would like to have automatically generated, optimised versions of this function for the cases:

- A is zero
- A is the identity
- b is zero

E.g. for the first case, I would ideally like the code to run as quickly as simply returning `b`.

It’s also worth noting that such optimisations will likely have ripple effects—one matrix being zero might cause another to be zero, and so on.

I appreciate that this is a lot to ask for (especially for a non-symbolic, eagerly executed language) so I’m not expecting a perfect solution to this. Instead, I would love to have an open discussion about this topic and hopefully find some ways I can get closer to my desired outcome.

**Initial Thoughts**

A manual way of approaching this would be to write variants of `linear_transformation` dispatching on zero values. Obviously, this is not a scalable solution, but I am also under the impression that value dispatch is discouraged in Julia and leads to inefficient code.

The `UniformScaling` matrix seems related to what I’m asking for, though I’m not completely sure it does exactly what I’m after. Perhaps, creating similar `Identity` and `Zero` matrices inspired by this is a good approach.

Can macros be used to achieve this task?

---

<div class="post-metadata">

**Author:** ![gdalle](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/gdalle/32/27854_2.png) [@gdalle](https://discourse.julialang.org/u/gdalle)\
**Post date:** [February 28, 2024, 12:10pm UTC](https://discourse.julialang.org/t/efficient-approach-to-generate-optimised-methods-for-special-cases-of-a-general-function/110892/2 "2024-02-28T12:10:16Z")

</div>

The natural reflex is to add an `if / else` to your function like so:

```julia
function linear_transformation(x, A, b)
    if all(iszero, A)
        return b
    elseif ...
        ...
    end
end

```

but that will be rather inefficient. First because of the `if`, and second because statements like `all(iszeros, A)` need to check every single coefficient.

The better method is to encode the matrix structure in the type. You have correctly spotted that our standard library `LinearAlgebra` does this, and already has an identity operator. Meanwhile, FillArrays.jl has types to represent arrays filled with zeros or ones. Both of these libraries overload the necessary operations to make `A * x + b` as fast as can be without additional intervention.

To sum up, the only thing you need to do is allow arbitrary matrix / vector subtypes in your function:

```julia
linear_transformation(x::AbstractVector, A::AbstractMatrix, b::AbstractVector) = A * x + b

```

Then you can apply it with the following varieties of arrays:

```julia
using LinearAlgebra, FillArrays
A = rand(2, 2); # dense matrix
A = I(2); # identity matrix
A = FillArrays.Zeros(2, 2); # zero matrix

```

Even better, you can combine types of `A` and `b` however you want, and the _multiple_ dispatch will still find the optimal method.

---

<div class="post-metadata">

**Author:** ![Dan](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/dan/32/42581_2.png) [@Dan](https://discourse.julialang.org/u/Dan)\
**Post date:** [February 28, 2024, 12:11pm UTC](https://discourse.julialang.org/t/efficient-approach-to-generate-optimised-methods-for-special-cases-of-a-general-function/110892/3 "2024-02-28T12:11:13Z")

</div>

> [@gdalle](#):
>
> `all(iszeros, A)` need to check every single coefficient

Not all coefficients for the `false` case.

---

<div class="post-metadata">

**Author:** ![TimHargreaves](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/timhargreaves/32/207284_2.png) [@TimHargreaves](https://discourse.julialang.org/u/TimHargreaves)\
**Post date:** [February 28, 2024, 12:27pm UTC](https://discourse.julialang.org/t/efficient-approach-to-generate-optimised-methods-for-special-cases-of-a-general-function/110892/4 "2024-02-28T12:27:27Z")

</div>

> [@gdalle](#):
>
> The better method is to encode the matrix structure in the type

Gotcha. That sounds great. Thank you very much.

One follow-up question:

This might be a case over over-engineering, but is `UniformScaling` as efficient as it can be for the unscaled identity matrix case? It is my understanding that `I` is indistinguishable from `1.0 * I` by the type system and so multiplying a matrix by `I` would involve multiplying each element (unnecessarily) by 1.0. Is this correct? If so, I assume I would just have to define something similar to `I` that is strictly for unscaled identity matrices.

---

<div class="post-metadata">

**Author:** ![gdalle](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/gdalle/32/27854_2.png) [@gdalle](https://discourse.julialang.org/u/gdalle)\
**Post date:** [February 28, 2024, 12:47pm UTC](https://discourse.julialang.org/t/efficient-approach-to-generate-optimised-methods-for-special-cases-of-a-general-function/110892/5 "2024-02-28T12:47:33Z")

</div>

Indeed you’re right and the performance difference is visible too, although most of the hit is actually due to the new allocation and not the multiplication by 1. If we use in-place multiplication with `LinearAlgebra.mul!`, the difference goes away, so I assume `1 * x` is optimized (at least for integer `1`):

```julia
using BenchmarkTools, LinearAlgebra

struct Identity end
Base.:*(::Identity, x) = x
LinearAlgebra.mul!(y, ::Identity, x) = y .= x

x = rand(1000)
y = zeros(1000)

```

```julia
julia> @btime (A * $x) setup=(A = Identity());
  2.455 ns (0 allocations: 0 bytes)

julia> @btime (A * $x) setup=(A = I);
  533.193 ns (1 allocation: 7.94 KiB)

julia> @btime mul!($y, A, $x) setup=(A = Identity());
  74.981 ns (0 allocations: 0 bytes)

julia> @btime mul!($y, A, $x) setup=(A = I);
  64.121 ns (0 allocations: 0 bytes)

```

---

<div class="post-metadata">

**Author:** ![PeterSimon](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/petersimon/32/25193_2.png) [@PeterSimon](https://discourse.julialang.org/u/PeterSimon)\
**Post date:** [March 2, 2024, 6:46pm UTC](https://discourse.julialang.org/t/efficient-approach-to-generate-optimised-methods-for-special-cases-of-a-general-function/110892/6 "2024-03-02T18:46:55Z")

</div>

> [@gdalle](#):
>
> ```julia
> julia> @btime (A * $x) setup=(A = I);
> 533.193 ns (1 allocation: 7.94 KiB)
> 
> ```

You mentioned that allocation is the main culprit here but I think it’s worth pointing out that the following definition

> [@gdalle](#):
>
> `Base.:*(::Identity, x) = x`

is not a fair way to compare against multiplication by `I`, since multiplication using `*` is supposed to create a new matrix, not provide a new name for the old matrix. If you instead define

```julia
Base.:*(::Identity, x) = copy(x)

```

then the first two timings become on my machine

```julia
julia> @btime (A * $x) setup=(A = Identity());
  502.525 ns (1 allocation: 7.94 KiB)

julia> @btime (A * $x) setup=(A = I);
  505.102 ns (1 allocation: 7.94 KiB)

```

so we can conclude that multiplication by the uniform scaling operator is maximally efficient even when using the `*` operator.
