-
Notifications
You must be signed in to change notification settings - Fork 200
Support std::pair constructors natively #1440
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Changes from all commits
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -26,7 +26,7 @@ double fn1(double x, double y) { | |
| argByVal g(x); | ||
| y = x; | ||
| return y + g.y; | ||
| } | ||
| } // x + x^2 | ||
|
|
||
| // CHECK: static void constructor_pullback(double val, argByVal *_d_this, double *_d_val) { | ||
| // CHECK-NEXT: argByVal *_this = (argByVal *)malloc(sizeof(argByVal)); | ||
|
|
@@ -285,6 +285,168 @@ double fn4(double i, double j) { | |
| // CHECK-NEXT: } | ||
| // CHECK-NEXT: } | ||
|
|
||
| struct argByValWrapper : public argByVal { | ||
| double z; | ||
| argByValWrapper(double v) : argByVal(v) {} | ||
| argByValWrapper(double v, double u) : argByValWrapper(v) { | ||
| z = y * u; | ||
| } | ||
| argByValWrapper(double v, bool) : argByVal(v) { | ||
| z = x * y; | ||
| } | ||
| }; | ||
|
|
||
| double fn5(double x, double y) { | ||
| argByValWrapper g(x); | ||
| y = x; | ||
| return y + g.y; | ||
| } // x + x^2 | ||
|
|
||
| // CHECK: static void constructor_pullback(double v, argByValWrapper *_d_this, double *_d_v) { | ||
| // CHECK-NEXT: { | ||
| // CHECK-NEXT: double _r0 = 0.; | ||
| // CHECK-NEXT: argByVal::constructor_pullback(v, &*_d_this, &_r0); | ||
| // CHECK-NEXT: *_d_v += _r0; | ||
| // CHECK-NEXT: } | ||
| // CHECK-NEXT: } | ||
|
|
||
| // CHECK: void fn5_grad(double x, double y, double *_d_x, double *_d_y) { | ||
| // CHECK-NEXT: argByValWrapper g(x); | ||
| // CHECK-NEXT: argByValWrapper _d_g(g); | ||
| // CHECK-NEXT: clad::zero_init(_d_g); | ||
| // CHECK-NEXT: double _t0 = y; | ||
| // CHECK-NEXT: y = x; | ||
| // CHECK-NEXT: { | ||
| // CHECK-NEXT: *_d_y += 1; | ||
| // CHECK-NEXT: _d_g.y += 1; | ||
| // CHECK-NEXT: } | ||
| // CHECK-NEXT: { | ||
| // CHECK-NEXT: y = _t0; | ||
| // CHECK-NEXT: double _r_d0 = *_d_y; | ||
| // CHECK-NEXT: *_d_y = 0.; | ||
| // CHECK-NEXT: *_d_x += _r_d0; | ||
| // CHECK-NEXT: } | ||
| // CHECK-NEXT: { | ||
| // CHECK-NEXT: double _r0 = 0.; | ||
| // CHECK-NEXT: argByValWrapper::constructor_pullback(x, &_d_g, &_r0); | ||
| // CHECK-NEXT: *_d_x += _r0; | ||
| // CHECK-NEXT: } | ||
| // CHECK-NEXT: } | ||
|
|
||
| double fn6(double x, double y) { | ||
| argByValWrapper g(x, y); | ||
| return g.z; | ||
| } // x^2 * y | ||
|
|
||
| // CHECK: static void constructor_pullback(double v, double u, argByValWrapper *_d_this, double *_d_v, double *_d_u) { | ||
| // CHECK-NEXT: argByValWrapper *_this = new argByValWrapper(v); | ||
| // CHECK-NEXT: double _t0 = _this->z; | ||
| // CHECK-NEXT: _this->z = _this->y * u; | ||
| // CHECK-NEXT: { | ||
| // CHECK-NEXT: _this->z = _t0; | ||
| // CHECK-NEXT: double _r_d0 = _d_this->z; | ||
| // CHECK-NEXT: _d_this->z = 0.; | ||
| // CHECK-NEXT: _d_this->y += _r_d0 * u; | ||
| // CHECK-NEXT: *_d_u += _this->y * _r_d0; | ||
| // CHECK-NEXT: } | ||
| // CHECK-NEXT: { | ||
| // CHECK-NEXT: double _r0 = 0.; | ||
| // CHECK-NEXT: argByValWrapper::constructor_pullback(v, &*_d_this, &_r0); | ||
| // CHECK-NEXT: *_d_v += _r0; | ||
| // CHECK-NEXT: } | ||
| // CHECK-NEXT: free(_this); | ||
| // CHECK-NEXT: } | ||
|
|
||
| // CHECK: void fn6_grad(double x, double y, double *_d_x, double *_d_y) { | ||
| // CHECK-NEXT: argByValWrapper g(x, y); | ||
| // CHECK-NEXT: argByValWrapper _d_g(g); | ||
| // CHECK-NEXT: clad::zero_init(_d_g); | ||
| // CHECK-NEXT: _d_g.z += 1; | ||
| // CHECK-NEXT: { | ||
| // CHECK-NEXT: double _r0 = 0.; | ||
| // CHECK-NEXT: double _r1 = 0.; | ||
| // CHECK-NEXT: argByValWrapper::constructor_pullback(x, y, &_d_g, &_r0, &_r1); | ||
| // CHECK-NEXT: *_d_x += _r0; | ||
| // CHECK-NEXT: *_d_y += _r1; | ||
| // CHECK-NEXT: } | ||
| // CHECK-NEXT: } | ||
|
|
||
| double fn7(double x, double y) { | ||
| argByValWrapper g(x, false); | ||
| return g.z; | ||
| } // x^3 | ||
|
|
||
| // CHECK: static void constructor_pullback(double v, bool arg, argByValWrapper *_d_this, double *_d_v, bool *_d_arg) { | ||
| // CHECK-NEXT: argByValWrapper *_this = (argByValWrapper *)malloc(sizeof(argByValWrapper)); | ||
| // CHECK-NEXT: new (static_cast<argByVal *>(_this)) argByVal(v); | ||
|
Comment on lines
+380
to
+381
Owner
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. What do we lose if we call
Collaborator
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Well, some constructor calls have side effects and shouldn't be called twice. A real-world example would be move-constructors. If that's better, we can check if the constructor doesn't have any side effects by examining its parameter types or other factors. Also, I think the question is more about the logic in constructor pullbacks in general than about this change in particular. |
||
| // CHECK-NEXT: double _t0 = _this->z; | ||
| // CHECK-NEXT: _this->z = _this->x * _this->y; | ||
| // CHECK-NEXT: { | ||
| // CHECK-NEXT: _this->z = _t0; | ||
| // CHECK-NEXT: double _r_d0 = _d_this->z; | ||
| // CHECK-NEXT: _d_this->z = 0.; | ||
| // CHECK-NEXT: _d_this->x += _r_d0 * _this->y; | ||
| // CHECK-NEXT: _d_this->y += _this->x * _r_d0; | ||
| // CHECK-NEXT: } | ||
| // CHECK-NEXT: { | ||
| // CHECK-NEXT: double _r0 = 0.; | ||
| // CHECK-NEXT: argByVal::constructor_pullback(v, &*_d_this, &_r0); | ||
| // CHECK-NEXT: *_d_v += _r0; | ||
| // CHECK-NEXT: } | ||
| // CHECK-NEXT: free(_this); | ||
| // CHECK-NEXT: } | ||
|
|
||
|
|
||
| // CHECK: void fn7_grad(double x, double y, double *_d_x, double *_d_y) { | ||
| // CHECK-NEXT: argByValWrapper g(x, false); | ||
| // CHECK-NEXT: argByValWrapper _d_g(g); | ||
| // CHECK-NEXT: clad::zero_init(_d_g); | ||
| // CHECK-NEXT: _d_g.z += 1; | ||
| // CHECK-NEXT: { | ||
| // CHECK-NEXT: double _r0 = 0.; | ||
| // CHECK-NEXT: bool _r1 = false; | ||
| // CHECK-NEXT: argByValWrapper::constructor_pullback(x, false, &_d_g, &_r0, &_r1); | ||
| // CHECK-NEXT: *_d_x += _r0; | ||
| // CHECK-NEXT: } | ||
| // CHECK-NEXT: } | ||
|
|
||
| double fn8(double u, double v) { | ||
| std::pair<double, double> p(u,v); | ||
| return p.first + p.second; | ||
| } | ||
|
|
||
| // 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}}) {{.*}}{ | ||
| // CHECK-NEXT: std::pair<double, double> *_this = (std::pair<double, double> *)malloc(sizeof(std::pair<double, double>)); | ||
| // CHECK: double _t0 = __{{u1|x}}; | ||
| // CHECK-NEXT: clad::ValueAndAdjoint<double &, double &> _t1 = clad::custom_derivatives::std::forward_reverse_forw(__{{u1|x}}, *_d___{{u1|x}}); | ||
| // CHECK-NEXT: _this->first = _t1.value; | ||
| // CHECK-NEXT: double _t2 = __{{u2|y}}; | ||
| // CHECK-NEXT: clad::ValueAndAdjoint<double &, double &> _t3 = clad::custom_derivatives::std::forward_reverse_forw(__{{u2|y}}, *_d___{{u2|y}}); | ||
| // CHECK-NEXT: _this->second = _t3.value; | ||
| // CHECK: { | ||
| // CHECK-NEXT: clad::custom_derivatives::std::forward_pullback(__{{u2|y}}, _d_this->second, &*_d___{{u2|y}}); | ||
| // CHECK-NEXT: __{{u2|y}} = _t2; | ||
| // CHECK-NEXT: _d_this->second = 0.; | ||
| // CHECK-NEXT: } | ||
| // CHECK-NEXT: { | ||
| // CHECK-NEXT: clad::custom_derivatives::std::forward_pullback(__{{u1|x}}, _d_this->first, &*_d___{{u1|x}}); | ||
| // CHECK-NEXT: __{{u1|x}} = _t0; | ||
| // CHECK-NEXT: _d_this->first = 0.; | ||
| // CHECK-NEXT: } | ||
| // CHECK-NEXT: free(_this); | ||
| // CHECK-NEXT: } | ||
|
|
||
| // CHECK: void fn8_grad(double u, double v, double *_d_u, double *_d_v) { | ||
| // CHECK-NEXT: std::pair<double, double> p(u, v); | ||
| // CHECK-NEXT: std::pair<double, double> _d_p(p); | ||
| // CHECK-NEXT: clad::zero_init(_d_p); | ||
| // CHECK-NEXT: { | ||
| // CHECK-NEXT: _d_p.first += 1; | ||
| // CHECK-NEXT: _d_p.second += 1; | ||
| // CHECK-NEXT: } | ||
| // CHECK-NEXT: pair::constructor_pullback(u, v, &_d_p, &*_d_u, &*_d_v); | ||
| // CHECK-NEXT: } | ||
|
|
||
| int main() { | ||
| double d_i, d_j; | ||
|
|
||
|
|
@@ -298,4 +460,16 @@ int main() { | |
|
|
||
| INIT_GRADIENT(fn4); | ||
| TEST_GRADIENT(fn4, /*numOfDerivativeArgs=*/2, 3, 4, &d_i, &d_j); // CHECK-EXEC: {1.00, 0.00} | ||
|
|
||
| INIT_GRADIENT(fn5); | ||
| TEST_GRADIENT(fn5, /*numOfDerivativeArgs=*/2, 3, 4, &d_i, &d_j); // CHECK-EXEC: {7.00, 0.00} | ||
|
|
||
| INIT_GRADIENT(fn6); | ||
| TEST_GRADIENT(fn6, /*numOfDerivativeArgs=*/2, 3, 4, &d_i, &d_j); // CHECK-EXEC: {24.00, 9.00} | ||
|
|
||
| INIT_GRADIENT(fn7); | ||
| TEST_GRADIENT(fn7, /*numOfDerivativeArgs=*/2, 2, 9, &d_i, &d_j); // CHECK-EXEC: {12.00, 0.00} | ||
|
|
||
| INIT_GRADIENT(fn8); | ||
| TEST_GRADIENT(fn8, /*numOfDerivativeArgs=*/2, 7, 2, &d_i, &d_j); // CHECK-EXEC: {1.00, 1.00} | ||
| } | ||
Uh oh!
There was an error while loading. Please reload this page.