# Error with defining customer gradients in Flux.jl

**URL:** <https://discourse.julialang.org/t/error-with-defining-customer-gradients-in-flux-jl/17769>\
**Category:** Machine Learning\
**Tags:** flux\
**Created:** [November 20, 2018, 12:56pm UTC](https://discourse.julialang.org/t/error-with-defining-customer-gradients-in-flux-jl/17769 "2018-11-20T12:56:11Z")\
**Posts on this page:** 10\
**Page:** 1

<div class="post-metadata">

**Author:** ![xiaodai](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/xiaodai/32/15937_2.png) [@xiaodai](https://discourse.julialang.org/u/xiaodai)\
**Post date:** [November 20, 2018, 12:56pm UTC](https://discourse.julialang.org/t/error-with-defining-customer-gradients-in-flux-jl/17769/1 "2018-11-20T12:56:11Z")

</div>

I am trying to define some simple custom graidents in Flux.jl but am running into issues. Here is an MWE

```julia
# start random parameters
W = param(rand(12))

model(x) = begin
    x*W
end

loss(x, y) = sum((model(x) .- y).^2)

x = rand(12)
y = rand(12)

params = Flux.Params([W])

grads = Flux.gradient(() -> loss(x,y), params)

```

gives error which I can’t seem to understand. I don’t se how I am doing anything different to the [documentation here][https://github.com/FluxML/Flux.jl/blob/master/docs/src/models/basics.md](https://github.com/FluxML/Flux.jl/blob/master/docs/src/models/basics.md))

> MethodError: \*(::LinearAlgebra.Transpose{Float64,Array{Float64,2}}, ::TrackedArray{…,Array{Float64,1}}) is ambiguous. Candidates:  
> \*(x::AbstractArray{T,2} where T, y::TrackedArray{T,1,A} where A where T) in Flux.Tracker at /home/jrun/.julia/packages/Flux/UHjNa/src/tracker/array.jl:281  
> \*(transA::LinearAlgebra.Transpose{#s549,#s548} where #s548\<:AbstractArray{T,2} where #s549, x::AbstractArray{S,1}) where {T, S} in LinearAlgebra at /buildworker/worker/package\_linux64/build/usr/share/julia/stdlib/v1.0/LinearAlgebra/src/matmul.jl:83  
> Possible fix, define  
> \*(::LinearAlgebra.Transpose{#s549,#s548} where #s548\<:AbstractArray{T,2} where #s549, ::TrackedArray{S,1,A} where A)
> 
> Stacktrace:  
> [1] (::getfield(Flux.Tracker, Symbol(“##326#327”)){Array{Float64,2},TrackedArray{…,Array{Float64,1}}})(::TrackedArray{…,Array{Float64,1}}) at /home/jrun/.julia/packages/Flux/UHjNa/src/tracker/array.jl:289  
> [2] back\_(::Flux.Tracker.Grads, ::Flux.Tracker.Call{getfield(Flux.Tracker, Symbol(“##326#327”)){Array{Float64,2},TrackedArray{…,Array{Float64,1}}},Tuple{Nothing,Flux.Tracker.Tracked{Array{Float64,1}}}}, ::TrackedArray{…,Array{Float64,1}}) at /home/jrun/.julia/packages/Flux/UHjNa/src/tracker/back.jl:103  
> [3] back(::Flux.Tracker.Grads, ::Flux.Tracker.Tracked{Array{Float64,1}}, ::TrackedArray{…,Array{Float64,1}}) at /home/jrun/.julia/packages/Flux/UHjNa/src/tracker/back.jl:118  
> [4] (::getfield(Flux.Tracker, Symbol(“##4#5”)){Flux.Tracker.Grads})(::Flux.Tracker.Tracked{Array{Float64,1}}, ::TrackedArray{…,Array{Float64,1}}) at /home/jrun/.julia/packages/Flux/UHjNa/src/tracker/back.jl:106  
> [5] foreach(::Function, ::Tuple{Nothing,Flux.Tracker.Tracked{Array{Float64,1}},Nothing,Nothing}, ::Tuple{Flux.Tracker.Tracked{Nothing},TrackedArray{…,Array{Float64,1}},TrackedArray{…,Array{Float64,1}},Flux.Tracker.Tracked{Nothing}}) at ./abstractarray.jl:1836  
> [6] back\_(::Flux.Tracker.Grads, ::Flux.Tracker.Call{getfield(Flux.Tracker, Symbol(“#back#353”)){4,getfield(Base.Broadcast, Symbol(“##26#28”)){getfield(Base.Broadcast, Symbol(“##5#6”)){getfield(Base.Broadcast, Symbol(“##27#29”)){typeof(-),getfield(Base.Broadcast, Symbol(“##9#10”)){getfield(Base.Broadcast, Symbol(“##9#10”)){getfield(Base.Broadcast, Symbol(“##11#12”))}},getfield(Base.Broadcast, Symbol(“##13#14”)){getfield(Base.Broadcast, Symbol(“##13#14”)){getfield(Base.Broadcast, Symbol(“##15#16”))}},getfield(Base.Broadcast, Symbol(“##5#6”)){getfield(Base.Broadcast, Symbol(“##5#6”)){getfield(Base.Broadcast, Symbol(“##5#6”)){getfield(Base.Broadcast, Symbol(“##3#4”))}}}}},typeof(Base.literal\_pow)},Tuple{Base.RefValue{typeof(^)},TrackedArray{…,Array{Float64,1}},Array{Float64,1},Base.RefValue{Val{2}}}},Tuple{Nothing,Flux.Tracker.Tracked{Array{Float64,1}},Nothing,Nothing}}, ::Array{Float64,1}) at /home/jrun/.julia/packages/Flux/UHjNa/src/tracker/back.jl:106  
> [7] back(::Flux.Tracker.Grads, ::Flux.Tracker.Tracked{Array{Float64,1}}, ::Array{Float64,1}) at /home/jrun/.julia/packages/Flux/UHjNa/src/tracker/back.jl:118  
> [8] (::getfield(Flux.Tracker, Symbol(“##4#5”)){Flux.Tracker.Grads})(::Flux.Tracker.Tracked{Array{Float64,1}}, ::Array{Float64,1}) at /home/jrun/.julia/packages/Flux/UHjNa/src/tracker/back.jl:106  
> [9] foreach at ./abstractarray.jl:1836 [inlined]  
> [10] back\_(::Flux.Tracker.Grads, ::Flux.Tracker.Call{getfield(Flux.Tracker, Symbol(“##299#300”)){TrackedArray{…,Array{Float64,1}}},Tuple{Flux.Tracker.Tracked{Array{Float64,1}}}}, ::Int64) at /home/jrun/.julia/packages/Flux/UHjNa/src/tracker/back.jl:106  
> [11] back(::Flux.Tracker.Grads, ::Flux.Tracker.Tracked{Float64}, ::Int64) at /home/jrun/.julia/packages/Flux/UHjNa/src/tracker/back.jl:118  
> [12] (::getfield(Flux.Tracker, Symbol(“##6#7”)){Flux.Tracker.Params,Flux.Tracker.TrackedReal{Float64}})(::Int64) at /home/jrun/.julia/packages/Flux/UHjNa/src/tracker/back.jl:131  
> [13] gradient(::Function, ::Flux.Tracker.Params) at /home/jrun/.julia/packages/Flux/UHjNa/src/tracker/back.jl:152  
> [14] top-level scope at In[139]:1

---

<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:** [November 20, 2018, 1:34pm UTC](https://discourse.julialang.org/t/error-with-defining-customer-gradients-in-flux-jl/17769/2 "2018-11-20T13:34:44Z")

</div>

Wildly guessing, but maybe `x*W` should be `W*x`?

---

<div class="post-metadata">

**Author:** ![xiaodai](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/xiaodai/32/15937_2.png) [@xiaodai](https://discourse.julialang.org/u/xiaodai)\
**Post date:** [November 20, 2018, 1:41pm UTC](https://discourse.julialang.org/t/error-with-defining-customer-gradients-in-flux-jl/17769/3 "2018-11-20T13:41:19Z")

</div>

Nope, cos `loss(x,y)` works

---

<div class="post-metadata">

**Author:** ![mohamed82008](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/mohamed82008/32/18171_2.png) [@mohamed82008](https://discourse.julialang.org/u/mohamed82008)\
**Post date:** [November 20, 2018, 1:50pm UTC](https://discourse.julialang.org/t/error-with-defining-customer-gradients-in-flux-jl/17769/5 "2018-11-20T13:50:43Z")

</div>

I can reproduce the error with the following:

```julia
k = 2
n = 12

W = param(rand(k))
model(x) = begin
    x*W
end

loss(x, y) = sum((model(x) .- y).^2)

x = rand(n, k)
y = rand(n)

params = Flux.Params([W])
grads = Flux.Tracker.gradient(() -> loss(x,y), params)

```

---

<div class="post-metadata">

**Author:** ![xiaodai](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/xiaodai/32/15937_2.png) [@xiaodai](https://discourse.julialang.org/u/xiaodai)\
**Post date:** [November 20, 2018, 2:00pm UTC](https://discourse.julialang.org/t/error-with-defining-customer-gradients-in-flux-jl/17769/6 "2018-11-20T14:00:02Z")

</div>

The key point is also that `loss(x,y)` works even in your example . But when backing out the gradient, it doesn’t work. I am scratching my head still.

---

<div class="post-metadata">

**Author:** ![xiaodai](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/xiaodai/32/15937_2.png) [@xiaodai](https://discourse.julialang.org/u/xiaodai)\
**Post date:** [November 20, 2018, 2:03pm UTC](https://discourse.julialang.org/t/error-with-defining-customer-gradients-in-flux-jl/17769/7 "2018-11-20T14:03:47Z")

</div>

Could be a genuine bug. I go report it

---

<div class="post-metadata">

**Author:** ![Rademcaher](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/rademcaher/32/7420_2.png) [@Rademcaher](https://discourse.julialang.org/u/Rademcaher)\
**Post date:** [November 20, 2018, 4:01pm UTC](https://discourse.julialang.org/t/error-with-defining-customer-gradients-in-flux-jl/17769/8 "2018-11-20T16:01:15Z")

</div>

Is that what you are looking to do?

```julia
W = param(rand(12))

predict(x) = W.*x

function loss(x, y)
   ŷ = predict(x)
    sum( (y - ŷ).^2 )
end

x = rand(12)
y = rand(12)

grads = Tracker.gradient(() -> loss(x,y), Params([W]))

```

---

<div class="post-metadata">

**Author:** ![crinders](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/crinders/32/8434_2.png) [@crinders](https://discourse.julialang.org/u/crinders)\
**Post date:** [November 20, 2018, 5:18pm UTC](https://discourse.julialang.org/t/error-with-defining-customer-gradients-in-flux-jl/17769/9 "2018-11-20T17:18:12Z")

</div>

Notice that the example is one-dimensional. `predict` is a single predictor. Multiplication of `W*x` for scalars is defined, but not for arrays.

```julia
julia> W*x
ERROR: MethodError: no method matching _forward(::typeof(*), ::TrackedArray{…,Array{Float64,1}}, ::Array{Float64,1})
Closest candidates are:
  _forward(::typeof(*), ::AbstractArray{T,2} where T, ::Union{AbstractArray{T,1}, AbstractArray{T,2}} where T) at /Users/berend/.julia/packages/Flux/UHjNa/src/tracker/array.jl:288
  _forward(::typeof(getindex), ::AbstractArray, ::Any...) at /Users/berend/.julia/packages/Flux/UHjNa/src/tracker/array.jl:74
  _forward(::typeof(vcat), ::Any...) at /Users/berend/.julia/packages/Flux/UHjNa/src/tracker/array.jl:136
  ...
Stacktrace:
 [1] #track#1(::Base.Iterators.Pairs{Union{},Union{},Tuple{},NamedTuple{(),Tuple{}}}, ::Function, ::typeof(*), ::TrackedArray{…,Array{Float64,1}}, ::Vararg{Any,N} where N) at /Users/berend/.julia/packages/Flux/UHjNa/src/tracker/Tracker.jl:50
 [2] track(::typeof(*), ::TrackedArray{…,Array{Float64,1}}, ::Array{Float64,1}) at /Users/berend/.julia/packages/Flux/UHjNa/src/tracker/Tracker.jl:50
 [3] *(::TrackedArray{…,Array{Float64,1}}, ::Array{Float64,1}) at /Users/berend/.julia/packages/Flux/UHjNa/src/tracker/array.jl:284
 [4] top-level scope at none:0

```

Instead use `W'*x`.

```julia
julia> W'*x
3.1105292208823814 (tracked)

```

---

<div class="post-metadata">

**Author:** ![crinders](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/crinders/32/8434_2.png) [@crinders](https://discourse.julialang.org/u/crinders)\
**Post date:** [November 20, 2018, 6:30pm UTC](https://discourse.julialang.org/t/error-with-defining-customer-gradients-in-flux-jl/17769/10 "2018-11-20T18:30:53Z")

</div>

Also your loss “worked” because the broadcast `x .+ y` adds a scalar constant to each element of an array independent of whether `x` or `y` is the array.

---

<div class="post-metadata">

**Author:** ![xiaodai](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/xiaodai/32/15937_2.png) [@xiaodai](https://discourse.julialang.org/u/xiaodai)\
**Post date:** [November 21, 2018, 7:21am UTC](https://discourse.julialang.org/t/error-with-defining-customer-gradients-in-flux-jl/17769/11 "2018-11-21T07:21:12Z")

</div>

> [@Rademcaher](#):
>
> W.\*x

Thanks but I want vector product not element wise product.
