Skip to content

Commit d0e1e2d

Browse files
fogsong233vgvassilev
authored andcommitted
Clone CXXBindTemporaryExpr nodes
1 parent 938d9c3 commit d0e1e2d

3 files changed

Lines changed: 50 additions & 0 deletions

File tree

include/clad/Differentiator/StmtClone.h

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -120,6 +120,7 @@ namespace utils {
120120
DECLARE_CLONE_FN(CXXThrowExpr)
121121
DECLARE_CLONE_FN(CXXConstructExpr)
122122
DECLARE_CLONE_FN(CXXTemporaryObjectExpr)
123+
DECLARE_CLONE_FN(CXXBindTemporaryExpr)
123124
DECLARE_CLONE_FN(MaterializeTemporaryExpr)
124125
DECLARE_CLONE_FN(PseudoObjectExpr)
125126
DECLARE_CLONE_FN(OpaqueValueExpr)

lib/Differentiator/StmtClone.cpp

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -221,6 +221,9 @@ Stmt* StmtClone::VisitCXXTemporaryObjectExpr(CXXTemporaryObjectExpr* Node) {
221221
return result;
222222
}
223223

224+
DEFINE_CREATE_EXPR(CXXBindTemporaryExpr,
225+
(Ctx, Node->getTemporary(), Clone(Node->getSubExpr())))
226+
224227
DEFINE_CLONE_EXPR(MaterializeTemporaryExpr,
225228
(CloneType(Node->getType()),
226229
Node->getSubExpr() ? Clone(Node->getSubExpr()) : nullptr,
Lines changed: 46 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,46 @@
1+
// RUN: %cladclang %s -I%S/../../include -o %t 2>&1 | %filecheck %s
2+
// RUN: %t | %filecheck_exec %s
3+
4+
#include "clad/Differentiator/Differentiator.h"
5+
#include "clad/Differentiator/STLBuiltins.h"
6+
7+
#include <cstdio>
8+
#include <vector>
9+
10+
struct Box {
11+
double value;
12+
// A non-trivial destructor makes Clang bind prvalues of this type in a
13+
// CXXBindTemporaryExpr.
14+
~Box() {}
15+
16+
double get() const { return value; }
17+
};
18+
19+
Box make_box(double x) { return {x * x}; }
20+
21+
double temporary_method(double x) { return make_box(x).get(); }
22+
23+
double vector_temporary_method(double x) {
24+
return x + static_cast<double>(std::vector<double>().size());
25+
}
26+
27+
// CHECK: void temporary_method_grad(double x, double *_d_x) {
28+
// CHECK: make_box(x).get_pullback(1, &{{.*}});
29+
// CHECK: make_box_pullback(x, {{.*}}, &{{.*}});
30+
// CHECK: }
31+
32+
// CHECK: void vector_temporary_method_grad(double x, double *_d_x) {
33+
// CHECK: *_d_x += 1;
34+
// CHECK: }
35+
36+
int main() {
37+
auto temporary_method_grad = clad::gradient(temporary_method);
38+
double d_x = 0.0;
39+
temporary_method_grad.execute(3.0, &d_x);
40+
std::printf("box=%.1f\n", d_x); // CHECK-EXEC: box=6.0
41+
42+
auto vector_temporary_method_grad = clad::gradient(vector_temporary_method);
43+
d_x = 0.0;
44+
vector_temporary_method_grad.execute(3.0, &d_x);
45+
std::printf("vector=%.1f\n", d_x); // CHECK-EXEC: vector=1.0
46+
}

0 commit comments

Comments
 (0)