Skip to content

Commit b769b40

Browse files
PetroZarytskyivgvassilev
authored andcommitted
Add support for ArrayInitLoopExpr/ArrayInitIndexExpr in reverse mode
Whenever Clang generates an implicit copy/move constructor of a class with a static array member, it uses ArrayInitLoopExpr to express the array copy. For example ``` struct arrWrapper { double arr[2]; }; ``` The implicit copy-constructor will look like ``` arrWrapper(const arrWrapper& other) : arr(ArrayInitLoopExpr(other[ArrayInitIndexExpr])) {} ``` The problem is that we cannot do the same with explicit functions. Instead, we have to replicate this behaviour using loops. This PR adds support for ``ArrayInitLoopExpr`` in constructors, including the ones that initialize high-dimensional arrays and arrays of copiable objects. Also, the last error in #791 is caused by this. Fixes #791.
1 parent 4f46756 commit b769b40

6 files changed

Lines changed: 210 additions & 2 deletions

File tree

include/clad/Differentiator/ReverseModeVisitor.h

Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -14,6 +14,7 @@
1414
#include "clad/Differentiator/ReverseModeVisitorDirectionKinds.h"
1515
#include "clad/Differentiator/VisitorBase.h"
1616

17+
#include "clang/AST/Decl.h"
1718
#include "clang/AST/DeclCXX.h"
1819
#include "clang/AST/Expr.h"
1920
#include "clang/AST/ExprCXX.h"
@@ -29,6 +30,7 @@
2930
#include <array>
3031
#include <limits>
3132
#include <memory>
33+
#include <queue>
3234
#include <stack>
3335
#include <unordered_map>
3436

@@ -397,6 +399,9 @@ namespace clad {
397399
StmtDiff VisitIntegerLiteral(const clang::IntegerLiteral* IL);
398400
StmtDiff VisitMemberExpr(const clang::MemberExpr* ME);
399401
StmtDiff VisitParenExpr(const clang::ParenExpr* PE);
402+
StmtDiff VisitArrayInitLoopExpr(const clang::ArrayInitLoopExpr* AILE);
403+
StmtDiff VisitArrayInitIndexExpr(const clang::ArrayInitIndexExpr* AIIE);
404+
StmtDiff VisitOpaqueValueExpr(const clang::OpaqueValueExpr* OVE);
400405
virtual StmtDiff VisitReturnStmt(const clang::ReturnStmt* RS);
401406
StmtDiff VisitStmt(const clang::Stmt* S);
402407
virtual StmtDiff VisitUnaryOperator(const clang::UnaryOperator* UnOp);
@@ -718,6 +723,13 @@ namespace clad {
718723
void PopSwitchStmtInfo() { m_SwitchStmtsData.pop_back(); }
719724

720725
private:
726+
// When differentiating ArrayInitLoopExpr, we need to replace
727+
// ArrayInitIndexExpr with real indices. We need to both add and pop them in
728+
// the right order, so we use std::queue. For example, for `arr[i][j]`, we
729+
// add `i` first, then `j`, and then pop them in the same order to generate
730+
// the subscript expr.
731+
std::queue<clang::VarDecl*> m_ArrayInitLoopIdx;
732+
721733
// FIXME: This variable is used to track
722734
// whether we're currently visiting an init of a var decl.
723735
// This is only necessary because we don't create constructors
@@ -726,6 +738,7 @@ namespace clad {
726738
// other cases, we have to use InitListExpr and change the constructor
727739
// style. Remove this once we generate constructors explicitly.
728740
bool m_TrackVarDeclConstructor = false;
741+
729742
/// A flag indicating if the Stmt is contained in a checkpointed loop.
730743
bool m_IsInsideCheckpointedLoop = false;
731744
};

include/clad/Differentiator/VisitorBase.h

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -13,6 +13,7 @@
1313

1414
#include "clang/AST/Expr.h"
1515
#include "clang/AST/RecursiveASTVisitor.h"
16+
#include "clang/AST/Stmt.h"
1617
#include "clang/AST/StmtVisitor.h"
1718
#include "clang/AST/Type.h"
1819
#include "clang/Basic/Diagnostic.h"
@@ -28,6 +29,7 @@
2829

2930
#include <array>
3031
#include <cassert>
32+
#include <cstddef>
3133
#include <stack>
3234
#include <unordered_map>
3335

@@ -294,6 +296,14 @@ namespace clad {
294296
clang::Expr* BuildOperatorCall(clang::OverloadedOperatorKind OOK,
295297
llvm::MutableArrayRef<clang::Expr*> ArgExprs,
296298
clang::SourceLocation OpLoc = noLoc);
299+
300+
/// A shorthand to generage a standard loop of form
301+
/// ```
302+
/// for (type loopCounter = 0; loopCounter < N; ++loopCounter)
303+
/// body;
304+
/// ```
305+
clang::ForStmt* BuildStandardForLoop(clang::VarDecl* loopCounter, size_t N,
306+
clang::Stmt* body);
297307
/// Function to resolve Unary Minus. If the leftmost operand
298308
/// has a Unary Minus then adds parens before adding the unary minus.
299309
/// \param[in] E Expression fed to the recursive call.

lib/Differentiator/DiffPlanner.cpp

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -709,6 +709,10 @@ static QualType GetDerivedFunctionType(const CallExpr* CE) {
709709
return false;
710710
return true;
711711
}
712+
// The sub-stmt of OpaqueValueExpr is not visited automatically
713+
bool VisitOpaqueValueExpr(const clang::OpaqueValueExpr* OVE) {
714+
return TraverseStmt(OVE->getSourceExpr());
715+
}
712716
// FIXME: This is a temporary measure until we add support for
713717
// `this` in varied analysis.
714718
bool VisitCXXThisExpr(const clang::CXXThisExpr* TE) { return false; }

lib/Differentiator/ReverseModeVisitor.cpp

Lines changed: 43 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -723,6 +723,49 @@ Expr* ReverseModeVisitor::getStdInitListSizeExpr(const Expr* E) {
723723
return Visit(ILE->getSubExpr(), dfdx());
724724
}
725725

726+
StmtDiff
727+
ReverseModeVisitor::VisitArrayInitLoopExpr(const ArrayInitLoopExpr* AILE) {
728+
// Since ArrayInitLoopExpr is not possible to express with regular syntax,
729+
// we have to replicate it with loops.
730+
// The code we're differentiated is of the form
731+
// res = ArrayInitLoopExpr(arr[ArrayInitIndexExpr])
732+
// We have to replace ArrayInitIndexExpr with an actual index `i`
733+
// and wrap the code in a for loop to compute the derivative as follows:
734+
// for (int i = 0; i < N; ++i)
735+
// _d_arr[i] += _d_res[i];
736+
beginScope(Scope::DeclScope);
737+
VarDecl* idxDecl = BuildVarDecl(m_Context.UnsignedIntTy, "i",
738+
getZeroInit(m_Context.IntTy));
739+
// Push the index to the queue so that we can replace ArrayInitIndexExpr
740+
// when we encounter it.
741+
m_ArrayInitLoopIdx.push(idxDecl);
742+
Expr* idx = BuildDeclRef(idxDecl);
743+
// Build `_d_res[i]`
744+
Expr* diff = BuildArraySubscript(dfdx(), {idx});
745+
beginBlock(direction::reverse);
746+
Visit(AILE->getSubExpr(), diff);
747+
Stmt* block = utils::unwrapIfSingleStmt(endBlock(direction::reverse));
748+
Stmt* loopDiff = BuildStandardForLoop(
749+
idxDecl, AILE->getArraySize().getZExtValue(), block);
750+
addToCurrentBlock(loopDiff, direction::reverse);
751+
endScope();
752+
// We cannot clone ArrayInitLoopExpr because it's not possible to express
753+
// with standard c++ syntax.
754+
return {};
755+
}
756+
757+
StmtDiff
758+
ReverseModeVisitor::VisitArrayInitIndexExpr(const ArrayInitIndexExpr* AIIE) {
759+
VarDecl* idxDecl = m_ArrayInitLoopIdx.front();
760+
m_ArrayInitLoopIdx.pop();
761+
return {BuildDeclRef(idxDecl)};
762+
}
763+
764+
StmtDiff
765+
ReverseModeVisitor::VisitOpaqueValueExpr(const OpaqueValueExpr* OVE) {
766+
return Visit(OVE->getSourceExpr(), dfdx());
767+
}
768+
726769
StmtDiff ReverseModeVisitor::VisitStmt(const Stmt* S) {
727770
diagUnsupported(S);
728771
// Unknown stmt, just clone it.

lib/Differentiator/VisitorBase.cpp

Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -23,6 +23,7 @@
2323
#include "clang/AST/Expr.h"
2424
#include "clang/AST/NestedNameSpecifier.h"
2525
#include "clang/AST/OperationKinds.h"
26+
#include "clang/AST/Stmt.h"
2627
#include "clang/AST/TemplateBase.h"
2728
#include "clang/Basic/OperatorKinds.h"
2829
#include "clang/Basic/SourceManager.h"
@@ -41,6 +42,7 @@
4142
#include "llvm/Support/Casting.h"
4243

4344
#include <algorithm>
45+
#include <cstddef>
4446
#include <numeric>
4547

4648
#include "clad/Differentiator/Compatibility.h"
@@ -264,6 +266,17 @@ namespace clad {
264266
return new (m_Context) DeclStmt(DGR, noLoc, noLoc);
265267
}
266268

269+
ForStmt* VisitorBase::BuildStandardForLoop(VarDecl* loopCounter, size_t N,
270+
Stmt* body) {
271+
Stmt* init = BuildDeclStmt(loopCounter);
272+
Expr* numExpr =
273+
ConstantFolder::synthesizeLiteral(m_Context.IntTy, m_Context, N);
274+
Expr* cond = BuildOp(BO_LT, BuildDeclRef(loopCounter), numExpr);
275+
Expr* inc = BuildOp(UO_PreInc, BuildDeclRef(loopCounter));
276+
return new (m_Context) ForStmt(m_Context, init, cond, /*CondVar=*/nullptr,
277+
inc, body, noLoc, noLoc, noLoc);
278+
}
279+
267280
DeclRefExpr* VisitorBase::BuildDeclRef(DeclaratorDecl* D,
268281
NestedNameSpecifier* NNS /*=nullptr*/,
269282
ExprValueKind VK /*=VK_LValue*/) {

test/Gradient/Constructors.C

Lines changed: 127 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -445,6 +445,123 @@ double fn9(S& s) {
445445
// CHECK-NEXT: }
446446
// CHECK-NEXT: }
447447

448+
struct arrWrapper {
449+
double arr[2];
450+
};
451+
452+
// CHECK: static inline constexpr void constructor_pullback(const arrWrapper &arg, arrWrapper *_d_this, arrWrapper *_d_arg) noexcept {
453+
// CHECK-NEXT: for (unsigned {{int|long}} i = 0; i < 2; ++i)
454+
// CHECK-NEXT: (*_d_arg).arr[i] += _d_this->arr[i];
455+
// CHECK-NEXT: }
456+
457+
double fn10(double x, double y) {
458+
arrWrapper a = {x, y};
459+
arrWrapper b = a;
460+
return b.arr[0] + b.arr[1];
461+
}
462+
463+
// CHECK: void fn10_grad(double x, double y, double *_d_x, double *_d_y) {
464+
// CHECK-NEXT: arrWrapper _d_a = {{.*0.*}};
465+
// CHECK-NEXT: arrWrapper a = {{.*x, y.*}};
466+
// CHECK-NEXT: arrWrapper _d_b = _d_a;
467+
// CHECK-NEXT: arrWrapper b = a;
468+
// CHECK-NEXT: {
469+
// CHECK-NEXT: _d_b.arr[0] += 1;
470+
// CHECK-NEXT: _d_b.arr[1] += 1;
471+
// CHECK-NEXT: }
472+
// CHECK-NEXT: arrWrapper::constructor_pullback(a, &_d_b, &_d_a);
473+
// CHECK-NEXT: {
474+
// CHECK-NEXT: *_d_x += _d_a.arr[0];
475+
// CHECK-NEXT: *_d_y += _d_a.arr[1];
476+
// CHECK-NEXT: }
477+
// CHECK-NEXT: }
478+
479+
struct arr2DWrapper {
480+
double arr[1][2];
481+
};
482+
483+
// CHECK: static inline constexpr void constructor_pullback(const arr2DWrapper &arg, arr2DWrapper *_d_this, arr2DWrapper *_d_arg) noexcept {
484+
// CHECK-NEXT: for (unsigned {{int|long}} i = 0; i < 1; ++i)
485+
// CHECK-NEXT: for (unsigned {{int|long}} i0 = 0; i0 < 2; ++i0)
486+
// CHECK-NEXT: (*_d_arg).arr[i][i0] += _d_this->arr[i][i0];
487+
// CHECK-NEXT: }
488+
489+
double fn11(double x, double y) {
490+
arr2DWrapper a = {x, y};
491+
arr2DWrapper b = a;
492+
return b.arr[0][0] + b.arr[0][1];
493+
}
494+
495+
// CHECK: void fn11_grad(double x, double y, double *_d_x, double *_d_y) {
496+
// CHECK-NEXT: arr2DWrapper _d_a = {{.*0.*}};
497+
// CHECK-NEXT: arr2DWrapper a = {{.*x, y.*}};
498+
// CHECK-NEXT: arr2DWrapper _d_b = _d_a;
499+
// CHECK-NEXT: arr2DWrapper b = a;
500+
// CHECK-NEXT: {
501+
// CHECK-NEXT: _d_b.arr[0][0] += 1;
502+
// CHECK-NEXT: _d_b.arr[0][1] += 1;
503+
// CHECK-NEXT: }
504+
// CHECK-NEXT: arr2DWrapper::constructor_pullback(a, &_d_b, &_d_a);
505+
// CHECK-NEXT: {
506+
// CHECK-NEXT: *_d_x += _d_a.arr[0][0];
507+
// CHECK-NEXT: *_d_y += _d_a.arr[0][1];
508+
// CHECK-NEXT: }
509+
// CHECK-NEXT: }
510+
511+
struct cust_double {
512+
cust_double(double x = 0): val(x) {}
513+
double val;
514+
};
515+
516+
// CHECK: static void constructor_pullback(double x, cust_double *_d_this, double *_d_x) {
517+
// CHECK-NEXT: {
518+
// CHECK-NEXT: *_d_x += _d_this->val;
519+
// CHECK-NEXT: _d_this->val = 0.;
520+
// CHECK-NEXT: }
521+
// CHECK-NEXT: }
522+
523+
// CHECK: static inline constexpr void constructor_pullback(const cust_double &arg, cust_double *_d_this, cust_double *_d_arg) noexcept {
524+
// CHECK-NEXT: {
525+
// CHECK-NEXT: (*_d_arg).val += _d_this->val;
526+
// CHECK-NEXT: _d_this->val = 0.;
527+
// CHECK-NEXT: }
528+
// CHECK-NEXT: }
529+
530+
struct arrStructWrapper {
531+
cust_double arr[2];
532+
};
533+
534+
// CHECK: static inline constexpr void constructor_pullback(const arrStructWrapper &arg, arrStructWrapper *_d_this, arrStructWrapper *_d_arg) noexcept {
535+
// CHECK-NEXT: for (unsigned {{int|long}} i = 0; i < 2; ++i)
536+
// CHECK-NEXT: cust_double::constructor_pullback(arg.arr[i], &_d_this->arr[i], &(*_d_arg).arr[i]);
537+
// CHECK-NEXT: }
538+
539+
double fn12(double x, double y) {
540+
arrStructWrapper a = {x, y};
541+
arrStructWrapper b = a;
542+
return b.arr[0].val + b.arr[1].val;
543+
}
544+
545+
// CHECK: void fn12_grad(double x, double y, double *_d_x, double *_d_y) {
546+
// CHECK-NEXT: arrStructWrapper _d_a = {{.*}};
547+
// CHECK-NEXT: arrStructWrapper a = {{.*x, y.*}};
548+
// CHECK-NEXT: arrStructWrapper _d_b = _d_a;
549+
// CHECK-NEXT: arrStructWrapper b = a;
550+
// CHECK-NEXT: {
551+
// CHECK-NEXT: _d_b.arr[0].val += 1;
552+
// CHECK-NEXT: _d_b.arr[1].val += 1;
553+
// CHECK-NEXT: }
554+
// CHECK-NEXT: arrStructWrapper::constructor_pullback(a, &_d_b, &_d_a);
555+
// CHECK-NEXT: {
556+
// CHECK-NEXT: double _r0 = 0.;
557+
// CHECK-NEXT: cust_double::constructor_pullback(x, &_d_a.arr[0], &_r0);
558+
// CHECK-NEXT: *_d_x += _r0;
559+
// CHECK-NEXT: double _r1 = 0.;
560+
// CHECK-NEXT: cust_double::constructor_pullback(y, &_d_a.arr[1], &_r1);
561+
// CHECK-NEXT: *_d_y += _r1;
562+
// CHECK-NEXT: }
563+
// CHECK-NEXT: }
564+
448565
int main() {
449566
double d_i, d_j;
450567

@@ -474,6 +591,14 @@ int main() {
474591
S s{new double[3]{5, 6, 7}}, _d_s{new double[3]{0}};
475592
auto dfn9 = clad::gradient(fn9);
476593
dfn9.execute(s, &_d_s);
477-
printf("{%.2f, %.2f, %.2f}\n", _d_s.a[0], _d_s.a[1], _d_s.a[2]);
478-
// TEST_GRADIENT(fn9, /*numOfDerivativeArgs=*/1, s, &d_s); // CHECK-EXEC: {0.00, 0.00, 1.00}
594+
printf("{%.2f, %.2f, %.2f}\n", _d_s.a[0], _d_s.a[1], _d_s.a[2]); // CHECK-EXEC: {0.00, 0.00, 1.00}
595+
596+
INIT_GRADIENT(fn10);
597+
TEST_GRADIENT(fn10, /*numOfDerivativeArgs=*/2, 7, 2, &d_i, &d_j); // CHECK-EXEC: {1.00, 1.00}
598+
599+
INIT_GRADIENT(fn11);
600+
TEST_GRADIENT(fn11, /*numOfDerivativeArgs=*/2, 9, -1, &d_i, &d_j); // CHECK-EXEC: {1.00, 1.00}
601+
602+
INIT_GRADIENT(fn12);
603+
TEST_GRADIENT(fn12, /*numOfDerivativeArgs=*/2, 3, 6, &d_i, &d_j); // CHECK-EXEC: {1.00, 1.00}
479604
}

0 commit comments

Comments
 (0)