# Custom differentiation with Flux

**URL:** <https://discourse.julialang.org/t/custom-differentiation-with-flux/130249>\
**Category:** Machine Learning\
**Created:** [June 26, 2025, 6:41pm UTC](https://discourse.julialang.org/t/custom-differentiation-with-flux/130249 "2025-06-26T18:41:59Z")\
**Posts on this page:** 3\
**Page:** 1

<div class="post-metadata">

**Author:** ![Fabrice\_Rosay](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/fabrice_rosay/32/15689_2.png) [@Fabrice\_Rosay](https://discourse.julialang.org/u/Fabrice_Rosay)\
**Post date:** [June 26, 2025, 6:41pm UTC](https://discourse.julialang.org/t/custom-differentiation-with-flux/130249/1 "2025-06-26T18:41:59Z")

</div>

Currently the derivative of `round`is 0. I would like it to be 1 instead and also that it works when broadcasted. So i wrote a custom `my_round`and tried to write the new rule with adjoint and `@scalar_rule` but it does not fully work in particular it never works with CUDA array. Is there any work around ? (for the record I’m trying to implement quantization aware training on simple dense network and miserably failing )

---

<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:** [June 26, 2025, 7:35pm UTC](https://discourse.julialang.org/t/custom-differentiation-with-flux/130249/2 "2025-06-26T19:35:06Z")

</div>

Assuming that you are using Zygote (still Flux’s default), then what you may be missing is that differentiates broadcasts uses ForwardDiff. This is not affected by ChainRules’s `@scalar_rule`. Instead, you would need `round(::Dual)`… something like this?

```julia
julia> using Zygote, ForwardDiff, JLArrays

julia> Zygote.gradient(x -> sum(round.(x ./ 10; digits=2)), randn(3)) # that's a bug
ERROR: MethodError: no method matching iterate(::Nothing)

julia> Zygote.gradient(x -> sum(round.(x ./ 10; digits=2)), jl(randn(3))) # with GPU array
ERROR: MethodError: no method matching round(::ForwardDiff.Dual{Nothing, Float64, 1}, ::RoundingMode{:Nearest})

julia> Base.round(x::ForwardDiff.Dual; kw...) = ForwardDiff.Dual(round(ForwardDiff.value(x); kw...), ForwardDiff.partials(x))

julia> Zygote.gradient(x -> sum(round.(x ./ 10; digits=2)), jl(randn(3)))
([0.1, 0.1, 0.1],)

```

---

<div class="post-metadata">

**Author:** ![Fabrice\_Rosay](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/fabrice_rosay/32/15689_2.png) [@Fabrice\_Rosay](https://discourse.julialang.org/u/Fabrice_Rosay)\
**Post date:** [June 26, 2025, 7:50pm UTC](https://discourse.julialang.org/t/custom-differentiation-with-flux/130249/3 "2025-06-26T19:50:32Z")

</div>

> [@mcabbott](#):
>
> ` Base.round(x::ForwardDiff.Dual; kw...) = ForwardDiff.Dual(round(ForwardDiff.value(x); kw...), ForwardDiff.partials(x))`

It works fine even on gpu by adding the last line, thank you.
