Skip to content

Commit 60831dc

Browse files
PetroZarytskyivgvassilev
authored andcommitted
Replace iterator-based std::move with clad::move
We use the iterator-based `std::move` to perform assignments of whole arrays in the reverse mode. This PR reimplements it with `clad::move` to make it compatible with CUDA. Fixes #1496
1 parent 26e2416 commit 60831dc

4 files changed

Lines changed: 38 additions & 39 deletions

File tree

include/clad/Differentiator/Differentiator.h

Lines changed: 17 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -197,6 +197,23 @@ CUDA_HOST_DEVICE void push(tape<T[N], SBO_SIZE, SLAB_SIZE>& to, const U& val) {
197197
for (std::size_t i = 0; i < N; ++i)
198198
zero_init(x[i]);
199199
}
200+
201+
// This function is similar to the iterator-based std::move but is designed to
202+
// work with CUDA.
203+
// NOLINTNEXTLINE(cppcoreguidelines-rvalue-reference-param-not-moved)
204+
template <class T, size_t N>
205+
CUDA_HOST_DEVICE void move(T (&&Input)[N], T* Output) {
206+
for (T& elem : Input) {
207+
*Output = std::forward<T>(elem);
208+
++Output;
209+
}
210+
}
211+
212+
// We cannot use forwarding references
213+
template <class T, size_t N>
214+
CUDA_HOST_DEVICE void move(T (&Input)[N], T* Output) {
215+
move(std::move(Input), Output);
216+
}
200217
// NOLINTEND(cppcoreguidelines-avoid-c-arrays)
201218

202219
/// Pad the args supplied with nullptr(s) or zeros to match the the num of

lib/Differentiator/ReverseModeVisitor.cpp

Lines changed: 3 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -3390,25 +3390,13 @@ Expr* ReverseModeVisitor::getStdInitListSizeExpr(const Expr* E) {
33903390

33913391
Expr* ReverseModeVisitor::BuildArrayAssignment(Expr* output, Expr* input,
33923392
direction d) {
3393-
// Build `type (&&_t0)[N] = input;`
3394-
QualType storeTy = input->getType();
3395-
if (input->isLValue())
3396-
storeTy = m_Context.getLValueReferenceType(storeTy);
3397-
else
3398-
storeTy = m_Context.getRValueReferenceType(storeTy);
3399-
input = StoreAndRef(input, storeTy, d);
3400-
llvm::SmallVector<Expr*, 1> argIn = {input};
3401-
// Build `std::begin(_t0)` and `std::end(_t0)`
3402-
Expr* beginIn = GetFunctionCall("begin", "std", argIn);
3403-
Expr* endIn = GetFunctionCall("end", "std", argIn);
3404-
34053393
// Build `std::begin(_output)`
34063394
llvm::SmallVector<Expr*, 1> argOut = {output};
34073395
Expr* beginOut = GetFunctionCall("begin", "std", argOut);
34083396

3409-
// Build `std::move(std::begin(_t0), std::end(_t0), std::begin(_output));`
3410-
llvm::SmallVector<Expr*, 1> moveArgs = {beginIn, endIn, beginOut};
3411-
return GetFunctionCall("move", "std", moveArgs);
3397+
// Build `clad::move(_input, std::begin(_output));`
3398+
llvm::SmallVector<Expr*, 1> moveArgs = {input, beginOut};
3399+
return GetFunctionCall("move", "clad", moveArgs);
34123400
}
34133401

34143402
Expr* ReverseModeVisitor::GlobalStoreAndRef(Expr* E, llvm::StringRef prefix,

test/Arrays/ArrayInputsReverseMode.C

Lines changed: 13 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -254,25 +254,24 @@ double func6(double seed) {
254254
//CHECK: void func6_grad(double seed, double *_d_seed) {
255255
//CHECK-NEXT: int _d_i = 0;
256256
//CHECK-NEXT: int i = 0;
257-
//CHECK-NEXT: clad::tape<double{{ ?}}[3]> _t2 = {};
257+
//CHECK-NEXT: clad::tape<double{{ ?}}[3]> _t1 = {};
258258
//CHECK-NEXT: double _d_arr[3] = {0};
259259
//CHECK-NEXT: double arr[3] = {0};
260260
//CHECK-NEXT: double _d_sum = 0.;
261261
//CHECK-NEXT: double sum = 0;
262262
//CHECK-NEXT: unsigned {{int|long|long long}} _t0 = 0;
263263
//CHECK-NEXT: for (i = 0; i < 3; i++) {
264264
//CHECK-NEXT: _t0++;
265-
//CHECK-NEXT: double (&&_t1)[3] = {seed, seed * i, seed + i};
266-
//CHECK-NEXT: clad::push(_t2, arr) , std::move(std::begin(_t1), std::end(_t1), std::begin(arr));
265+
//CHECK-NEXT: clad::push(_t1, arr) , clad::move({seed, seed * i, seed + i}, std::begin(arr));
267266
//CHECK-NEXT: sum += addArr(arr, 3);
268267
//CHECK-NEXT: }
269268
//CHECK-NEXT: _d_sum += 1;
270269
//CHECK-NEXT: for (; _t0; _t0--) {
271270
//CHECK-NEXT: i--;
272271
//CHECK-NEXT: {
273272
//CHECK-NEXT: double _r_d0 = _d_sum;
274-
//CHECK-NEXT: int _r1 = 0;
275-
//CHECK-NEXT: addArr_pullback(arr, 3, _r_d0, _d_arr, &_r1);
273+
//CHECK-NEXT: int _r0 = 0;
274+
//CHECK-NEXT: addArr_pullback(arr, 3, _r_d0, _d_arr, &_r0);
276275
//CHECK-NEXT: }
277276
//CHECK-NEXT: {
278277
//CHECK-NEXT: *_d_seed += _d_arr[0];
@@ -281,9 +280,8 @@ double func6(double seed) {
281280
//CHECK-NEXT: *_d_seed += _d_arr[2];
282281
//CHECK-NEXT: _d_i += _d_arr[2];
283282
//CHECK-NEXT: clad::zero_init(_d_arr);
284-
//CHECK-NEXT: double &_r0[3] = clad::back(_t2);
285-
//CHECK-NEXT: std::move(std::begin(_r0), std::end(_r0), std::begin(arr));
286-
//CHECK-NEXT: clad::pop(_t2);
283+
//CHECK-NEXT: clad::move(clad::back(_t1), std::begin(arr));
284+
//CHECK-NEXT: clad::pop(_t1);
287285
//CHECK-NEXT: }
288286
//CHECK-NEXT: }
289287
//CHECK-NEXT: }
@@ -318,14 +316,13 @@ double func7(double *params) {
318316
//CHECK-NEXT: double _d_out = 0.;
319317
//CHECK-NEXT: double out = 0.;
320318
//CHECK-NEXT: unsigned {{int|long|long long}} _t0 = 0;
321-
// CHECK-NEXT: for (i = 0; i < 1; ++i) {
322-
// CHECK-NEXT: _t0++;
323-
// CHECK-NEXT: double (&&_t1)[1] = {params[0]};
324-
// CHECK-NEXT: std::move(std::begin(_t1), std::end(_t1), std::begin(paramsPrime));
325-
// CHECK-NEXT: out = out + inv_square(paramsPrime);
326-
// CHECK-NEXT: }
327-
// CHECK-NEXT: _d_out += 1;
328-
// CHECK-NEXT: for (; _t0; _t0--) {
319+
//CHECK-NEXT: for (i = 0; i < 1; ++i) {
320+
//CHECK-NEXT: _t0++;
321+
//CHECK-NEXT: clad::move({params[0]}, std::begin(paramsPrime));
322+
//CHECK-NEXT: out = out + inv_square(paramsPrime);
323+
//CHECK-NEXT: }
324+
//CHECK-NEXT: _d_out += 1;
325+
//CHECK-NEXT: for (; _t0; _t0--) {
329326
//CHECK-NEXT: {
330327
//CHECK-NEXT: double _r_d0 = _d_out;
331328
//CHECK-NEXT: _d_out = 0.;

test/Gradient/Loops.C

Lines changed: 5 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -1384,8 +1384,7 @@ double fn21(double x) {
13841384
// CHECK-NEXT: unsigned {{int|long|long long}} _t0 = 0;
13851385
// CHECK-NEXT: for (i = 0; i < 5; ++i) {
13861386
// CHECK-NEXT: _t0++;
1387-
// CHECK-NEXT: double (&&_t1)[3] = {1, x, 2};
1388-
// CHECK-NEXT: std::move(std::begin(_t1), std::end(_t1), std::begin(arr));
1387+
// CHECK-NEXT: clad::move({1, x, 2}, std::begin(arr));
13891388
// CHECK-NEXT: res += arr[0] + arr[1];
13901389
// CHECK-NEXT: }
13911390
// CHECK-NEXT: _d_res += 1;
@@ -1415,16 +1414,15 @@ double fn22(double param) {
14151414
// CHECK: void fn22_grad(double param, double *_d_param) {
14161415
// CHECK-NEXT: int _d_i = 0;
14171416
// CHECK-NEXT: int i = 0;
1418-
// CHECK-NEXT: clad::tape<double{{ ?}}[1]> _t2 = {};
1417+
// CHECK-NEXT: clad::tape<double{{ ?}}[1]> _t1 = {};
14191418
// CHECK-NEXT: double _d_arr[1] = {0};
14201419
// CHECK-NEXT: double arr[1] = {0};
14211420
// CHECK-NEXT: double _d_out = 0.;
14221421
// CHECK-NEXT: double out = 0.;
14231422
// CHECK-NEXT: unsigned {{int|long|long long}} _t0 = 0;
14241423
// CHECK-NEXT: for (i = 0; i < 1; i++) {
14251424
// CHECK-NEXT: _t0++;
1426-
// CHECK-NEXT: double (&&_t1)[1] = {1.};
1427-
// CHECK-NEXT: clad::push(_t2, arr) , std::move(std::begin(_t1), std::end(_t1), std::begin(arr));
1425+
// CHECK-NEXT: clad::push(_t1, arr) , clad::move({1.}, std::begin(arr));
14281426
// CHECK-NEXT: out += arr[0] * param;
14291427
// CHECK-NEXT: }
14301428
// CHECK-NEXT: _d_out += 1;
@@ -1436,9 +1434,8 @@ double fn22(double param) {
14361434
// CHECK-NEXT: }
14371435
// CHECK-NEXT: {
14381436
// CHECK-NEXT: clad::zero_init(_d_arr);
1439-
// CHECK-NEXT: double &_r0[1] = clad::back(_t2);
1440-
// CHECK-NEXT: std::move(std::begin(_r0), std::end(_r0), std::begin(arr));
1441-
// CHECK-NEXT: clad::pop(_t2);
1437+
// CHECK-NEXT: clad::move(clad::back(_t1), std::begin(arr));
1438+
// CHECK-NEXT: clad::pop(_t1);
14421439
// CHECK-NEXT: }
14431440
// CHECK-NEXT: }
14441441
// CHECK-NEXT: }

0 commit comments

Comments
 (0)