Skip to content

Commit c31d12f

Browse files
Add custom derivatives for std::forward.
Fixes #1274.
1 parent d670b20 commit c31d12f

3 files changed

Lines changed: 48 additions & 1 deletion

File tree

include/clad/Differentiator/Differentiator.h

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -120,7 +120,7 @@ CUDA_HOST_DEVICE T push(tape<T>& to, ArgsT... val) {
120120
// copying the sequence of characters to the destination [27.5.1(3)].
121121
// (C++ has deprecated the volatile qualifiers. However, we drop them here
122122
// to make sure things still work with codebases which still have them)
123-
std::memcpy(const_cast<T*>(&t), tmp, sizeof(T));
123+
std::memcpy((void*)const_cast<T*>(&t), tmp, sizeof(T));
124124
}
125125

126126
template <class T, typename std::enable_if<is_range<T>::value, int>::type = 0>

include/clad/Differentiator/STLBuiltins.h

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -809,6 +809,15 @@ template <typename... Args> auto make_tuple_pushforward(Args... args) noexcept {
809809
second_half_tuple(t));
810810
}
811811

812+
template <class T> constexpr void forward_pullback(T& t, T dy, T* dt) noexcept {
813+
*dt += dy;
814+
}
815+
816+
template <class T>
817+
constexpr void forward_pullback(T&& t, T dy, T* dt) noexcept {
818+
*dt += dy;
819+
}
820+
812821
} // namespace std
813822

814823
} // namespace custom_derivatives

test/Gradient/Constructors.C

Lines changed: 38 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -411,6 +411,41 @@ double fn7(double x, double y) {
411411
// CHECK-NEXT: }
412412
// CHECK-NEXT: }
413413

414+
double fn8(double u, double v) {
415+
std::pair<double, double> p(u,v);
416+
return p.first + p.second;
417+
}
418+
419+
// CHECK: static constexpr void constructor_pullback(double &__{{u1|x}}, double &__{{u2|y}}, std::pair<double, double> *_d_this, double *_d___{{u1|x}}, double *_d___{{u2|y}}) {{.*}}{
420+
// CHECK-NEXT: std::pair<double, double> *_this = (std::pair<double, double> *)malloc(sizeof(std::pair<double, double>));
421+
// CHECK: double _t0 = __{{u1|x}};
422+
// CHECK-NEXT: _this->first = std::forward<double &>(__{{u1|x}});
423+
// CHECK-NEXT: double _t1 = __{{u2|y}};
424+
// CHECK-NEXT: _this->second = std::forward<double &>(__{{u2|y}});
425+
// CHECK-NEXT: {
426+
// CHECK-NEXT: clad::custom_derivatives::std::forward_pullback(__{{u2|y}}, _d_this->second, &*_d___{{u2|y}});
427+
// CHECK-NEXT: __{{u2|y}} = _t1;
428+
// CHECK-NEXT: _d_this->second = 0.;
429+
// CHECK-NEXT: }
430+
// CHECK-NEXT: {
431+
// CHECK-NEXT: clad::custom_derivatives::std::forward_pullback(__{{u1|x}}, _d_this->first, &*_d___{{u1|x}});
432+
// CHECK-NEXT: __{{u1|x}} = _t0;
433+
// CHECK-NEXT: _d_this->first = 0.;
434+
// CHECK-NEXT: }
435+
// CHECK-NEXT: free(_this);
436+
// CHECK-NEXT: }
437+
438+
// CHECK: void fn8_grad(double u, double v, double *_d_u, double *_d_v) {
439+
// CHECK-NEXT: std::pair<double, double> p(u, v);
440+
// CHECK-NEXT: std::pair<double, double> _d_p(p);
441+
// CHECK-NEXT: clad::zero_init(_d_p);
442+
// CHECK-NEXT: {
443+
// CHECK-NEXT: _d_p.first += 1;
444+
// CHECK-NEXT: _d_p.second += 1;
445+
// CHECK-NEXT: }
446+
// CHECK-NEXT: pair::constructor_pullback(u, v, &_d_p, &*_d_u, &*_d_v);
447+
// CHECK-NEXT: }
448+
414449
int main() {
415450
double d_i, d_j;
416451

@@ -433,4 +468,7 @@ int main() {
433468

434469
INIT_GRADIENT(fn7);
435470
TEST_GRADIENT(fn7, /*numOfDerivativeArgs=*/2, 2, 9, &d_i, &d_j); // CHECK-EXEC: {12.00, 0.00}
471+
472+
INIT_GRADIENT(fn8);
473+
TEST_GRADIENT(fn8, /*numOfDerivativeArgs=*/2, 7, 2, &d_i, &d_j); // CHECK-EXEC: {1.00, 1.00}
436474
}

0 commit comments

Comments
 (0)