Skip to content

Commit a92b3c4

Browse files
PetroZarytskyivgvassilev
authored andcommitted
Make the clad::restore_tracker param in reverse_forw request-based
1 parent b765488 commit a92b3c4

7 files changed

Lines changed: 31 additions & 38 deletions

File tree

include/clad/Differentiator/CladUtils.h

Lines changed: 5 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -339,12 +339,11 @@ namespace clad {
339339
///
340340
/// \param[in] forCustomDerv If true, turns member functions into regular
341341
/// functions by moving the base to the parameters.
342-
clang::QualType
343-
GetDerivativeType(clang::Sema& S, const clang::FunctionDecl* FD,
344-
DiffMode mode,
345-
llvm::ArrayRef<const clang::ValueDecl*> diffParams,
346-
bool forCustomDerv = false,
347-
llvm::ArrayRef<clang::QualType> customParams = {});
342+
clang::QualType GetDerivativeType(
343+
clang::Sema& S, const clang::FunctionDecl* FD, DiffMode mode,
344+
llvm::ArrayRef<const clang::ValueDecl*> diffParams,
345+
bool forCustomDerv = false, bool shouldUseRestoreTracker = false,
346+
llvm::ArrayRef<clang::QualType> customParams = {});
348347
/// Find declaration of clad::class templated type
349348
///
350349
/// \param[in] className name of the class to be found

include/clad/Differentiator/DiffPlanner.h

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -90,8 +90,13 @@ struct DiffRequest {
9090
bool VerboseDiags = false;
9191
/// A flag to enable TBR analysis during reverse-mode differentiation.
9292
bool EnableTBRAnalysis = false;
93+
/// A flag to enable varied analysis during reverse-mode differentiation.
9394
bool EnableVariedAnalysis = false;
95+
/// A flag to enable useful analysis during reverse-mode differentiation.
9496
bool EnableUsefulAnalysis = false;
97+
/// A flag to request a clad::restore_tracker parameter in the generated
98+
/// _reverse_forw function.
99+
bool UseRestoreTracker = false;
95100
/// A flag specifying whether this differentiation is to be used
96101
/// in immediate contexts.
97102
bool ImmediateMode = false;

lib/Differentiator/CladUtils.cpp

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1142,7 +1142,7 @@ namespace clad {
11421142
QualType
11431143
GetDerivativeType(Sema& S, const clang::FunctionDecl* FD, DiffMode mode,
11441144
llvm::ArrayRef<const clang::ValueDecl*> diffParams,
1145-
bool forCustomDerv,
1145+
bool forCustomDerv, bool shouldUseRestoreTracker,
11461146
llvm::ArrayRef<QualType> customParams) {
11471147
ASTContext& C = S.getASTContext();
11481148
if (mode == DiffMode::forward)
@@ -1266,7 +1266,7 @@ namespace clad {
12661266
FnTypes.insert(FnTypes.begin(), typeTag);
12671267
}
12681268

1269-
if (shouldUseRestoreTracker(FD) && !forCustomDerv) {
1269+
if (shouldUseRestoreTracker) {
12701270
QualType trackerTy = GetRestoreTrackerType(S);
12711271
trackerTy = C.getLValueReferenceType(trackerTy);
12721272
FnTypes.push_back(trackerTy);

lib/Differentiator/DiffPlanner.cpp

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -968,7 +968,8 @@ static QualType GetDerivedFunctionType(const CallExpr* CE) {
968968
for (const DiffInputVarInfo& VarInfo : R.DVI)
969969
diffParams.push_back(VarInfo.param);
970970
QualType dTy = utils::GetDerivativeType(S, R.Function, R.Mode, diffParams,
971-
/*forCustomDerv=*/true);
971+
/*forCustomDerv=*/true,
972+
/*shouldUseRestoreTracker=*/false);
972973
// We disable diagnostics for methods and operators because they often have
973974
// ideantical names: `constructor_pullback`, `operator_star_pushforward`,
974975
// etc. If we turn it on, every such operator will trigger diagnostics
@@ -1424,6 +1425,7 @@ static QualType GetDerivedFunctionType(const CallExpr* CE) {
14241425
forwPassRequest.BaseFunctionName = request.BaseFunctionName;
14251426
forwPassRequest.Mode = DiffMode::reverse_mode_forward_pass;
14261427
forwPassRequest.CallContext = request.CallContext;
1428+
forwPassRequest.UseRestoreTracker = shouldUseRestoreTracker;
14271429
QualType returnType = request->getReturnType();
14281430
if (LookupCustomDerivativeDecl(forwPassRequest) ||
14291431
utils::isMemoryType(returnType) || shouldUseRestoreTracker)

lib/Differentiator/ReverseModeForwPassVisitor.cpp

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -158,7 +158,7 @@ ReverseModeForwPassVisitor::BuildParams(DiffParams& diffParams) {
158158
m_Variables[*it] = BuildDeclRef(dPVD), m_DiffReq->getLocation();
159159
}
160160
}
161-
if (utils::shouldUseRestoreTracker(m_DiffReq.Function)) {
161+
if (m_DiffReq.UseRestoreTracker) {
162162
QualType trackerTy = utils::GetRestoreTrackerType(m_Sema);
163163
trackerTy = m_Sema.getASTContext().getLValueReferenceType(trackerTy);
164164
ParmVarDecl* trackerPVD = utils::BuildParmVarDecl(

lib/Differentiator/VisitorBase.cpp

Lines changed: 4 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1006,9 +1006,10 @@ namespace clad {
10061006
llvm::SmallVector<const ValueDecl*, 4> diffParams{};
10071007
for (const DiffInputVarInfo& VarInfo : m_DiffReq.DVI)
10081008
diffParams.push_back(VarInfo.param);
1009-
return utils::GetDerivativeType(m_Sema, m_DiffReq.Function, m_DiffReq.Mode,
1010-
diffParams, /*forCustomDerv=*/false,
1011-
customParams);
1009+
return utils::GetDerivativeType(
1010+
m_Sema, m_DiffReq.Function, m_DiffReq.Mode, diffParams,
1011+
/*forCustomDerv=*/false,
1012+
/*shouldUseRestoreTracker=*/m_DiffReq.UseRestoreTracker, customParams);
10121013
}
10131014

10141015
FunctionDecl* VisitorBase::FindDerivedFunction(DiffRequest& request) {

test/Gradient/UserDefinedTypes.C

Lines changed: 11 additions & 25 deletions
Original file line numberDiff line numberDiff line change
@@ -959,7 +959,7 @@ namespace class_functions {
959959
constructor_reverse_forw(::clad::Tag<ptrClass>, double* mptr, double* d_mptr) elidable_reverse_forw;
960960
}}}
961961

962-
// CHECK: clad::ValueAndAdjoint<double &, double &> operator_star_reverse_forw(ptrClass *_d_this, clad::restore_tracker &_tracker0) {
962+
// CHECK: clad::ValueAndAdjoint<double &, double &> operator_star_reverse_forw(ptrClass *_d_this) {
963963
// CHECK-NEXT: return {*this->ptr, *_d_this->ptr};
964964
// CHECK-NEXT: }
965965

@@ -974,12 +974,8 @@ double fn26(double x, double y) {
974974
// CHECK: void fn26_grad(double x, double y, double *_d_x, double *_d_y) {
975975
// CHECK-NEXT: ptrClass p(&x);
976976
// CHECK-NEXT: ptrClass _d_p(_d_x);
977-
// CHECK-NEXT: clad::restore_tracker _tracker0 = {};
978-
// CHECK-NEXT: clad::ValueAndAdjoint<double &, double &> _t0 = p.operator_star_reverse_forw(&_d_p, _tracker0);
979-
// CHECK-NEXT: {
980-
// CHECK-NEXT: _t0.adjoint += 1;
981-
// CHECK-NEXT: _tracker0.restore();
982-
// CHECK-NEXT: }
977+
// CHECK-NEXT: clad::ValueAndAdjoint<double &, double &> _t0 = p.operator_star_reverse_forw(&_d_p);
978+
// CHECK-NEXT: _t0.adjoint += 1;
983979
// CHECK-NEXT: }
984980

985981
struct MyStructWrapper {
@@ -1120,7 +1116,7 @@ MyStruct& objRef(MyStruct& s) {
11201116
return s;
11211117
}
11221118

1123-
// CHECK: clad::ValueAndAdjoint<MyStruct &, MyStruct &> objRef_reverse_forw(MyStruct &s, MyStruct &_d_s, clad::restore_tracker &_tracker0) {
1119+
// CHECK: clad::ValueAndAdjoint<MyStruct &, MyStruct &> objRef_reverse_forw(MyStruct &s, MyStruct &_d_s) {
11241120
// CHECK-NEXT: return {s, _d_s};
11251121
// CHECK-NEXT: }
11261122

@@ -1135,12 +1131,8 @@ double fn30(double x, double y) {
11351131
// CHECK: void fn30_grad(double x, double y, double *_d_x, double *_d_y) {
11361132
// CHECK-NEXT: MyStruct _d_a = {0., 0.};
11371133
// CHECK-NEXT: MyStruct a{x, y};
1138-
// CHECK-NEXT: clad::restore_tracker _tracker0 = {};
1139-
// CHECK-NEXT: clad::ValueAndAdjoint<MyStruct &, MyStruct &> _t0 = objRef_reverse_forw(a, _d_a, _tracker0);
1140-
// CHECK-NEXT: {
1141-
// CHECK-NEXT: _tracker0.restore();
1142-
// CHECK-NEXT: _t0.adjoint.b += 1;
1143-
// CHECK-NEXT: }
1134+
// CHECK-NEXT: clad::ValueAndAdjoint<MyStruct &, MyStruct &> _t0 = objRef_reverse_forw(a, _d_a);
1135+
// CHECK-NEXT: _t0.adjoint.b += 1;
11441136
// CHECK-NEXT: {
11451137
// CHECK-NEXT: *_d_x += _d_a.a;
11461138
// CHECK-NEXT: *_d_y += _d_a.b;
@@ -1248,7 +1240,7 @@ struct structToConvert {
12481240
return data;
12491241
}
12501242

1251-
// CHECK: clad::ValueAndAdjoint<otherStruct, otherStruct> conversion_operator_reverse_forw(clad::Tag<otherStruct>, structToConvert *_d_this, clad::restore_tracker &_tracker0) {
1243+
// CHECK: clad::ValueAndAdjoint<otherStruct, otherStruct> conversion_operator_reverse_forw(clad::Tag<otherStruct>, structToConvert *_d_this) {
12521244
// CHECK-NEXT: return {{[{][{]}}2 * this->data, &this->data}, {0., &_d_this->data{{[}][}]}};
12531245
// CHECK-NEXT:}
12541246

@@ -1260,7 +1252,7 @@ struct structToConvert {
12601252
return {2 * data, &data};
12611253
}
12621254

1263-
// CHECK: clad::ValueAndAdjoint<double &, double &> conversion_operator_reverse_forw(clad::Tag<double &>, structToConvert *_d_this, clad::restore_tracker &_tracker0) {
1255+
// CHECK: clad::ValueAndAdjoint<double &, double &> conversion_operator_reverse_forw(clad::Tag<double &>, structToConvert *_d_this) {
12641256
// CHECK-NEXT: return {this->data, _d_this->data};
12651257
// CHECK-NEXT:}
12661258

@@ -1280,21 +1272,15 @@ double fn34(double x, double y) {
12801272
// CHECK-NEXT: structToConvert obj_x{x};
12811273
// CHECK-NEXT: structToConvert _d_obj_y = {0.};
12821274
// CHECK-NEXT: structToConvert obj_y{y};
1283-
// CHECK-NEXT: clad::restore_tracker _tracker0 = {};
1284-
// CHECK-NEXT: clad::ValueAndAdjoint<otherStruct, otherStruct> _t0 = obj_y.conversion_operator_reverse_forw(clad::Tag<otherStruct>(), &_d_obj_y, _tracker0);
1275+
// CHECK-NEXT: clad::ValueAndAdjoint<otherStruct, otherStruct> _t0 = obj_y.conversion_operator_reverse_forw(clad::Tag<otherStruct>(), &_d_obj_y);
12851276
// CHECK-NEXT: otherStruct _d_conv = (otherStruct)_t0.adjoint;
12861277
// CHECK-NEXT: otherStruct conv = (otherStruct)_t0.value;
1287-
// CHECK-NEXT: clad::restore_tracker _tracker1 = {};
1288-
// CHECK-NEXT: clad::ValueAndAdjoint<double &, double &> _t1 = obj_x.conversion_operator_reverse_forw(clad::Tag<double &>(), &_d_obj_x, _tracker1);
1278+
// CHECK-NEXT: clad::ValueAndAdjoint<double &, double &> _t1 = obj_x.conversion_operator_reverse_forw(clad::Tag<double &>(), &_d_obj_x);
12891279
// CHECK-NEXT: {
12901280
// CHECK-NEXT: _t1.adjoint += 1;
1291-
// CHECK-NEXT: _tracker1.restore();
12921281
// CHECK-NEXT: _d_conv.val += 1;
12931282
// CHECK-NEXT: }
1294-
// CHECK-NEXT: {
1295-
// CHECK-NEXT: _tracker0.restore();
1296-
// CHECK-NEXT: obj_y.conversion_operator_pullback(_d_conv, &_d_obj_y);
1297-
// CHECK-NEXT: }
1283+
// CHECK-NEXT: obj_y.conversion_operator_pullback(_d_conv, &_d_obj_y);
12981284
// CHECK-NEXT: *_d_y += _d_obj_y.data;
12991285
// CHECK-NEXT: *_d_x += _d_obj_x.data;
13001286
// CHECK-NEXT:}

0 commit comments

Comments
 (0)