Skip to content

Commit b8933b9

Browse files
Eliminates the reverse_forw pass when the custom pullback exists and it removes unwanted tape overhead (#1810)
Fixes issue #1808
1 parent 8b8d22a commit b8933b9

3 files changed

Lines changed: 13 additions & 18 deletions

File tree

lib/Differentiator/DiffPlanner.cpp

Lines changed: 7 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1420,9 +1420,14 @@ static QualType GetDerivedFunctionType(const CallExpr* CE) {
14201420
forwPassRequest.CallContext = request.CallContext;
14211421
forwPassRequest.UseRestoreTracker = shouldUseRestoreTracker;
14221422
QualType returnType = request->getReturnType();
1423-
if (LookupCustomDerivativeDecl(forwPassRequest) ||
1424-
utils::isMemoryType(returnType) || shouldUseRestoreTracker)
1423+
bool hasCustomPullback = request.CustomDerivative != nullptr;
1424+
bool hasCustomReverseForw = LookupCustomDerivativeDecl(forwPassRequest);
1425+
1426+
if (hasCustomReverseForw ||
1427+
(!hasCustomPullback &&
1428+
(utils::isMemoryType(returnType) || shouldUseRestoreTracker))) {
14251429
m_DiffRequestGraph.addNode(forwPassRequest, /*isSource=*/true);
1430+
}
14261431
}
14271432

14281433
if (!nonDiff && request.Mode != DiffMode::unknown)

test/Gradient/FunctionCalls.C

Lines changed: 0 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -1098,16 +1098,6 @@ void inner_function(double *out, bool flag, const double *C) {
10981098
if (flag)
10991099
out[0] = C[0];
11001100
}
1101-
// CHECK: void inner_function_reverse_forw(double *out, bool flag, const double *C, double *_d_out, bool _d_flag, {{(const )?}}double *_d_C, clad::restore_tracker &_tracker0) {
1102-
// CHECK-NEXT: {
1103-
// CHECK-NEXT: bool _cond0 = flag;
1104-
// CHECK-NEXT: if (_cond0) {
1105-
// CHECK-NEXT: _tracker0.store(out[0]);
1106-
// CHECK-NEXT: out[0] = C[0];
1107-
// CHECK-NEXT: }
1108-
// CHECK-NEXT: }
1109-
// CHECK-NEXT: }
1110-
11111101
double fn31(double *variables) {
11121102
double out = 0.;
11131103
inner_function(&out, true, variables);

test/Gradient/STLCustomDerivatives.C

Lines changed: 6 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -714,23 +714,23 @@ int main() {
714714
// CHECK-NEXT: _t2++;
715715
// CHECK-NEXT: res += v.at(i0);
716716
// CHECK-NEXT: }
717-
// CHECK-NEXT: clad::restore_tracker _tracker0 = {};
718-
// CHECK-NEXT: v.assign_reverse_forw(3, 0, &_d_v, 0, 0, _tracker0);
719-
// CHECK-NEXT: clad::restore_tracker _tracker1 = {};
720-
// CHECK-NEXT: v.assign_reverse_forw(2, y, &_d_v, 0, *_d_y, _tracker1);
717+
// CHECK-NEXT: std::vector<double> _t3 = v;
718+
// CHECK-NEXT: v.assign(3, 0);
719+
// CHECK-NEXT: std::vector<double> _t4 = v;
720+
// CHECK-NEXT: v.assign(2, y);
721721
// CHECK-NEXT: {
722722
// CHECK-NEXT: _d_res += 1;
723723
// CHECK-NEXT: _d_v[0] += 1;
724724
// CHECK-NEXT: _d_v[1] += 1;
725725
// CHECK-NEXT: _d_v[2] += 1;
726726
// CHECK-NEXT: }
727727
// CHECK-NEXT: {
728-
// CHECK-NEXT: _tracker1.restore();
728+
// CHECK-NEXT: v = _t4;
729729
// CHECK-NEXT: {{.*size_type|size_t}} _r2 = {{0U|0UL|0}};
730730
// CHECK-NEXT: {{.*}}assign_pullback(&v, 2, y, &_d_v, &_r2, _d_y);
731731
// CHECK-NEXT: }
732732
// CHECK-NEXT: {
733-
// CHECK-NEXT: _tracker0.restore();
733+
// CHECK-NEXT: v = _t3;
734734
// CHECK-NEXT: {{.*size_type|size_t}} _r0 = {{0U|0UL|0}};
735735
// CHECK-NEXT: {{.*}}value_type _r1 = 0.;
736736
// CHECK-NEXT: {{.*}}assign_pullback(&v, 3, 0, &_d_v, &_r0, &_r1);

0 commit comments

Comments
 (0)