@@ -92,8 +92,10 @@ function _reverse_mode(d::NLPEvaluator, x)
9292 end
9393 # Phase I
9494 for k in d. subexpression_order
95+ _forward_eval (d. subexpressions[k], d, x)
96+ # FIXME this assumes scalar output
9597 d. subexpression_forward_values[k] =
96- _forward_eval ( d. subexpressions[k], d, x)
98+ d. subexpressions[k]. forward_storage[ 1 ]
9799 end
98100 if d. objective != = nothing
99101 _forward_eval (something (d. objective). expr, d, x)
@@ -158,7 +160,7 @@ function _forward_eval(
158160 f:: _SubexpressionStorage ,
159161 d:: NLPEvaluator ,
160162 x:: AbstractVector{T} ,
161- ):: T where {T}
163+ ) where {T}
162164 @assert length (f. forward_storage) >= length (f. nodes)
163165 @assert length (f. partials_storage) >= length (f. nodes)
164166 operators = d. data. operators
@@ -633,10 +635,7 @@ function _forward_eval(
633635 f. partials_storage[rhs] = zero (T)
634636 end
635637 end
636- # Caller is responsible for reading the right range of `f.forward_storage`
637- # for vector-valued roots (use `_storage_range(f.sizes, 1)`); the scalar
638- # return is only meaningful when the root is scalar.
639- return f. forward_storage[1 ]
638+ return
640639end
641640
642641"""
@@ -823,9 +822,9 @@ Reverse-mode evaluation of an expression tree given in `f`.
823822 * This function assumes that `f.reverse_storage` has been initialized with 0.0.
824823"""
825824function _reverse_eval (
826- f:: _SubexpressionStorage ,
827- seed:: Union{Nothing,AbstractVector{Float64 }} = nothing ,
828- )
825+ f:: _SubexpressionStorage{T} ,
826+ seed:: Union{Nothing,AbstractVector{T }} = nothing ,
827+ ) where {T}
829828 @assert length (f. reverse_storage) >= _length (f. sizes)
830829 @assert length (f. partials_storage) >= _length (f. sizes)
831830 # f.nodes is already in order such that parents always appear before
@@ -834,14 +833,10 @@ function _reverse_eval(
834833 children_arr = SparseArrays. rowvals (f. adj)
835834 root_range = _storage_range (f. sizes, 1 )
836835 if seed === nothing
837- for i in root_range
838- f. reverse_storage[i] = one (Float64)
839- end
836+ f. reverse_storage[root_range] .= one (T)
840837 else
841838 @assert length (seed) == length (root_range)
842- for (j, i) in enumerate (root_range)
843- f. reverse_storage[i] = seed[j]
844- end
839+ f. reverse_storage[root_range] .= seed
845840 end
846841 for k in 1 : length (f. nodes)
847842 node = f. nodes[k]
@@ -1108,17 +1103,20 @@ function _reverse_eval(
11081103 elseif op == :^
11091104 # We start with just .^2 to simplify
11101105 @assert f. sizes. ndims[rhs] == 0 " Broadcasted ^ requires scalar exponent"
1111- exp = _getscalar (f. forward_storage, f. sizes, rhs)
1112- # To simplify, so we don't need to compute its derivative
1106+ # If it is a constant, we can just read it from the `const_values` and avoid a GPU->CPU communication
1107+ # We also don't need to compute its derivative
1108+ @assert f. nodes[rhs]. type == NODE_VALUE
1109+ exponent = f. const_values[f. nodes[rhs]. index]
1110+
11131111 @assert f. nodes[rhs]. type == NODE_VALUE
11141112 rev_parent = _view_linear (f. reverse_storage, f. sizes, k)
11151113 rev_child =
11161114 _view_linear (f. reverse_storage, f. sizes, lhs)
1117- if exp == 2
1115+ if exponent == 2
11181116 child =
11191117 _view_linear (f. forward_storage, f. sizes, lhs)
11201118 rev_child .= 2 .* child .* rev_parent
1121- elseif exp == 1
1119+ elseif exponent == 1
11221120 rev_child .= rev_parent
11231121 else
11241122 partial =
0 commit comments