Skip to content

Commit ffa3f5f

Browse files
PetroZarytskyivgvassilev
authored andcommitted
Add support for init lists of different length in clad::move
After #1626, we started initializing arrays inside loops with clad::move. However, since the size of the array was deduced from the size of the init list, this change silently dropped support for decls like `double arr[10] = {0}`.
1 parent 9f62f71 commit ffa3f5f

4 files changed

Lines changed: 72 additions & 20 deletions

File tree

include/clad/Differentiator/Differentiator.h

Lines changed: 13 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -24,6 +24,7 @@
2424
#include <cassert>
2525
#include <cstddef>
2626
#include <cstring>
27+
#include <initializer_list>
2728
#include <iterator>
2829
#include <type_traits>
2930
#include <utility>
@@ -201,18 +202,22 @@ CUDA_HOST_DEVICE void push(tape<T[N], SBO_SIZE, SLAB_SIZE>& to, const U& val) {
201202
// This function is similar to the iterator-based std::move but is designed to
202203
// work with CUDA.
203204
// NOLINTNEXTLINE(cppcoreguidelines-rvalue-reference-param-not-moved)
205+
// An overload to initialize arrays from buffers, e.g., `clad::move(t0, arr)`
204206
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-
}
207+
CUDA_HOST_DEVICE void move(T* Input, T (&Output)[N]) {
208+
for (size_t i = 0; i < N; ++i)
209+
Output[i] = std::move(Input[i]);
210210
}
211211

212-
// We cannot use forwarding references
212+
// An overload to initialize arrays with init lists, e.g., `clad::move({1, 2},
213+
// arr)`
213214
template <class T, size_t N>
214-
CUDA_HOST_DEVICE void move(T (&Input)[N], T* Output) {
215-
move(std::move(Input), Output);
215+
CUDA_HOST_DEVICE void move(std::initializer_list<T> Input, T (&Output)[N]) {
216+
size_t i = 0;
217+
for (auto it = Input.begin(); it != Input.end() && i < N; ++it, ++i)
218+
Output[i] = *it;
219+
for (; i < N; ++i)
220+
Output[i] = T();
216221
}
217222
// NOLINTEND(cppcoreguidelines-avoid-c-arrays)
218223

lib/Differentiator/ReverseModeVisitor.cpp

Lines changed: 2 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -3438,12 +3438,8 @@ Expr* ReverseModeVisitor::getStdInitListSizeExpr(const Expr* E) {
34383438

34393439
Expr* ReverseModeVisitor::BuildArrayAssignment(Expr* output, Expr* input,
34403440
direction d) {
3441-
// Build `std::begin(_output)`
3442-
llvm::SmallVector<Expr*, 1> argOut = {output};
3443-
Expr* beginOut = GetFunctionCall("begin", "std", argOut);
3444-
3445-
// Build `clad::move(_input, std::begin(_output));`
3446-
llvm::SmallVector<Expr*, 1> moveArgs = {input, beginOut};
3441+
// Build `clad::move(_input, _output);`
3442+
llvm::SmallVector<Expr*, 1> moveArgs = {input, output};
34473443
return GetFunctionCall("move", "clad", moveArgs);
34483444
}
34493445

test/Arrays/ArrayInputsReverseMode.C

Lines changed: 54 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -262,7 +262,7 @@ double func6(double seed) {
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: clad::push(_t1, arr) , clad::move({seed, seed * i, seed + i}, std::begin(arr));
265+
//CHECK-NEXT: clad::push(_t1, arr) , clad::move({seed, seed * i, seed + i}, arr);
266266
//CHECK-NEXT: sum += addArr(arr, 3);
267267
//CHECK-NEXT: }
268268
//CHECK-NEXT: _d_sum += 1;
@@ -280,7 +280,7 @@ double func6(double seed) {
280280
//CHECK-NEXT: *_d_seed += _d_arr[2];
281281
//CHECK-NEXT: _d_i += _d_arr[2];
282282
//CHECK-NEXT: clad::zero_init(_d_arr);
283-
//CHECK-NEXT: clad::move(clad::back(_t1), std::begin(arr));
283+
//CHECK-NEXT: clad::move(clad::back(_t1), arr);
284284
//CHECK-NEXT: clad::pop(_t1);
285285
//CHECK-NEXT: }
286286
//CHECK-NEXT: }
@@ -318,7 +318,7 @@ double func7(double *params) {
318318
//CHECK-NEXT: unsigned {{int|long|long long}} _t0 = 0;
319319
//CHECK-NEXT: for (i = 0; i < 1; ++i) {
320320
//CHECK-NEXT: _t0++;
321-
//CHECK-NEXT: clad::move({params[0]}, std::begin(paramsPrime));
321+
//CHECK-NEXT: clad::move({params[0]}, paramsPrime);
322322
//CHECK-NEXT: out = out + inv_square(paramsPrime);
323323
//CHECK-NEXT: }
324324
//CHECK-NEXT: _d_out += 1;
@@ -565,6 +565,52 @@ double func12(double x[3], double y[3]) {
565565
//CHECK-NEXT: }
566566
//CHECK-NEXT: }
567567

568+
double func13(double* x, double y) {
569+
double prod = 0;
570+
for (int i = 0; i < 2; ++i) {
571+
double arr[4] = {1. + i, 0., y};
572+
prod += arr[0] * x[0] + arr[1] * x[1] + arr[2] * x[2] + arr[3] * x[3];
573+
}
574+
return prod; // 3 * x[0] + 2 * x[2] * y
575+
}
576+
577+
// CHECK: void func13_grad(double *x, double y, double *_d_x, double *_d_y) {
578+
// CHECK-NEXT: int _d_i = 0;
579+
// CHECK-NEXT: int i = 0;
580+
// CHECK-NEXT: clad::tape<double{{ ?}}[4]> _t1 = {};
581+
// CHECK-NEXT: double _d_arr[4] = {0};
582+
// CHECK-NEXT: double arr[4] = {0};
583+
// CHECK-NEXT: double _d_prod = 0.;
584+
// CHECK-NEXT: double prod = 0;
585+
// CHECK-NEXT: unsigned {{int|long}} _t0 = 0;
586+
// CHECK-NEXT: for (i = 0; i < 2; ++i) {
587+
// CHECK-NEXT: _t0++;
588+
// CHECK-NEXT: clad::push(_t1, arr) , clad::move({1. + i, 0., y}, arr);
589+
// CHECK-NEXT: prod += arr[0] * x[0] + arr[1] * x[1] + arr[2] * x[2] + arr[3] * x[3];
590+
// CHECK-NEXT: }
591+
// CHECK-NEXT: _d_prod += 1;
592+
// CHECK-NEXT: for (; _t0; _t0--) {
593+
// CHECK-NEXT: {
594+
// CHECK-NEXT: double _r_d0 = _d_prod;
595+
// CHECK-NEXT: _d_arr[0] += _r_d0 * x[0];
596+
// CHECK-NEXT: _d_x[0] += arr[0] * _r_d0;
597+
// CHECK-NEXT: _d_arr[1] += _r_d0 * x[1];
598+
// CHECK-NEXT: _d_x[1] += arr[1] * _r_d0;
599+
// CHECK-NEXT: _d_arr[2] += _r_d0 * x[2];
600+
// CHECK-NEXT: _d_x[2] += arr[2] * _r_d0;
601+
// CHECK-NEXT: _d_arr[3] += _r_d0 * x[3];
602+
// CHECK-NEXT: _d_x[3] += arr[3] * _r_d0;
603+
// CHECK-NEXT: }
604+
// CHECK-NEXT: {
605+
// CHECK-NEXT: _d_i += _d_arr[0];
606+
// CHECK-NEXT: *_d_y += _d_arr[2];
607+
// CHECK-NEXT: clad::zero_init(_d_arr);
608+
// CHECK-NEXT: clad::move(clad::back(_t1), arr);
609+
// CHECK-NEXT: clad::pop(_t1);
610+
// CHECK-NEXT: }
611+
// CHECK-NEXT: }
612+
// CHECK-NEXT: }
613+
568614
int main() {
569615
double arr[] = {1, 2, 3};
570616
auto f_dx = clad::gradient(f);
@@ -639,4 +685,9 @@ int main() {
639685
auto func12grad = clad::gradient(func12, "x");
640686
func12grad.execute(x, y, dx);
641687
printf("{%.2f, %.2f, %.2f}\n", dx[0], dx[1], dx[2]); // CHECK-EXEC: {5.00, 6.00, 7.00}
688+
689+
auto func13grad = clad::gradient(func13);
690+
double x1[] = {1, 2, 3, 4}, dx1[4] = {0}, dy = 0;
691+
func13grad.execute(x1, 5, dx1, &dy);
692+
printf("{%.2f, %.2f, %.2f, %.2f, %.2f}\n", dx1[0], dx1[1], dx1[2], dx1[3], dy); // CHECK-EXEC: {3.00, 0.00, 10.00, 0.00, 6.00}
642693
}

test/Gradient/Loops.C

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1384,7 +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: clad::move({1, x, 2}, std::begin(arr));
1387+
// CHECK-NEXT: clad::move({1, x, 2}, arr);
13881388
// CHECK-NEXT: res += arr[0] + arr[1];
13891389
// CHECK-NEXT: }
13901390
// CHECK-NEXT: _d_res += 1;
@@ -1422,7 +1422,7 @@ double fn22(double param) {
14221422
// CHECK-NEXT: unsigned {{int|long|long long}} _t0 = 0;
14231423
// CHECK-NEXT: for (i = 0; i < 1; i++) {
14241424
// CHECK-NEXT: _t0++;
1425-
// CHECK-NEXT: clad::push(_t1, arr) , clad::move({1.}, std::begin(arr));
1425+
// CHECK-NEXT: clad::push(_t1, arr) , clad::move({1.}, arr);
14261426
// CHECK-NEXT: out += arr[0] * param;
14271427
// CHECK-NEXT: }
14281428
// CHECK-NEXT: _d_out += 1;
@@ -1434,7 +1434,7 @@ double fn22(double param) {
14341434
// CHECK-NEXT: }
14351435
// CHECK-NEXT: {
14361436
// CHECK-NEXT: clad::zero_init(_d_arr);
1437-
// CHECK-NEXT: clad::move(clad::back(_t1), std::begin(arr));
1437+
// CHECK-NEXT: clad::move(clad::back(_t1), arr);
14381438
// CHECK-NEXT: clad::pop(_t1);
14391439
// CHECK-NEXT: }
14401440
// CHECK-NEXT: }

0 commit comments

Comments
 (0)