Linux GNU 11.4.0 Code Coverage Report


Directory: ./
Coverage: low: ≥ 0% medium: ≥ 75.0% high: ≥ 90.0%
Coverage Exec / Excl / Total
Lines: 57.2% 676 / 0 / 1182
Functions: -% 0 / 1 / 1
Branches: 50.0% 342 / 0 / 684

OMCompiler/Compiler/NBackEnd/Util/NBDifferentiate.mo
Line Branch Exec Source
1 /*
2 * This file is part of OpenModelica.
3 *
4 * Copyright (c) 1998-2026, Open Source Modelica Consortium (OSMC),
5 * c/o Linköpings universitet, Department of Computer and Information Science,
6 * SE-58183 Linköping, Sweden.
7 *
8 * All rights reserved.
9 *
10 * THIS PROGRAM IS PROVIDED UNDER THE TERMS OF AGPL VERSION 3 LICENSE OR
11 * THIS OSMC PUBLIC LICENSE (OSMC-PL) VERSION 1.8.
12 * ANY USE, REPRODUCTION OR DISTRIBUTION OF THIS PROGRAM CONSTITUTES
13 * RECIPIENT'S ACCEPTANCE OF THE OSMC PUBLIC LICENSE OR THE GNU AGPL
14 * VERSION 3, ACCORDING TO RECIPIENTS CHOICE.
15 *
16 * The OpenModelica software and the OSMC (Open Source Modelica Consortium)
17 * Public License (OSMC-PL) are obtained from OSMC, either from the above
18 * address, from the URLs:
19 * http://www.openmodelica.org or
20 * https://github.com/OpenModelica/ or
21 * http://www.ida.liu.se/projects/OpenModelica,
22 * and in the OpenModelica distribution.
23 *
24 * GNU AGPL version 3 is obtained from:
25 * https://www.gnu.org/licenses/licenses.html#GPL
26 *
27 * This program is distributed WITHOUT ANY WARRANTY; without
28 * even the implied warranty of MERCHANTABILITY or FITNESS
29 * FOR A PARTICULAR PURPOSE, EXCEPT AS EXPRESSLY SET FORTH
30 * IN THE BY RECIPIENT SELECTED SUBSIDIARY LICENSE CONDITIONS OF OSMC-PL.
31 *
32 * See the full OSMC Public License conditions for more details.
33 *
34 */
35
36 encapsulated package NBDifferentiate
37 "file: NBDifferentiate.mo
38 package: NBDifferentiate
39 description: This file contains the functions to differentiate equations and
40 expressions symbolically.
41 "
42 public
43 // OF imports
44 import Absyn.Path;
45 import AbsynUtil;
46
47 // NF imports
48 import Algorithm = NFAlgorithm;
49 import Binding = NFBinding;
50 import BuiltinFuncs = NFBuiltinFuncs;
51 import Call = NFCall;
52 import Class = NFClass;
53 import Restriction = NFRestriction;
54 import NFClassTree.ClassTree;
55 import Component = NFComponent;
56 import ComponentRef = NFComponentRef;
57 import Dimension = NFDimension;
58 import Expression = NFExpression;
59 import InstContext = NFInstContext;
60 import NFInstNode.{InstNode, CachedData};
61 import NFFunction.{Function, Slot};
62 import FunctionDerivative = NFFunctionDerivative;
63 import Operator = NFOperator;
64 import Prefixes = NFPrefixes;
65 import Sections = NFSections;
66 import SimplifyExp = NFSimplifyExp;
67 import Statement = NFStatement;
68 import Subscript = NFSubscript;
69 import Type = NFType;
70 import NFPrefixes.Variability;
71 import Variable = NFVariable;
72
73 // Backend imports
74 import NFBackendExtension.BackendInfo;
75 import NBEquation.{Equation, EquationAttributes, EquationPointer, EquationPointers, IfEquationBody, WhenEquationBody, WhenStatement};
76 import NBVariable.{VariablePointer};
77 import BVariable = NBVariable;
78 import Replacements = NBReplacements;
79 import StrongComponent = NBStrongComponent;
80 import Tearing = NBTearing;
81
82 // Util imports
83 import Array;
84 import BackendUtil = NBBackendUtil;
85 import Error;
86 import UnorderedMap;
87 import Slice = NBSlice;
88
89 protected
90 import NFFunction;
91 import NFPrefixes;
92
93 public
94 // ================================
95 // TYPES AND UNIONTYPES
96 // ================================
97 type DifferentiationType = enumeration(TIME, SIMPLE, FUNCTION, JACOBIAN);
98
99 uniontype DifferentiationArguments
100 record DIFFERENTIATION_ARGUMENTS
101 ComponentRef diffCref "The input will be differentiated w.r.t. this cref (only SIMPLE).";
102 list<Pointer<Variable>> new_vars "contains all new variables that need to be added to the system";
103 Option<UnorderedMap<ComponentRef, ComponentRef>> diff_map "seed and temporary cref map x --> $SEED.MATRIX.x, y --> $pDer.MATRIX.y. Can be used for any differentiation rules";
104 DifferentiationType diffType "Differentiation use case (time, simple, function, jacobian)";
105 UnorderedMap<Path, Function> funcMap "Function tree containing all functions and their known derivatives";
106 Boolean scalarized "true if the variables are scalarized";
107 Option<UnorderedMap<ComponentRef, list<Expression>>> adjoint_map "map for accumulating adjoint gradients for component refs";
108 Expression current_grad "current gradient expression, used in reverse mode";
109 Boolean collectAdjoints "If false, skip writing into adjoint_map (used for LHS traversal in reverse/Jacobian).";
110 end DIFFERENTIATION_ARGUMENTS;
111
112 function default
113 input DifferentiationType ty = DifferentiationType.TIME;
114 input UnorderedMap<Path, Function> funcMap = UnorderedMap.new<Function>(AbsynUtil.pathHash, AbsynUtil.pathEqual);
115 output DifferentiationArguments diffArgs = DIFFERENTIATION_ARGUMENTS(
116 diffCref = ComponentRef.EMPTY(),
117 new_vars = {},
118 diff_map = NONE(),
119 diffType = ty,
120 funcMap = funcMap,
121 scalarized = false,
122 adjoint_map = NONE(),
123 current_grad= Expression.EMPTY(Type.REAL()),
124 collectAdjoints = false
125 );
126 end default;
127
128 function simpleCref "Differentiate w.r.t. cref"
129 input ComponentRef cref;
130 input UnorderedMap<Path, Function> funcMap = UnorderedMap.new<Function>(AbsynUtil.pathHash, AbsynUtil.pathEqual);
131 output DifferentiationArguments diffArgs = DIFFERENTIATION_ARGUMENTS(
132 diffCref = cref,
133 new_vars = {},
134 diff_map = NONE(),
135 diffType = DifferentiationType.SIMPLE,
136 funcMap = funcMap,
137 scalarized = false,
138 adjoint_map = NONE(),
139 current_grad = Expression.EMPTY(Type.REAL()),
140 collectAdjoints = false
141 );
142 end simpleCref;
143
144 function toString
145 input DifferentiationArguments diffArgs;
146 output String str = "[" + diffTypeStr(diffArgs.diffType) + "]";
147 algorithm
148 ✗ if diffArgs.diffType == DifferentiationType.SIMPLE then
149 ✗ str := str + " " + ComponentRef.toString(diffArgs.diffCref);
150 end if;
151 end toString;
152
153 function diffTypeStr
154 input DifferentiationType diffType;
155 output String str;
156 algorithm
157 str := match diffType
158 case DifferentiationType.TIME then "TIME";
159 case DifferentiationType.SIMPLE then "SIMPLE";
160 case DifferentiationType.FUNCTION then "FUNCTION";
161 case DifferentiationType.JACOBIAN then "JACOBIAN";
162 else "FAIL";
163 end match;
164 end diffTypeStr;
165 end DifferentiationArguments;
166
167 // ================================
168 // FUNCTIONS
169 // ================================
170
171 function differentiateStrongComponentList
172 "author: kabdelhak
173 Differentiates a list of strong components."
174 input output list<StrongComponent> comps;
175 input output DifferentiationArguments diffArguments;
176 input Pointer<Integer> idx;
177 input String context;
178 input String name;
179 protected
180 Pointer<DifferentiationArguments> diffArguments_ptr = Pointer.create(diffArguments);
181 algorithm
182 86 comps := List.map(comps, function differentiateStrongComponent(diffArguments_ptr = diffArguments_ptr, idx = idx, context = context, name = name));
183 85 diffArguments := Pointer.access(diffArguments_ptr);
184 end differentiateStrongComponentList;
185
186 function differentiateStrongComponent
187 input output StrongComponent comp;
188 input Pointer<DifferentiationArguments> diffArguments_ptr;
189 input Pointer<Integer> idx;
190 input String context;
191 input String name;
192 algorithm
193 comp := match comp
194 local
195 Pointer<Variable> new_var;
196 Pointer<Equation> new_eqn;
197 list<Slice<VariablePointer>> new_var_slices;
198 ComponentRef new_cref;
199 Slice<VariablePointer> new_var_slice;
200 Slice<EquationPointer> new_eqn_slice;
201 DifferentiationArguments diffArguments;
202 Tearing strict;
203 Option<Tearing> casual;
204 Boolean linear;
205
206 case StrongComponent.SINGLE_COMPONENT() algorithm
207 854 new_var := differentiateVariablePointer(comp.var, diffArguments_ptr);
208 854 new_eqn := differentiateEquationPointer(comp.eqn, diffArguments_ptr, name);
209 853 Equation.createName(new_eqn, idx, context);
210 853 then StrongComponent.SINGLE_COMPONENT(new_var, new_eqn, comp.status);
211
212 case StrongComponent.MULTI_COMPONENT() algorithm
213 ✗ new_var_slices := list(Slice.apply(var, function differentiateVariablePointer(diffArguments_ptr = diffArguments_ptr)) for var in comp.vars);
214 ✗ new_eqn_slice := Slice.apply(comp.eqn, function differentiateEquationPointer(diffArguments_ptr = diffArguments_ptr, name = name));
215 ✗ Equation.createName(Slice.getT(new_eqn_slice), idx = idx, context = context);
216 ✗ then StrongComponent.MULTI_COMPONENT(new_var_slices, new_eqn_slice, comp.status);
217
218 case StrongComponent.SLICED_COMPONENT() algorithm
219 // Map the subscripted LHS cref without collecting into the adjoint_map if one exists
220
1/2
✗ Branch 3 not taken.
✓ Branch 4 taken 214 times.
214 (Expression.CREF(cref = new_cref), diffArguments) := differentiateComponentRefNoCollect(Expression.fromCref(comp.var_cref), Pointer.access(diffArguments_ptr));
221 214 Pointer.update(diffArguments_ptr, diffArguments);
222 214 new_var_slice := Slice.apply(comp.var, function differentiateVariablePointer(diffArguments_ptr = diffArguments_ptr));
223 214 new_eqn_slice := Slice.apply(comp.eqn, function differentiateEquationPointer(diffArguments_ptr = diffArguments_ptr, name = name));
224 214 Slice.applyMutable(new_eqn_slice, function Equation.createName(idx = idx, context = context));
225 214 then StrongComponent.SLICED_COMPONENT(new_cref, new_var_slice, new_eqn_slice, comp.status);
226
227 case StrongComponent.RESIZABLE_COMPONENT() algorithm
228
1/2
✗ Branch 3 not taken.
✓ Branch 4 taken 25 times.
25 (Expression.CREF(cref = new_cref), diffArguments) := differentiateComponentRef(Expression.fromCref(comp.var_cref), Pointer.access(diffArguments_ptr));
229 25 Pointer.update(diffArguments_ptr, diffArguments);
230 25 new_var_slice := Slice.apply(comp.var, function differentiateVariablePointer(diffArguments_ptr = diffArguments_ptr));
231 25 new_eqn_slice := Slice.apply(comp.eqn, function differentiateEquationPointer(diffArguments_ptr = diffArguments_ptr, name = name));
232 25 Slice.applyMutable(new_eqn_slice, function Equation.createName(idx = idx, context = context));
233 25 then StrongComponent.RESIZABLE_COMPONENT(new_cref, new_var_slice, new_eqn_slice, comp.order, comp.status);
234
235 case StrongComponent.GENERIC_COMPONENT() algorithm
236
1/2
✗ Branch 3 not taken.
✓ Branch 4 taken 37 times.
37 (Expression.CREF(cref = new_cref), diffArguments) := differentiateComponentRef(Expression.fromCref(comp.var_cref), Pointer.access(diffArguments_ptr));
237 37 Pointer.update(diffArguments_ptr, diffArguments);
238 37 new_var_slice := Slice.apply(comp.var, function differentiateVariablePointer(diffArguments_ptr = diffArguments_ptr));
239 37 new_eqn_slice := Slice.apply(comp.eqn, function differentiateEquationPointer(diffArguments_ptr = diffArguments_ptr, name = name));
240 37 Slice.applyMutable(new_eqn_slice, function Equation.createName(idx = idx, context = context));
241 37 then StrongComponent.GENERIC_COMPONENT(new_cref, new_var_slice, new_eqn_slice);
242
243 case StrongComponent.ALGEBRAIC_LOOP() algorithm
244 1 strict := differentiateTearing(comp.strict, diffArguments_ptr, idx, context, name);
245 1 casual := Util.applyOption(comp.casual, function differentiateTearing(diffArguments_ptr=diffArguments_ptr, idx=idx, context=context, name=name));
246 // if we differentiate for jacobian, the algebraic loops will always be linear
247 ✗ linear := match Pointer.access(diffArguments_ptr) case DIFFERENTIATION_ARGUMENTS(diffType = NBDifferentiate.DifferentiationType.JACOBIAN) then true; else comp.linear; end match;
248
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 1 time.
1 then StrongComponent.ALGEBRAIC_LOOP(-1, strict, casual, linear, false, comp.homotopy, comp.status, comp.implicitlyCreated);
249
250 case StrongComponent.ENTWINED_COMPONENT() algorithm
251 ✗ Error.addMessage(Error.INTERNAL_ERROR,{getInstanceName() + " not implemented for entwined equation:\n" + StrongComponent.toString(comp)});
252 ✗ then fail();
253
254 10 case StrongComponent.ALIAS() then differentiateStrongComponent(comp.original, diffArguments_ptr, idx, context, name);
255
256 else algorithm
257 ✗ Error.addMessage(Error.INTERNAL_ERROR,{getInstanceName() + " not implemented for unknown strong component:\n" + StrongComponent.toString(comp)});
258 ✗ then fail();
259 end match;
260 end differentiateStrongComponent;
261
262 function differentiateTearing
263 input Tearing tearing;
264 input Pointer<DifferentiationArguments> diffArguments_ptr;
265 input Pointer<Integer> idx;
266 input String context;
267 input String name;
268 output Tearing diff_tearing;
269 protected
270 list<Slice<VariablePointer>> ite_vars;
271 list<Slice<EquationPointer>> res_eqns;
272 array<StrongComponent> inner_eqns;
273 algorithm
274
4/4
✓ Branch 0 taken 2 times.
✓ Branch 1 taken 1 time.
✓ Branch 2 taken 2 times.
✓ Branch 3 taken 1 time.
3 ite_vars := list(Slice.apply(var, function differentiateVariablePointer(diffArguments_ptr = diffArguments_ptr)) for var in tearing.iteration_vars);
275
4/4
✓ Branch 0 taken 2 times.
✓ Branch 1 taken 1 time.
✓ Branch 2 taken 2 times.
✓ Branch 3 taken 1 time.
3 res_eqns := list(Slice.apply(eqn, function differentiateEquationPointer(diffArguments_ptr = diffArguments_ptr, name = name)) for eqn in tearing.residual_eqns);
276 // Only differentiate continuous inner equations; discrete ones contribute zero to the Jacobian.
277
2/6
✗ Branch 2 not taken.
✗ Branch 3 not taken.
✗ Branch 4 not taken.
✓ Branch 5 taken 1 time.
✗ Branch 6 not taken.
✓ Branch 7 taken 1 time.
1 inner_eqns := listArray(list(differentiateStrongComponent(ie, diffArguments_ptr, idx, context, name) for ie guard(not StrongComponent.isDiscrete(ie)) in arrayList(tearing.innerEquations)));
278
279 1 diff_tearing := Tearing.TEARING_SET(ite_vars, res_eqns, inner_eqns, NONE());
280 end differentiateTearing;
281
282 function differentiateEquationPointerList
283 "author: kabdelhak
284 Differentiates a list of equations wrapped in pointers."
285 input output list<Pointer<Equation>> equations;
286 input output DifferentiationArguments diffArguments;
287 input Pointer<Integer> idx;
288 input String context;
289 input String name;
290 protected
291 Pointer<DifferentiationArguments> diffArguments_ptr = Pointer.create(diffArguments);
292 algorithm
293 ✗ equations := List.map(equations, function differentiateEquationPointer(diffArguments_ptr = diffArguments_ptr, name = name));
294 ✗ for eqn in equations loop
295 ✗ Equation.createName(eqn, idx, context);
296 end for;
297 ✗ diffArguments := Pointer.access(diffArguments_ptr);
298 end differentiateEquationPointerList;
299
300 function differentiateEquationPointer
301 input Pointer<Equation> eq_ptr;
302 input Pointer<DifferentiationArguments> diffArguments_ptr;
303 input String name = "";
304 output Pointer<Equation> derivative_ptr;
305 protected
306 Equation eq, diffedEq;
307 DifferentiationArguments old_diffArguments, new_diffArguments;
308 algorithm
309 1350 eq := Pointer.access(eq_ptr);
310 1350 old_diffArguments := Pointer.access(diffArguments_ptr);
311
312 derivative_ptr := match Equation.getAttributes(eq)
313
314 // we differentiate w.r.t time and there already is a derivative saved
315 case EquationAttributes.EQUATION_ATTRIBUTES(derivative = SOME(derivative_ptr))
316 guard(old_diffArguments.diffType == DifferentiationType.TIME)
317 then derivative_ptr;
318
319 // else differentiate the equation
320 else algorithm
321 1350 (diffedEq, new_diffArguments) := differentiateEquation(eq, old_diffArguments, name);
322 1349 derivative_ptr := Pointer.create(diffedEq);
323 // save the derivative if we derive w.r.t. time
324
2/2
✓ Branch 0 taken 218 times.
✓ Branch 1 taken 1131 times.
1349 if new_diffArguments.diffType == DifferentiationType.TIME then
325 218 Pointer.update(eq_ptr, Equation.setDerivative(eq, derivative_ptr));
326 end if;
327
2/2
✓ Branch 0 taken 147 times.
✓ Branch 1 taken 1202 times.
1349 if not referenceEq(new_diffArguments, old_diffArguments) then
328 147 Pointer.update(diffArguments_ptr, new_diffArguments);
329 end if;
330 then derivative_ptr;
331 end match;
332 end differentiateEquationPointer;
333
334 function differentiateEquation
335 input output Equation eq;
336 input output DifferentiationArguments diffArguments;
337 input String name = "";
338 algorithm
339
3/6
✓ Branch 1 taken 7 times.
✓ Branch 2 taken 1461 times.
✗ Branch 3 not taken.
✓ Branch 4 taken 7 times.
✗ Branch 6 not taken.
✗ Branch 7 not taken.
1468 if Flags.isSet(Flags.DEBUG_DIFFERENTIATION) and not stringEqual(name, "") then
340 7 print("### debugDifferentiation | " + name + " ###\n");
341 7 print("[BEFORE] " + Equation.toString(eq) + "\n");
342 end if;
343 (eq, diffArguments) := match eq
344 local
345 Expression lhs, rhs;
346 list<Equation> forBody = {};
347 IfEquationBody ifBody;
348 WhenEquationBody whenBody;
349 Pointer<DifferentiationArguments> diffArguments_ptr;
350 EquationAttributes attr;
351 Algorithm alg;
352
353 // ToDo: Element source stuff (see old backend)
354 case Equation.SCALAR_EQUATION() algorithm
355 1280 (lhs, diffArguments) := differentiateExpressionNoCollect(eq.lhs, diffArguments);
356 1280 (rhs, diffArguments) := differentiateExpression(eq.rhs, diffArguments);
357 1279 attr := differentiateEquationAttributes(eq.attr, diffArguments);
358 1279 then (Equation.SCALAR_EQUATION(eq.ty, lhs, rhs, eq.source, attr), diffArguments);
359
360 case Equation.ARRAY_EQUATION() algorithm
361 70 (lhs, diffArguments) := differentiateExpressionNoCollect(eq.lhs, diffArguments);
362 70 (rhs, diffArguments) := differentiateExpression(eq.rhs, diffArguments);
363 70 attr := differentiateEquationAttributes(eq.attr, diffArguments);
364 70 then (Equation.ARRAY_EQUATION(eq.ty, lhs, rhs, eq.source, attr, eq.recordSize), diffArguments);
365
366 case Equation.RECORD_EQUATION() algorithm
367 ✗ (lhs, diffArguments) := differentiateExpressionNoCollect(eq.lhs, diffArguments);
368 ✗ (rhs, diffArguments) := differentiateExpression(eq.rhs, diffArguments);
369 ✗ attr := differentiateEquationAttributes(eq.attr, diffArguments);
370 ✗ then (Equation.RECORD_EQUATION(eq.ty, lhs, rhs, eq.source, attr, eq.recordSize), diffArguments);
371
372 case Equation.IF_EQUATION() algorithm
373 ✗ (ifBody, diffArguments_ptr) := differentiateIfEquationBody(eq.body, Pointer.create(diffArguments));
374 ✗ attr := differentiateEquationAttributes(eq.attr, diffArguments);
375 ✗ then (Equation.IF_EQUATION(eq.size, ifBody, eq.source, attr), Pointer.access(diffArguments_ptr));
376
377 case Equation.FOR_EQUATION() algorithm
378
2/2
✓ Branch 0 taken 118 times.
✓ Branch 1 taken 118 times.
236 for body_eqn in eq.body loop
379 118 (body_eqn, diffArguments) := differentiateEquation(body_eqn, diffArguments);
380 forBody := body_eqn :: forBody;
381 end for;
382 118 attr := differentiateEquationAttributes(eq.attr, diffArguments);
383 118 then (Equation.FOR_EQUATION(eq.size,
384 eq.iter,
385 listReverse(forBody),
386 eq.source,
387 attr),
388 diffArguments);
389
390 case Equation.WHEN_EQUATION() algorithm
391 ✗ (whenBody, diffArguments) := differentiateWhenEquationBody(eq.body, diffArguments);
392 ✗ attr := differentiateEquationAttributes(eq.attr, diffArguments);
393 ✗ then (Equation.WHEN_EQUATION(eq.size, whenBody, eq.source, attr), diffArguments);
394
395 case Equation.ALGORITHM() algorithm
396 ✗ (alg, diffArguments) := differentiateAlgorithm(eq.alg, diffArguments); // may need differentiateAlgorithmAdjoint
397 ✗ then (Equation.ALGORITHM(eq.size, alg, eq.source, eq.expand, eq.attr), diffArguments);
398
399 else algorithm
400 // maybe add failtrace here and allow failing
401 ✗ Error.addMessage(Error.INTERNAL_ERROR,{getInstanceName() + " failed for: " + Equation.toString(eq)});
402 ✗ then fail();
403
404 end match;
405
406 /* ToDo
407 record AUX_EQUATION
408 "Auxiliary equations are generated when auxiliary variables are generated
409 that are known to always be solved in this specific equation. E.G. $CSE
410 The variable binding contains the equation, but this equation is also
411 allowed to have a body for special cases."
412 Pointer<Variable> auxiliary "Corresponding auxiliary variable";
413 Option<Equation> body "Optional body equation"; // -> Expression
414 end AUX_EQUATION;
415
416 record DUMMY_EQUATION
417 end DUMMY_EQUATION;
418
419 */
420
3/6
✓ Branch 1 taken 7 times.
✓ Branch 2 taken 1460 times.
✗ Branch 3 not taken.
✓ Branch 4 taken 7 times.
✗ Branch 6 not taken.
✗ Branch 7 not taken.
1467 if Flags.isSet(Flags.DEBUG_DIFFERENTIATION) and not stringEqual(name, "") then
421 7 eq := Equation.simplify(eq, name, "\t");
422 7 print("[AFTER ] " + Equation.toString(eq) + "\n\n");
423 else
424 1460 eq := Equation.simplify(eq, name);
425 end if;
426 end differentiateEquation;
427
428 function differentiateEquationAdjoint
429 "Adjoint-mode equation differentiation.
430 Given an equation and a DifferentiationArguments with a fresh adjoint_map,
431 populates the map via reverse-mode differentiation of the RHS, then emits
432 accumulation statements (v := v + sum(M[v])) and a seed reset (seed := 0).
433 Returns the list of adjoint statements and updated diffArguments."
434 input Equation eq;
435 input output DifferentiationArguments diffArguments;
436 output list<Statement> adjointStatements;
437 algorithm
438 (diffArguments, adjointStatements) := match eq
439
440 local
441 Expression lhs;
442 ComponentRef lhsCref, seedCref;
443 UnorderedMap<ComponentRef, ComponentRef> dm;
444 list<Statement> stmts;
445
446 // For-equation locals
447 list<Statement> bodyStmts, allStmts;
448
449 // If-equation locals
450 list<tuple<Expression, list<Statement>>> ifBranches;
451 Option<list<tuple<Expression, list<Statement>>>> elseIfBranches;
452
453 // Array equation locals
454 ComponentRef lhs_base;
455 ComponentRef seed_base;
456 Expression seed_subscripted;
457
458 // For-equation iterator locals
459 list<ComponentRef> iterNames;
460 list<Expression> iterRanges;
461 list<Option<NBEquation.Iterator>> iterMaps;
462 ComponentRef iterName;
463 Expression iterRange;
464 Option<NBEquation.Iterator> iterMap;
465 list<tuple<ComponentRef, array<Expression>>> sub_iters_stmt;
466 NBEquation.Iterator revIter;
467 ComponentRef iter_name;
468 array<Expression> iter_elems;
469
470 // ===================== SCALAR_EQUATION (Assignment) =====================
471 case Equation.SCALAR_EQUATION() algorithm
472
2/4
✗ Branch 0 not taken.
✓ Branch 1 taken 16 times.
✗ Branch 2 not taken.
✓ Branch 3 taken 16 times.
16 SOME(dm) := diffArguments.diff_map;
473 16 lhsCref := Expression.toCref(eq.lhs);
474
475 // Check if LHS variable is in the diff_map; if so, get seed cref and differentiate RHS, else skip
476
2/4
✓ Branch 1 taken 16 times.
✗ Branch 2 not taken.
✓ Branch 5 taken 16 times.
✗ Branch 6 not taken.
16 if (not ComponentRef.isEmpty(lhsCref)) and UnorderedMap.contains(ComponentRef.stripSubscriptsAll(lhsCref), dm) then
477 16 seedCref := UnorderedMap.getOrFail(ComponentRef.stripSubscriptsAll(lhsCref), dm);
478
1/2
✓ Branch 0 taken 16 times.
✗ Branch 1 not taken.
16 if not diffArguments.scalarized then
479 // Strip subscripts from base seed before copying: diff_map[base] may store a
480 // subscripted element seed (partial-slice NLS). Stripping lets copySubscripts
481 // place the origin subscripts onto an unsubscripted template without conflict.
482 16 seedCref := ComponentRef.copySubscripts(lhsCref, ComponentRef.stripSubscriptsAll(seedCref));
483 end if;
484
485 // Set seed in diffArguments and differentiate RHS
486 16 diffArguments.current_grad := Expression.fromCref(seedCref);
487 16 diffArguments.collectAdjoints := true;
488 16 (_, diffArguments) := differentiateExpression(eq.rhs, diffArguments);
489
490 // After differentiating RHS, emit accumulation statements from adjoint_map
491 16 (diffArguments, stmts) := makeAdjointAccumulationStatements(diffArguments);
492 // unneccesary reverse: stmts := listReverse(stmts);
493 else
494 ✗ stmts := {};
495 end if;
496 16 then (diffArguments, stmts);
497
498 // ===================== ARRAY_EQUATION (Assignment) =====================
499 case Equation.ARRAY_EQUATION() algorithm
500
2/4
✗ Branch 0 not taken.
✓ Branch 1 taken 1 time.
✗ Branch 2 not taken.
✓ Branch 3 taken 1 time.
1 SOME(dm) := diffArguments.diff_map;
501 1 lhs_base := Expression.toCref(eq.lhs);
502
503
2/4
✓ Branch 1 taken 1 time.
✗ Branch 2 not taken.
✓ Branch 5 taken 1 time.
✗ Branch 6 not taken.
1 if (not ComponentRef.isEmpty(lhs_base)) and UnorderedMap.contains(ComponentRef.stripSubscriptsAll(lhs_base), dm) then
504 1 seed_base := UnorderedMap.getOrFail(ComponentRef.stripSubscriptsAll(lhs_base), dm);
505 // Strip subscripts from base seed before applying: diff_map[base] may store
506 // a subscripted element seed (partial-slice NLS). applySubscripts then
507 // merges the lhs subscripts onto an unsubscripted template correctly.
508 1 seed_subscripted := Expression.applySubscripts(
509 ComponentRef.subscriptsAllFlat(lhs_base),
510 Expression.fromCref(ComponentRef.stripSubscriptsAll(seed_base)),
511 true);
512 1 diffArguments.current_grad := seed_subscripted;
513 1 diffArguments.collectAdjoints := true;
514 1 (_, diffArguments) := differentiateExpression(eq.rhs, diffArguments);
515 1 (diffArguments, stmts) := makeAdjointAccumulationStatements(diffArguments);
516 else
517 ✗ stmts := {};
518 end if;
519 1 then (diffArguments, stmts);
520
521 // ===================== RECORD_EQUATION (same as Scalar) =====================
522 case Equation.RECORD_EQUATION() algorithm
523 ✗ SOME(dm) := diffArguments.diff_map;
524 ✗ lhsCref := Expression.toCref(eq.lhs);
525
526 ✗ if (not ComponentRef.isEmpty(lhsCref)) and UnorderedMap.contains(ComponentRef.stripSubscriptsAll(lhsCref), dm) then
527 ✗ seedCref := UnorderedMap.getOrFail(ComponentRef.stripSubscriptsAll(lhsCref), dm);
528 ✗ if not diffArguments.scalarized then
529 // Strip subscripts from base seed before copying (same reason as SCALAR_EQUATION).
530 ✗ seedCref := ComponentRef.copySubscripts(lhsCref, ComponentRef.stripSubscriptsAll(seedCref));
531 end if;
532
533 ✗ diffArguments.current_grad := Expression.fromCref(seedCref);
534 ✗ diffArguments.collectAdjoints := true;
535 ✗ (_, diffArguments) := differentiateExpression(eq.rhs, diffArguments);
536
537 ✗ (diffArguments, stmts) := makeAdjointAccumulationStatements(diffArguments);
538 // stmts := listReverse(stmts);
539 else
540 ✗ stmts := {};
541 end if;
542 ✗ then (diffArguments, stmts);
543
544 // ===================== IF_EQUATION =====================
545 case Equation.IF_EQUATION() algorithm
546 ✗ (diffArguments, ifBranches, elseIfBranches) := differentiateIfEquationBodyAdjoint(eq.body, diffArguments);
547
548 // Wrap into a single IF statement
549 ✗ stmts := {Statement.IF(ifBranches, DAE.emptyElementSource)};
550 ✗ then (diffArguments, stmts);
551
552 // ===================== FOR_EQUATION =====================
553 case Equation.FOR_EQUATION() algorithm
554 1 stmts := {};
555
2/2
✓ Branch 0 taken 1 time.
✓ Branch 1 taken 1 time.
2 for bodyEqn in eq.body loop
556 1 (diffArguments, bodyStmts) := differentiateEquationAdjoint(bodyEqn, diffArguments);
557 1 stmts := listAppend(bodyStmts, stmts);
558 end for;
559
560 // Wrap in nested FOR statement with reversed iterator range
561 1 revIter := reverseEquationIterator(eq.iter);
562 1 (iterNames, iterRanges, iterMaps) := NBEquation.Iterator.getFrames(revIter);
563
2/2
✓ Branch 2 taken 1 time.
✓ Branch 3 taken 1 time.
2 for tpl in listReverse(List.zip3(iterNames, iterRanges, iterMaps)) loop
564 1 (iterName, iterRange, iterMap) := tpl;
565 sub_iters_stmt := match iterMap
566 case SOME(NBEquation.Iterator.SINGLE(name = iter_name, range = Expression.ARRAY(elements = iter_elems), map = NONE()))
567 ✗ then {(iter_name, iter_elems)};
568 else {};
569 end match;
570 2 stmts := {Statement.FOR(
571 ComponentRef.node(iterName),
572 SOME(iterRange),
573 stmts,
574 Statement.ForType.NORMAL(),
575 DAE.emptyElementSource,
576 sub_iters_stmt
577 )};
578 end for;
579 1 then (diffArguments, stmts);
580
581 // ===================== ALGORITHM =====================
582 case Equation.ALGORITHM() algorithm
583 allStmts := {};
584
2/2
✓ Branch 0 taken 3 times.
✓ Branch 1 taken 1 time.
4 for s in eq.alg.statements loop
585 3 (diffArguments, bodyStmts) := differentiateStatementAdjoint(s, diffArguments);
586
2/2
✓ Branch 1 taken 5 times.
✓ Branch 2 taken 3 times.
8 for bs in bodyStmts loop
587 allStmts := bs :: allStmts;
588 end for;
589 end for;
590 stmts := allStmts;
591 1 then (diffArguments, stmts);
592
593 else algorithm
594 ✗ Error.addMessage(Error.INTERNAL_ERROR,{getInstanceName() + " failed for: " + Equation.toString(eq)});
595 ✗ then fail();
596 end match;
597 end differentiateEquationAdjoint;
598
599 function differentiateStatementAdjoint
600 "Adjoint-mode statement differentiation.
601 Given a statement and DifferentiationArguments with a fresh adjoint_map,
602 returns adjoint statements and updated diffArguments."
603 input Statement stmt;
604 input output DifferentiationArguments diffArguments;
605 output list<Statement> adjointStatements;
606 algorithm
607 (diffArguments, adjointStatements) := match stmt
608 local
609 Expression lhs;
610 ComponentRef lhsCref;
611 list<Statement> stmts, bodyStmts, allStmts;
612 list<tuple<Expression, list<Statement>>> adjBranches;
613 Expression cond;
614
615 // Real assignment statement
616 case Statement.ASSIGNMENT() guard(Type.isReal(Type.arrayElementType(Expression.typeOf(stmt.lhs)))) algorithm
617 // Differentiate the LHS to get seed variable cref without collecting into the adjoint_map (avoid duplicates)
618 3 (lhs, diffArguments) := differentiateExpressionNoCollect(stmt.lhs, diffArguments);
619 lhsCref := match lhs
620 3 case Expression.CREF() then lhs.cref;
621 else ComponentRef.EMPTY();
622 end match;
623
624
1/2
✓ Branch 1 taken 3 times.
✗ Branch 2 not taken.
3 if not ComponentRef.isEmpty(lhsCref) then
625 // Set seed to the differentiated LHS variable
626 3 diffArguments.current_grad := lhs;
627 3 diffArguments.collectAdjoints := true;
628
629 // Differentiate RHS to accumulate into adjoint_map
630 3 (_, diffArguments) := differentiateExpression(stmt.rhs, diffArguments);
631
632 // Emit accumulation statements
633 3 (diffArguments, stmts) := makeAdjointAccumulationStatements(diffArguments);
634 else
635 ✗ stmts := {};
636 end if;
637 3 then (diffArguments, stmts);
638
639 // FOR statement
640 case Statement.FOR() algorithm
641 allStmts := {};
642 ✗ for s in stmt.body loop
643 ✗ (diffArguments, bodyStmts) := differentiateStatementAdjoint(s, diffArguments);
644 ✗ for bs in bodyStmts loop
645 allStmts := bs :: allStmts;
646 end for;
647 end for;
648 ✗ stmts := {Statement.FOR(
649 stmt.iterator,
650 reverseForRange(stmt.range),
651 allStmts,
652 stmt.forType,
653 stmt.source,
654 stmt.sub_iters
655 )};
656 ✗ then (diffArguments, stmts);
657
658 // IF statement
659 case Statement.IF() algorithm
660 adjBranches := {};
661 ✗ for branch in stmt.branches loop
662 ✗ (cond, bodyStmts) := branch;
663 allStmts := {};
664 ✗ for s in bodyStmts loop
665 ✗ (diffArguments, stmts) := differentiateStatementAdjoint(s, diffArguments);
666 ✗ for bs in stmts loop
667 allStmts := bs :: allStmts;
668 end for;
669 end for;
670 ✗ adjBranches := (cond, allStmts) :: adjBranches;
671 end for;
672 ✗ stmts := {Statement.IF(listReverse(adjBranches), stmt.source)};
673 ✗ then (diffArguments, stmts);
674
675 // Non-Real assignments pass through unchanged
676 ✗ case Statement.ASSIGNMENT() then (diffArguments, {stmt});
677
678 ✗ else (diffArguments, {stmt});
679 end match;
680 end differentiateStatementAdjoint;
681
682 function differentiateIfEquationBodyAdjoint
683 "Adjoint-mode differentiation of an IfEquationBody.
684 Returns branches as (condition, adjointStatements) tuples and optional else branches."
685 input IfEquationBody body;
686 input output DifferentiationArguments diffArguments;
687 output list<tuple<Expression, list<Statement>>> branches;
688 output Option<list<tuple<Expression, list<Statement>>>> elseIfBranches;
689 protected
690 list<Statement> allStmts, bodyStmts;
691 Equation bodyEqn;
692 IfEquationBody elseBody;
693 list<tuple<Expression, list<Statement>>> elseBranches;
694 Option<list<tuple<Expression, list<Statement>>>> nestedElse;
695 algorithm
696 // Process then-equations in LIFO order
697 allStmts := {};
698 ✗ for eqPtr in body.then_eqns loop
699 ✗ bodyEqn := Pointer.access(eqPtr);
700 ✗ (diffArguments, bodyStmts) := differentiateEquationAdjoint(bodyEqn, diffArguments);
701 ✗ for s in bodyStmts loop
702 allStmts := s :: allStmts;
703 end for;
704 end for;
705 // allStmts is in LIFO order
706
707 ✗ branches := {(body.condition, allStmts)};
708
709 // Recurse into else-if
710 ✗ if isSome(body.else_if) then
711 ✗ SOME(elseBody) := body.else_if;
712 ✗ (diffArguments, elseBranches, nestedElse) := differentiateIfEquationBodyAdjoint(elseBody, diffArguments);
713 // Flatten: append elseBranches and nestedElse
714 ✗ for b in elseBranches loop
715 branches := b :: branches;
716 end for;
717 ✗ elseIfBranches := nestedElse;
718 else
719 elseIfBranches := NONE();
720 end if;
721 ✗ branches := listReverse(branches);
722 end differentiateIfEquationBodyAdjoint;
723
724 function makeAdjointAccumulationStatements
725 "Read the adjoint_map and generate accumulation statements:
726 For each key v in the map with entries [(_, e1), (_, e2), ...]:
727 v := v + e1 + e2 + ...
728 Clears the adjoint_map after reading."
729 input output DifferentiationArguments diffArguments;
730 output list<Statement> stmts;
731 protected
732 UnorderedMap<ComponentRef, list<Expression>> amap;
733 list<ComponentRef> keys;
734 list<Expression> taggedTerms;
735 Expression accRhs;
736 Type vty;
737 NFOperator.SizeClassification sc;
738 Operator addOp;
739 ComponentRef key;
740 algorithm
741 stmts := {};
742 // Only generate accumulation statements if there is an adjoint_map
743
2/4
✗ Branch 0 not taken.
✓ Branch 1 taken 20 times.
✓ Branch 2 taken 20 times.
✗ Branch 3 not taken.
20 if isSome(diffArguments.adjoint_map) then
744 20 SOME(amap) := diffArguments.adjoint_map;
745 20 keys := UnorderedMap.keyList(amap);
746 // the keys in the adjoint_map are the variables we need to accumulate into; the values are the terms to accumulate
747
2/2
✓ Branch 0 taken 30 times.
✓ Branch 1 taken 20 times.
50 for key in keys loop
748 // each adjoint term is a tuple of (original variable cref, differentiated expression); we only need the expression for accumulation
749 // TODO: and could probably/surely remove the original variable cref from the map entirely
750 30 taggedTerms := UnorderedMap.getOrFail(key, amap);
751
1/2
✓ Branch 0 taken 30 times.
✗ Branch 1 not taken.
30 if not listEmpty(taggedTerms) then
752 // Build RHS: key + sum(terms)
753 // TODO: Turn this into a single multary construction
754 // Use subscripted type so indexed crefs in FOR bodies become scalar assignments.
755 30 vty := ComponentRef.getSubscriptedType(key, true);
756 30 sc := sizeClassificationFromType(vty);
757 30 addOp := Operator.fromClassification((NFOperator.MathClassification.ADDITION, sc), vty);
758 // First sum the terms if there are more than one; if only one term, use it directly
759
1/2
✓ Branch 1 taken 30 times.
✗ Branch 2 not taken.
30 if List.hasOneElement(taggedTerms) then
760 30 accRhs := listHead(taggedTerms);
761 else
762 ✗ accRhs := SimplifyExp.simplify(Expression.MULTARY(taggedTerms, {}, addOp));
763 end if;
764
765 // Basic accumulation only: v := v + sum(terms)
766 60 accRhs := SimplifyExp.simplify(Expression.MULTARY({Expression.fromCref(key), accRhs}, {}, addOp));
767 30 accRhs := Expression.map(accRhs, Expression.repairOperator);
768 // Emit accumulation statement because we put this into an algorithm body
769 30 stmts := Statement.ASSIGNMENT(
770 Expression.fromCref(key),
771 accRhs,
772 vty,
773 DAE.emptyElementSource
774 ) :: stmts;
775 end if;
776 end for;
777
778 // Clear the adjoint_map for next use
779 20 UnorderedMap.clear(amap);
780 20 diffArguments.adjoint_map := SOME(amap);
781 end if;
782 end makeAdjointAccumulationStatements;
783
784 function differentiateIfEquationBody
785 input output IfEquationBody body;
786 input output Pointer<DifferentiationArguments> diffArguments_ptr;
787 protected
788 list<Pointer<Equation>> then_eqns;
789 IfEquationBody else_if;
790 algorithm
791 // ToDo: this is a little ugly
792 // 1. why are the then_eqns Pointers? no need for that
793 // 2. we could just traverse it regularly without creating a pointer for diffArguments
794 ✗ then_eqns := List.map(body.then_eqns, function differentiateEquationPointer(diffArguments_ptr = diffArguments_ptr, name = ""));
795 ✗ if isSome(body.else_if) then
796 ✗ (else_if, diffArguments_ptr) := differentiateIfEquationBody(Util.getOption(body.else_if), diffArguments_ptr);
797 ✗ body := IfEquationBody.IF_EQUATION_BODY(body.condition, then_eqns, SOME(else_if));
798 else
799 ✗ body := IfEquationBody.IF_EQUATION_BODY(body.condition, then_eqns, NONE());
800 end if;
801 end differentiateIfEquationBody;
802
803 function differentiateWhenEquationBody
804 input output WhenEquationBody body;
805 input output DifferentiationArguments diffArguments;
806 protected
807 list<WhenStatement> when_stmts;
808 WhenEquationBody else_when;
809 algorithm
810 ✗ (when_stmts, diffArguments) := List.mapFold(body.when_stmts, function differentiateWhenStatement(), diffArguments);
811 ✗ if isSome(body.else_when) then
812 ✗ (else_when, diffArguments) := differentiateWhenEquationBody(Util.getOption(body.else_when), diffArguments);
813 ✗ body := WhenEquationBody.WHEN_EQUATION_BODY(body.condition, when_stmts, SOME(else_when));
814 else
815 ✗ body := WhenEquationBody.WHEN_EQUATION_BODY(body.condition, when_stmts, NONE());
816 end if;
817 end differentiateWhenEquationBody;
818
819 function differentiateWhenStatement
820 input output WhenStatement stmt;
821 input output DifferentiationArguments diffArguments;
822 algorithm
823 (stmt, diffArguments) := match stmt
824 local
825 Expression lhs, rhs;
826 // Only differentiate assignments
827 case WhenStatement.ASSIGN() algorithm
828 ✗ (lhs, diffArguments) := differentiateExpression(stmt.lhs, diffArguments);
829 ✗ (rhs, diffArguments) := differentiateExpression(stmt.rhs, diffArguments);
830 ✗ then (WhenStatement.ASSIGN(lhs, rhs, stmt.source), diffArguments);
831 else (stmt, diffArguments);
832 end match;
833 end differentiateWhenStatement;
834
835 function differentiateExpressionDump
836 "wrapper function for differentiation to allow dumping before and afterwards"
837 input output Expression exp;
838 input output DifferentiationArguments diffArguments;
839 input String name = "";
840 input String indent = "";
841 algorithm
842
1/2
✗ Branch 1 not taken.
✓ Branch 2 taken 3718 times.
3718 if Flags.isSet(Flags.DEBUG_DIFFERENTIATION) then
843 ✗ print(indent + "### debugDifferentiation | " + name + " ###\n");
844 ✗ print(indent + "[BEFORE] " + Expression.toString(exp) + "\n");
845 ✗ (exp, diffArguments) := differentiateExpression(exp, diffArguments);
846 ✗ print(indent + "[AFTER ] " + Expression.toString(exp) + "\n\n");
847 else
848 3718 (exp, diffArguments) := differentiateExpression(exp, diffArguments);
849 end if;
850 end differentiateExpressionDump;
851
852 function differentiateExpression
853 input output Expression exp;
854 input output DifferentiationArguments diffArguments;
855 algorithm
856 (exp, diffArguments) := match exp
857 local
858 Expression elem1, elem2, current_grad, gradTrue, gradFalse;
859 list<Expression> new_elements = {};
860 list<list<Expression>> new_matrix_elements = {};
861 array<Expression> arr;
862 ComponentRef d_fn;
863 Boolean isReverse = isSome(diffArguments.adjoint_map);
864
865 // differentiation of constant expressions results in zero
866 755 case Expression.INTEGER() then (Expression.INTEGER(0), diffArguments);
867 1763 case Expression.REAL() then (Expression.REAL(0.0), diffArguments);
868 // leave boolean and string expressions as is
869 3 case Expression.STRING() then (exp, diffArguments);
870 1 case Expression.BOOLEAN() then (exp, diffArguments);
871
872 // differentiate cref
873 17897 case Expression.CREF() then differentiateComponentRef(exp, diffArguments);
874
875 // [a, b, c, ...]' = [a', b', c', ...]
876 case Expression.ARRAY() algorithm
877 91 (arr, diffArguments) := Array.mapFold(exp.elements, differentiateExpression, diffArguments);
878 91 exp.elements := arr;
879 91 then (exp, diffArguments);
880
881 // |a, b, c|' |a', b', c'|
882 // |d, e, f| = |d', e', f'|
883 // |g, h, i| |g', h', i'|
884 case Expression.MATRIX() algorithm
885 ✗ for element_lst in exp.elements loop
886 new_elements := {};
887 ✗ for element in element_lst loop
888 ✗ (element, diffArguments) := differentiateExpression(element, diffArguments);
889 new_elements := element :: new_elements;
890 end for;
891 ✗ new_matrix_elements := listReverse(new_elements) :: new_matrix_elements;
892 end for;
893 ✗ then (Expression.MATRIX(listReverse(new_matrix_elements)), diffArguments);
894
895 // (a, b, c, ...)' = (a', b', c', ...)
896 case Expression.TUPLE() algorithm
897
2/2
✓ Branch 0 taken 18 times.
✓ Branch 1 taken 9 times.
27 for element in exp.elements loop
898 18 (element, diffArguments) := differentiateExpression(element, diffArguments);
899 new_elements := element :: new_elements;
900 end for;
901 9 then (Expression.TUPLE(exp.ty, listReverse(new_elements)), diffArguments);
902
903 // REC(a, b, c, ...)' = REC(a', b', c', ...)
904 case Expression.RECORD() algorithm
905
2/2
✓ Branch 0 taken 30 times.
✓ Branch 1 taken 3 times.
33 for element in exp.elements loop
906 30 (element, diffArguments) := differentiateExpression(element, diffArguments);
907 new_elements := element :: new_elements;
908 end for;
909 3 then (Expression.RECORD(exp.path, exp.ty, listReverse(new_elements)), diffArguments);
910
911 // e.g. (f(x))' = f'(x) * x' (more rules in differentiateCall)
912 609 case Expression.CALL() then differentiateCall(exp, diffArguments);
913
914 // Forward: (if c then a else b)' = if c then a' else b'
915 // Reverse: upstream G is only sent to taken branch:
916 // grad_a = if c then G else 0
917 // grad_b = if c then 0 else G
918 // Then recurse with those masked gradients.
919 case Expression.IF() algorithm
920
2/2
✓ Branch 0 taken 1 time.
✓ Branch 1 taken 21 times.
22 if isReverse then
921 // Keep original upstream
922 1 current_grad := diffArguments.current_grad;
923
924 // Masked gradients
925 1 gradTrue := Expression.IF(Expression.typeOf(current_grad), exp.condition, current_grad, Expression.makeZero(Expression.typeOf(current_grad)));
926 1 gradFalse := Expression.IF(Expression.typeOf(current_grad), exp.condition, Expression.makeZero(Expression.typeOf(current_grad)), current_grad);
927
928 // Recurse true branch
929 1 diffArguments.current_grad := gradTrue;
930 1 (elem1, diffArguments) := differentiateExpression(exp.trueBranch, diffArguments);
931
932 // Recurse false branch
933 1 diffArguments.current_grad := gradFalse;
934 1 (elem2, diffArguments) := differentiateExpression(exp.falseBranch, diffArguments);
935
936 // Restore upstream
937 1 diffArguments.current_grad := current_grad;
938 else
939 21 (elem1, diffArguments) := differentiateExpression(exp.trueBranch, diffArguments);
940 21 (elem2, diffArguments) := differentiateExpression(exp.falseBranch, diffArguments);
941 end if;
942 22 then (Expression.IF(exp.ty, exp.condition, elem1, elem2), diffArguments);
943
944 // e.g. (fg)' = fg' + f'g (more rules in differentiateBinary)
945 1252 case Expression.BINARY() then differentiateBinary(exp, diffArguments);
946
947 // e.g. (fgh)' = f'gh + fg'h + fgh' (more rules in differentiateMultary)
948 9573 case Expression.MULTARY() then differentiateMultary(exp, diffArguments);
949
950 // (-x)' = -(x')
951 case Expression.UNARY() algorithm
952
2/2
✓ Branch 0 taken 1 time.
✓ Branch 1 taken 1146 times.
1147 if isReverse then
953 1 current_grad := diffArguments.current_grad;
954
955 // apply same unary operator to current_grad
956 2 diffArguments.current_grad := Expression.UNARY(exp.operator, current_grad);
957 1 (elem1, diffArguments) := differentiateExpression(exp.exp, diffArguments);
958
959 1 diffArguments.current_grad := current_grad;
960 else
961 1146 (elem1, diffArguments) := differentiateExpression(exp.exp, diffArguments);
962 end if;
963 1147 then (Expression.UNARY(exp.operator, elem1), diffArguments);
964
965 // ((Real) x)' = (Real) x'
966 case Expression.CAST() algorithm
967 108 (elem1, diffArguments) := differentiateExpression(exp.exp, diffArguments);
968 108 then (Expression.CAST(exp.ty, elem1), diffArguments);
969
970 // BOX(x)' = BOX(x')
971 case Expression.BOX() algorithm
972 ✗ (elem1, diffArguments) := differentiateExpression(exp.exp, diffArguments);
973 ✗ then (Expression.BOX(elem1), diffArguments);
974
975 // UNBOX(x)' = UNBOX(x')
976 case Expression.UNBOX() algorithm
977 ✗ (elem1, diffArguments) := differentiateExpression(exp.exp, diffArguments);
978 ✗ then (Expression.UNBOX(elem1, exp.ty), diffArguments);
979
980 // (x(1))' = x'(1)
981 case Expression.SUBSCRIPTED_EXP() algorithm
982 11 (elem1, diffArguments) := differentiateExpression(exp.exp, diffArguments);
983 11 then (Expression.SUBSCRIPTED_EXP(elem1, exp.subscripts, exp.ty, exp.split), diffArguments);
984
985 // (..., a_i,...)' = (..., a'_i, ...)
986 case Expression.TUPLE_ELEMENT() algorithm
987 6 (elem1, diffArguments) := differentiateExpression(exp.tupleExp, diffArguments);
988 6 then (Expression.TUPLE_ELEMENT(elem1, exp.index, exp.ty), diffArguments);
989
990 // REC(i, ...)' = REC(i', ...)
991 case Expression.RECORD_ELEMENT() algorithm
992 // check if differentiating for simple cref and if it contains it
993
1/4
✗ Branch 0 not taken.
✓ Branch 1 taken 2 times.
✗ Branch 3 not taken.
✗ Branch 4 not taken.
2 if diffArguments.diffType == DifferentiationType.SIMPLE and not Expression.containsCref(exp.recordExp, diffArguments.diffCref) then
994 ✗ elem1 := Expression.makeZero(Expression.typeOf(exp));
995 else
996 2 (elem1, diffArguments) := differentiateExpression(exp.recordExp, diffArguments);
997 2 elem1 := Expression.RECORD_ELEMENT(elem1, exp.index, exp.fieldName, exp.ty);
998 end if;
999 2 then (elem1, diffArguments);
1000
1001 // differentiate a passed function pointer
1002 case Expression.PARTIAL_FUNCTION_APPLICATION() algorithm
1003 ✗ d_fn := BVariable.makeFDerVar(exp.fn);
1004 ✗ for element in exp.args loop
1005 ✗ (element, diffArguments) := differentiateExpression(element, diffArguments);
1006 new_elements := element :: new_elements;
1007 end for;
1008 ✗ then (Expression.PARTIAL_FUNCTION_APPLICATION(d_fn, listAppend(exp.args, listReverse(new_elements)),
1009 listAppend(exp.argNames, list(BackendUtil.makeFDerString(name) for name in exp.argNames)), exp.ty), diffArguments);
1010
1011 // Binary expressions, conditions and placeholders are not differentiated and left as they are
1012 ✗ case Expression.LBINARY() then (exp, diffArguments);
1013 ✗ case Expression.LUNARY() then (exp, diffArguments);
1014 ✗ case Expression.RELATION() then (exp, diffArguments);
1015 ✗ case Expression.SIZE() then (exp, diffArguments);
1016 ✗ case Expression.RANGE() then (exp, diffArguments);
1017 ✗ case Expression.END() then (exp, diffArguments);
1018 ✗ case Expression.EMPTY() then (exp, diffArguments);
1019 ✗ case Expression.ENUM_LITERAL() then (exp, diffArguments);
1020 ✗ case Expression.TYPENAME() then (exp, diffArguments);
1021
1022 else algorithm
1023 // maybe add failtrace here and allow failing
1024 ✗ Error.addMessage(Error.INTERNAL_ERROR,{getInstanceName() + " failed for: " + Expression.toString(exp)});
1025 ✗ then fail();
1026 end match;
1027 end differentiateExpression;
1028
1029 function differentiateExpressionNoCollect
1030 input output Expression expr;
1031 input output DifferentiationArguments diffArguments;
1032 protected
1033 Boolean oldCollect;
1034 algorithm
1035
3/4
✗ Branch 0 not taken.
✓ Branch 1 taken 1353 times.
✓ Branch 2 taken 1350 times.
✓ Branch 3 taken 3 times.
1353 if isSome(diffArguments.adjoint_map) then
1036 3 oldCollect := diffArguments.collectAdjoints;
1037 3 diffArguments.collectAdjoints := false;
1038 3 (expr, diffArguments) := differentiateExpression(expr, diffArguments);
1039
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 3 times.
3 diffArguments.collectAdjoints := oldCollect;
1040 else
1041 1350 (expr, diffArguments) := differentiateExpression(expr, diffArguments);
1042 end if;
1043 end differentiateExpressionNoCollect;
1044
1045 function differentiateComponentRef
1046 input output Expression exp "Has to be Expression.CREF()";
1047 input output DifferentiationArguments diffArguments;
1048 protected
1049 Pointer<Variable> var_ptr, der_ptr;
1050 ComponentRef derCref, strippedCref;
1051 algorithm
1052 // extract var pointer first to have following code more readable
1053 var_ptr := match exp
1054 // function body expressions, empty and wild crefs are not lowered (maybe do it?)
1055 1205 case _ guard(diffArguments.diffType == DifferentiationType.FUNCTION) then Pointer.create(NBVariable.DUMMY_VARIABLE);
1056 28 case Expression.CREF(cref = ComponentRef.EMPTY()) then Pointer.create(NBVariable.DUMMY_VARIABLE);
1057 ✗ case Expression.CREF(cref = ComponentRef.WILD()) then Pointer.create(NBVariable.DUMMY_VARIABLE);
1058 18228 case Expression.CREF() then BVariable.getVarPointer(exp.cref, sourceInfo());
1059 else algorithm
1060 ✗ Error.addMessage(Error.INTERNAL_ERROR,{getInstanceName() + " failed for: " + Expression.toString(exp)});
1061 ✗ then fail();
1062 end match;
1063
1064 // Debug entry summary
1065
3/4
✓ Branch 6 taken 19461 times.
✗ Branch 7 not taken.
✓ Branch 10 taken 19419 times.
✓ Branch 11 taken 42 times.
58341 dbg("[dCREF] exp=" + Expression.toString(exp)
1066 + " | diffType=" + DifferentiationArguments.diffTypeStr(diffArguments.diffType)
1067 + " | scalarized=" + boolString(diffArguments.scalarized)
1068 + " | collectAdjoints=" + boolString(diffArguments.collectAdjoints));
1069
3/4
✗ Branch 0 not taken.
✓ Branch 1 taken 19461 times.
✓ Branch 2 taken 45 times.
✓ Branch 3 taken 19416 times.
19461 if isSome(diffArguments.adjoint_map) then
1070 45 dbg("[dCREF] current_grad=" + Expression.toString(diffArguments.current_grad));
1071 end if;
1072
1073 (exp, diffArguments) := match (exp, diffArguments.diffType, diffArguments.diff_map)
1074 local
1075 Expression res;
1076 UnorderedMap<ComponentRef,ComponentRef> diff_map;
1077 list<Subscript> expCrefSubscripts;
1078 ComponentRef adjointKey;
1079 list<ComponentRef> elem_crefs;
1080 list<Expression> elem_exps;
1081 Expression elem_res;
1082 Boolean hasSetSub;
1083 // -------------------------------------
1084 // EMPTY and WILD crefs do nothing
1085 // -------------------------------------
1086 28 case (Expression.CREF(cref = ComponentRef.EMPTY()), _, _) then (exp, diffArguments);
1087 ✗ case (Expression.CREF(cref = ComponentRef.WILD()), _, _) then (exp, diffArguments);
1088
1089 // -------------------------------------
1090 // Special rules for Type: FUNCTION
1091 // (needs to be first because var_ptr is DUMMY)
1092 // -------------------------------------
1093
1094 // Types: (FUNCTION)
1095 // Any variable that is in the HT will be differentiated accordingly. 0 otherwise
1096 case (Expression.CREF(), DifferentiationType.FUNCTION, SOME(diff_map)) algorithm
1097 1205 strippedCref := ComponentRef.stripSubscriptsAll(exp.cref);
1098 // discrete variables (e.g. for-loop iterators with the name of a local) have no derivative
1099
4/4
✓ Branch 2 taken 1161 times.
✓ Branch 3 taken 44 times.
✓ Branch 5 taken 991 times.
✓ Branch 6 taken 170 times.
1205 if not Type.isDiscrete(Type.arrayElementType(exp.ty)) and UnorderedMap.contains(strippedCref, diff_map) then
1100 // get the derivative and reapply subscripts
1101 991 derCref := UnorderedMap.getOrFail(strippedCref, diff_map);
1102 991 derCref := ComponentRef.copySubscripts(exp.cref, derCref);
1103 991 res := Expression.fromCref(derCref);
1104 elseif not Type.isDiscrete(Type.arrayElementType(exp.ty)) and isSome(derivativeOfPrefix(strippedCref, diff_map)) then
1105 // a field of a record input, e.g. s.T -> $Ds.T
1106 ✗ SOME(derCref) := derivativeOfPrefix(strippedCref, diff_map);
1107 ✗ derCref := ComponentRef.copySubscripts(exp.cref, derCref);
1108 ✗ res := Expression.fromCref(derCref);
1109 elseif Type.isArray(exp.ty) and not Type.hasKnownSize(exp.ty) then
1110 // an input of unknown size, e.g. a[:], has the zero derivative fill(0, size(a, 1), ...)
1111 ✗ res := Expression.CALL(Call.makeTypedCall(NFBuiltinFuncs.FILL_FUNC,
1112 makeZero(Type.arrayElementType(exp.ty)) :: list(Expression.SIZE(exp, SOME(Expression.INTEGER(i))) for i in 1:Type.dimensionCount(exp.ty)),
1113 Variability.CONTINUOUS, NFPrefixes.Purity.PURE, exp.ty));
1114 else
1115 214 res := makeZero(exp.ty);
1116 end if;
1117 1205 then (res, diffArguments);
1118
1119 // Types: (SIMPLE, TIME)
1120 // a record variable is differentiated fieldwise, D(r)/dr.x => R(1, 0, ...)
1121 case (Expression.CREF(), _, _)
1122 guard((diffArguments.diffType == DifferentiationType.SIMPLE or diffArguments.diffType == DifferentiationType.TIME)
1123 and Type.isRecord(exp.ty) and not ComponentRef.isEqual(exp.cref, diffArguments.diffCref)
1124 and BVariable.checkCref(exp.cref, BVariable.isRecord, sourceInfo()))
1125 6 then differentiateRecordCref(exp, diffArguments);
1126
1127 // -------------------------------------
1128 // Generic Rules
1129 // -------------------------------------
1130
1131 // Types: (TIME)
1132 // differentiate time cref => 1
1133 case (Expression.CREF(), DifferentiationType.TIME, _)
1134 guard(ComponentRef.isTime(exp.cref))
1135 7 then (Expression.makeOne(exp.ty), diffArguments);
1136
1137 // Types: not (TIME)
1138 // differentiate time cref => 0
1139 case (Expression.CREF(), _, _)
1140 guard(ComponentRef.isTime(exp.cref))
1141 79 then (Expression.makeZero(exp.ty), diffArguments);
1142
1143 // Types: (ALL)
1144 // differentiate start cref => 0
1145 case (Expression.CREF(), _, _)
1146 guard(BVariable.isStart(var_ptr))
1147 ✗ then (Expression.makeZero(exp.ty), diffArguments);
1148
1149 // ToDo: Records, Arrays, WILD (?)
1150
1151 // Types: (SIMPLE)
1152 // D(x)/dx => 1
1153 case (Expression.CREF(), DifferentiationType.SIMPLE, _)
1154 guard(ComponentRef.isEqual(exp.cref, diffArguments.diffCref))
1155 4470 then (makeOne(exp.ty), diffArguments);
1156
1157 // Types: (SIMPLE)
1158 // D(y)/dx => 0
1159 case (Expression.CREF(), DifferentiationType.SIMPLE, _)
1160 7490 then (Expression.makeZero(exp.ty), diffArguments);
1161
1162 // Types: (ALL)
1163 // Known variables, except for top level inputs have a 0-derivative
1164 case (Expression.CREF(), _, _)
1165 guard(BVariable.isParamOrConst(var_ptr) and
1166 not (ComponentRef.isTopLevel(exp.cref) and BVariable.isInput(var_ptr))
1167 and not BVariable.isOptimizable(var_ptr) /* TODO? */ )
1168 165 then (Expression.makeZero(exp.ty), diffArguments);
1169
1170 // -------------------------------------
1171 // Special rules for Type: TIME
1172 // -------------------------------------
1173
1174 // Types: (TIME)
1175 // D(discrete)/d(x) = 0
1176 case (Expression.CREF(), DifferentiationType.TIME, _)
1177 guard(BVariable.isDiscrete(var_ptr) or BVariable.isDiscreteState(var_ptr))
1178 ✗ then (Expression.makeZero(exp.ty), diffArguments);
1179
1180 // Types: (TIME)
1181 // known derivatives by state order
1182 case (Expression.CREF(), DifferentiationType.TIME, SOME(diff_map))
1183 guard(UnorderedMap.contains(ComponentRef.stripSubscriptsAll(exp.cref), diff_map)) algorithm
1184 // get the derivative and reapply subscripts
1185 41 derCref := UnorderedMap.getOrFail(ComponentRef.stripSubscriptsAll(exp.cref), diff_map);
1186 41 derCref := ComponentRef.copySubscripts(exp.cref, derCref);
1187 41 res := Expression.fromCref(derCref);
1188 41 then (res, diffArguments);
1189
1190 // Types: (TIME)
1191 // DUMMY_STATES => DUMMY_DER
1192 case (Expression.CREF(), DifferentiationType.TIME, _)
1193 guard(BVariable.isDummyState(var_ptr))
1194 40 then (Expression.fromCref(BVariable.getPartnerCref(exp.cref, BVariable.getVarDummyDer)), diffArguments);
1195
1196 // Types: (TIME)
1197 // D(x)/dtime --> der(x) --> $DER.x
1198 // STATE => STATE_DER
1199 case (Expression.CREF(), DifferentiationType.TIME, _)
1200 guard(BVariable.isState(var_ptr))
1201 187 then (Expression.fromCref(BVariable.getPartnerCref(exp.cref, BVariable.getVarDer)), diffArguments);
1202
1203 // Types: (TIME)
1204 // D(y)/dtime --> der(y) --> $DER.y
1205 // ALGEBRAIC => STATE_DER
1206 // make y a state and add new STATE_DER
1207 case (Expression.CREF(), DifferentiationType.TIME, _)
1208 guard(BVariable.isContinuous(var_ptr, false))
1209 algorithm
1210 // create derivative
1211 98 (derCref, der_ptr) := BVariable.makeDerVar(exp.cref);
1212 // add derivative to new_vars
1213 196 diffArguments.new_vars := der_ptr :: diffArguments.new_vars;
1214 // update algebraic variable to be a state
1215 98 BVariable.setStateDerivativeVar(var_ptr, der_ptr);
1216 98 then (Expression.fromCref(derCref), diffArguments);
1217
1218 // -------------------------------------
1219 // Special rules for Type: JACOBIAN
1220 // -------------------------------------
1221
1222 // Types: (JACOBIAN)
1223 // cref in diff_map => get $SEED or $pDER variable from hash table
1224 case (Expression.CREF(), DifferentiationType.JACOBIAN, SOME(diff_map))
1225 guard(diffArguments.scalarized)
1226 algorithm
1227 ✗ if Type.isRecord(exp.ty) and isMixedRecordDerivative(ComponentRef.stripSubscriptsAll(exp.cref), diff_map) then
1228 // a record with a seed of its own whose fields are not all seeds is differentiated fieldwise
1229 ✗ (res, diffArguments) := differentiateRecordCref(exp, diffArguments);
1230 elseif UnorderedMap.contains(exp.cref, diff_map) then
1231 ✗ res := Expression.fromCref(UnorderedMap.getOrFail(exp.cref, diff_map));
1232
1233 // Accumulate adjoint contribution: append current_grad to list at key exp.cref.
1234 ✗ if diffArguments.collectAdjoints then
1235 ✗ UnorderedMap.tryAddUpdate(exp.cref, function updateAdjointList(current_grad = diffArguments.current_grad), Util.getOption(diffArguments.adjoint_map));
1236 end if;
1237 else
1238 // an array whose elements are in diff_map (e.g. a partially torn array) is differentiated
1239 // elementwise, everything else that is not in diff_map gets differentiated to zero
1240 hasSetSub := false;
1241 elem_crefs := {};
1242 ✗ if Type.isArray(exp.ty) and Type.sizeOf(exp.ty) <= 256 then
1243 ✗ elem_crefs := listReverse(ComponentRef.scalarizeAll(exp.cref, false));
1244 ✗ for c in elem_crefs loop
1245 ✗ if UnorderedMap.contains(c, diff_map) then
1246 hasSetSub := true;
1247 break;
1248 end if;
1249 end for;
1250 end if;
1251 ✗ if hasSetSub then
1252 elem_exps := {};
1253 ✗ for c in elem_crefs loop
1254 ✗ (elem_res, diffArguments) := differentiateComponentRef(Expression.fromCref(c), diffArguments);
1255 elem_exps := elem_res :: elem_exps;
1256 end for;
1257 ✗ res := makeShapedArray(exp.ty, listReverse(elem_exps));
1258 elseif Type.isRecord(exp.ty) then
1259 ✗ (res, diffArguments) := differentiateRecordCref(exp, diffArguments);
1260 else
1261 ✗ res := differentiateIteratorElement(exp, diffArguments, diff_map);
1262 end if;
1263 end if;
1264 ✗ then (res, diffArguments);
1265
1266 // Types: (JACOBIAN)
1267 // cref in diff_map => get $SEED or $pDER variable from hash table
1268 case (Expression.CREF(), DifferentiationType.JACOBIAN, SOME(diff_map))
1269 guard(not diffArguments.scalarized)
1270 algorithm
1271 5645 strippedCref := ComponentRef.stripSubscriptsAll(exp.cref);
1272 5645 expCrefSubscripts := ComponentRef.subscriptsAllFlat(exp.cref);
1273 5645 dbg("[dCREF:JAC] cref=" + ComponentRef.toString(exp.cref)
1274 + " | stripped=" + ComponentRef.toString(strippedCref)
1275 + " | subs=" + Subscript.toStringList(expCrefSubscripts));
1276
3/4
✓ Branch 1 taken 6 times.
✓ Branch 2 taken 5639 times.
✓ Branch 4 taken 6 times.
✗ Branch 5 not taken.
5645 if Type.isRecord(exp.ty) and isMixedRecordDerivative(strippedCref, diff_map) then
1277 6 (res, diffArguments) := differentiateRecordCref(exp, diffArguments);
1278 elseif UnorderedMap.contains(exp.cref, diff_map) then
1279 // exp.cref is itself one of this Jacobian's own registered unknowns:
1280 // use it directly rather than falling through to the base-cref template,
1281 // which may belong to an unrelated element sharing the same base cref.
1282 4124 derCref := UnorderedMap.getOrFail(exp.cref, diff_map);
1283 4124 dbg("[dCREF:JAC] exact match -> " + ComponentRef.toString(derCref));
1284 4124 res := Expression.fromCref(derCref);
1285
2/2
✓ Branch 0 taken 31 times.
✓ Branch 1 taken 4093 times.
4124 if diffArguments.collectAdjoints then
1286 // Accumulate into the derivative (pDER/SEED) cref's own adjoint slot, not
1287 // the source variable's - matches the base-cref-fallback branch below, whose
1288 // adjointKey is likewise derived from derCref, never from exp.cref directly.
1289
2/2
✓ Branch 2 taken 28 times.
✓ Branch 3 taken 3 times.
31 if not UnorderedMap.contains(derCref, Util.getOption(diffArguments.adjoint_map)) then
1290 28 UnorderedMap.tryAdd(derCref, {}, Util.getOption(diffArguments.adjoint_map));
1291 end if;
1292 31 UnorderedMap.tryAddUpdate(derCref, function updateAdjointList(current_grad = diffArguments.current_grad), Util.getOption(diffArguments.adjoint_map));
1293 end if;
1294 elseif UnorderedMap.contains(strippedCref, diff_map) then
1295 // get the derivative and reapply subscripts
1296 904 derCref := UnorderedMap.getOrFail(strippedCref, diff_map);
1297 904 dbg("[dCREF:JAC] mapped -> " + ComponentRef.toString(derCref));
1298 // Strip subscripts from derCref before copying: diff_map[base] may store a
1299 // subscripted element seed (partial-slice NLS iter vars). Stripping ensures
1300 // exp.cref subscripts (including iterators) merge onto an unsubscripted template.
1301 904 res := Expression.fromCref(ComponentRef.copySubscripts(exp.cref, ComponentRef.stripSubscriptsAll(derCref)));
1302 904 dbg("[dCREF:JAC] get variable for derivative cref: " + NBVariable.pointerToString(NBVariable.getVarPointer(derCref, sourceInfo())));
1303
2/2
✓ Branch 0 taken 2 times.
✓ Branch 1 taken 902 times.
904 if diffArguments.collectAdjoints then // if derCref is on the rhs then collect adjoint (collectAdjoints is false when differentiating lhs)
1304 2 adjointKey := ComponentRef.copySubscripts(exp.cref, ComponentRef.stripSubscriptsAll(derCref));
1305
1/2
✓ Branch 2 taken 2 times.
✗ Branch 3 not taken.
2 if not UnorderedMap.contains(adjointKey, Util.getOption(diffArguments.adjoint_map)) then
1306 2 UnorderedMap.tryAdd(adjointKey, {}, Util.getOption(diffArguments.adjoint_map));
1307 end if;
1308 2 UnorderedMap.tryAddUpdate(adjointKey, function updateAdjointList(current_grad = diffArguments.current_grad), Util.getOption(diffArguments.adjoint_map));
1309 else
1310 902 dbg("[dCREF:JAC] collectAdjoints=false, skip append");
1311 end if;
1312 else
1313 // a cref with a set/array-valued subscript whose individual scalar elements
1314 // are each registered in diff_map (e.g. i_s[{1, 2}] when the Jacobian's seeds
1315 // are the individual i_s[1]/i_s[2], as produced for a torn slice's residual
1316 // equation -- see NBTearing.scalarSlices) has no single matching diff_map
1317 // entry of its own: neither the exact nor the whole-base-stripped lookup
1318 // above can find it, so without this the symbolic derivative fell through to
1319 // a hardcoded zero, silently producing a zero column in the analytical
1320 // Jacobian for those seeds. Differentiate it elementwise instead, matching how
1321 // dependency collection resolves the same shape of cref (see
1322 // NBAdjacency.collectDependenciesCref).
1323 hasSetSub := false;
1324
2/2
✓ Branch 1 taken 255 times.
✓ Branch 2 taken 611 times.
866 for s in ComponentRef.subscriptsAllFlat(exp.cref) loop
1325 // WHOLE (":") and SLICE (e.g. "1:3") are ordinary range subscripts, not
1326 // the literal/array-valued INDEX subscript case (e.g. "{1, 2}") this is
1327 // meant to catch -- see NBAdjacency.collectDependenciesCref.
1328
1/4
✗ Branch 1 not taken.
✓ Branch 2 taken 255 times.
✗ Branch 4 not taken.
✗ Branch 5 not taken.
255 if not Subscript.isScalar(s) and not Subscript.isSliced(s) then
1329 hasSetSub := true;
1330 end if;
1331 end for;
1332 // a slice (e.g. i[1:2]) of variables whose elements are the seeds needs to be expanded as well
1333
4/6
✓ Branch 0 taken 611 times.
✗ Branch 1 not taken.
✓ Branch 3 taken 47 times.
✓ Branch 4 taken 564 times.
✓ Branch 6 taken 47 times.
✗ Branch 7 not taken.
611 if not hasSetSub and Type.isArray(exp.ty) and Type.sizeOf(exp.ty) <= 256 then
1334
2/2
✓ Branch 2 taken 84 times.
✓ Branch 3 taken 29 times.
113 for c in listReverse(ComponentRef.scalarizeAll(exp.cref, false)) loop
1335
2/2
✓ Branch 1 taken 66 times.
✓ Branch 2 taken 18 times.
84 if UnorderedMap.contains(c, diff_map) then
1336 hasSetSub := true;
1337 break;
1338 end if;
1339 end for;
1340 end if;
1341
2/2
✓ Branch 0 taken 593 times.
✓ Branch 1 taken 18 times.
611 if not hasSetSub then
1342 // no set-valued subscript to expand (e.g. a fully bare/unsubscripted
1343 // matrix cref like Rot_dq): keep the original whole-type zero, since
1344 // building it element-by-element would flatten its shape and break
1345 // codegen for multi-dimensional types.
1346
1/2
✗ Branch 1 not taken.
✓ Branch 2 taken 593 times.
593 if Type.isRecord(exp.ty) then
1347 ✗ (res, diffArguments) := differentiateRecordCref(exp, diffArguments);
1348 else
1349 593 res := differentiateIteratorElement(exp, diffArguments, diff_map);
1350 end if;
1351 else
1352 18 elem_crefs := listReverse(ComponentRef.scalarizeAll(exp.cref, false));
1353 elem_exps := {};
1354
2/2
✓ Branch 0 taken 108 times.
✓ Branch 1 taken 18 times.
126 for c in elem_crefs loop
1355 108 (elem_res, diffArguments) := differentiateComponentRef(Expression.fromCref(c), diffArguments);
1356 elem_exps := elem_res :: elem_exps;
1357 end for;
1358 18 res := makeShapedArray(exp.ty, listReverse(elem_exps));
1359 end if;
1360 end if;
1361 5645 then (res, diffArguments);
1362
1363 else algorithm
1364 // maybe add failtrace here and allow failing
1365 ✗ Error.addMessage(Error.INTERNAL_ERROR,{getInstanceName() + " failed for: " + Expression.toString(exp)});
1366 ✗ then fail();
1367
1368 end match;
1369 end differentiateComponentRef;
1370
1371 function makeShapedArray
1372 "Builds an array expression of type ty from its scalar elements given in row-major
1373 order. A multi-dimensional type gets one nested array per leading dimension, e.g.
1374 Real[2, 1] with {a, b} becomes {{a}, {b}}, so that the shape matches the type. A flat
1375 array of scalars typed as a matrix breaks the code generation."
1376 input Type ty;
1377 input list<Expression> elems "row-major";
1378 output Expression res;
1379 protected
1380 list<Dimension> dims = Type.arrayDims(ty);
1381 Type row_ty;
1382 Integer row_size, n_rows;
1383 list<Expression> rows = {}, row, remaining = elems;
1384 algorithm
1385
1/2
✓ Branch 1 taken 20 times.
✗ Branch 2 not taken.
20 if listLength(dims) < 2 then
1386 20 res := Expression.makeArray(ty, listArray(elems));
1387 else
1388 ✗ row_ty := Type.ARRAY(Type.arrayElementType(ty), listRest(dims));
1389 ✗ row_size := Type.sizeOf(row_ty);
1390 ✗ if row_size < 1 or listLength(elems) <> row_size * Dimension.size(listHead(dims)) then
1391 // sizes that do not add up, keep the previous (flat) result rather than guess
1392 ✗ res := Expression.makeArray(ty, listArray(elems));
1393 else
1394 ✗ n_rows := Dimension.size(listHead(dims));
1395 ✗ for i in 1:n_rows loop
1396 ✗ (row, remaining) := List.split(remaining, row_size);
1397 ✗ rows := makeShapedArray(row_ty, row) :: rows;
1398 end for;
1399 ✗ res := Expression.makeArray(ty, listArray(listReverse(rows)));
1400 end if;
1401 end if;
1402 end makeShapedArray;
1403
1404 function derivativeOfPrefix
1405 "The derivative of a cref whose prefix has a derivative, e.g. s.T -> $Ds.T if s -> $Ds."
1406 input ComponentRef cref "without subscripts";
1407 input UnorderedMap<ComponentRef, ComponentRef> diff_map;
1408 output Option<ComponentRef> derCref;
1409 algorithm
1410 derCref := match cref
1411 local
1412 ComponentRef rest, der_rest;
1413 case ComponentRef.CREF(restCref = rest as ComponentRef.CREF()) algorithm
1414
1/2
✗ Branch 1 not taken.
✓ Branch 2 taken 160 times.
160 if UnorderedMap.contains(rest, diff_map) then
1415 ✗ derCref := SOME(ComponentRef.prepend(UnorderedMap.getOrFail(rest, diff_map), cref));
1416 else
1417 derCref := match derivativeOfPrefix(rest, diff_map)
1418 ✗ case SOME(der_rest) then SOME(ComponentRef.prepend(der_rest, cref));
1419 else NONE();
1420 end match;
1421 end if;
1422 then derCref;
1423 else NONE();
1424 end match;
1425 end derivativeOfPrefix;
1426
1427 function isMixedRecordDerivative
1428 "true if the fields of a record do not all have derivatives of the same kind as the record itself,
1429 e.g. a torn record with seeds and inner variables. It has to be differentiated fieldwise then."
1430 input ComponentRef cref;
1431 input UnorderedMap<ComponentRef, ComponentRef> diff_map;
1432 output Boolean b = false;
1433 protected
1434 Option<ComponentRef> der_opt = UnorderedMap.get(cref, diff_map);
1435 String root;
1436 algorithm
1437
3/6
✗ Branch 0 not taken.
✓ Branch 1 taken 6 times.
✓ Branch 2 taken 6 times.
✗ Branch 3 not taken.
✗ Branch 5 not taken.
✓ Branch 6 taken 6 times.
6 if isSome(der_opt) and BVariable.checkCref(cref, BVariable.isRecord, sourceInfo()) then
1438 6 root := crefRoot(Util.getOption(der_opt));
1439
1/2
✓ Branch 2 taken 6 times.
✗ Branch 3 not taken.
6 for child in BVariable.getRecordChildrenCref(cref) loop
1440 b := match UnorderedMap.get(ComponentRef.stripSubscriptsAll(child), diff_map)
1441 local
1442 ComponentRef child_der;
1443 ✗ case SOME(child_der) then crefRoot(child_der) <> root;
1444 else true;
1445 end match;
1446 if b then break; end if;
1447 end for;
1448 end if;
1449 end isMixedRecordDerivative;
1450
1451 function crefRoot
1452 input ComponentRef cref;
1453 output String root = listHead(Util.stringSplitAtChar(ComponentRef.toString(cref), "."));
1454 end crefRoot;
1455
1456 function makeZero
1457 "Expression.makeZero, but records without a '0' operator are zero field by field"
1458 input Type ty;
1459 output Expression zero;
1460 protected
1461 InstNode node;
1462 list<Expression> fields = {};
1463 algorithm
1464 zero := match ty
1465 case Type.COMPLEX() guard(Type.isRecord(ty) and not Restriction.isOperatorRecord(Class.restriction(InstNode.getClass(Type.complexNode(ty))))) algorithm
1466 ✗ node := Type.complexNode(ty);
1467 ✗ for comp in Class.getComponents(InstNode.getClass(node)) loop
1468 ✗ fields := makeZero(InstNode.getType(comp)) :: fields;
1469 end for;
1470 ✗ then Expression.makeRecord(InstNode.fullPath(node), ty, listReverse(fields));
1471 case Type.ARRAY() guard(Type.isRecord(Type.arrayElementType(ty)))
1472 ✗ then Expression.fillType(ty, makeZero(Type.arrayElementType(ty)));
1473 // strings of a record have no derivative
1474 case Type.STRING() then Expression.STRING("");
1475 214 else Expression.makeZero(ty);
1476 end match;
1477 end makeZero;
1478
1479 function makeOne
1480 "Expression.makeOne, but records without a '1' operator are one field by field"
1481 input Type ty;
1482 output Expression one;
1483 protected
1484 InstNode node;
1485 list<Expression> fields = {};
1486 algorithm
1487 one := match ty
1488 case Type.COMPLEX() guard(Type.isRecord(ty) and not Restriction.isOperatorRecord(Class.restriction(InstNode.getClass(Type.complexNode(ty))))) algorithm
1489 ✗ node := Type.complexNode(ty);
1490 ✗ for comp in Class.getComponents(InstNode.getClass(node)) loop
1491 ✗ fields := makeOne(InstNode.getType(comp)) :: fields;
1492 end for;
1493 ✗ then Expression.makeRecord(InstNode.fullPath(node), ty, listReverse(fields));
1494 case Type.ARRAY() guard(Type.isRecord(Type.arrayElementType(ty)))
1495 ✗ then Expression.fillType(ty, makeOne(Type.arrayElementType(ty)));
1496 // strings of a record have no derivative
1497 case Type.STRING() then Expression.STRING("");
1498 4470 else Expression.makeOne(ty);
1499 end match;
1500 end makeOne;
1501
1502 function differentiateRecordCref
1503 "A record variable whose fields are differentiated on their own: Record(der(field1), ...)."
1504 input output Expression exp;
1505 input output DifferentiationArguments diffArguments;
1506 protected
1507 list<ComponentRef> children;
1508 list<Expression> elements = {};
1509 Expression elem;
1510 ComponentRef cref;
1511 Type ty;
1512 algorithm
1513 (cref, ty) := match exp
1514 12 case Expression.CREF() then (exp.cref, exp.ty);
1515 else (ComponentRef.EMPTY(), Type.UNKNOWN());
1516 end match;
1517
2/4
✓ Branch 1 taken 12 times.
✗ Branch 2 not taken.
✓ Branch 4 taken 12 times.
✗ Branch 5 not taken.
12 if Type.isRecord(ty) and BVariable.checkCref(cref, BVariable.isRecord, sourceInfo()) then
1518 12 children := BVariable.getRecordChildrenCref(cref);
1519
1/2
✓ Branch 2 taken 12 times.
✗ Branch 3 not taken.
12 if List.compareLength(children, Type.recordFields(ty)) == 0 then
1520
2/2
✓ Branch 0 taken 36 times.
✓ Branch 1 taken 12 times.
48 for child in children loop
1521 // strings have no derivative, keep them
1522
1/2
✗ Branch 2 not taken.
✓ Branch 3 taken 36 times.
36 if Type.isString(ComponentRef.getSubscriptedType(child)) then
1523 ✗ elem := Expression.fromCref(child);
1524 else
1525 36 (elem, diffArguments) := differentiateComponentRef(Expression.fromCref(child), diffArguments);
1526 end if;
1527 elements := elem :: elements;
1528 end for;
1529 end if;
1530 end if;
1531
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 12 times.
12 if listEmpty(elements) then
1532 ✗ exp := Expression.makeZero(Expression.typeOf(exp));
1533 else
1534 12 exp := Expression.makeRecord(InstNode.fullPath(Type.complexNode(ty)), ty, listReverse(elements));
1535 end if;
1536 end differentiateRecordCref;
1537
1538 function differentiateIteratorElement
1539 "An element of an array with iterator subscripts (e.g. x[i] in a reduction) whose
1540 elements are seeds on their own: {der(x[1]), ..., der(x[n])}[i]. Zero otherwise."
1541 input Expression exp;
1542 input DifferentiationArguments diffArguments;
1543 input UnorderedMap<ComponentRef, ComponentRef> diff_map;
1544 output Expression res;
1545 protected
1546 ComponentRef base;
1547 list<Subscript> subs;
1548 list<ComponentRef> elem_crefs;
1549 list<Expression> elem_exps = {};
1550 Expression elem_res;
1551 Type base_ty;
1552 DifferentiationArguments args = diffArguments;
1553 Boolean found;
1554 algorithm
1555 593 res := Expression.makeZero(Expression.typeOf(exp));
1556 // subscripts of all parts, e.g. module[i].x for a component array
1557 (base, subs) := match exp
1558 593 case Expression.CREF() then (ComponentRef.stripSubscriptsAll(exp.cref), ComponentRef.subscriptsAllFlat(exp.cref));
1559 else (ComponentRef.EMPTY(), {});
1560 end match;
1561
4/4
✓ Branch 0 taken 233 times.
✓ Branch 1 taken 360 times.
✓ Branch 3 taken 210 times.
✓ Branch 4 taken 23 times.
593 if listEmpty(subs) or List.all(subs, Subscript.isLiteral) then return; end if;
1562 23 base_ty := ComponentRef.getSubscriptedType(base, true);
1563
2/4
✓ Branch 1 taken 23 times.
✗ Branch 2 not taken.
✗ Branch 4 not taken.
✓ Branch 5 taken 23 times.
23 if not Type.isArray(base_ty) or Type.sizeOf(base_ty) > 256 then return; end if;
1564 23 elem_crefs := listReverse(ComponentRef.scalarizeAll(base, false));
1565 found := false;
1566
2/2
✓ Branch 0 taken 240 times.
✓ Branch 1 taken 21 times.
261 for c in elem_crefs loop
1567
2/2
✓ Branch 1 taken 238 times.
✓ Branch 2 taken 2 times.
240 if UnorderedMap.contains(c, diff_map) then
1568 found := true;
1569 break;
1570 end if;
1571 end for;
1572
2/2
✓ Branch 0 taken 2 times.
✓ Branch 1 taken 21 times.
23 if not found then return; end if;
1573
2/2
✓ Branch 0 taken 12 times.
✓ Branch 1 taken 2 times.
14 for c in elem_crefs loop
1574 12 (elem_res, args) := differentiateComponentRef(Expression.fromCref(c), args);
1575 elem_exps := elem_res :: elem_exps;
1576 end for;
1577 2 res := Expression.applySubscripts(subs, makeShapedArray(base_ty, listReverse(elem_exps)));
1578 end differentiateIteratorElement;
1579
1580 function differentiateComponentRefNoCollect
1581 input output Expression exp;
1582 input output DifferentiationArguments diffArguments;
1583 protected
1584 Boolean oldCollect;
1585 algorithm
1586
2/4
✗ Branch 0 not taken.
✓ Branch 1 taken 1346 times.
✓ Branch 2 taken 1346 times.
✗ Branch 3 not taken.
1346 if isSome(diffArguments.adjoint_map) then
1587 ✗ oldCollect := diffArguments.collectAdjoints;
1588 ✗ diffArguments.collectAdjoints := false;
1589 ✗ (exp, diffArguments) := differentiateComponentRef(exp, diffArguments);
1590 ✗ diffArguments.collectAdjoints := oldCollect;
1591 else
1592 1346 (exp, diffArguments) := differentiateComponentRef(exp, diffArguments);
1593 end if;
1594 end differentiateComponentRefNoCollect;
1595
1596 function differentiateVariablePointer
1597 input Pointer<Variable> var_ptr;
1598 input Pointer<DifferentiationArguments> diffArguments_ptr;
1599 output Pointer<Variable> diff_ptr;
1600 protected
1601 DifferentiationArguments diffArguments = Pointer.access(diffArguments_ptr);
1602 Variable var = Pointer.access(var_ptr);
1603 Expression crefExp;
1604 algorithm
1605 1132 (crefExp, diffArguments) := differentiateComponentRefNoCollect(Expression.fromCref(var.name), diffArguments);
1606 diff_ptr := match crefExp
1607 14 case Expression.CREF(cref = ComponentRef.EMPTY()) then Pointer.create(NBVariable.DUMMY_VARIABLE);
1608 ✗ case Expression.CREF(cref = ComponentRef.WILD()) then Pointer.create(NBVariable.DUMMY_VARIABLE);
1609 1118 case Expression.CREF() then BVariable.getVarPointer(crefExp.cref, sourceInfo());
1610 else algorithm
1611 ✗ Error.addMessage(Error.INTERNAL_ERROR,{getInstanceName() + " failed for " + Variable.toString(var)
1612 + " because the result is expected to be a variable but turned out to be " + Expression.toString(crefExp) + "."});
1613 ✗ then fail();
1614 end match;
1615 1132 Pointer.update(diffArguments_ptr, diffArguments);
1616 end differentiateVariablePointer;
1617
1618 function differentiateCall
1619 "Differentiate builtin function calls
1620 1. if the function is builtin -> use hardcoded logic
1621 2. if the function is not builtin -> check if there is a 'fitting' derivative defined.
1622 - 'fitting' means that all the zeroDerivative annotations have to hold
1623 2.1 fitting function found -> use it
1624 2.2 fitting function not found -> differentiate the body of the function
1625 ToDo: respect the 'order' of the derivative when differentiating!"
1626 input output Expression exp "Has to be Expression.CALL()";
1627 input output DifferentiationArguments diffArguments;
1628 protected
1629 constant Boolean debug = false;
1630 algorithm
1631 if debug then
1632 print("\nDifferentiate Exp-Call: "+ Expression.toString(exp) + "\n");
1633 end if;
1634
1635 (exp, diffArguments) := match exp
1636 local
1637 Expression ret, arg;
1638 Call call;
1639 Option<Function> func_opt, der_func_opt;
1640 Function func, der_func;
1641 list<Expression> arguments = {};
1642 list<tuple<Expression, InstNode>> arguments_inputs;
1643 InstNode inp;
1644 Boolean isCont, isReal, isFunc, isSkipped, skippedVarying = false;
1645 // interface map. If the map contains a variable it has a zero derivative
1646 // if the value is "true" it has to be stripped from the interface
1647 // (it is possible that a variable has a zero derivative, but still appears in the interface)
1648 UnorderedMap<String, Boolean> interface_map;
1649
1650 // for array constructors only differentiate the argument
1651 case ret as Expression.CALL(call = call as Call.TYPED_ARRAY_CONSTRUCTOR()) algorithm
1652 24 (arg, diffArguments) := differentiateExpression(call.exp, diffArguments);
1653 24 call.exp := arg;
1654 24 ret.call := call;
1655 24 then (ret, diffArguments);
1656
1657 // handle reductions
1658 case Expression.CALL(call = call as Call.TYPED_REDUCTION()) algorithm
1659 17 (ret, diffArguments) := differentiateReduction(AbsynUtil.pathString(Function.nameConsiderBuiltin(call.fn)), exp, diffArguments);
1660 then (ret, diffArguments);
1661
1662 // builtin functions
1663 case Expression.CALL(call = call as Call.TYPED_CALL()) guard(Function.isBuiltin(call.fn)) algorithm
1664 495 (ret, diffArguments) := differentiateBuiltinCall(AbsynUtil.pathString(Function.nameConsiderBuiltin(call.fn)), exp, diffArguments);
1665 then (ret, diffArguments);
1666
1667 // user defined functions
1668 case Expression.CALL(call = call as Call.TYPED_CALL()) algorithm
1669 73 func_opt := UnorderedMap.get(call.fn.path, diffArguments.funcMap);
1670
2/4
✗ Branch 0 not taken.
✓ Branch 1 taken 73 times.
✗ Branch 2 not taken.
✓ Branch 3 taken 73 times.
73 if isSome(func_opt) then
1671 // The function is in the function tree
1672 73 SOME(func) := func_opt;
1673 73 interface_map := UnorderedMap.new<Boolean>(stringHashDjb2, stringEqual);
1674
1675 // build interface map to check if a function fits
1676 // save all inputs that would end up in a zero derivative in a map
1677 73 arguments_inputs := List.zip(call.arguments, func.inputs);
1678
2/2
✓ Branch 0 taken 376 times.
✓ Branch 1 taken 73 times.
449 for tpl in arguments_inputs loop
1679 376 (arg, inp) := tpl;
1680 // check if it is (or contains) something continuous -- do not check for functions
1681 // (differentiating a function inside a function) since crefs are not lowered
1682 // there, assume continuous. Use an OR-fold (does this expression contain AT LEAST
1683 // ONE continuous variable), not isContinuous's ALL-fold (are ALL crefs in it
1684 // continuous): a mixed expression like a constant parameter times a genuinely
1685 // time-varying variable (e.g. "e * a", a unit-vector parameter times an
1686 // acceleration) genuinely needs differentiating -- isContinuous's ALL-fold wrongly
1687 // judged it "not continuous" purely because the parameter "e" is present, silently
1688 // dropping every argument built this way from the derivative computation entirely.
1689
4/4
✓ Branch 0 taken 231 times.
✓ Branch 1 taken 145 times.
✓ Branch 3 taken 187 times.
✓ Branch 4 taken 44 times.
376 isCont := ((diffArguments.diffType == DifferentiationType.FUNCTION) or BackendUtil.containsContinuousVar(arg));
1690
1691 // input type has to be real value or a function pointer, skip if its in the interface diff info
1692 // records are differentiated fieldwise
1693
4/4
✓ Branch 3 taken 42 times.
✓ Branch 4 taken 334 times.
✓ Branch 8 taken 16 times.
✓ Branch 9 taken 26 times.
376 isReal := Type.isReal(Type.arrayElementType(Expression.typeOf(arg))) or Type.isRecord(Type.arrayElementType(Expression.typeOf(arg)));
1694 376 isFunc := InstNode.isFunction(inp);
1695 376 isSkipped := Util.applyOptionOrDefault(func.interfaceDiffInfo, function UnorderedSet.contains(key = inp), false);
1696
6/6
✓ Branch 0 taken 364 times.
✓ Branch 1 taken 12 times.
✓ Branch 2 taken 361 times.
✓ Branch 3 taken 3 times.
✓ Branch 4 taken 48 times.
✓ Branch 5 taken 313 times.
376 if isSkipped or not (isFunc or (isCont and isReal)) then
1697 // add to map; if it is not Real also already set to true (always removed from interface)
1698
2/2
✓ Branch 0 taken 37 times.
✓ Branch 1 taken 23 times.
97 UnorderedMap.add(InstNode.name(inp), not (isFunc or isReal), interface_map);
1699 end if;
1700 end for;
1701
1702 // try to get a fitting function from derivatives -> if none is found, differentiate
1703 73 der_func_opt := Function.getDerivative(func, interface_map);
1704
3/4
✗ Branch 0 not taken.
✓ Branch 1 taken 73 times.
✓ Branch 2 taken 27 times.
✓ Branch 3 taken 46 times.
73 if isSome(der_func_opt) then
1705 46 SOME(der_func) := der_func_opt;
1706 46 der_func := addDiffInfo(func, der_func, diffArguments);
1707 elseif List.any(func.inputs, InstNode.isFunction) then
1708 // the body calls the function input, which has no derivative (e.g. solveOneNonlinearEquation)
1709 3 fail();
1710 elseif Function.isExternal(func) then
1711 // external functions without a derivative annotation have no body to differentiate
1712 ✗ fail();
1713 else
1714 24 (der_func, diffArguments) := differentiateFunction(func, interface_map, diffArguments);
1715 end if;
1716
1717
2/2
✓ Branch 1 taken 367 times.
✓ Branch 2 taken 70 times.
437 for tpl in listReverse(arguments_inputs) loop
1718 367 (arg, inp) := tpl;
1719 367 isSkipped := Util.applyOptionOrDefault(func.interfaceDiffInfo, function UnorderedSet.contains(key = inp), false);
1720 // only keep the arguments which are not in the map or have value false
1721
4/4
✓ Branch 0 taken 355 times.
✓ Branch 1 taken 12 times.
✓ Branch 4 taken 314 times.
✓ Branch 5 taken 41 times.
367 if not (isSkipped or UnorderedMap.getOrDefault(InstNode.name(inp), interface_map, false)) then
1722 314 arguments := arg :: arguments;
1723 elseif isSkipped and diffArguments.diffType <> DifferentiationType.FUNCTION and BackendUtil.containsContinuousVar(arg) then
1724 // inputs of a derivative function are not differentiated again, but it still depends on them
1725 skippedVarying := true;
1726 end if;
1727 end for;
1728
1729 // differentiate type arguments and append to original ones
1730 70 (arguments, diffArguments) := List.mapFold(arguments, differentiateExpression, diffArguments);
1731
11/14
✓ Branch 0 taken 48 times.
✓ Branch 1 taken 22 times.
✓ Branch 2 taken 44 times.
✓ Branch 3 taken 4 times.
✓ Branch 5 taken 8 times.
✓ Branch 6 taken 36 times.
✓ Branch 9 taken 8 times.
✗ Branch 10 not taken.
✓ Branch 14 taken 8 times.
✗ Branch 15 not taken.
✓ Branch 18 taken 2 times.
✓ Branch 19 taken 6 times.
✓ Branch 22 taken 2 times.
✗ Branch 23 not taken.
70 if diffArguments.diffType <> DifferentiationType.FUNCTION and not skippedVarying and List.all(arguments, isZeroDerivative)
1732 and not Type.isTuple(Expression.typeOf(exp)) and not Type.isComplex(Type.arrayElementType(Expression.typeOf(exp)))
1733 and (not Type.isArray(Expression.typeOf(exp)) or Type.hasKnownSize(Expression.typeOf(exp))) then
1734 // no argument depends on the differentiation variable (keeps the arguments out of the derivative)
1735 8 ret := Expression.makeZero(Expression.typeOf(exp));
1736 else
1737 62 arguments := listAppend(call.arguments, arguments);
1738 62 ret := Expression.CALL(Call.makeTypedCall(der_func, arguments, call.var, call.purity));
1739 end if;
1740 else
1741 // The function is not in the function tree and not builtin -> error
1742 ✗ Error.addMessage(Error.INTERNAL_ERROR,{getInstanceName()
1743 + " failed because the function is not a builtin function and could not be found in the function tree: "
1744 + Expression.toString(exp)});
1745 ✗ fail();
1746 end if;
1747 70 then (ret, diffArguments);
1748
1749 // If the call was not typed correctly by the frontend
1750 else algorithm
1751 ✗ Error.addMessage(Error.INTERNAL_ERROR,{getInstanceName() + " failed for: " + Expression.toString(exp)});
1752 ✗ then fail();
1753 end match;
1754
1755 if debug then
1756 print("Differentiate-ExpCall-result: " + Expression.toString(exp) + "\n");
1757 end if;
1758 end differentiateCall;
1759
1760 function differentiateReduction
1761 "This function differentiates reduction expressions with respect to a given variable.
1762 Also creates and multiplies inner derivatives."
1763 input String name;
1764 input output Expression exp;
1765 input output DifferentiationArguments diffArguments;
1766 algorithm
1767 exp := match exp
1768 local
1769 Call call;
1770 Expression arg;
1771
1772 case Expression.CALL(call = call as Call.TYPED_REDUCTION()) guard(name == "sum") algorithm
1773 17 (arg, diffArguments) := differentiateExpression(call.exp, diffArguments);
1774 17 call.exp := arg;
1775
1/2
✓ Branch 0 taken 17 times.
✗ Branch 1 not taken.
17 exp.call := call;
1776 then exp;
1777
1778 // ToDo: product, min, max
1779
1780 else algorithm
1781 ✗ Error.addMessage(Error.INTERNAL_ERROR,{getInstanceName() + " failed because of non-call expression: " + Expression.toString(exp)});
1782 ✗ then fail();
1783 end match;
1784 end differentiateReduction;
1785
1786 function differentiateBuiltinCall
1787 "This function differentiates built-in call expressions with respect to a given variable.
1788 Also creates and multiplies inner derivatives."
1789 input String name;
1790 input output Expression exp;
1791 input output DifferentiationArguments diffArguments;
1792 protected
1793 // these need to be adapted to size and type of exp
1794 Operator.SizeClassification sizeClass = NFOperator.SizeClassification.SCALAR;
1795 Operator addOp = Operator.fromClassification((NFOperator.MathClassification.ADDITION, sizeClass), Type.REAL());
1796 Operator mulOp = Operator.fromClassification((NFOperator.MathClassification.MULTIPLICATION, sizeClass), Type.REAL());
1797 algorithm
1798 // math functions that trigger events have the index of their event values as last argument
1799 495 exp := stripMathEventIndex(name, exp);
1800 exp := match exp
1801 local
1802 Integer i;
1803 Expression ret, ret1, ret2, arg1, arg2, arg3, diffArg1, diffArg2, diffArg3, current_grad = diffArguments.current_grad, cond1, cond2, cond, zero1, zero2, grad_x, grad_y, old_grad;
1804 list<Expression> rest, diffRest;
1805 Type ty;
1806 DifferentiationType diffType;
1807 Integer rY, rX;
1808 Boolean isReverse = isSome(diffArguments.adjoint_map);
1809
1810 Type elTy;
1811 // sumG = G + Gᵀ
1812 Operator addM, subM;
1813 Expression sumG, triuG;
1814
1815 // diagG = G .* I(n), I(n) from diagonal(ones(n))
1816 Integer nExp;
1817 Expression eyeNN;
1818 Operator mulEW;
1819 Expression diagG;
1820
1821 // d/dz delay(x, delta) = (dt/dz - d delta/dz) * delay(der(x), delta)
1822 case Expression.CALL() guard(name == "delay")
1823 algorithm
1824 (arg1, arg2, arg3) := match Call.arguments(exp.call)
1825 case {arg1, arg2, arg3} then (arg1, arg2, arg3);
1826 else algorithm
1827 ✗ Error.addMessage(Error.INTERNAL_ERROR,{getInstanceName() + " failed for: " + Expression.toString(exp) + "."});
1828 ✗ then fail();
1829 end match;
1830 // if z = t then dt/dz = 1 else dt/dz = 0
1831 ✗ ret1 := Expression.REAL(if diffArguments.diffType == DifferentiationType.TIME then 1.0 else 0.0);
1832 // d delta/dz
1833 ✗ (ret2, diffArguments) := differentiateExpression(arg2, diffArguments);
1834 // dt/dz - d delta/dz
1835 ✗ ret2 := SimplifyExp.simplifyDump(Expression.MULTARY({ret1}, {ret2}, addOp), true, getInstanceName());
1836 ✗ if Expression.isZero(ret2) then
1837 ✗ ret := Expression.makeZero(Expression.typeOf(arg1));
1838 else
1839 ✗ diffType := diffArguments.diffType;
1840 ✗ diffArguments.diffType := DifferentiationType.TIME;
1841 ✗ (ret1, diffArguments) := differentiateExpression(arg1, diffArguments);
1842 ✗ diffArguments.diffType := diffType;
1843 ✗ exp.call := Call.setArguments(exp.call, {ret1, arg2, arg3});
1844 ✗ ret := Expression.MULTARY({ret2, exp}, {}, mulOp);
1845 end if;
1846 then ret;
1847
1848 // SMOOTH
1849 case Expression.CALL() guard(name == "smooth")
1850 algorithm
1851 ret := match Call.arguments(exp.call)
1852 case {arg1 as Expression.INTEGER(i), arg2} guard(i > 0) algorithm
1853 ✗ (ret2, diffArguments) := differentiateExpression(arg2, diffArguments);
1854 ✗ exp.call := Call.setArguments(exp.call, {Expression.INTEGER(i-1), ret2});
1855 then exp;
1856 case {arg1 as Expression.INTEGER(i), arg2} algorithm
1857 2 (ret2, diffArguments) := differentiateExpression(arg2, diffArguments);
1858 2 exp := Expression.CALL(Call.makeTypedCall(
1859 fn = NFBuiltinFuncs.NO_EVENT,
1860 args = {ret2},
1861 variability = Expression.variability(ret2),
1862 purity = NFPrefixes.Purity.PURE
1863 ));
1864 then exp;
1865 else algorithm
1866 ✗ Error.addMessage(Error.INTERNAL_ERROR,{getInstanceName() + " failed for: " + Expression.toString(exp) + "."});
1867 ✗ then fail();
1868 end match;
1869 then ret;
1870
1871 case Expression.CALL() guard(name == "sum")
1872 algorithm
1873 arg1 := match Call.arguments(exp.call)
1874 case {arg1} then arg1;
1875 else algorithm
1876 ✗ Error.addMessage(Error.INTERNAL_ERROR,{getInstanceName() + " failed for: " + Expression.toString(exp) + "."});
1877 ✗ then fail();
1878 end match;
1879
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 22 times.
22 if isReverse then
1880 ✗ current_grad := diffArguments.current_grad;
1881 // sum is linear: propagate the scalar upstream adjoint to the array elements
1882 // by differentiating the argument with the upstream gradient already set.
1883 ✗ diffArguments.current_grad := current_grad;
1884
1885 ✗ (ret1, diffArguments) := differentiateExpression(arg1, diffArguments);
1886
1887 // restore upstream
1888 ✗ diffArguments.current_grad := current_grad;
1889 else
1890 22 (ret1, diffArguments) := differentiateExpression(arg1, diffArguments);
1891 end if;
1892
1893 22 exp.call := Call.setArguments(exp.call, {ret1});
1894 then exp;
1895
1896 // symmetric(A):
1897 // Forward: symmetric(dA/dz)
1898 // Reverse: grad_A = triu(G + Gᵀ) - diag(G)
1899 case Expression.CALL() guard(name == "symmetric")
1900 algorithm
1901 arg1 := match Call.arguments(exp.call)
1902 case {arg1} then arg1;
1903 else algorithm
1904 ✗ Error.addMessage(Error.INTERNAL_ERROR, {getInstanceName() + " failed for: " + Expression.toString(exp) + "."});
1905 ✗ then fail();
1906 end match;
1907
1908 ✗ if isReverse then
1909 ✗ current_grad := diffArguments.current_grad;
1910
1911 // upstream gradient type (matrix)
1912 ✗ ty := Expression.typeOf(current_grad);
1913 // element type
1914 ✗ elTy := if Type.isArray(ty) then Type.arrayElementType(ty) else ty;
1915 // matrix dimension (assume square)
1916 ✗ nExp := Dimension.size(listHead(Type.arrayDims(Expression.typeOf(arg1))));
1917
1918 // element-wise add / mul operators with full matrix type (not element type)
1919 ✗ addM := Operator.fromClassification(
1920 (NFOperator.MathClassification.ADDITION, NFOperator.SizeClassification.ELEMENT_WISE),
1921 ty);
1922 ✗ subM := Operator.fromClassification(
1923 (NFOperator.MathClassification.SUBTRACTION, NFOperator.SizeClassification.ELEMENT_WISE),
1924 ty);
1925 ✗ mulEW := Operator.fromClassification(
1926 (NFOperator.MathClassification.MULTIPLICATION, NFOperator.SizeClassification.ELEMENT_WISE),
1927 ty);
1928
1929 // sumG = G + Gᵀ (binary)
1930 ✗ sumG := Expression.BINARY(
1931 current_grad,
1932 addM,
1933 typeTransposeCall(current_grad));
1934
1935 // triu(sumG) = sumG .* triu(ones(n,n)) (binary)
1936 ✗ triuG := Expression.BINARY(
1937 sumG,
1938 mulEW,
1939 Expression.makeTriuMask(nExp, elTy));
1940
1941 // I(n)
1942 ✗ eyeNN := Expression.makeIdentityMatrix(nExp, elTy);
1943
1944 // diagG = G .* I (binary)
1945 ✗ diagG := Expression.BINARY(
1946 current_grad,
1947 mulEW,
1948 eyeNN);
1949
1950 // triu(G + Gᵀ) - diag(G) (binary)
1951 ✗ diffArguments.current_grad := Expression.BINARY(
1952 triuG,
1953 subM,
1954 diagG);
1955 end if;
1956
1957 // Forward: symmetric(dA/dz)
1958 ✗ (ret1, diffArguments) := differentiateExpression(arg1, diffArguments);
1959
1960 ✗ if isReverse then
1961 // restore upstream
1962 ✗ diffArguments.current_grad := current_grad;
1963 end if;
1964 ✗ exp.call := Call.setArguments(exp.call, {ret1});
1965 then exp;
1966
1967 // diagonal(v):
1968 // Forward: diagonal(dv/dz)
1969 // Reverse: grad_v = diag(G) (extract diagonal of upstream matrix)
1970 case Expression.CALL() guard(name == "diagonal")
1971 algorithm
1972 arg1 := match Call.arguments(exp.call)
1973 case {arg1} then arg1;
1974 else algorithm
1975 ✗ Error.addMessage(Error.INTERNAL_ERROR, {getInstanceName() + " failed for: " + Expression.toString(exp) + "."});
1976 ✗ then fail();
1977 end match;
1978
1979 ✗ if isReverse then
1980 ✗ current_grad := diffArguments.current_grad;
1981 // number of elements in v and in diagonal of G
1982 ✗ nExp := Dimension.size(listHead(Type.arrayDims(Expression.typeOf(arg1))));
1983 // Literal: [ G[1,1], G[2,2], ..., G[n,n] ]
1984 ✗ diffArguments.current_grad := extractDiagonalVector(current_grad, nExp, Expression.typeOf(arg1));
1985 end if;
1986
1987 // Forward: diagonal(dv/dz)
1988 ✗ (ret1, diffArguments) := differentiateExpression(arg1, diffArguments);
1989
1990 ✗ if isReverse then
1991 // Restore upstream and return updated call
1992 ✗ diffArguments.current_grad := current_grad;
1993 end if;
1994 ✗ exp.call := Call.setArguments(exp.call, {ret1});
1995 then exp;
1996
1997 // matrix(A)
1998 // Forward: matrix(dA/dz)
1999 // Reverse: let rX = ndims(A), G the upstream matrix:
2000 // - if rX < 2: dropLastDimIndex1(G) (2-rX times)
2001 // - if rX = 2: G
2002 // - if rX > 2: promote(G, rX)
2003 case Expression.CALL() guard(name == "matrix")
2004 algorithm
2005 arg1 := match Call.arguments(exp.call)
2006 case {arg1} then arg1;
2007 else algorithm
2008 ✗ Error.addMessage(Error.INTERNAL_ERROR, {getInstanceName() + " failed for: " + Expression.toString(exp) + "."});
2009 ✗ then fail();
2010 end match;
2011
2012 ✗ if isReverse then
2013 ✗ current_grad := diffArguments.current_grad;
2014 // Rank of input A
2015 ✗ ty := Expression.typeOf(arg1);
2016 ✗ rX := if Type.isArray(ty) then Type.dimensionCount(ty) else 0;
2017
2018 // Map upstream gradient back to A's shape
2019 grad_x := current_grad;
2020
2021 // If A has rank < 2, drop trailing dims by indexing with 1
2022 ✗ if rX < 2 then
2023 ✗ for i in 1:(2 - rX) loop
2024 ✗ grad_x := dropLastDimIndex1(grad_x);
2025 end for;
2026 elseif rX > 2 then
2027 // If A has rank > 2 (with trailing singleton dims), promote G to rank rX
2028 ✗ grad_x := typePromoteCall(grad_x, rX);
2029 end if;
2030
2031 // Recurse into A with mapped upstream gradient
2032 ✗ diffArguments.current_grad := grad_x;
2033
2034 ✗ (ret1, diffArguments) := differentiateExpression(arg1, diffArguments);
2035
2036 // restore upstream
2037 ✗ diffArguments.current_grad := current_grad;
2038 else
2039 ✗ (ret1, diffArguments) := differentiateExpression(arg1, diffArguments);
2040 end if;
2041 // Forward: matrix(dA/dz)
2042 ✗ exp.call := Call.setArguments(exp.call, {ret1});
2043 then exp;
2044
2045 // Functions with one argument that differentiate "through"
2046 // through means that the derivative of the function wrt. its input is equal to the function of derivative of input
2047 // d/dz f(x) -> f(dx/dz)
2048 case Expression.CALL() guard(List.contains({"pre", "noEvent", "scalar", "vector", "transpose", "skew"}, name, stringEqual))
2049 algorithm
2050 arg1 := match Call.arguments(exp.call)
2051 case {arg1} then arg1;
2052 else algorithm
2053 ✗ Error.addMessage(Error.INTERNAL_ERROR,{getInstanceName() + " failed for: " + Expression.toString(exp) + "."});
2054 ✗ then fail();
2055 end match;
2056 2 (ret1, diffArguments) := differentiateExpression(arg1, diffArguments);
2057 2 exp.call := Call.setArguments(exp.call, {ret1});
2058 then exp;
2059
2060 // Functions with two arguments that differentiate "through"
2061 // df(x,y)/dz = f(dx/dz, dy/dz)
2062 case Expression.CALL() guard(List.contains({"homotopy", "$OMC$inStreamDiv"}, name, stringEqual))
2063 algorithm
2064 (arg1, arg2) := match Call.arguments(exp.call)
2065 case {arg1, arg2} then (arg1, arg2);
2066 else algorithm
2067 ✗ Error.addMessage(Error.INTERNAL_ERROR,{getInstanceName() + " failed for: " + Expression.toString(exp) + "."});
2068 ✗ then fail();
2069 end match;
2070 9 (ret1, diffArguments) := differentiateExpression(arg1, diffArguments);
2071 9 (ret2, diffArguments) := differentiateExpression(arg2, diffArguments);
2072 9 exp.call := Call.setArguments(exp.call, {ret1, ret2});
2073 then exp;
2074
2075 // d/dz cat(k, A, B, C, ...) = cat(k, dA/dz, dB/dz, dC/dz, ...)
2076 case Expression.CALL() guard name == "cat"
2077 algorithm
2078 ✗ if isReverse then
2079 ✗ Error.addInternalError(getInstanceName() + " failed for: " + Expression.toString(exp) + "\nReverse Mode not implemented for `cat()`.", sourceInfo());
2080 ✗ fail();
2081 end if;
2082
2083 ✗ arg1 :: rest := Call.arguments(exp.call);
2084 diffRest := {};
2085 ✗ for arg in listReverse(rest) loop
2086 ✗ (ret, diffArguments) := differentiateExpression(arg, diffArguments);
2087 diffRest := ret :: diffRest;
2088 end for;
2089 ✗ exp.call := Call.setArguments(exp.call, arg1 :: diffRest);
2090 then exp;
2091
2092 // d/dz promote(A, n) = promote(dA/dz, n)
2093 case Expression.CALL() guard(name == "promote")
2094 algorithm
2095 (arg1, arg2) := match Call.arguments(exp.call)
2096 case {arg1, arg2} then (arg1, arg2);
2097 else algorithm
2098 ✗ Error.addMessage(Error.INTERNAL_ERROR,{getInstanceName() + " failed for: " + Expression.toString(exp) + "."});
2099 ✗ then fail();
2100 end match;
2101 ✗ if isReverse then
2102 ✗ rY := if Type.isArray(Expression.typeOf(exp)) then Type.dimensionCount(Expression.typeOf(exp)) else 0;
2103 ✗ rX := if Type.isArray(Expression.typeOf(arg1)) then Type.dimensionCount(Expression.typeOf(arg1)) else 0;
2104 ✗ current_grad := diffArguments.current_grad;
2105 old_grad := current_grad;
2106 ✗ for i in 1:max(0, rY - rX) loop
2107 ✗ current_grad := dropLastDimIndex1(current_grad);
2108 end for;
2109 ✗ diffArguments.current_grad := current_grad;
2110 ✗ (ret1, diffArguments) := differentiateExpression(arg1, diffArguments);
2111 ✗ diffArguments.current_grad := old_grad;
2112 else
2113 ✗ (ret1, diffArguments) := differentiateExpression(arg1, diffArguments);
2114 end if;
2115 ✗ exp.call := Call.setArguments(exp.call, {ret1, arg2});
2116 then exp;
2117
2118 // d/dz identity(n) = zeros(n, n)
2119 case Expression.CALL() guard(name == "identity")
2120 algorithm
2121 // diffArguments.current_grad := Expression.makeZero(Expression.typeOf(exp));?
2122 arg1 := match Call.arguments(exp.call)
2123 case {arg1} then arg1;
2124 else algorithm
2125 ✗ Error.addMessage(Error.INTERNAL_ERROR,{getInstanceName() + " failed for: " + Expression.toString(exp) + "."});
2126 ✗ then fail();
2127 end match;
2128 ✗ then Expression.CALL(Call.makeTypedCall(
2129 fn = NFBuiltinFuncs.FILL_FUNC,
2130 args = {Expression.INTEGER(0), arg1, arg1},
2131 variability = Variability.CONSTANT,
2132 purity = NFPrefixes.Purity.PURE
2133 ));
2134
2135 // d/dz fill(x, n1, n2, ...) = fill(dx/dz, n1, n2, ...)
2136 case Expression.CALL() guard(name == "fill")
2137 algorithm
2138 // only differentiate 1st input
2139 ✗ arg1 :: rest := Call.arguments(exp.call);
2140 ✗ if isReverse then
2141 ✗ rY := if Type.isArray(Expression.typeOf(exp)) then Type.dimensionCount(Expression.typeOf(exp)) else 0;
2142 ✗ rX := if Type.isArray(Expression.typeOf(arg1)) then Type.dimensionCount(Expression.typeOf(arg1)) else 0;
2143 ✗ current_grad := diffArguments.current_grad;
2144 old_grad := current_grad;
2145 ✗ for i in 1:max(0, rY - rX) loop // reduce over all added dimensions with sum (TODO: change to only sum over added dimensions)
2146 ✗ current_grad := typeSumCall(current_grad); // sum over first (or last?) dimension
2147 end for;
2148 ✗ diffArguments.current_grad := current_grad;
2149 ✗ (ret1, diffArguments) := differentiateExpression(arg1, diffArguments);
2150 ✗ diffArguments.current_grad := old_grad;
2151 else
2152 ✗ (ret1, diffArguments) := differentiateExpression(arg1, diffArguments);
2153 end if;
2154 ✗ exp.call := Call.setArguments(exp.call, ret1 :: rest);
2155 then exp;
2156
2157 // SEMI LINEAR
2158 // d sL(x, m1, m2)/dz = sL(x, dm1/dz, dm2/dz) + dx/dz * (if x >= 0 then m1 else m2)
2159 case Expression.CALL() guard(name == "semiLinear")
2160 algorithm
2161 (arg1, arg2, arg3) := match Call.arguments(exp.call)
2162 case {arg1, arg2, arg3} then (arg1, arg2, arg3);
2163 else algorithm
2164 ✗ Error.addMessage(Error.INTERNAL_ERROR,{getInstanceName() + " failed for: " + Expression.toString(exp) + "."});
2165 ✗ then fail();
2166 end match;
2167 10 current_grad := diffArguments.current_grad;
2168
2169
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 10 times.
10 if isReverse then
2170 ✗ cond := Expression.RELATION(
2171 arg1, // x
2172 Operator.makeGreaterEq(Expression.typeOf(arg1)),
2173 Expression.makeZero(Expression.typeOf(arg1)),
2174 -1);
2175
2176 ✗ grad_x := Expression.IF(
2177 Expression.typeOf(arg1),
2178 cond,
2179 Expression.MULTARY({arg2, current_grad}, {}, mulOp), // d(positive_slope * x)/dx = positive_slope * current_grad
2180 Expression.MULTARY({arg3, current_grad}, {}, mulOp) // d(negative_slope * x)/dx = negative_slope * current_grad
2181 );
2182 ✗ diffArguments.current_grad := grad_x;
2183 end if;
2184
2185 // dx/dz, dm1/dz, dm2/dz
2186 10 (diffArg1, diffArguments) := differentiateExpression(arg1, diffArguments);
2187 10 diffArguments.current_grad := current_grad; // restore upstream
2188 10 (diffArg2, diffArguments) := differentiateExpression(arg2, diffArguments);
2189 10 (diffArg3, diffArguments) := differentiateExpression(arg3, diffArguments);
2190
2191 // sL(x, dm1/dz, dm2/dz)
2192 10 exp.call := Call.setArguments(exp.call, {arg1, diffArg2, diffArg3});
2193 ret := exp;
2194
2195 // only add second part if dx/dz is nonzero
2196
1/2
✗ Branch 1 not taken.
✓ Branch 2 taken 10 times.
10 if not Expression.isZero(diffArg1) then
2197 ✗ ty := Expression.typeOf(diffArg1);
2198 // x >= 0
2199 ✗ ret1 := Expression.RELATION(arg1, Operator.makeGreaterEq(ty), Expression.makeZero(ty), -1);
2200 // if x >= 0 then m1 else m2
2201 ✗ ret1 := Expression.IF(ty, ret1, arg2, arg3);
2202 // dx/dz * (if x >= 0 then m1 else m2)
2203 ✗ ret2 := Expression.MULTARY({diffArg1, ret1}, {}, mulOp);
2204 // sL(x, dm1/dz, dm2/dz) + dx/dz * (if x >= 0 then m1 else m2)
2205 ✗ ret := Expression.MULTARY({ret, ret2}, {}, addOp);
2206 end if;
2207 then ret;
2208
2209 // d/dz min(X) = (dX/dz)[argmin(X)]
2210 // d/dz max(X) = (dX/dz)[argmax(X)]
2211 // d/dz min(x,y) = if x < y then dx/dz else dy/dz
2212 // d/dz max(x,y) = if x > y then dx/dz else dy/dz
2213 case Expression.CALL() guard(name == "min" or name == "max")
2214 algorithm
2215 ret := match Call.arguments(exp.call)
2216 case {arg1} algorithm
2217 // dX/dz
2218 ✗ (diffArg1, diffArguments) := differentiateExpression(arg1, diffArguments);
2219 ✗ ty := Expression.typeOf(diffArg1);
2220 ✗ if Expression.isZero(diffArg1) then
2221 // make 0 of reduced type
2222 ✗ ret := Expression.makeZero(Type.arrayElementType(ty));
2223 else
2224 ✗ ret1 := Expression.CALL(Call.makeTypedCall(
2225 fn = if name == "min" then NFBuiltinFuncs.ARG_MIN_ARR_REAL else NFBuiltinFuncs.ARG_MAX_ARR_REAL,
2226 args = {arg1},
2227 variability = Expression.variability(arg1),
2228 purity = NFPrefixes.Purity.PURE));
2229 ✗ ret := Expression.applySubscripts({Subscript.INDEX(ret1)}, diffArg1, true);
2230 end if;
2231 then ret;
2232
2233 case {arg1, arg2} algorithm
2234
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 292 times.
292 if isReverse then
2235 ✗ current_grad := diffArguments.current_grad;
2236 // Relation: for min use x<y; for max use x>y
2237 ✗ cond1 := Expression.RELATION(
2238 arg1,
2239 if name == "min" then Operator.makeLess(Expression.typeOf(arg1))
2240 else Operator.makeGreater(Expression.typeOf(arg1)),
2241 arg2,
2242 -1);
2243 ✗ cond2 := Expression.RELATION(
2244 arg2,
2245 if name == "min" then Operator.makeLess(Expression.typeOf(arg2))
2246 else Operator.makeGreater(Expression.typeOf(arg2)),
2247 arg1,
2248 -1);
2249
2250 // Reverse local masks:
2251 // For min: grad_x = upstream if x<y else 0; grad_y = upstream if x>=y else 0
2252 // For max: grad_x = upstream if x>y else 0; grad_y = upstream if x<=y else 0
2253 ✗ zero1 := Expression.makeZero(Expression.typeOf(arg1));
2254 ✗ zero2 := Expression.makeZero(Expression.typeOf(arg2));
2255
2256 ✗ grad_x := Expression.IF(
2257 Expression.typeOf(arg1),
2258 cond1,
2259 current_grad,
2260 zero1);
2261
2262 ✗ grad_y := Expression.IF(
2263 Expression.typeOf(arg2),
2264 cond2,
2265 current_grad,
2266 zero2);
2267
2268 // Reverse recurse arg1 with grad_x
2269 ✗ old_grad := diffArguments.current_grad;
2270 ✗ diffArguments.current_grad := grad_x;
2271 // dx/dz
2272 ✗ (diffArg1, diffArguments) := differentiateExpression(arg1, diffArguments);
2273
2274 // Reverse recurse arg2 with grad_y
2275 ✗ diffArguments.current_grad := grad_y;
2276 // dy/dz
2277 ✗ (diffArg2, diffArguments) := differentiateExpression(arg2, diffArguments);
2278
2279 // Restore upstream
2280 ✗ diffArguments.current_grad := old_grad;
2281 else
2282 // Forward: dx/dz and dy/dz
2283 292 (diffArg1, diffArguments) := differentiateExpression(arg1, diffArguments);
2284 292 (diffArg2, diffArguments) := differentiateExpression(arg2, diffArguments);
2285 end if;
2286
2287 292 ty := Expression.typeOf(diffArg1);
2288
3/4
✓ Branch 1 taken 144 times.
✓ Branch 2 taken 148 times.
✓ Branch 4 taken 144 times.
✗ Branch 5 not taken.
292 if Expression.isZero(diffArg1) and Expression.isZero(diffArg2) then
2289 144 ret := Expression.makeZero(ty);
2290 else
2291 // condition x < y or x > y
2292
3/4
✓ Branch 0 taken 148 times.
✗ Branch 1 not taken.
✓ Branch 3 taken 2 times.
✓ Branch 4 taken 146 times.
148 ret1 := Expression.RELATION(arg1, if name == "min" then Operator.makeLess(ty) else Operator.makeGreater(ty), arg2, -1);
2293 // if condition then dx/dz else dy/dz
2294 148 ret := Expression.IF(ty, ret1, diffArg1, diffArg2);
2295 end if;
2296 then ret;
2297 else algorithm
2298 ✗ Error.addMessage(Error.INTERNAL_ERROR,{getInstanceName() + " failed for: " + Expression.toString(exp) + "."});
2299 ✗ then fail();
2300 end match;
2301 then ret;
2302
2303 // Builtin function call with one argument
2304 // df(x)/dz = df/dx * dx/dz
2305 case Expression.CALL() guard List.hasOneElement(Call.arguments(exp.call))
2306 algorithm
2307 arg1 := match Call.arguments(exp.call)
2308 case {arg1} then arg1;
2309 else algorithm
2310 ✗ Error.addMessage(Error.INTERNAL_ERROR,{getInstanceName() + " failed for: " + Expression.toString(exp) + "."});
2311 ✗ then fail();
2312 end match;
2313 // differentiate the call df/dx
2314 158 ret := differentiateBuiltinCall1Arg(name, arg1);
2315
2/2
✓ Branch 1 taken 156 times.
✓ Branch 2 taken 2 times.
158 if not Expression.isZero(ret) then
2316 156 current_grad := diffArguments.current_grad;
2317
2318 312 diffArguments.current_grad := Expression.MULTARY({current_grad, ret}, {}, mulOp);
2319 // differentiate the argument (inner derivative) dx/dz
2320 156 (diffArg1, diffArguments) := differentiateExpression(arg1, diffArguments);
2321
2322 156 diffArguments.current_grad := current_grad;
2323 156 ret := Expression.MULTARY({ret, diffArg1}, {}, mulOp);
2324 end if;
2325 then ret;
2326
2327 // Builtin function call with two arguments
2328 // df(x,y)/dz = df/dx * dx/dz + df/dy * dy/dz
2329 case Expression.CALL() guard(listLength(Call.arguments(exp.call)) == 2)
2330 algorithm
2331 (arg1, arg2) := match Call.arguments(exp.call)
2332 case {arg1, arg2} then (arg1, arg2);
2333 else algorithm
2334 ✗ Error.addMessage(Error.INTERNAL_ERROR,{getInstanceName() + " failed for: " + Expression.toString(exp) + "."});
2335 ✗ then fail();
2336 end match;
2337 // differentiate the call
2338 ✗ (ret1, ret2) := differentiateBuiltinCall2Arg(name, arg1, arg2); // df/dx and df/dy
2339 ✗ current_grad := diffArguments.current_grad;
2340
2341 ✗ diffArguments.current_grad := Expression.MULTARY({current_grad, ret1}, {}, mulOp);
2342 ✗ (diffArg1, diffArguments) := differentiateExpression(arg1, diffArguments); // dx/dz
2343
2344 ✗ diffArguments.current_grad := Expression.MULTARY({current_grad, ret2}, {}, mulOp);
2345 ✗ (diffArg2, diffArguments) := differentiateExpression(arg2, diffArguments); // dy/dz
2346
2347 ✗ diffArguments.current_grad := current_grad;
2348 ✗ ret1 := Expression.MULTARY({ret1, diffArg1}, {}, mulOp); // df/dx * dx/dz
2349 ✗ ret2 := Expression.MULTARY({ret2, diffArg2}, {}, mulOp); // df/dy * dy/dz
2350 ✗ ret := Expression.MULTARY({ret1,ret2}, {}, addOp); // df/dx * dx/dz + df/dy * dy/dz
2351 then ret;
2352
2353 // try some simple known cases
2354 case Expression.CALL() algorithm
2355 ret := match Call.functionNameLast(exp.call)
2356 case "sample" then Expression.BOOLEAN(false);
2357 else algorithm
2358 ✗ Error.addMessage(Error.INTERNAL_ERROR,{getInstanceName() + " failed for: " + Expression.toString(exp)});
2359 ✗ then fail();
2360 end match;
2361 then ret;
2362
2363 else algorithm
2364 ✗ Error.addMessage(Error.INTERNAL_ERROR,{getInstanceName() + " failed because of non-call expression: " + Expression.toString(exp)});
2365 ✗ then fail();
2366 end match;
2367 end differentiateBuiltinCall;
2368
2369 function stripMathEventIndex
2370 "integer(x, index), floor(x, index), ceil(x, index), div(x, y, index) and mod(x, y, index)
2371 are differentiated like the functions without the index"
2372 input String name;
2373 input output Expression exp;
2374 protected
2375 list<Expression> args;
2376 Integer n;
2377 algorithm
2378 exp := match exp
2379 case Expression.CALL() algorithm
2380 495 args := Call.arguments(exp.call);
2381 495 n := listLength(args);
2382
15/24
✓ Branch 0 taken 2 times.
✓ Branch 1 taken 493 times.
✓ Branch 3 taken 2 times.
✗ Branch 4 not taken.
✓ Branch 5 taken 15 times.
✓ Branch 6 taken 480 times.
✓ Branch 8 taken 15 times.
✗ Branch 9 not taken.
✓ Branch 10 taken 22 times.
✓ Branch 11 taken 473 times.
✗ Branch 13 not taken.
✓ Branch 14 taken 22 times.
✗ Branch 15 not taken.
✗ Branch 16 not taken.
✓ Branch 17 taken 435 times.
✓ Branch 18 taken 60 times.
✓ Branch 20 taken 435 times.
✗ Branch 21 not taken.
✓ Branch 22 taken 435 times.
✓ Branch 23 taken 60 times.
✗ Branch 25 not taken.
✓ Branch 26 taken 435 times.
✗ Branch 27 not taken.
✗ Branch 28 not taken.
495 if ((name == "integer" or name == "floor" or name == "ceil") and n == 2) or ((name == "div" or name == "mod") and n == 3) then
2383 ✗ exp.call := Call.setArguments(exp.call, List.firstN(args, n - 1));
2384 end if;
2385 then exp;
2386 else exp;
2387 end match;
2388 end stripMathEventIndex;
2389
2390 function differentiateBuiltinCall1Arg
2391 "differentiate a builtin call with one argument."
2392 input String name;
2393 input Expression arg;
2394 output Expression derFuncCall;
2395 protected
2396 // these probably need to be adapted to the size and type of arg
2397 Operator.SizeClassification sizeClass = NFOperator.SizeClassification.SCALAR;
2398 Operator powOp = Operator.fromClassification((NFOperator.MathClassification.POWER, sizeClass), Type.REAL());
2399 Operator addOp = Operator.fromClassification((NFOperator.MathClassification.ADDITION, sizeClass), Type.REAL());
2400 Operator mulOp = Operator.fromClassification((NFOperator.MathClassification.MULTIPLICATION, sizeClass), Type.REAL());
2401 algorithm
2402 derFuncCall := match name
2403 local
2404 Expression ret;
2405
2406 // all these have integer values and therefore zero derivative
2407 case "sign" then Expression.INTEGER(0);
2408 case "ceil" then Expression.REAL(0.0);
2409 case "floor" then Expression.REAL(0.0);
2410 case "integer" then Expression.INTEGER(0);
2411
2412 // abs(arg) -> sign(arg)
2413 18 case "abs" then Expression.CAST(
2414 Expression.typeOf(arg),
2415 Expression.CALL(Call.makeTypedCall(
2416 fn = NFBuiltinFuncs.SIGN,
2417 args = {arg},
2418 variability = Expression.variability(arg),
2419 purity = NFPrefixes.Purity.PURE
2420 )));
2421
2422 // sqrt(arg) -> 0.5/arg^(0.5)
2423 case "sqrt" algorithm
2424 14 ret := Expression.BINARY(arg, powOp, Expression.REAL(0.5)); // arg^0.5
2425 14 ret := Expression.MULTARY({Expression.REAL(0.5)}, {ret}, mulOp); // 1/(2*arg^0.5)
2426 then ret;
2427
2428 // sin(arg) -> cos(arg)
2429 100 case "sin" then Expression.CALL(Call.makeTypedCall(
2430 fn = NFBuiltinFuncs.COS_REAL,
2431 args = {arg},
2432 variability = Expression.variability(arg),
2433 purity = NFPrefixes.Purity.PURE
2434 ));
2435
2436 // cos(arg) -> -sin(arg)
2437 38 case "cos" then Expression.negate(Expression.CALL(Call.makeTypedCall(
2438 fn = NFBuiltinFuncs.SIN_REAL,
2439 args = {arg},
2440 variability = Expression.variability(arg),
2441 purity = NFPrefixes.Purity.PURE
2442 )));
2443
2444 // tan(arg) -> 1/cos(arg)^2
2445 // kabdelhak: ToDo - investigate numerical properties: 1+tan(arg)^2 maybe better?
2446 case "tan" algorithm
2447 ✗ ret := Expression.CALL(Call.makeTypedCall(
2448 fn = NFBuiltinFuncs.COS_REAL,
2449 args = {arg},
2450 variability = Expression.variability(arg),
2451 purity = NFPrefixes.Purity.PURE)); // cos(arg)
2452 ✗ ret := Expression.BINARY(ret, powOp, Expression.REAL(2.0)); // cos(arg)^2
2453 ✗ ret := Expression.MULTARY({Expression.REAL(1.0)}, {ret}, mulOp); // 1/cos(arg)^2
2454 then ret;
2455
2456 // asin(arg) -> 1/sqrt(1-arg^2)
2457 case "asin" algorithm
2458 ✗ ret := Expression.BINARY(arg, powOp, Expression.REAL(2.0)); // arg^2
2459 ✗ ret := Expression.MULTARY({Expression.REAL(1.0)}, {ret}, addOp); // 1-arg^2
2460 ✗ ret := Expression.BINARY(ret, powOp, Expression.REAL(0.5)); // sqrt(1-arg^2)
2461 ✗ ret := Expression.MULTARY({Expression.REAL(1.0)}, {ret}, mulOp); // 1/sqrt(1-arg^2)
2462 then ret;
2463
2464 // acos(arg) -> -1/sqrt(1-arg^2)
2465 case "acos" algorithm
2466 ✗ ret := Expression.BINARY(arg, powOp, Expression.REAL(2.0)); // arg^2
2467 ✗ ret := Expression.MULTARY({Expression.REAL(1.0)}, {ret}, addOp); // 1-arg^2
2468 ✗ ret := Expression.BINARY(ret, powOp, Expression.REAL(0.5)); // sqrt(1-arg^2)
2469 ✗ ret := Expression.MULTARY({Expression.REAL(-1.0)}, {ret}, mulOp); // -1/sqrt(1-arg^2)
2470 then ret;
2471
2472 // atan(arg) -> 1/(1+arg^2)
2473 case "atan" algorithm
2474 6 ret := Expression.BINARY(arg, powOp, Expression.REAL(2.0)); // arg^2
2475 6 ret := Expression.MULTARY({Expression.REAL(1.0), ret}, {}, addOp);// 1+arg^2
2476 6 ret := Expression.MULTARY({Expression.REAL(1.0)}, {ret}, mulOp); // 1/(1+arg^2)
2477 then ret;
2478
2479 // sinh(arg) -> cosh(arg)
2480 ✗ case "sinh" then Expression.CALL(Call.makeTypedCall(
2481 fn = NFBuiltinFuncs.COSH_REAL,
2482 args = {arg},
2483 variability = Expression.variability(arg),
2484 purity = NFPrefixes.Purity.PURE
2485 ));
2486
2487 // cosh(arg) -> sinh(arg)
2488 ✗ case "cosh" then Expression.CALL(Call.makeTypedCall(
2489 fn = NFBuiltinFuncs.SINH_REAL,
2490 args = {arg},
2491 variability = Expression.variability(arg),
2492 purity = NFPrefixes.Purity.PURE
2493 ));
2494
2495 // tanh(arg) -> 1-tanh(arg)^2
2496 case "tanh" algorithm
2497 ✗ ret := Expression.CALL(Call.makeTypedCall(
2498 fn = NFBuiltinFuncs.TANH_REAL,
2499 args = {arg},
2500 variability = Expression.variability(arg),
2501 purity = NFPrefixes.Purity.PURE)); // tanh(arg)
2502 ✗ ret := Expression.BINARY(ret, powOp, Expression.REAL(2.0)); // tanh(arg)^2
2503 ✗ ret := Expression.MULTARY({Expression.REAL(1.0)}, {ret}, addOp); // 1-tanh(arg)^2
2504 then ret;
2505
2506 // acosh(arg) -> 1/sqrt(arg^2-1)
2507 case "acosh" algorithm
2508 ✗ ret := Expression.BINARY(arg, powOp, Expression.REAL(2.0)); // arg^2
2509 ✗ ret := Expression.MULTARY({ret}, {Expression.REAL(1.0)}, addOp); // arg^2-1
2510 ✗ ret := Expression.BINARY(ret, powOp, Expression.REAL(0.5)); // sqrt(arg^2-1)
2511 ✗ ret := Expression.MULTARY({Expression.REAL(1.0)}, {ret}, mulOp); // 1/sqrt(arg^2-1)
2512 then ret;
2513
2514 // asinh(arg) -> 1/sqrt(arg^2+1)
2515 case "asinh" algorithm
2516 ✗ ret := Expression.BINARY(arg, powOp, Expression.REAL(2.0)); // arg^2
2517 ✗ ret := Expression.MULTARY({ret, Expression.REAL(1.0)}, {}, addOp); // arg^2+1
2518 ✗ ret := Expression.BINARY(ret, powOp, Expression.REAL(0.5)); // sqrt(arg^2+1)
2519 ✗ ret := Expression.MULTARY({Expression.REAL(1.0)}, {ret}, mulOp); // 1/sqrt(arg^2+1)
2520 then ret;
2521
2522 // atanh(arg) -> 1/(1-arg^2)
2523 case "atanh" algorithm
2524 ✗ ret := Expression.BINARY(arg, powOp, Expression.REAL(2.0)); // arg^2
2525 ✗ ret := Expression.MULTARY({Expression.REAL(1.0)}, {ret}, addOp); // 1-arg^2
2526 ✗ ret := Expression.MULTARY({Expression.REAL(1.0)}, {ret}, mulOp); // 1/(1-arg^2)
2527 then ret;
2528
2529 // exp(arg) -> exp(arg)
2530 48 case "exp" then Expression.CALL(Call.makeTypedCall(
2531 fn = NFBuiltinFuncs.EXP_REAL,
2532 args = {arg},
2533 variability = Expression.variability(arg),
2534 purity = NFPrefixes.Purity.PURE
2535 ));
2536
2537 // log(arg) -> 1/arg
2538 19 case "log" then Expression.MULTARY({Expression.REAL(1.0)}, {arg}, mulOp);
2539
2540 // log10(arg) -> 1/(arg*log(10))
2541 case "log10" algorithm
2542 15 ret := Expression.CALL(Call.makeTypedCall(
2543 fn = NFBuiltinFuncs.LOG_REAL,
2544 args = {Expression.REAL(10.0)},
2545 variability = Variability.CONSTANT,
2546 purity = NFPrefixes.Purity.PURE)); // log(10)
2547 15 ret := Expression.MULTARY({Expression.REAL(1.0)}, {arg, ret}, mulOp); // 1/(arg*log(10))
2548 then ret;
2549
2550 else algorithm
2551 ✗ Error.addMessage(Error.INTERNAL_ERROR,{getInstanceName() + " failed for: " + name});
2552 ✗ then fail();
2553 end match;
2554 end differentiateBuiltinCall1Arg;
2555
2556 function differentiateBuiltinCall2Arg
2557 "differentiate a builtin call with two arguments."
2558 input String name;
2559 input Expression arg1;
2560 input Expression arg2;
2561 output Expression derFuncCall1;
2562 output Expression derFuncCall2;
2563 protected
2564 // these probably need to be adapted to the size and type of arg
2565 Operator.SizeClassification sizeClass = NFOperator.SizeClassification.SCALAR;
2566 Operator powOp = Operator.fromClassification((NFOperator.MathClassification.POWER, sizeClass), Type.REAL());
2567 Operator addOp = Operator.fromClassification((NFOperator.MathClassification.ADDITION, sizeClass), Type.REAL());
2568 Operator mulOp = Operator.fromClassification((NFOperator.MathClassification.MULTIPLICATION, sizeClass), Type.REAL());
2569 algorithm
2570 (derFuncCall1, derFuncCall2) := match name
2571 local
2572 Expression exp1, exp2, ret1, ret2;
2573
2574 // div(arg1, arg2) truncates the fractional part of arg1/arg2 so it has discrete values
2575 // therefore it has zero derivative where it's defined
2576 case "div" then (Expression.INTEGER(0), Expression.INTEGER(0));
2577
2578 // d/darg1 mod(arg1, arg2) -> 1
2579 // d/darg2 mod(arg1, arg2) -> -floor(arg1/arg2)
2580 case "mod" algorithm
2581 ✗ exp2 := Expression.CALL(Call.makeTypedCall(
2582 fn = NFBuiltinFuncs.FLOOR,
2583 args = {Expression.MULTARY({arg1}, {arg2}, mulOp)}, // arg1/arg2
2584 variability = Prefixes.variabilityMax(Expression.variability(arg1), Expression.variability(arg2)),
2585 purity = NFPrefixes.Purity.PURE
2586 )); // floor(arg1/arg2)
2587 ✗ ret2 := Expression.negate(exp2); // -floor(arg1/arg2)
2588 then (Expression.REAL(1), ret2);
2589
2590 // d/darg1 rem(arg1, arg2) -> 1
2591 // d/darg2 rem(arg1, arg2) -> -div(arg1, arg2)
2592 case "rem" algorithm
2593 ✗ exp2 := Expression.CALL(Call.makeTypedCall(
2594 fn = NFBuiltinFuncs.DIV_REAL,
2595 args = {arg1, arg2},
2596 variability = Prefixes.variabilityMax(Expression.variability(arg1), Expression.variability(arg2)),
2597 purity = NFPrefixes.Purity.PURE
2598 )); // div(arg1, arg2)
2599 ✗ ret2 := Expression.negate(exp2); // -div(arg1, arg2)
2600 then (Expression.REAL(1), ret2);
2601
2602 // d/darg1 atan2(arg1, arg2) -> -arg2/(arg1^2+arg2^2)
2603 // d/darg2 atan2(arg1, arg2) -> arg1/(arg1^2+arg2^2)
2604 case "atan2" algorithm
2605 ✗ exp1 := Expression.BINARY(arg1, powOp, Expression.REAL(2.0)); // arg1^2
2606 ✗ exp2 := Expression.BINARY(arg2, powOp, Expression.REAL(2.0)); // arg2^2
2607 ✗ exp1 := Expression.MULTARY({exp1, exp2}, {}, addOp); // arg1^2+arg2^2
2608 ✗ ret1 := Expression.MULTARY({Expression.negate(arg2)}, {exp1}, mulOp); // -arg2/(arg1^2+arg2^2)
2609 ✗ ret2 := Expression.MULTARY({arg1}, {exp1}, mulOp); // arg1/(arg1^2+arg2^2)
2610 then (ret1, ret2);
2611
2612 else algorithm
2613 ✗ Error.addMessage(Error.INTERNAL_ERROR,{getInstanceName() + " failed for: " + name});
2614 ✗ then fail();
2615 end match;
2616 end differentiateBuiltinCall2Arg;
2617
2618 function addDiffInfo
2619 "adds differentiation info to a pre-defined derivative function
2620 ToDo: do this generally when creating the function tree instead.
2621 Current approach only works if differentiated in the proper order."
2622 input Function func;
2623 input output Function der_func;
2624 input output DifferentiationArguments diffArguments;
2625 protected
2626 UnorderedSet<InstNode> diffInfo;
2627 algorithm
2628 // use the previous differentiation info and extend upon it
2629 diffInfo := match func.interfaceDiffInfo
2630 3 case SOME(diffInfo) then UnorderedSet.copy(diffInfo);
2631 43 else UnorderedSet.new(InstNode.hash, InstNode.nameEqual);
2632 end match;
2633
2634 // add all interface nodes of func (inputs, locals, outputs) since all have been
2635 // differentiated to produce der_func; this prevents re-differentiation of func.outputs
2636 // that became locals in der_func when creating higher-order derivatives
2637
2/2
✓ Branch 0 taken 257 times.
✓ Branch 1 taken 46 times.
303 for node in func.inputs loop
2638 257 UnorderedSet.add(node, diffInfo);
2639 end for;
2640
2/2
✓ Branch 0 taken 258 times.
✓ Branch 1 taken 46 times.
304 for node in func.locals loop
2641 258 UnorderedSet.add(node, diffInfo);
2642 end for;
2643
2/2
✓ Branch 0 taken 53 times.
✓ Branch 1 taken 46 times.
99 for o in func.outputs loop
2644 53 UnorderedSet.add(InstNode.fromHandle(o), diffInfo);
2645 end for;
2646
2647 46 der_func.interfaceDiffInfo := SOME(diffInfo);
2648
2649 // add function to function tree
2650 46 UnorderedMap.add(der_func.path, der_func, diffArguments.funcMap);
2651 end addDiffInfo;
2652
2653 function differentiateFunction
2654 input Function func;
2655 output Function der_func;
2656 input UnorderedMap<String, Boolean> interface_map;
2657 input output DifferentiationArguments diffArguments;
2658 algorithm
2659 der_func := match func
2660 local
2661 InstNode node;
2662 Pointer<Class> cls;
2663 Class new_cls;
2664 DifferentiationArguments funcDiffArgs;
2665 UnorderedMap<ComponentRef, ComponentRef> diff_map = UnorderedMap.new<ComponentRef>(ComponentRef.hash, ComponentRef.isEqual);
2666 UnorderedSet<InstNode> diffInfo;
2667 list<Algorithm> algorithms;
2668 FunctionDerivative funcDer;
2669 Function dummy_func;
2670 CachedData cachedData;
2671 String der_func_name;
2672 list<InstNode> inputs, locals, outputs, local_outputs, uninitialized;
2673 list<Slot> slots;
2674
2675 case der_func as Function.FUNCTION() algorithm
2676 24 node := InstNode.fromHandle(der_func.node);
2677
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 24 times.
24 InstNode.CLASS_NODE(cls = cls) := node;
2678 new_cls := match Pointer.access(cls)
2679 case new_cls as Class.INSTANCED_CLASS() algorithm
2680 // prepare outputs that become locals
2681
4/4
✓ Branch 0 taken 32 times.
✓ Branch 1 taken 24 times.
✓ Branch 2 taken 32 times.
✓ Branch 3 taken 24 times.
56 local_outputs := list(InstNode.setComponentDirection(NFPrefixes.Direction.NONE, InstNode.fromHandle(lout)) for lout in der_func.outputs);
2682
4/4
✓ Branch 0 taken 32 times.
✓ Branch 1 taken 24 times.
✓ Branch 2 taken 32 times.
✓ Branch 3 taken 24 times.
56 local_outputs := list(InstNode.protect(lout) for lout in local_outputs);
2683
2684 // prepare differentiation arguments
2685 24 funcDiffArgs := DifferentiationArguments.default();
2686 24 funcDiffArgs.diffType := DifferentiationType.FUNCTION;
2687 24 funcDiffArgs.funcMap := diffArguments.funcMap;
2688 // prepare interface diff info if the function
2689 diffInfo := match der_func.interfaceDiffInfo
2690 2 case SOME(diffInfo) then UnorderedSet.copy(diffInfo);
2691 22 else UnorderedSet.new(InstNode.hash, InstNode.nameEqual);
2692 end match;
2693
2694 24 createInterfaceDerivatives(der_func.inputs, interface_map, diff_map);
2695 24 createInterfaceDerivatives(der_func.locals, interface_map, diff_map);
2696
4/4
✓ Branch 0 taken 32 times.
✓ Branch 1 taken 24 times.
✓ Branch 2 taken 32 times.
✓ Branch 3 taken 24 times.
56 createInterfaceDerivatives(list(InstNode.fromHandle(o) for o in der_func.outputs), interface_map, diff_map);
2697 24 funcDiffArgs.diff_map := SOME(diff_map);
2698
2699 // differentiate interface arguments
2700 24 (inputs, funcDiffArgs) := differentiateFunctionInterfaceNodes(der_func.inputs, interface_map, diff_map, funcDiffArgs, diffInfo, true);
2701 24 (locals, funcDiffArgs) := differentiateFunctionInterfaceNodes(der_func.locals, interface_map, diff_map, funcDiffArgs, diffInfo, false);
2702
4/4
✓ Branch 0 taken 32 times.
✓ Branch 1 taken 24 times.
✓ Branch 2 taken 32 times.
✓ Branch 3 taken 24 times.
56 (outputs, funcDiffArgs) := differentiateFunctionInterfaceNodes(list(InstNode.fromHandle(o) for o in der_func.outputs), interface_map, diff_map, funcDiffArgs, diffInfo, false);
2703
2704 // update inputs, outputs and locals, add old outputs to locals as they might still be used as temporary variables
2705 24 der_func.inputs := inputs;
2706 48 der_func.locals := List.flatten({der_func.locals, locals, local_outputs});
2707
4/4
✓ Branch 0 taken 32 times.
✓ Branch 1 taken 24 times.
✓ Branch 2 taken 32 times.
✓ Branch 3 taken 24 times.
80 der_func.outputs := list(NFInstNode.NodeHandle.VALUE(o) for o in outputs);
2708 // also add the new locals to the class
2709 24 new_cls.elements := ClassTree.appendComponentsToFlatTree(locals, new_cls.elements);
2710
2711 // differentiate slots
2712 24 (slots, funcDiffArgs) := createSlotDerivatives(der_func.slots, interface_map, diff_map, funcDiffArgs);
2713 24 der_func.slots := listAppend(der_func.slots, slots);
2714
2715 // create "fake" function with correct interface to have the interface
2716 // in the case of recursive differentiation (e.g. function calls itself)
2717 24 dummy_func := func;
2718 24 node := InstNode.replaceClass(new_cls, node);
2719 24 der_func_name := NBVariable.FUNCTION_DERIVATIVE_STR + intString(listLength(func.derivatives));
2720 // A copy of the differentiated function, not an update of it: it
2721 // needs its own identity, or both nodes publish into one cell and
2722 // the derivative reads back the function it was derived from.
2723 24 node := InstNode.rename(der_func_name + "." + InstNode.name(node), node);
2724 24 node := InstNode.setDefinition(
2725 SCodeUtil.setElementName(InstNode.definition(node), InstNode.name(node)), node);
2726 // create "fake" function from new node, update cache to get correct derivative name
2727 24 der_func.path := AbsynUtil.prefixPath(der_func_name, der_func.path);
2728 24 der_func.derivatives := {};
2729 24 der_func.derivedInputs := {};
2730 24 der_func.interfaceDiffInfo := SOME(diffInfo);
2731 24 cachedData := CachedData.FUNCTION({der_func}, true, false);
2732 48 der_func.node := NFInstNode.NodeHandle.VALUE(InstNode.newFuncCache(node, cachedData));
2733
2734 // create fake derivative
2735 24 funcDer := FunctionDerivative.FUNCTION_DER(
2736 derivativeFn = InstNode.identityCell(InstNode.fromHandle(der_func.node)),
2737 derivedFn = InstNode.identityCell(InstNode.fromHandle(dummy_func.node)),
2738 order = Expression.INTEGER(1),
2739 conditions = FunctionDerivative.conditionsFromMap(interface_map),
2740 lowerOrderDerivatives = {} // possibly needs updating
2741 );
2742
2743 // add fake derivative to function tree
2744 48 dummy_func.derivatives := funcDer :: dummy_func.derivatives;
2745 24 UnorderedMap.add(dummy_func.path, dummy_func, funcDiffArgs.funcMap);
2746
2747 // differentiate function statements (if there are any. empty for function pointer arguments)
2748 funcDiffArgs := match new_cls.sections
2749 local
2750 Sections sections;
2751 case sections as Sections.SECTIONS() algorithm
2752 try
2753 24 (algorithms, funcDiffArgs) := List.mapFold(sections.algorithms, differentiateAlgorithm, funcDiffArgs);
2754 else
2755 // remove the fake derivative, it has the undifferentiated body
2756 ✗ UnorderedMap.add(func.path, func, funcDiffArgs.funcMap);
2757 ✗ fail();
2758 end try;
2759
2760 // add them to new node
2761 24 sections.algorithms := algorithms;
2762 24 new_cls.sections := sections;
2763 24 then funcDiffArgs;
2764 ✗ else funcDiffArgs;
2765 end match;
2766
2767 // update the class pointer in place; the fake node created above for
2768 // recursive differentiation shares it and reaches codegen via the cache
2769
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 24 times.
24 InstNode.CLASS_NODE(cls = cls) := node;
2770 24 Pointer.update(cls, new_cls);
2771 24 der_func.derivatives := {};
2772 24 der_func.derivedInputs := {};
2773 24 der_func.interfaceDiffInfo := SOME(diffInfo);
2774 24 cachedData := CachedData.FUNCTION({der_func}, true, false);
2775 48 der_func.node := NFInstNode.NodeHandle.VALUE(InstNode.newFuncCache(node, cachedData));
2776
2777 // check the generated body for use-before-assign and initialize
2778 // variables not provably assigned (the frontend check is skipped here)
2779 24 uninitialized := Function.checkUseBeforeAssignGenerated(der_func);
2780
2/2
✓ Branch 0 taken 6 times.
✓ Branch 1 taken 18 times.
24 if not listEmpty(uninitialized) then
2781 6 new_cls.sections := Function.initializeUninitialized(new_cls.sections, uninitialized, AbsynUtil.pathString(der_func.path));
2782
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 6 times.
6 InstNode.CLASS_NODE(cls = cls) := node;
2783 6 Pointer.update(cls, new_cls);
2784 end if;
2785
2786 // save the function tree
2787 24 diffArguments.funcMap := funcDiffArgs.funcMap;
2788 24 then new_cls;
2789
2790 else algorithm
2791 ✗ Error.addMessage(Error.INTERNAL_ERROR,{getInstanceName() + " failed for class " + Class.toFlatString(Pointer.access(cls), InstNode.fromHandle(func.node)) + "."});
2792 ✗ then fail();
2793 end match;
2794
2795 // add function to function tree
2796 24 UnorderedMap.add(der_func.path, der_func, diffArguments.funcMap);
2797 // add new function as derivative to original function
2798 24 funcDer := FunctionDerivative.FUNCTION_DER(
2799 derivativeFn = InstNode.identityCell(InstNode.fromHandle(der_func.node)),
2800 derivedFn = InstNode.identityCell(InstNode.fromHandle(func.node)),
2801 order = Expression.INTEGER(1),
2802 conditions = FunctionDerivative.conditionsFromMap(interface_map),
2803 lowerOrderDerivatives = {} // possibly needs updating
2804 );
2805 24 func.derivatives := List.appendElt(funcDer, func.derivatives);
2806 24 UnorderedMap.add(func.path, func, diffArguments.funcMap);
2807 then der_func;
2808
2809 else algorithm
2810 ✗ Error.addMessage(Error.INTERNAL_ERROR,{getInstanceName() + " failed for uninstantiated function " + Function.signatureString(func) + "."});
2811 ✗ then fail();
2812 end match;
2813
2/2
✓ Branch 1 taken 23 times.
✓ Branch 2 taken 1 time.
24 if Flags.isSet(Flags.DEBUG_DIFFERENTIATION) then
2814 1 print("\n[BEFORE] " + Function.toFlatString(func) + "\n");
2815 1 print("\n[AFTER ] " + Function.toFlatString(der_func) + "\n\n");
2816 end if;
2817 end differentiateFunction;
2818
2819 function differentiateFunctionInterfaceNodes
2820 "differentiates function interface nodes (inputs, outputs, locals) and
2821 adds them to the diff_map used for differentiation. Also returns the new
2822 interface node lists for the differentiated function.
2823 Note1: outputs only have the differentiated and not the original interface nodes
2824 Note2: for derivatives of higher orders skip the previously differentiated interface nodes"
2825 input output list<InstNode> interface_nodes;
2826 input UnorderedMap<String, Boolean> interface_map;
2827 input UnorderedMap<ComponentRef, ComponentRef> diff_map;
2828 input output DifferentiationArguments diffArgs;
2829 input UnorderedSet<InstNode> diffInfo;
2830 input Boolean keepOld;
2831 protected
2832 list<InstNode> new_nodes;
2833 InstNode d_node;
2834 algorithm
2835
2/2
✓ Branch 0 taken 24 times.
✓ Branch 1 taken 48 times.
72 new_nodes := if keepOld then listReverse(interface_nodes) else {};
2836
2/2
✓ Branch 0 taken 251 times.
✓ Branch 1 taken 72 times.
323 for node in interface_nodes loop
2837 // check if its part of the interface
2838
2/2
✓ Branch 2 taken 234 times.
✓ Branch 3 taken 17 times.
251 if not UnorderedMap.contains(InstNode.name(node), interface_map) then
2839 // check if derivative of higher order skips this
2840
2/2
✓ Branch 1 taken 232 times.
✓ Branch 2 taken 2 times.
234 if not UnorderedSet.contains(node, diffInfo) then
2841 232 (d_node, diffArgs) := differentiateFunctionInterfaceNode(node, diff_map, diffArgs);
2842 new_nodes := d_node :: new_nodes;
2843 // add to skipped nodes if differentiated again because the derivative now already exists
2844 232 UnorderedSet.add(node, diffInfo);
2845 else
2846 end if;
2847 end if;
2848 end for;
2849 72 interface_nodes := listReverse(new_nodes);
2850 end differentiateFunctionInterfaceNodes;
2851
2852 function differentiateFunctionInterfaceNode
2853 input InstNode node;
2854 output InstNode d_node;
2855 input UnorderedMap<ComponentRef, ComponentRef> diff_map;
2856 input output DifferentiationArguments diffArgs;
2857 protected
2858 ComponentRef cref, diff_cref;
2859 Component comp;
2860 Binding binding;
2861 Function func, d_func;
2862 algorithm
2863 325 cref := ComponentRef.fromNode(node, InstNode.getType(node));
2864 325 diff_cref := UnorderedMap.getSafe(cref, diff_map, sourceInfo());
2865 diff_cref := match diff_cref
2866 case ComponentRef.CREF() guard InstNode.isComponent(ComponentRef.node(diff_cref)) algorithm
2867 325 d_node := ComponentRef.node(diff_cref);
2868 // differentiate bindings
2869 325 comp := InstNode.component(d_node);
2870 comp := match comp
2871 case comp as Component.COMPONENT() algorithm
2872 325 (binding, diffArgs) := differentiateBinding(comp.binding, diffArgs);
2873 325 comp.binding := binding;
2874 then comp;
2875 else comp;
2876 end match;
2877 325 d_node := InstNode.replaceComponent(comp, d_node);
2878 325 diff_cref.node := ComponentRef.storeNode(d_node, update = true);
2879 then diff_cref;
2880 else diff_cref;
2881 end match;
2882
2883 // if the node is a function, its a function pointer argument
2884
1/2
✗ Branch 1 not taken.
✓ Branch 2 taken 325 times.
325 if InstNode.isFunction(node) then
2885 ✗ func := listHead(Function.getCachedFuncs(node));
2886 ✗ (d_func, diffArgs) := differentiateFunction(func, UnorderedMap.new<Boolean>(stringHashDjb2, stringEqual), diffArgs);
2887 end if;
2888
2889 325 d_node := ComponentRef.node(diff_cref);
2890 end differentiateFunctionInterfaceNode;
2891
2892 function createInterfaceDerivatives
2893 input list<InstNode> interface_nodes;
2894 input UnorderedMap<String, Boolean> interface_map;
2895 input UnorderedMap<ComponentRef, ComponentRef> diff_map;
2896 protected
2897 ComponentRef cref;
2898
2899 function addCref
2900 input ComponentRef cref;
2901 input UnorderedMap<ComponentRef, ComponentRef> diff_map;
2902 protected
2903 ComponentRef diff_cref;
2904 list<ComponentRef> children;
2905 algorithm
2906 291 diff_cref := BVariable.makeFDerVar(cref);
2907 291 UnorderedMap.add(cref, diff_cref, diff_map);
2908
2909 291 children := ComponentRef.getRecordChildren(cref);
2910
2/2
✓ Branch 0 taken 57 times.
✓ Branch 1 taken 291 times.
348 for child in children loop
2911 57 addCref(child, diff_map);
2912 end for;
2913 end addCref;
2914 algorithm
2915
2/2
✓ Branch 0 taken 251 times.
✓ Branch 1 taken 72 times.
323 for node in interface_nodes loop
2916
2/2
✓ Branch 2 taken 234 times.
✓ Branch 3 taken 17 times.
251 if not UnorderedMap.contains(InstNode.name(node), interface_map) then
2917 234 cref := ComponentRef.fromNode(node, InstNode.getType(node));
2918 234 addCref(cref, diff_map);
2919 end if;
2920 end for;
2921 end createInterfaceDerivatives;
2922
2923 function createSlotDerivatives
2924 input list<Slot> slots;
2925 output list<Slot> new_slots = {};
2926 input UnorderedMap<String, Boolean> interface_map;
2927 input UnorderedMap<ComponentRef, ComponentRef> diff_map;
2928 input output DifferentiationArguments diffArgs;
2929 protected
2930 InstNode d_node;
2931 Integer local_index = listLength(slots) + 1;
2932 algorithm
2933
2/2
✓ Branch 0 taken 110 times.
✓ Branch 1 taken 24 times.
134 for slot in slots loop
2934
2/2
✓ Branch 2 taken 93 times.
✓ Branch 3 taken 17 times.
110 if not UnorderedMap.contains(InstNode.name(slot.node), interface_map) then
2935 93 (d_node, diffArgs) := differentiateFunctionInterfaceNode(slot.node, diff_map, diffArgs);
2936 93 slot.node := d_node;
2937 slot.index := local_index;
2938 new_slots := slot :: new_slots;
2939 93 local_index := local_index + 1;
2940 end if;
2941 end for;
2942 24 new_slots := listReverse(new_slots);
2943 end createSlotDerivatives;
2944
2945 function resolvePartialDerivatives
2946 input output Function func;
2947 input UnorderedMap<Path, Function> funcMap;
2948 protected
2949 Function der_func;
2950 InstNode node;
2951 Pointer<Class> cls, tmp_cls;
2952 Class new_cls, wrap_cls;
2953 Sections sections;
2954 UnorderedMap<ComponentRef, ComponentRef> diff_map = UnorderedMap.new<ComponentRef>(ComponentRef.hash, ComponentRef.isEqual);
2955 UnorderedMap<String, Boolean> interface_map;
2956 DifferentiationArguments diffArgs = DifferentiationArguments.default();
2957 UnorderedSet<InstNode> diffInfo;
2958 list<Algorithm> algorithms;
2959 CachedData cachedData;
2960 ComponentRef diffCref;
2961 list<InstNode> locals, outputs, local_outputs;
2962 Boolean changed = false;
2963
2964 algorithm
2965 func := match func
2966 case der_func as Function.FUNCTION() algorithm
2967
1/2
✗ Branch 1 not taken.
✓ Branch 2 taken 406 times.
406 InstNode.CLASS_NODE(cls = cls) := InstNode.fromHandle(der_func.node);
2968 406 wrap_cls := Pointer.access(cls);
2969 new_cls := match wrap_cls
2970 case wrap_cls as Class.TYPED_DERIVED(baseClass = node as InstNode.CLASS_NODE(cls = tmp_cls)) algorithm
2971 new_cls := match Pointer.access(tmp_cls)
2972 case new_cls as Class.INSTANCED_CLASS(sections = sections as Sections.SECTIONS(algorithms = algorithms)) algorithm
2973 // prepare differentiation arguments
2974 ✗ diffArgs.diffType := DifferentiationType.FUNCTION;
2975 ✗ diffArgs.funcMap := funcMap;
2976 // prepare interface diff info if the function
2977 diffInfo := match der_func.interfaceDiffInfo
2978 ✗ case SOME(diffInfo) then UnorderedSet.copy(diffInfo);
2979 ✗ else UnorderedSet.new(InstNode.hash, InstNode.nameEqual);
2980 end match;
2981
2982 ✗ interface_map := UnorderedMap.fromLists(list(InstNode.name(var) for var in der_func.inputs), List.fill(false, listLength(der_func.inputs)), stringHashDjb2, stringEqual);
2983
2984 // add all differentiated inputs to the interface map
2985 ✗ for var in List.getAtIndexLst(der_func.inputs, der_func.derivedInputs) loop
2986 ✗ UnorderedMap.remove(InstNode.name(var), interface_map);
2987
2988 // prepare outputs that become locals
2989 ✗ local_outputs := list(InstNode.setComponentDirection(NFPrefixes.Direction.NONE, InstNode.fromHandle(node)) for node in der_func.outputs);
2990 ✗ local_outputs := list(InstNode.protect(node) for node in local_outputs);
2991
2992 // differentiate interface arguments
2993 ✗ createInterfaceDerivatives({var}, interface_map, diff_map);
2994 ✗ createInterfaceDerivatives(der_func.locals, interface_map, diff_map);
2995 ✗ createInterfaceDerivatives(list(InstNode.fromHandle(o) for o in der_func.outputs), interface_map, diff_map);
2996 ✗ diffArgs.diff_map := SOME(diff_map);
2997
2998 ✗ (locals, diffArgs) := differentiateFunctionInterfaceNodes(der_func.locals, interface_map, diff_map, diffArgs, diffInfo, true);
2999 ✗ (outputs, diffArgs) := differentiateFunctionInterfaceNodes(list(InstNode.fromHandle(o) for o in der_func.outputs), interface_map, diff_map, diffArgs, diffInfo, false);
3000
3001 ✗ diffCref := UnorderedMap.getSafe(ComponentRef.fromNode(var, InstNode.getType(var)), diff_map, sourceInfo());
3002 ✗ der_func.locals := listAppend(locals, local_outputs);
3003 ✗ der_func.outputs := list(NFInstNode.NodeHandle.VALUE(o) for o in outputs);
3004 ✗ der_func.interfaceDiffInfo := SOME(diffInfo);
3005
3006 // differentiate function statements
3007 ✗ (algorithms, diffArgs) := List.mapFold(algorithms, differentiateAlgorithm, diffArgs);
3008 ✗ algorithms := Algorithm.mapExpList(algorithms, function Replacements.single(old = Expression.fromCref(diffCref), new = Expression.makeOne(ComponentRef.getSubscriptedType(diffCref))));
3009
3010 ✗ UnorderedMap.add(InstNode.name(var), false, interface_map);
3011 end for;
3012
3013 // add them to new node
3014 ✗ sections.algorithms := algorithms;
3015 ✗ new_cls.sections := sections;
3016 ✗ new_cls.ty := wrap_cls.ty;
3017 ✗ new_cls.restriction := wrap_cls.restriction;
3018 ✗ node.cls := Pointer.create(new_cls);
3019 ✗ der_func.derivatives := {};
3020 ✗ der_func.derivedInputs := {};
3021 ✗ der_func.interfaceDiffInfo := SOME(diffInfo);
3022 ✗ cachedData := CachedData.FUNCTION({der_func}, true, false);
3023 ✗ der_func.node := NFInstNode.NodeHandle.VALUE(InstNode.newFuncCache(node, cachedData));
3024
3025
3026 changed := true;
3027 then new_cls;
3028
3029 else wrap_cls;
3030 end match;
3031 then new_cls;
3032 else wrap_cls;
3033 end match;
3034
3035
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 406 times.
406 if changed then
3036 ✗ if Flags.isSet(Flags.DEBUG_DIFFERENTIATION) then
3037 ✗ print("\n[BEFORE] " + Function.toFlatString(func) + "\n");
3038 ✗ print("\n[AFTER ] " + Function.toFlatString(der_func) + "\n\n");
3039 end if;
3040 ✗ UnorderedMap.add(der_func.path, der_func, funcMap);
3041 end if;
3042 then der_func;
3043
3044 else func;
3045 end match;
3046 end resolvePartialDerivatives;
3047
3048 function differentiateAlgorithm
3049 input output Algorithm alg;
3050 input output DifferentiationArguments diffArguments;
3051 protected
3052 list<list<Statement>> statements;
3053 list<Statement> statements_flat;
3054 list<ComponentRef> inputs, outputs;
3055 UnorderedSet<Statement> diffInfo;
3056 algorithm
3057 // store which statements are differentiated so they wont be differentiated again
3058 diffInfo := match alg.stmtDiffInfo
3059 2 case SOME(diffInfo) then UnorderedSet.copy(diffInfo);
3060 22 else UnorderedSet.new(Statement.hash, Statement.isEqual);
3061 end match;
3062
3063 // differentiate the statements
3064 24 (statements, diffArguments) := List.mapFold(alg.statements, function differentiateStatement(diffInfo = diffInfo), diffArguments);
3065
3066 // add all original statements to the set of statements that should not be differentiated
3067
2/2
✓ Branch 0 taken 103 times.
✓ Branch 1 taken 24 times.
127 for stmt in alg.statements loop
3068 103 UnorderedSet.add(stmt, diffInfo);
3069 end for;
3070
3071 24 statements_flat := List.flatten(statements);
3072 24 (inputs, outputs) := Algorithm.getInputsOutputs(statements_flat);
3073 24 alg := Algorithm.ALGORITHM(statements_flat, inputs, outputs, SOME(diffInfo), alg.scope, alg.source);
3074 end differentiateAlgorithm;
3075
3076 function wildIfNotCref
3077 "replaces direct non-cref tuple elements by a wildcard"
3078 input output Expression exp;
3079 algorithm
3080 exp := match exp
3081 case Expression.TUPLE()
3082
6/6
✓ Branch 0 taken 18 times.
✓ Branch 1 taken 9 times.
✓ Branch 2 taken 18 times.
✓ Branch 3 taken 9 times.
✓ Branch 5 taken 1 time.
✓ Branch 6 taken 17 times.
27 then Expression.TUPLE(exp.ty, list(if Expression.isCref(e) then e else Expression.CREF(Expression.typeOf(e), ComponentRef.WILD()) for e in exp.elements));
3083 else exp;
3084 end match;
3085 end wildIfNotCref;
3086
3087 function differentiateStatement
3088 input Statement stmt;
3089 input UnorderedSet<Statement> diffInfo;
3090 output list<Statement> diff_stmts "two statements for 'Real' assignments (diff; original) and else one";
3091 input output DifferentiationArguments diffArguments;
3092 algorithm
3093 diff_stmts := match stmt
3094 local
3095 Statement diff_stmt;
3096 Expression exp, lhs, rhs;
3097 list<Statement> branch_stmts_flat;
3098 list<list<Statement>> branch_stmts;
3099 list<tuple<Expression, list<Statement>>> branches = {};
3100 Boolean isReverse = isSome(diffArguments.adjoint_map);
3101
3102 // 0. do not differentiate if it already exists differentiated due to previous differentiation
3103 case _ guard(UnorderedSet.contains(stmt, diffInfo)) then {stmt};
3104
3105 // I. differentiate 'Real' assignment and return differentiated and original statement
3106 case diff_stmt as Statement.ASSIGNMENT() guard(Type.isReal(Type.arrayElementType(Expression.typeOf(diff_stmt.lhs)))) algorithm
3107 // In reverse mode the assignment LHS is the destination; traverse it without
3108 // collecting into adjoint_map to avoid artificial self-contributions.
3109 158 (lhs, diffArguments) := differentiateExpression(diff_stmt.lhs, diffArguments);
3110 158 (rhs, diffArguments) := differentiateExpression(diff_stmt.rhs, diffArguments);
3111 158 diff_stmt.lhs := lhs;
3112 158 diff_stmt.rhs := SimplifyExp.simplifyDump(rhs, true, getInstanceName());
3113
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 158 times.
158 then if isReverse then {diff_stmt} else {diff_stmt, stmt};
3114
3115 // I-b. differentiate record-type assignment from a function call
3116 // e.g. f := Helmholtz(d, T) where f is a record — propagate seeds through
3117 // the called function so the derivative record gets populated correctly.
3118 // Without this, the derivative variable is left zero-initialised and the
3119 // analytical Jacobian for any NLS that calls the outer function is wrong.
3120 case diff_stmt as Statement.ASSIGNMENT() guard(
3121 Type.isComplex(Expression.typeOf(diff_stmt.lhs)) and
3122 Expression.isCall(diff_stmt.rhs)
3123 ) algorithm
3124 1 (lhs, diffArguments) := differentiateExpression(diff_stmt.lhs, diffArguments);
3125 1 (rhs, diffArguments) := differentiateExpression(diff_stmt.rhs, diffArguments);
3126 1 diff_stmt.lhs := lhs;
3127 1 diff_stmt.rhs := SimplifyExp.simplifyDump(rhs, true, getInstanceName());
3128
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 1 time.
1 then if isReverse then {diff_stmt} else {diff_stmt, stmt};
3129
3130 // I-c. differentiate tuple assignment from a function call
3131 // (a, b) := f(x) -> (a', b') := f'(x, x')
3132 case diff_stmt as Statement.ASSIGNMENT(lhs = Expression.TUPLE()) guard(Expression.isCall(diff_stmt.rhs)) algorithm
3133 9 (lhs, diffArguments) := differentiateExpression(diff_stmt.lhs, diffArguments);
3134 9 (rhs, diffArguments) := differentiateExpression(diff_stmt.rhs, diffArguments);
3135 // outputs without a derivative variable (e.g. Integer) are ignored
3136 9 diff_stmt.lhs := wildIfNotCref(lhs);
3137 9 diff_stmt.rhs := SimplifyExp.simplifyDump(rhs, true, getInstanceName());
3138
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 9 times.
9 then if isReverse then {diff_stmt} else {diff_stmt, stmt};
3139
3140 // II. delegate differentiation to body and only return differentiated statement
3141 case diff_stmt as Statement.FOR() algorithm
3142 22 (branch_stmts, diffArguments) := List.mapFold(diff_stmt.body, function differentiateStatement(diffInfo = diffInfo), diffArguments);
3143 22 diff_stmt.body := List.flatten(branch_stmts);
3144 then {diff_stmt};
3145
3146 case diff_stmt as Statement.WHILE() algorithm
3147 ✗ (branch_stmts, diffArguments) := List.mapFold(diff_stmt.body, function differentiateStatement(diffInfo = diffInfo), diffArguments);
3148 ✗ diff_stmt.body := List.flatten(branch_stmts);
3149 then {diff_stmt};
3150
3151 case diff_stmt as Statement.FAILURE() algorithm
3152 ✗ (branch_stmts, diffArguments) := List.mapFold(diff_stmt.body, function differentiateStatement(diffInfo = diffInfo), diffArguments);
3153 ✗ diff_stmt.body := List.flatten(branch_stmts);
3154 then {diff_stmt};
3155
3156 case diff_stmt as Statement.IF() algorithm
3157
2/2
✓ Branch 0 taken 36 times.
✓ Branch 1 taken 18 times.
54 for branch in diff_stmt.branches loop
3158 36 (exp, branch_stmts_flat) := branch;
3159 36 (branch_stmts, diffArguments) := List.mapFold(branch_stmts_flat, function differentiateStatement(diffInfo = diffInfo), diffArguments);
3160 36 branches := (exp, List.flatten(branch_stmts)) :: branches;
3161 end for;
3162 18 diff_stmt.branches := listReverse(branches);
3163 then {diff_stmt};
3164
3165 case diff_stmt as Statement.WHEN() algorithm
3166 ✗ for branch in diff_stmt.branches loop
3167 ✗ (exp, branch_stmts_flat) := branch;
3168 ✗ (branch_stmts, diffArguments) := List.mapFold(branch_stmts_flat, function differentiateStatement(diffInfo = diffInfo), diffArguments);
3169 ✗ branches := (exp, List.flatten(branch_stmts)) :: branches;
3170 end for;
3171 ✗ diff_stmt.branches := listReverse(branches);
3172 then {diff_stmt};
3173
3174 // III. assignments of non-Real are not differentiated, as well as empty statements
3175 case Statement.ASSIGNMENT() then {stmt};
3176 case Statement.FUNCTION_ARRAY_INIT() then {stmt};
3177 case Statement.ASSERT() then {stmt};
3178 case Statement.TERMINATE() then {stmt};
3179 case Statement.NORETCALL() then {stmt};
3180 case Statement.RETURN() then {stmt};
3181 case Statement.BREAK() then {stmt};
3182
3183 else algorithm
3184 ✗ Error.addMessage(Error.INTERNAL_ERROR,{getInstanceName() + " failed for:" + Statement.toString(stmt)});
3185 ✗ then fail();
3186 end match;
3187 end differentiateStatement;
3188
3189 function reverseForRange
3190 "Reverse a for-loop range for adjoint sweeps:
3191 start[:step]:stop -> stop[:-step]:start
3192 If no step is given, use -1 by default (typical forward loops like 1:N)."
3193 input Option<Expression> rangeIn;
3194 output Option<Expression> rangeOut;
3195 protected
3196 Expression startExp, stopExp, stepExp;
3197 algorithm
3198 rangeOut := match rangeIn
3199 case SOME(Expression.RANGE(start = startExp, step = SOME(stepExp), stop = stopExp))
3200 ✗ then SOME(Expression.makeRange(stopExp, SOME(Expression.negate(stepExp)), startExp));
3201
3202 case SOME(Expression.RANGE(start = startExp, step = NONE(), stop = stopExp))
3203 1 then SOME(Expression.makeRange(stopExp, SOME(Expression.INTEGER(-1)), startExp));
3204
3205 else rangeIn;
3206 end match;
3207 end reverseForRange;
3208
3209 function reverseEquationIterator
3210 "Reverse all iterator ranges of an equation iterator for reverse sweeps."
3211 input NBEquation.Iterator iterIn;
3212 output NBEquation.Iterator iterOut;
3213 protected
3214 list<ComponentRef> names;
3215 list<Expression> ranges;
3216 list<Option<NBEquation.Iterator>> maps;
3217 list<Expression> revRanges = {};
3218 Option<Expression> o_range;
3219 algorithm
3220 1 (names, ranges, maps) := NBEquation.Iterator.getFrames(iterIn);
3221
2/2
✓ Branch 1 taken 1 time.
✓ Branch 2 taken 1 time.
2 for range in ranges loop
3222 1 o_range := reverseForRange(SOME(range));
3223 1 revRanges := Util.getOption(o_range) :: revRanges;
3224 end for;
3225 1 iterOut := NBEquation.Iterator.fromFrames(List.zip3(names, listReverse(revRanges), maps));
3226 end reverseEquationIterator;
3227
3228 function bothZero
3229 "true if both derivatives of a product or quotient are zero, then the derivative is zero as well
3230 (keeps the operands out of it, they would make it look nonlinear)"
3231 input Expression diffExp1;
3232 input Expression diffExp2;
3233 input Operator operator;
3234 output Boolean b = isZeroDerivative(diffExp1) and isZeroDerivative(diffExp2)
3235 and (not Type.isArray(Operator.typeOf(operator)) or Type.hasKnownSize(Operator.typeOf(operator)));
3236 end bothZero;
3237
3238 function isZeroDerivative
3239 "zero after simplification, e.g. a subscripted array of zeros"
3240 input Expression exp;
3241 output Boolean b = Expression.isZero(exp) or Expression.isZero(SimplifyExp.simplify(exp));
3242 end isZeroDerivative;
3243
3244 function differentiateBinary
3245 "Some of this is depcreated because of Expression.MULTARY().
3246 Will always try to convert to MULTARY whenever possible. (commutativity)"
3247 input output Expression exp "Has to be Expression.BINARY()";
3248 input output DifferentiationArguments diffArguments;
3249 algorithm
3250
1/2
✗ Branch 1 not taken.
✓ Branch 2 taken 1252 times.
1252 if Flags.isSet(Flags.DEBUG_ADJOINT) then
3251 ✗ print("differentiateBinary: " + Expression.toString(exp) + "\n");
3252 end if;
3253 (exp, diffArguments) := match exp
3254 local
3255 Expression exp1, exp2, diffExp1, diffExp2, e1, e2, e3, res;
3256 Operator operator, addOp, mulOp, powOp, divOp;
3257 Operator.SizeClassification sizeClass, powSizeClass;
3258 Expression current_grad = diffArguments.current_grad;
3259 // Local reverse grads (to assign before recursing)
3260 Expression grad_exp1, grad_exp2, denom2, numUF;
3261 Boolean isVec1, isVec2, isMat1, isMat2;
3262 Type ty1, ty2;
3263 Integer r1, r2;
3264 list<Integer> dim1, dim2;
3265 Boolean isReverse = isSome(diffArguments.adjoint_map);
3266
3267 // Addition calculations (ADD, ADD_EW, ...)
3268 // (f + g)' = f' + g'
3269 // Adjoint rule: ∂(f + g)/∂f = 1, ∂(f + g)/∂g = 1
3270 // diffArguments.current_grad = ∂Out/∂(f + g) * ∂(f + g)/∂f = current_grad * 1 = current_grad
3271 case Expression.BINARY(exp1 = exp1, operator = operator, exp2 = exp2)
3272 guard(Operator.getMathClassification(operator) == NFOperator.MathClassification.ADDITION)
3273 algorithm
3274 //current_grad := diffArguments.current_grad;
3275
3276 //diffArguments.current_grad := current_grad; // not needed, but for clarity
3277 212 (diffExp1, diffArguments) := differentiateExpression(exp1, diffArguments);
3278
3279 //diffArguments.current_grad := current_grad; // not needed, but for clarity
3280 212 (diffExp2, diffArguments) := differentiateExpression(exp2, diffArguments);
3281
3282 //diffArguments.current_grad := current_grad;
3283 212 then (Expression.MULTARY({diffExp1, diffExp2}, {}, operator), diffArguments);
3284
3285 // Subtraction calculations (SUB, SUB_EW, ...)
3286 // (f - g)' = f' - g'
3287 // ∂(f - g)/∂f = 1, ∂(f - g)/∂g = -1
3288 case Expression.BINARY(exp1 = exp1, operator = operator, exp2 = exp2)
3289 guard(Operator.getMathClassification(operator) == NFOperator.MathClassification.SUBTRACTION)
3290 algorithm
3291 157 current_grad := diffArguments.current_grad;
3292
3293 // differentiate first argument
3294 //diffArguments.current_grad := current_grad; // not needed, but for clarity
3295 157 (diffExp1, diffArguments) := differentiateExpression(exp1, diffArguments);
3296
3297 // differentiate second argument
3298 157 diffArguments.current_grad := Expression.negate(current_grad);
3299 157 (diffExp2, diffArguments) := differentiateExpression(exp2, diffArguments);
3300
3301 157 diffArguments.current_grad := current_grad;
3302 // create addition operator from the size classification of original multiplication operator
3303 157 (_, sizeClass) := Operator.classify(operator);
3304 157 addOp := Operator.fromClassification((NFOperator.MathClassification.ADDITION, sizeClass), operator.ty);
3305 157 then (Expression.MULTARY({diffExp1}, {diffExp2}, addOp), diffArguments);
3306
3307 // Multiplication (MUL, MUL_EW, ...)
3308 // (f * g)' = f'g + fg'
3309 // ∂(f * g)/∂f = g, ∂(f * g)/∂g = f
3310 case Expression.BINARY(exp1 = exp1, operator = operator, exp2 = exp2)
3311 guard(Operator.getMathClassification(operator) == NFOperator.MathClassification.MULTIPLICATION)
3312 algorithm
3313
2/2
✓ Branch 0 taken 2 times.
✓ Branch 1 taken 537 times.
539 if isReverse then
3314 // Upstream gradient
3315 2 current_grad := diffArguments.current_grad;
3316
3317 // Type / rank info
3318 2 ty1 := Expression.typeOf(exp1);
3319 2 ty2 := Expression.typeOf(exp2);
3320
1/2
✓ Branch 1 taken 2 times.
✗ Branch 2 not taken.
2 r1 := if Type.isArray(ty1) then Type.dimensionCount(ty1) else 0;
3321
1/2
✓ Branch 1 taken 2 times.
✗ Branch 2 not taken.
2 r2 := if Type.isArray(ty2) then Type.dimensionCount(ty2) else 0;
3322
1/2
✓ Branch 0 taken 2 times.
✗ Branch 1 not taken.
2 dim1 := if r1 > 0 then NFDimension.sizes(Type.arrayDims(ty1)) else {};
3323
1/2
✓ Branch 0 taken 2 times.
✗ Branch 1 not taken.
2 dim2 := if r2 > 0 then NFDimension.sizes(Type.arrayDims(ty2)) else {};
3324
3325 2 isVec1 := (r1 == 1);
3326 2 isVec2 := (r2 == 1);
3327 2 isMat1 := (r1 == 2);
3328 2 isMat2 := (r2 == 2);
3329
3330 // Original size classification (kept for forward combination)
3331 2 (_, sizeClass) := Operator.classify(operator);
3332 // Decide shape case
3333 // Inner product
3334
1/4
✗ Branch 0 not taken.
✓ Branch 1 taken 2 times.
✗ Branch 2 not taken.
✗ Branch 3 not taken.
2 if isVec1 and isVec2 and sizeClass == NFOperator.SizeClassification.SCALAR then
3335 ✗ grad_exp1 := Expression.BINARY(
3336 current_grad,
3337 Operator.fromClassification(
3338 (NFOperator.MathClassification.MULTIPLICATION, NFOperator.SizeClassification.SCALAR_ARRAY),
3339 operator.ty),
3340 exp2); // G * y
3341 ✗ grad_exp2 := Expression.BINARY(
3342 current_grad,
3343 Operator.fromClassification(
3344 (NFOperator.MathClassification.MULTIPLICATION, NFOperator.SizeClassification.SCALAR_ARRAY),
3345 operator.ty),
3346 exp1); // G * x
3347 // outer product
3348 elseif isMat1 and isMat2 and sizeClass == NFOperator.SizeClassification.MATRIX and listGet(dim1, 1) > 1 and listGet(dim1, 2) == 1 and listGet(dim2, 1) == 1 and listGet(dim2, 2) > 1 then
3349 ✗ grad_exp1 := Expression.BINARY(
3350 current_grad,
3351 Operator.fromClassification(
3352 (NFOperator.MathClassification.MULTIPLICATION, NFOperator.SizeClassification.MATRIX),
3353 operator.ty),
3354 exp2); // G * y
3355 ✗ grad_exp2 := Expression.BINARY(
3356 typeTransposeCall(current_grad),
3357 Operator.fromClassification(
3358 (NFOperator.MathClassification.MULTIPLICATION, NFOperator.SizeClassification.MATRIX),
3359 operator.ty),
3360 exp1); // G^T * x
3361 // Matrix * Vector
3362 elseif isMat1 and isVec2 then
3363 1 grad_exp1 := Expression.BINARY(
3364 current_grad,
3365 Operator.fromClassification(
3366 (NFOperator.MathClassification.MULTIPLICATION, NFOperator.SizeClassification.MATRIX),
3367 operator.ty),
3368 typeTransposeCall(exp2)); // G * xᵀ
3369 1 grad_exp2 := Expression.BINARY(
3370 typeTransposeCall(exp1),
3371 Operator.fromClassification(
3372 (NFOperator.MathClassification.MULTIPLICATION, NFOperator.SizeClassification.MATRIX_VECTOR),
3373 operator.ty),
3374 current_grad); // Aᵀ * G
3375 // Vector * Matrix
3376 elseif isVec1 and isMat2 then
3377 // grad w.r.t exp1 (x): B * Gᵀ -> treat Gᵀ via transpose(current_grad)
3378 ✗ grad_exp1 := Expression.BINARY(
3379 exp2,
3380 Operator.fromClassification(
3381 (NFOperator.MathClassification.MULTIPLICATION, NFOperator.SizeClassification.MATRIX_VECTOR),
3382 operator.ty),
3383 typeTransposeCall(current_grad)); // B * Gᵀ (shape n)
3384 // grad w.r.t exp2 (B): xᵀ * G
3385 ✗ grad_exp2 := Expression.BINARY(
3386 typeTransposeCall(exp1),
3387 Operator.fromClassification(
3388 (NFOperator.MathClassification.MULTIPLICATION, NFOperator.SizeClassification.MATRIX),
3389 operator.ty),
3390 current_grad); // xᵀ * G (outer product)
3391 // Matrix * Matrix
3392 elseif isMat1 and isMat2 then
3393 1 grad_exp1 := Expression.BINARY(
3394 current_grad,
3395 Operator.fromClassification(
3396 (NFOperator.MathClassification.MULTIPLICATION, NFOperator.SizeClassification.MATRIX),
3397 operator.ty),
3398 typeTransposeCall(exp2)); // G * Bᵀ
3399 1 grad_exp2 := Expression.BINARY(
3400 typeTransposeCall(exp1),
3401 Operator.fromClassification(
3402 (NFOperator.MathClassification.MULTIPLICATION, NFOperator.SizeClassification.MATRIX),
3403 operator.ty),
3404 current_grad); // Aᵀ * G
3405 else
3406 ✗ grad_exp1 := Expression.MULTARY({current_grad, exp2}, {}, makeMulFromOperator(operator));
3407 ✗ grad_exp2 := Expression.MULTARY({current_grad, exp1}, {}, makeMulFromOperator(operator));
3408 end if;
3409
3410 // Reverse recurse: exp1
3411 2 diffArguments.current_grad := grad_exp1;
3412 2 (diffExp1, diffArguments) := differentiateExpression(exp1, diffArguments);
3413 // Reverse recurse: exp2
3414 2 diffArguments.current_grad := grad_exp2;
3415 2 (diffExp2, diffArguments) := differentiateExpression(exp2, diffArguments);
3416 // Restore upstream
3417 2 diffArguments.current_grad := current_grad;
3418 else
3419 // only forward differentiation
3420 537 (diffExp1, diffArguments) := differentiateExpression(exp1, diffArguments);
3421 537 (diffExp2, diffArguments) := differentiateExpression(exp2, diffArguments);
3422 end if;
3423 // Forward derivative assembly: f*g' + f'*g
3424 539 sizeClass := Operator.classifyAddition(operator);
3425 539 addOp := Operator.fromClassification(
3426 (NFOperator.MathClassification.ADDITION, sizeClass),
3427 operator.ty);
3428
4/4
✓ Branch 0 taken 537 times.
✓ Branch 1 taken 2 times.
✓ Branch 3 taken 140 times.
✓ Branch 4 taken 397 times.
938 then (if not isReverse and bothZero(diffExp1, diffExp2, operator) then Expression.makeZero(Operator.typeOf(operator))
3429 else Expression.MULTARY(
3430 {Expression.BINARY(diffExp1, operator, exp2), // f'g
3431 Expression.BINARY(exp1, operator, diffExp2)}, // fg'
3432 {},
3433 addOp
3434 ),
3435 diffArguments);
3436
3437 // Division (DIV, DIV_EW, ...)
3438 // (f / g)' = (f'g - fg') / g^2
3439 case Expression.BINARY(exp1 = exp1, operator = operator, exp2 = exp2)
3440 guard(Operator.getMathClassification(operator) == NFOperator.MathClassification.DIVISION)
3441 algorithm
3442 powSizeClass := NFOperator.SizeClassification.SCALAR;
3443 121 powOp := Operator.fromClassification((NFOperator.MathClassification.POWER, powSizeClass), Type.REAL());
3444
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 121 times.
121 if isReverse then
3445 ✗ current_grad := diffArguments.current_grad; // upstream gradient
3446 ✗ diffArguments.current_grad := Expression.MULTARY({current_grad}, {exp2}, Operator.fromClassification(
3447 (NFOperator.MathClassification.MULTIPLICATION, if Type.isArray(Expression.typeOf(current_grad)) then NFOperator.SizeClassification.ARRAY_SCALAR else NFOperator.SizeClassification.SCALAR),
3448 operator.ty)); // z = f/g going into f
3449 end if;
3450 121 (diffExp1, diffArguments) := differentiateExpression(exp1, diffArguments);
3451
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 121 times.
121 if isReverse then
3452 // Reverse local grad for denominator g: G_g = - ( (upstream .* f) / g^2 )
3453 // Build g^2
3454 ✗ denom2 := Expression.BINARY(exp2, powOp, Expression.REAL(2.0));
3455
3456 // Build numerator = upstream .* f with proper size classification
3457 ✗ numUF := Expression.BINARY(current_grad, if Type.isArray(Expression.typeOf(exp1)) then Operator.makeScalarProduct(operator.ty) else Operator.fromClassification(
3458 (NFOperator.MathClassification.MULTIPLICATION, NFOperator.SizeClassification.SCALAR),
3459 Type.REAL()), exp1);
3460
3461 // Divide by g^2 (array/scalar-safe)
3462 ✗ divOp := Operator.fromClassification(
3463 (NFOperator.MathClassification.DIVISION, NFOperator.SizeClassification.SCALAR),
3464 Type.REAL());
3465 ✗ diffArguments.current_grad := Expression.negate(
3466 Expression.BINARY(numUF, divOp, denom2));
3467 end if;
3468 121 (diffExp2, diffArguments) := differentiateExpression(exp2, diffArguments);
3469
3470
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 121 times.
121 if isReverse then
3471 // Restore upstream
3472 ✗ diffArguments.current_grad := current_grad;
3473 end if;
3474 // create subtraction and multiplication operator from the size classification of original division operator
3475 121 (_, sizeClass) := Operator.classify(operator);
3476 // the frontend treats multiplication equally for element and non-elementwise, but pow needs to have the correct operator
3477 // the addition in the numerator f'g +/- fg' must be element-wise when the result is an array (same as multiplication case)
3478 121 addOp := Operator.fromClassification((NFOperator.MathClassification.ADDITION, Operator.classifyAddition(operator)), operator.ty);
3479 121 mulOp := Operator.fromClassification((NFOperator.MathClassification.MULTIPLICATION, sizeClass), operator.ty);
3480
3/4
✓ Branch 0 taken 121 times.
✗ Branch 1 not taken.
✓ Branch 3 taken 105 times.
✓ Branch 4 taken 16 times.
541 then (if not isReverse and bothZero(diffExp1, diffExp2, operator) then Expression.makeZero(Operator.typeOf(operator))
3481 else Expression.MULTARY(
3482 {Expression.MULTARY(
3483 {Expression.BINARY(diffExp1, mulOp, exp2)}, // f'g
3484 {Expression.BINARY(exp1, mulOp, diffExp2)}, // - fg'
3485 addOp
3486 )},
3487 {Expression.BINARY(exp2, powOp, Expression.REAL(2.0))}, // / g^2
3488 mulOp
3489 ),
3490 diffArguments);
3491
3492 // Power (POW, POW_EW, ...) with base zero
3493 // (0^r)' = 0
3494 case Expression.BINARY(exp1 = exp1, operator = operator)
3495 guard((Operator.getMathClassification(operator) == NFOperator.MathClassification.POWER) and
3496 Expression.isZero(exp1))
3497 ✗ then (Expression.makeZero(operator.ty), diffArguments);
3498
3499 // Power (POW, POW_EW, ...) general case
3500 case Expression.BINARY(exp1 = exp1, operator = operator, exp2 = exp2)
3501 guard((Operator.getMathClassification(operator) == NFOperator.MathClassification.POWER))
3502 algorithm
3503 223 (_, sizeClass) := Operator.classify(operator);
3504 223 addOp := Operator.fromClassification((NFOperator.MathClassification.ADDITION, sizeClass), operator.ty);
3505 223 current_grad := diffArguments.current_grad; // upstream gradient
3506
3507 669 diffArguments.current_grad := Expression.MULTARY({current_grad, exp2, Expression.BINARY(exp1, operator, minusOne(exp2, addOp))}, {}, makeMulFromOperator(operator));
3508 223 (diffExp1, diffArguments) := differentiateExpression(exp1, diffArguments);
3509
3510 669 diffArguments.current_grad := Expression.MULTARY({current_grad, exp, expLog(exp1)}, {}, makeMulFromOperator(operator));
3511 223 (diffExp2, diffArguments) := differentiateExpression(exp2, diffArguments);
3512
3513 223 diffArguments.current_grad := current_grad;
3514 223 diffExp1 := SimplifyExp.simplifyDump(diffExp1, true, getInstanceName());
3515 223 diffExp2 := SimplifyExp.simplifyDump(diffExp2, true, getInstanceName());
3516 223 mulOp := Operator.fromClassification((NFOperator.MathClassification.MULTIPLICATION, sizeClass), operator.ty);
3517
3518 res := match (Expression.isZero(diffExp1), Expression.isZero(diffExp2))
3519 // Power (POW, POW_EW, ...) with constant exponent and constant base
3520 // (r1^r2)' = 0
3521 45 case (true, true) then Expression.makeZero(operator.ty);
3522 // Power (POW, POW_EW, ...) with constant exponent
3523 // (x^r)' = r*(x^(r-1))*x'
3524 326 case (false, true) then Expression.MULTARY({exp2, Expression.BINARY(exp1, operator, minusOne(exp2, addOp)), diffExp1}, {}, mulOp);
3525 // Power (POW, POW_EW, ...) with constant base
3526 // (r^x)' = r^x*ln(r)*x'
3527 4 case (true, false) then Expression.MULTARY({exp, expLog(exp1), diffExp2}, {}, mulOp);
3528 // Power (POW, POW_EW, ...) regular case
3529 // (x^y)' = x^(y-1) * (x*ln(x)*y'+(y*x'))
3530 else algorithm
3531 // x^(y-1)
3532 13 e1 := Expression.BINARY(exp1, operator, minusOne(exp2, addOp));
3533 // x * ln(x) * y'
3534 26 e2 := Expression.MULTARY({exp1, expLog(exp1), diffExp2}, {}, mulOp);
3535 // y * x'
3536 13 e3 := Expression.MULTARY({exp2, diffExp1}, {}, mulOp);
3537 26 then Expression.MULTARY({e1, Expression.MULTARY({e2, e3}, {}, addOp)}, {}, mulOp);
3538 end match;
3539 223 then (res, diffArguments);
3540
3541 // Logical and Comparing operators => just return as is
3542 case Expression.BINARY(operator = operator)
3543 guard((Operator.getMathClassification(operator) == NFOperator.MathClassification.LOGICAL) or
3544 (Operator.getMathClassification(operator) == NFOperator.MathClassification.RELATION))
3545 ✗ then (exp, diffArguments);
3546
3547 else algorithm
3548 // maybe add failtrace here and allow failing
3549 ✗ Error.addMessage(Error.INTERNAL_ERROR,{getInstanceName() + " failed for: " + Expression.toString(exp)});
3550 ✗ then fail();
3551
3552 end match;
3553 // simplify?
3554 end differentiateBinary;
3555
3556 function differentiateMultary
3557 "Differentiates a multary expression. Expression.MULTARY()
3558 Note: these can only contain commutative operators"
3559 input output Expression exp "Has to be Expression.MULTARY()";
3560 input output DifferentiationArguments diffArguments;
3561 protected
3562 Boolean isReverse = isSome(diffArguments.adjoint_map);
3563 algorithm
3564
1/2
✗ Branch 1 not taken.
✓ Branch 2 taken 9573 times.
9573 if Flags.isSet(Flags.DEBUG_ADJOINT) then
3565 ✗ print("differentiateMultary: " + Expression.toString(exp) + "\n");
3566 end if;
3567 exp := match exp
3568 local
3569 Expression diff_arg, divisor, diff_enumerator, diff_divisor;
3570 list<Expression> arguments, new_arguments = {};
3571 list<Expression> inv_arguments, new_inv_arguments = {};
3572 list<Expression> diff_arguments, diff_inv_arguments;
3573 Operator operator, addOp, powOp, mulEWOp;
3574 Operator.SizeClassification sizeClass, powSizeClass;
3575 Expression current_grad = diffArguments.current_grad, upstream, e_over_f, e_over_g, numProd, denomProd;
3576 List<Expression> arg_rest;
3577 Boolean hasArray = false, hasArrayNum;
3578 Expression local_grad, localUpF, localUpG;
3579 Integer i;
3580 Type powTy;
3581
3582 // Dash calculations (ADD, SUB, ADD_EW, SUB_EW, ...)
3583 // NOTE: Multary always contains ADDITION
3584 // (sum(f_i))' = sum(f_i')
3585 // e.g. (f + g + h - p - q)' = f' + g' + h' - p' - q'
3586 // Reverse-mode note:
3587 // - If an argument is scalar but at least one other argument is an array,
3588 // its local upstream must be sum-reduced to a scalar before recursion.
3589 case Expression.MULTARY(arguments = arguments, inv_arguments = inv_arguments, operator = operator)
3590 guard(Operator.getMathClassification(operator) == NFOperator.MathClassification.ADDITION)
3591 algorithm
3592
2/2
✓ Branch 0 taken 10 times.
✓ Branch 1 taken 7183 times.
7193 if isReverse then
3593 // Detect if any term is an array (for mixed scalar/array broadcasting)
3594
2/4
✓ Branch 1 taken 10 times.
✗ Branch 2 not taken.
✗ Branch 4 not taken.
✓ Branch 5 taken 10 times.
10 hasArray := List.any(arguments, Expression.hasArrayType) or List.any(inv_arguments, Expression.hasArrayType);
3595 end if;
3596 // go over addition arguments
3597
2/2
✓ Branch 1 taken 9121 times.
✓ Branch 2 taken 7190 times.
16311 for arg in listReverse(arguments) loop
3598
2/2
✓ Branch 0 taken 16 times.
✓ Branch 1 taken 9105 times.
9121 if isReverse then
3599 16 current_grad := diffArguments.current_grad;
3600 // For scalar arg in mixed case: sum-reduce upstream to scalar
3601
2/4
✓ Branch 1 taken 16 times.
✗ Branch 2 not taken.
✗ Branch 3 not taken.
✓ Branch 4 taken 16 times.
16 if Expression.isScalar(arg) and hasArray then
3602 ✗ diffArguments.current_grad := typeSumCall(current_grad);
3603 else
3604 16 diffArguments.current_grad := current_grad;
3605 end if;
3606 end if;
3607
3608 9121 (diff_arg, diffArguments) := differentiateExpression(arg, diffArguments);
3609
3610
2/2
✓ Branch 0 taken 16 times.
✓ Branch 1 taken 9102 times.
9118 if isReverse then
3611 16 diffArguments.current_grad := current_grad;
3612 else
3613 new_arguments := diff_arg :: new_arguments;
3614 end if;
3615 end for;
3616 // go over subtraction arguments
3617
2/2
✓ Branch 1 taken 6354 times.
✓ Branch 2 taken 7190 times.
13544 for arg in listReverse(inv_arguments) loop
3618
2/2
✓ Branch 0 taken 5 times.
✓ Branch 1 taken 6349 times.
6354 if isReverse then
3619 5 current_grad := diffArguments.current_grad;
3620
3621 5 local_grad := Expression.negate(current_grad);
3622
2/4
✓ Branch 1 taken 5 times.
✗ Branch 2 not taken.
✗ Branch 3 not taken.
✓ Branch 4 taken 5 times.
5 if Expression.isScalar(arg) and hasArray then
3623 ✗ local_grad := typeSumCall(local_grad);
3624 end if;
3625 5 diffArguments.current_grad := local_grad;
3626 end if;
3627
3628 6354 (diff_arg, diffArguments) := differentiateExpression(arg, diffArguments);
3629
3630
2/2
✓ Branch 0 taken 5 times.
✓ Branch 1 taken 6349 times.
6354 if isReverse then
3631 5 diffArguments.current_grad := current_grad;
3632 else
3633 new_inv_arguments := diff_arg :: new_inv_arguments;
3634 end if;
3635 end for;
3636 7190 then Expression.MULTARY(new_arguments, new_inv_arguments, operator);
3637
3638 // Dot calculations (MUL, DIV, MUL_EW, DIV_EW, ...)
3639 // NOTE: Multary always contains MULTIPLICATION
3640 // no inverse arguments so single product rule:
3641 // prod(f_i)) = sum((f_i)' * prod(f_k | k <> i))
3642 // e.g. (fgh)' = f'gh + fg'h + fgh'
3643 case Expression.MULTARY(arguments = arguments, inv_arguments = {}, operator = operator)
3644 guard(Operator.getMathClassification(operator) == NFOperator.MathClassification.MULTIPLICATION)
3645 algorithm
3646 // create addition operator
3647 2172 sizeClass := Operator.classifyAddition(operator);
3648 2172 addOp := Operator.fromClassification((NFOperator.MathClassification.ADDITION, sizeClass), operator.ty);
3649 // the adjoint is handled inside here
3650 2172 (new_arguments, diffArguments) := differentiateMultaryMultiplicationArgs(arguments, diffArguments, operator);
3651 2172 then Expression.MULTARY(new_arguments, {}, addOp);
3652
3653 // Dot calculations (MUL, DIV, MUL_EW, DIV_EW, ...)
3654 // NOTE: Multary always contains MULTIPLICATION
3655 // (prod(f_i)) / prod(g_j))'
3656 // makes use of single product rule:
3657 // prod(f_i)) = sum((f_i)' * prod(f_k | k <> i))
3658 // e.g. (abc)' = a'bc + ab'c + abc'
3659 // and binary division rule
3660 // (f / g)' = (f'g - g'f) / g^2
3661 // this is implemented like so:
3662 // E = (prod arguments) / (prod inv_arguments)
3663 // dE = Σ_i f_i' * (E / f_i) - Σ_j g_j' * (E / g_j)
3664 // Reverse mode local grads (used via current_grad):
3665 // for f_i: G_i = G * (E / f_i)
3666 // for g_j: G_j = -G * (E / g_j)
3667 // Reverse assumptions:
3668 // - Broadcasting only happens in the numerator.
3669 // - All denominators are scalar.
3670 // - If the numerator is an array then the division by the denominator is elementwise.
3671 // Sum reduction is needed:
3672 // - For scalar f_i in numerator if any other numerator factor is an array.
3673 // - For denominator g_j (scalar) if numerator is an array.
3674 case Expression.MULTARY(arguments = arguments, inv_arguments = inv_arguments, operator = operator)
3675 guard(Operator.getMathClassification(operator) == NFOperator.MathClassification.MULTIPLICATION
3676 and (not listEmpty(inv_arguments)) and isReverse)
3677 algorithm
3678 1 (_, sizeClass) := Operator.classify(operator);
3679 // Determine operators
3680 1 Operator.fromClassification(
3681 (NFOperator.MathClassification.ADDITION, sizeClass),
3682 operator.ty);
3683 1 makeMulFromOperator(operator);
3684
3685 // Use element-wise mul for reverse local upstream assembly to avoid array*scalar miscodegen
3686 // when upstream and partial products are arrays.
3687 // We keep forward terms using mulOp as before.
3688 1 mulEWOp := Operator.fromClassification(
3689 (NFOperator.MathClassification.MULTIPLICATION, NFOperator.SizeClassification.ELEMENT_WISE),
3690 operator.ty);
3691 1 Operator.fromClassification(
3692 (NFOperator.MathClassification.ADDITION, NFOperator.SizeClassification.ELEMENT_WISE),
3693 operator.ty);
3694
3695 // Does the numerator contain any arrays?
3696 1 hasArrayNum := List.any(arguments, Expression.hasArrayType);
3697
3698 1 numProd := Expression.MULTARY(arguments, {}, operator);
3699 1 denomProd := Expression.MULTARY(inv_arguments, {}, operator);
3700
3701 // Forward derivative term accumulator
3702 1 upstream := diffArguments.current_grad;
3703 // Differentiate numerator factors
3704 i := 1;
3705
2/2
✓ Branch 1 taken 1 time.
✓ Branch 2 taken 1 time.
2 for f in arguments loop
3706 // Remove first occurrence of f from numerator list using List.deleteMemberOnTrue
3707 // this may be an issue if f occurs multiple times
3708 1 arg_rest := listDelete(arguments, i);
3709 1 e_over_f := Expression.MULTARY(arg_rest, {denomProd}, operator);
3710
3711 // Reverse local upstream for f: G_f = upstream .* (exp / f)
3712 1 localUpF := Expression.MULTARY({upstream, e_over_f}, {}, mulEWOp);
3713
3714 // If f is scalar but numerator has arrays -> sum-reduce to scalar
3715
2/4
✓ Branch 1 taken 1 time.
✗ Branch 2 not taken.
✗ Branch 3 not taken.
✓ Branch 4 taken 1 time.
1 if Expression.isScalar(f) and hasArrayNum then
3716 ✗ localUpF := typeSumCall(localUpF);
3717 end if;
3718
3719 // Recurse into f with G_f
3720 1 diffArguments.current_grad := localUpF;
3721 1 (diff_arg, diffArguments) := differentiateExpression(f, diffArguments);
3722
3723 // Forward term: f' * (exp / f)
3724 1 i := i + 1;
3725 end for;
3726
3727 // Differentiate denominator factors
3728 i := 1;
3729
1/2
✓ Branch 2 taken 1 time.
✗ Branch 3 not taken.
1 powSizeClass := if Expression.hasArrayType(listHead(inv_arguments)) then NFOperator.SizeClassification.ARRAY_SCALAR else NFOperator.SizeClassification.SCALAR;
3730 1 Operator.fromClassification((NFOperator.MathClassification.POWER, powSizeClass), Type.REAL());
3731
2/2
✓ Branch 1 taken 1 time.
✓ Branch 2 taken 1 time.
2 for g in inv_arguments loop
3732 1 listDelete(inv_arguments, i);
3733 // exp / g : add one more g to denominator list
3734 1 e_over_g := Expression.MULTARY({numProd}, g :: inv_arguments, operator);
3735
3736 // Reverse local upstream for g: G_g = - upstream .* (exp / g)
3737 1 localUpG := Expression.negate(Expression.MULTARY({upstream, e_over_g}, {}, mulEWOp));
3738
3739 // If numerator has arrays -> sum-reduce scalar denominator upstream
3740
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 1 time.
1 if hasArrayNum then
3741 ✗ localUpG := typeSumCall(localUpG);
3742 end if;
3743
3744 1 diffArguments.current_grad := localUpG;
3745 1 (diff_arg, diffArguments) := differentiateExpression(g, diffArguments);
3746
3747 // Forward term: - g' * (exp / g)
3748 1 Expression.negate(Expression.MULTARY({diff_arg, e_over_g}, {}, mulEWOp));
3749 1 i := i + 1;
3750 end for;
3751 // Restore upstream gradient
3752 1 diffArguments.current_grad := upstream;
3753 then (Expression.END());
3754
3755 case Expression.MULTARY(arguments = arguments, inv_arguments = inv_arguments, operator = operator)
3756 guard(Operator.getMathClassification(operator) == NFOperator.MathClassification.MULTIPLICATION
3757 and (not listEmpty(inv_arguments)))
3758 algorithm
3759 // the frontend treats multiplication equally for elementwise and non-elementwise, but pow needs to have the correct operator
3760
1/2
✗ Branch 3 not taken.
✓ Branch 4 taken 207 times.
207 if not listEmpty(inv_arguments) and Type.isArray(Expression.typeOf(listHead(inv_arguments))) then
3761 powSizeClass := NFOperator.SizeClassification.ARRAY_SCALAR;
3762 ✗ powTy := operator.ty;
3763 else
3764 powSizeClass := NFOperator.SizeClassification.SCALAR;
3765 powTy := Type.REAL();
3766 end if;
3767
3768 // check if the addition size class has to be element wise
3769
3/4
✓ Branch 0 taken 199 times.
✓ Branch 1 taken 8 times.
✓ Branch 5 taken 199 times.
✗ Branch 6 not taken.
207 if not listEmpty(arguments) and Type.isArray(Expression.typeOf(listHead(arguments))) then
3770 sizeClass := NFOperator.SizeClassification.ELEMENT_WISE;
3771 else
3772 207 (_, sizeClass) := Operator.classify(operator);
3773 end if;
3774
3775 207 addOp := Operator.fromClassification((NFOperator.MathClassification.ADDITION, sizeClass), operator.ty);
3776 207 powOp := Operator.fromClassification((NFOperator.MathClassification.POWER, powSizeClass), powTy);
3777 // f'
3778 207 (diff_arguments, diffArguments) := differentiateMultaryMultiplicationArgs(arguments, diffArguments, operator);
3779 207 diff_enumerator := Expression.MULTARY(diff_arguments, {}, addOp);
3780 // g'
3781 207 (diff_inv_arguments, diffArguments) := differentiateMultaryMultiplicationArgs(inv_arguments, diffArguments, operator);
3782 207 diff_divisor := Expression.MULTARY(diff_inv_arguments, {}, addOp);
3783 // g
3784 207 divisor := Expression.MULTARY(inv_arguments, {}, operator);
3785 1035 then Expression.MULTARY(
3786 {Expression.MULTARY(
3787 {Expression.MULTARY(diff_enumerator :: inv_arguments, {}, operator)}, // f'g
3788 {Expression.MULTARY(diff_divisor :: arguments, {}, operator)}, // -g'f
3789 addOp
3790 )},
3791 {Expression.BINARY(divisor, powOp, Expression.REAL(2.0))},
3792 operator
3793 );
3794
3795 else algorithm
3796 // maybe add failtrace here and allow failing
3797 ✗ Error.addMessage(Error.INTERNAL_ERROR,{getInstanceName() + " failed for: " + Expression.toString(exp)});
3798 ✗ then fail();
3799 end match;
3800 end differentiateMultary;
3801
3802 function differentiateMultaryMultiplicationArgs
3803 "prod_i(f_i)' = sum_i((f_i)' * prod(f_k | k <> i))
3804 e.g. (fgh)' = f'gh + fg'h + fgh'"
3805 input list<Expression> arguments;
3806 output list<Expression> new_arguments = {};
3807 input output DifferentiationArguments diffArguments;
3808 input Operator operator;
3809 protected
3810 Expression diff_arg, current_grad = diffArguments.current_grad, localUp, restProd;
3811 Array<List<Expression>> diff_lists = listArray({});
3812 List<Expression> arg_products = {}, restArgs;
3813 Integer idx = 1;
3814 Boolean isReverse = isSome(diffArguments.adjoint_map);
3815 Operator mulEWOp = Operator.fromClassification(
3816 (NFOperator.MathClassification.MULTIPLICATION, NFOperator.SizeClassification.ELEMENT_WISE),
3817 operator.ty);
3818 algorithm
3819
2/2
✓ Branch 0 taken 12 times.
✓ Branch 1 taken 2574 times.
2586 if isReverse then
3820 12 arg_products := Expression.productOfListExceptSelf(arguments, makeMulFromOperator(operator));
3821 else
3822 2574 diff_lists := arrayCreate(listLength(arguments), {});
3823 end if;
3824
2/2
✓ Branch 0 taken 4949 times.
✓ Branch 1 taken 2586 times.
7535 for arg in arguments loop
3825
2/2
✓ Branch 0 taken 24 times.
✓ Branch 1 taken 4925 times.
4949 if isReverse then
3826 24 current_grad := diffArguments.current_grad;
3827
3828 // product of remaining factors (k <> i)
3829 24 restProd := listGet(arg_products, idx);
3830
3831 // Build local upstream = current_grad .* restProd, but flatten if restProd is also a MULTARY product.
3832 restArgs := match restProd
3833 local
3834 Operator mOp;
3835 list<Expression> rA;
3836 case Expression.MULTARY(operator = mOp, arguments = rA)
3837 guard Operator.getMathClassification(mOp) == NFOperator.MathClassification.MULTIPLICATION
3838 then rA;
3839 else {restProd};
3840 end match;
3841
3842 24 localUp := Expression.MULTARY(
3843 listAppend({current_grad}, restArgs),
3844 {},
3845 mulEWOp); // may need to adapt this aswell to scalar when scalar
3846
3847 // If current argument is scalar but the rest-product is array-shaped,
3848 // sum-reduce the local upstream to a scalar before recursing.
3849
2/4
✓ Branch 1 taken 24 times.
✗ Branch 2 not taken.
✗ Branch 4 not taken.
✓ Branch 5 taken 24 times.
24 if Expression.isScalar(arg) and Expression.hasArrayType(restProd) then
3850 ✗ localUp := typeSumCall(localUp);
3851 end if;
3852 24 diffArguments.current_grad := localUp;
3853 end if;
3854
3855 4949 (diff_arg, diffArguments) := differentiateExpression(arg, diffArguments);
3856
3857
2/2
✓ Branch 0 taken 24 times.
✓ Branch 1 taken 4925 times.
4949 if isReverse then
3858 24 diffArguments.current_grad := current_grad;
3859 else
3860
1/2
✓ Branch 0 taken 4925 times.
✗ Branch 1 not taken.
14864 for i in 1:arrayLength(diff_lists) loop
3861
2/2
✓ Branch 0 taken 4925 times.
✓ Branch 1 taken 5014 times.
19878 diff_lists[i] := if i == idx then diff_arg :: diff_lists[i] else arg :: diff_lists[i];
3862 end for;
3863 end if;
3864 4949 idx := idx + 1;
3865 end for;
3866
2/2
✓ Branch 0 taken 12 times.
✓ Branch 1 taken 2574 times.
2586 if not isReverse then
3867
2/2
✓ Branch 0 taken 8 times.
✓ Branch 1 taken 2566 times.
7499 for i in arrayLength(diff_lists):-1:1 loop
3868 4925 new_arguments := Expression.MULTARY(listReverse(diff_lists[i]), {}, operator) :: new_arguments;
3869 end for;
3870 end if;
3871 end differentiateMultaryMultiplicationArgs;
3872
3873 function differentiateEquationAttributes
3874 "Differentiates the residual variable for diffType JACOBIAN, if it exists.
3875 The cref has to be saved in the diff_map for this to work.
3876 ToDo: needs to be adapted for torn/inner equations"
3877 input output EquationAttributes attr;
3878 input DifferentiationArguments diffArguments;
3879 algorithm
3880 attr := match (attr, diffArguments)
3881 local
3882 Pointer<Variable> residualVar, diffedResidualVar;
3883 UnorderedMap<ComponentRef,ComponentRef> diff_map;
3884
3885 case (EquationAttributes.EQUATION_ATTRIBUTES(residualVar = SOME(residualVar)),
3886 DIFFERENTIATION_ARGUMENTS(diff_map = SOME(diff_map), diffType = DifferentiationType.JACOBIAN))
3887 guard(UnorderedMap.contains(BVariable.getVarName(residualVar), diff_map))
3888 algorithm
3889 724 diffedResidualVar := BVariable.getVarPointer(UnorderedMap.getOrFail(BVariable.getVarName(residualVar), diff_map), sourceInfo());
3890 724 attr.residualVar := SOME(diffedResidualVar);
3891 then attr;
3892
3893 else attr;
3894
3895 end match;
3896 end differentiateEquationAttributes;
3897
3898 function differentiateBinding
3899 input output Binding binding;
3900 input output DifferentiationArguments diffArgs;
3901 protected
3902 Option<Expression> opt_exp;
3903 Expression exp;
3904 algorithm
3905 325 opt_exp := Binding.getExpOpt(binding);
3906
3/4
✗ Branch 0 not taken.
✓ Branch 1 taken 325 times.
✓ Branch 2 taken 56 times.
✓ Branch 3 taken 269 times.
325 if isSome(opt_exp) then
3907 56 (exp, diffArgs) := differentiateExpression(Util.getOption(opt_exp), diffArgs);
3908 56 binding := Binding.setExp(exp, binding);
3909 end if;
3910 end differentiateBinding;
3911
3912 protected
3913 function sizeClassificationFromType
3914 input Type ty;
3915 output Operator.SizeClassification sc;
3916 algorithm
3917 sc := match Type.dimensionCount(ty)
3918 case 0 then NFOperator.SizeClassification.SCALAR;
3919 case 1 then NFOperator.SizeClassification.ELEMENT_WISE;
3920 case 2 then NFOperator.SizeClassification.MATRIX;
3921 else NFOperator.SizeClassification.ELEMENT_WISE;
3922 end match;
3923 end sizeClassificationFromType;
3924
3925 function minusOne
3926 input output Expression exp;
3927 input Operator op;
3928 algorithm
3929 exp := match exp
3930 local
3931 Real r;
3932 Integer i;
3933 304 case Expression.REAL(value = r) then Expression.REAL(r - 1.0);
3934 ✗ case Expression.INTEGER(value = i) then Expression.INTEGER(i - 1);
3935 190 else Expression.MULTARY({exp}, {Expression.makeOne(op.ty)}, op);
3936 end match;
3937 end minusOne;
3938
3939 function expLog
3940 input output Expression exp;
3941 algorithm
3942 exp := match exp
3943 local
3944 Real r;
3945 Integer i;
3946
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 4 times.
4 case Expression.REAL(value = r) then Expression.REAL(log(r));
3947 ✗ case Expression.INTEGER(value = i) then Expression.REAL(log(i));
3948 468 else Expression.CALL(Call.makeTypedCall(
3949 fn = NFBuiltinFuncs.LOG_REAL,
3950 args = {exp},
3951 variability = Expression.variability(exp),
3952 purity = NFPrefixes.Purity.PURE
3953 ));
3954 end match;
3955 end expLog;
3956
3957 function makeMulFromOperator
3958 input Operator operator;
3959 output Operator mulOp;
3960 algorithm
3961 459 mulOp := Operator.fromClassification((NFOperator.MathClassification.MULTIPLICATION, Operator.getSizeClassification(operator)), operator.ty);
3962 end makeMulFromOperator;
3963
3964 function typeTransposeCall
3965 "Create a typed builtin transpose(mat) call without expanding mat.
3966 Returns mat if it is not an array with at least 2 dimensions."
3967 input Expression mat;
3968 output Expression tr;
3969 protected
3970 Type inTy = Expression.typeOf(mat);
3971 list<Type.Dimension> dims;
3972 Type elTy;
3973 Type resTy;
3974 NFCall call;
3975 NFPrefixes.Variability var = Expression.variability(mat);
3976 NFPrefixes.Purity pur = Expression.purity(mat);
3977 algorithm
3978 // Only handle array types
3979
1/2
✗ Branch 1 not taken.
✓ Branch 2 taken 4 times.
4 if not Type.isArray(inTy) then
3980 tr := mat;
3981 ✗ return;
3982 end if;
3983
3984 4 elTy := Type.arrayElementType(inTy);
3985 4 dims := Type.arrayDims(inTy);
3986
3987 // Need at least 2 dimensions to transpose
3988
2/2
✓ Branch 1 taken 1 time.
✓ Branch 2 taken 3 times.
4 if listLength(dims) < 2 then
3989 tr := mat;
3990 1 return;
3991 end if;
3992
3993 // Swap first two dimensions; keep the rest
3994 9 resTy := Type.ARRAY(
3995 elTy,
3996 listAppend({listGet(dims,2), listGet(dims,1)}, listRest(listRest(dims)))
3997 );
3998
3999 6 call := NFCall.makeTypedCall(NFBuiltinFuncs.TRANSPOSE, {mat}, var, pur, resTy);
4000 3 tr := Expression.CALL(call);
4001 end typeTransposeCall;
4002
4003 // Helper: build a typed builtin promote(A, n) call that appends (n - ndims(A)) singleton dims.
4004 function typePromoteCall
4005 input Expression arr; // A (scalar or array)
4006 input Integer n; // desired rank
4007 output Expression promoted;
4008 protected
4009 Type inTy = Expression.typeOf(arr);
4010 Type elTy;
4011 list<Type.Dimension> inDims;
4012 Integer m, k;
4013 list<Type.Dimension> ones = {};
4014 list<Type.Dimension> resDims;
4015 Type resTy;
4016 NFCall call;
4017 NFPrefixes.Variability var = Expression.variability(arr);
4018 NFPrefixes.Purity pur = Expression.purity(arr);
4019 algorithm
4020 ✗ elTy := if Type.isArray(inTy) then Type.arrayElementType(inTy) else inTy;
4021 ✗ inDims := if Type.isArray(inTy) then Type.arrayDims(inTy) else {};
4022 ✗ m := listLength(inDims);
4023
4024 // Append singleton dims to the right until rank n
4025 ✗ for k in 1:max(0, n - m) loop
4026 ✗ ones := Dimension.fromInteger(1) :: ones;
4027 end for;
4028 ✗ resDims := List.append_reverse(ones, inDims);
4029 ✗ resTy := if n > 0 then Type.ARRAY(elTy, resDims) else elTy;
4030
4031 ✗ call := NFCall.makeTypedCall(NFBuiltinFuncs.PROMOTE, {arr, Expression.INTEGER(n)}, var, pur, resTy);
4032 ✗ promoted := Expression.CALL(call);
4033 end typePromoteCall;
4034
4035
4036 function typeSumCall
4037 "
4038 Create a typed builtin sum(A) call without expanding A.
4039 Semantics:
4040 - If A is not an array => return A (defensive fallback).
4041 - If A is an array => return sum over all elements, resulting in a scalar of element type.
4042 "
4043 input Expression arr;
4044 output Expression s;
4045 protected
4046 Type inTy = Expression.typeOf(arr);
4047 list<Type.Dimension> dims;
4048 Type elTy;
4049 Type resTy;
4050 NFCall call;
4051 NFPrefixes.Variability var = Expression.variability(arr);
4052 NFPrefixes.Purity pur = Expression.purity(arr);
4053 algorithm
4054 // Not an array: just return expression (sum(x) == x)
4055 ✗ if not Type.isArray(inTy) then
4056 s := arr;
4057 ✗ return;
4058 end if;
4059
4060 ✗ elTy := Type.arrayElementType(inTy);
4061 ✗ dims := Type.arrayDims(inTy);
4062 resTy := elTy; // always reduce to scalar of element type
4063
4064 ✗ call := NFCall.makeTypedCall(NFBuiltinFuncs.SUM, {arr}, var, pur, resTy);
4065 ✗ s := Expression.CALL(call);
4066 end typeSumCall;
4067
4068 // Helper: build matrix * vector (or matrix * matrix) MULTARY with a proper mul operator
4069 function makeMul
4070 input Expression a;
4071 input Expression b;
4072 input Operator.SizeClassification sc;
4073 input Type ty;
4074 output Expression res;
4075 algorithm
4076 ✗ res := Expression.BINARY(
4077 a,
4078 Operator.fromClassification((NFOperator.MathClassification.MULTIPLICATION, sc), ty),
4079 b);
4080 end makeMul;
4081
4082 // Drop the last array dimension by indexing it with 1:
4083 // arr[..., 1]. If arr is not an array, return it unchanged.
4084 function dropLastDimIndex1
4085 input Expression arr;
4086 output Expression res;
4087 protected
4088 Type ty = Expression.typeOf(arr);
4089 list<Type.Dimension> dims;
4090 Integer m, i;
4091 list<Subscript> subs = {};
4092 algorithm
4093 ✗ if not Type.isArray(ty) then
4094 ✗ res := arr; return;
4095 end if;
4096
4097 ✗ dims := Type.arrayDims(ty);
4098 ✗ m := listLength(dims);
4099 ✗ if m <= 0 then
4100 ✗ res := arr; return;
4101 end if;
4102
4103 // Build subscripts: WHOLE for first m-1 dims, INDEX(1) for last
4104 ✗ for i in 1:(m-1) loop
4105 subs := Subscript.WHOLE() :: subs;
4106 end for;
4107 subs := Subscript.INDEX(Expression.INTEGER(1)) :: subs;
4108 ✗ subs := listReverse(subs);
4109
4110 ✗ res := Expression.applySubscripts(subs, arr, true);
4111 end dropLastDimIndex1;
4112
4113 // Build vector[n] with elements A[i,i], i=1..n (literal array).
4114 function extractDiagonalVector
4115 input Expression A; // matrix
4116 input Integer n;
4117 input Type vecTy; // vector[n] type
4118 output Expression v;
4119 protected
4120 list<Expression> elems = {};
4121 Integer i;
4122 algorithm
4123 ✗ for i in 1:n loop
4124 ✗ elems := Expression.applySubscripts(
4125 { Subscript.INDEX(Expression.INTEGER(i)), Subscript.INDEX(Expression.INTEGER(i)) },
4126 A, true) :: elems;
4127 end for;
4128 ✗ v := Expression.ARRAY(vecTy, listArray(listReverse(elems)), false);
4129 end extractDiagonalVector;
4130
4131 function dbg
4132 input String s;
4133 algorithm
4134
1/2
✓ Branch 1 taken 31985 times.
✗ Branch 2 not taken.
31985 if Flags.isSet(Flags.DEBUG_ADJOINT) then
4135 ✗ print(s + "\n");
4136 end if;
4137 end dbg;
4138
4139 function expressionHasIteratorCref
4140 input Expression exp;
4141 output Boolean hasIter;
4142
4143 function foldIter
4144 input Expression e;
4145 input output Boolean b;
4146 algorithm
4147 b := match e
4148 ✗ case Expression.CREF() then b or ComponentRef.isIterator(e.cref);
4149 else b;
4150 end match;
4151 end foldIter;
4152 algorithm
4153 ✗ hasIter := Expression.fold(exp, foldIter, false);
4154 end expressionHasIteratorCref;
4155
4156 function subscriptHasIterator
4157 input Subscript sub;
4158 output Boolean hasIter;
4159 algorithm
4160 hasIter := match sub
4161 ✗ case Subscript.INDEX() then expressionHasIteratorCref(sub.index);
4162 ✗ case Subscript.SLICE() then expressionHasIteratorCref(sub.slice);
4163 else false;
4164 end match;
4165 end subscriptHasIterator;
4166
4167 function subscriptsHaveIterator
4168 input list<Subscript> subs;
4169 output Boolean hasIter = false;
4170 algorithm
4171 ✗ for sub in subs loop
4172 ✗ if subscriptHasIterator(sub) then
4173 hasIter := true;
4174 break;
4175 end if;
4176 end for;
4177 end subscriptsHaveIterator;
4178
4179 function updateAdjointList
4180 input Option<list<Expression>> oldOpt;
4181 input Expression current_grad;
4182 output list<Expression> newList;
4183 protected
4184 list<Expression> oldList;
4185 algorithm
4186 newList := match oldOpt
4187 // probably the only case since empty list is used to initialize
4188 case SOME(oldList) then (current_grad :: oldList);
4189 else {current_grad};
4190 end match;
4191 end updateAdjointList;
4192
4193 annotation(__OpenModelica_Interface="nbackend");
4194 end NBDifferentiate;
4195