# Zygote Update with Parametric Type

**URL:** <https://discourse.julialang.org/t/zygote-update-with-parametric-type/98588>\
**Category:** Machine Learning\
**Created:** [May 10, 2023, 2:50am UTC](https://discourse.julialang.org/t/zygote-update-with-parametric-type/98588 "2023-05-10T02:50:36Z")\
**Posts on this page:** 6\
**Page:** 1

<div class="post-metadata">

**Author:** ![Ian\_L](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/ian_l/32/49509_2.png) [@Ian\_L](https://discourse.julialang.org/u/Ian_L)\
**Post date:** [May 10, 2023, 2:50am UTC](https://discourse.julialang.org/t/zygote-update-with-parametric-type/98588/1 "2023-05-10T02:50:36Z")

</div>

Hi. I’m trying to update my flux model which is parameterized by two types:

```julia
abstract type C end

struct A <: C end
struct B <: C end

struct Block{T<:C}
    C_inout::Int
    time::Vector{Float64}
end
Flux.@functor Block
Flux.trainable(m::Block) = (diffusion_time = m.time,)

function (model::Block{A})(x)
    0.0
end
function (model::Block{B})(x)
    0.0
end

```

Here is the gradient and update

```julia
grad = gradient(loss, m, x, y)
Flux.update!(opt_state, m, grad[1]) # Breaks!

```

The update step attempts to construct an unparameterized Block:

```julia
MethodError: no method matching Block(::Int64, ::Vector{Float64})

```

What would be the best way about solving this/organizing the codes?

---

<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:** [May 10, 2023, 5:23am UTC](https://discourse.julialang.org/t/zygote-update-with-parametric-type/98588/2 "2023-05-10T05:23:00Z")

</div>

The problem here is that a constructor `Block(C_inout, time)` doesn’t make sense because there is an additional type parameter `T` which cannot be inferred from the attributes. What is it for?

---

<div class="post-metadata">

**Author:** ![Ian\_L](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/ian_l/32/49509_2.png) [@Ian\_L](https://discourse.julialang.org/u/Ian_L)\
**Post date:** [May 10, 2023, 5:55am UTC](https://discourse.julialang.org/t/zygote-update-with-parametric-type/98588/3 "2023-05-10T05:55:12Z")

</div>

I have two possible forward modes - one fast and one slow. In the code above, these correspond to types `A` and `B`. I was hoping to dispatch based of the type `Block{T}` so that the code would avoid branching. I was worried that if I used something like an if-else statement in my forward pass, this would hurt performance… but maybe it isn’t too bad?

---

<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:** [May 10, 2023, 5:59am UTC](https://discourse.julialang.org/t/zygote-update-with-parametric-type/98588/4 "2023-05-10T05:59:02Z")

</div>

If you cannot change the struct, you will have to define a method of `functor` instead of having the macro `@functor` do it for you – the macro does not know about this parameter.

Alternatively, store an instance & the type parameter will take care of itself:

```julia
struct Block{T<:C}
    C_instance::T
    C_inout::Int
    time::Vector{Float64}
end

Flux.@functor Block

Block(A(), 1, [2.0])

```

(Either way, no need to overload `trainable`.)

---

<div class="post-metadata">

**Author:** ![Ian\_L](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/ian_l/32/49509_2.png) [@Ian\_L](https://discourse.julialang.org/u/Ian_L)\
**Post date:** [May 10, 2023, 6:18am UTC](https://discourse.julialang.org/t/zygote-update-with-parametric-type/98588/5 "2023-05-10T06:18:46Z")

</div>

Interesting. How would I dispatch against `C_instance` for the forward pass?

---

<div class="post-metadata">

**Author:** ![Ian\_L](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/ian_l/32/49509_2.png) [@Ian\_L](https://discourse.julialang.org/u/Ian_L)\
**Post date:** [May 10, 2023, 6:26am UTC](https://discourse.julialang.org/t/zygote-update-with-parametric-type/98588/6 "2023-05-10T06:26:22Z")

</div>

Ah nevermind. I see what I missed. Thanks!
