Skip to content

Commit 9f62f71

Browse files
PetroZarytskyivgvassilev
authored andcommitted
Support std::min/max with custom comparators
1 parent 1a94ffa commit 9f62f71

3 files changed

Lines changed: 83 additions & 6 deletions

File tree

include/clad/Differentiator/BuiltinDerivatives.h

Lines changed: 33 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -12,6 +12,7 @@
1212

1313
#include <algorithm>
1414
#include <cmath>
15+
#include <functional>
1516

1617
#define elidable_reverse_forw __attribute__((annotate("elidable_reverse_forw")))
1718

@@ -962,36 +963,62 @@ CUDA_HOST_DEVICE void pow_pullback(T1 x, T2 exponent, T3 d_y, T1* d_x,
962963
*d_exponent += t.pushforward * d_y;
963964
}
964965

966+
template <typename T, class Compare>
967+
CUDA_HOST_DEVICE ValueAndPushforward<const T&, const T&>
968+
min_pushforward(const T& a, const T& b, Compare comp, const T& d_a,
969+
const T& d_b, Compare /*dcomp*/) {
970+
return {::std::min(a, b, comp), comp(a, b) ? d_a : d_b};
971+
}
972+
965973
template <typename T>
966974
CUDA_HOST_DEVICE ValueAndPushforward<const T&, const T&>
967975
min_pushforward(const T& a, const T& b, const T& d_a, const T& d_b) {
968976
return {::std::min(a, b), a < b ? d_a : d_b};
969977
}
970978

979+
template <typename T, class Compare>
980+
CUDA_HOST_DEVICE ValueAndPushforward<const T&, const T&>
981+
max_pushforward(const T& a, const T& b, Compare comp, const T& d_a,
982+
const T& d_b, Compare /*dcomp*/) {
983+
return {::std::max(a, b, comp), comp(a, b) ? d_b : d_a};
984+
}
985+
971986
template <typename T>
972987
CUDA_HOST_DEVICE ValueAndPushforward<const T&, const T&>
973988
max_pushforward(const T& a, const T& b, const T& d_a, const T& d_b) {
974989
return {::std::max(a, b), a < b ? d_b : d_a};
975990
}
976991

977-
template <typename T, typename U>
978-
CUDA_HOST_DEVICE void min_pullback(const T& a, const T& b, U d_y, T* d_a,
979-
T* d_b) {
980-
if (a < b)
992+
template <typename T, typename U, class Compare>
993+
CUDA_HOST_DEVICE void min_pullback(const T& a, const T& b, Compare comp, U d_y,
994+
T* d_a, T* d_b, Compare* /*dcomp*/) {
995+
if (comp(a, b))
981996
*d_a += d_y;
982997
else
983998
*d_b += d_y;
984999
}
9851000

9861001
template <typename T, typename U>
987-
CUDA_HOST_DEVICE void max_pullback(const T& a, const T& b, U d_y, T* d_a,
1002+
CUDA_HOST_DEVICE void min_pullback(const T& a, const T& b, U d_y, T* d_a,
9881003
T* d_b) {
989-
if (a < b)
1004+
min_pullback(a, b, ::std::less<T>{}, d_y, d_a, d_b, (::std::less<T>*)nullptr);
1005+
}
1006+
1007+
template <typename T, typename U, class Compare>
1008+
CUDA_HOST_DEVICE void max_pullback(const T& a, const T& b, Compare comp, U d_y,
1009+
T* d_a, T* d_b, Compare* /*dcomp*/) {
1010+
if (comp(a, b))
9901011
*d_b += d_y;
9911012
else
9921013
*d_a += d_y;
9931014
}
9941015

1016+
template <typename T, typename U>
1017+
CUDA_HOST_DEVICE void max_pullback(const T& a, const T& b, U d_y, T* d_a,
1018+
T* d_b) {
1019+
max_pullback(a, b, ::std::less<T>{}, d_y, d_a, d_b, (::std::less<T>*)nullptr);
1020+
}
1021+
9951022
#if __cplusplus >= 201703L
9961023
template <typename T>
9971024
CUDA_HOST_DEVICE ValueAndPushforward<T, T>

test/FirstDerivative/BuiltinDerivatives.C

Lines changed: 22 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -499,6 +499,22 @@ double f_pow_zero(double x, double y) { return std::pow(x, y); }
499499
// CHECK-NEXT: return _t0.pushforward;
500500
// CHECK-NEXT: }
501501

502+
double f_custom_max(double x, double y) { return std::max(x, y, std::greater<double>()); }
503+
// CHECK: double f_custom_max_darg0(double x, double y) {
504+
// CHECK-NEXT: double _d_x = 1;
505+
// CHECK-NEXT: double _d_y = 0;
506+
// CHECK-NEXT: ValueAndPushforward<const double &, const double &> _t0 = clad::custom_derivatives::std::max_pushforward(x, y, std::greater<double>(), _d_x, _d_y, std::greater<double>());
507+
// CHECK-NEXT: return _t0.pushforward;
508+
// CHECK-NEXT: }
509+
510+
double f_custom_min(double x, double y) { return std::min(x, y, std::greater<double>()); }
511+
// CHECK: double f_custom_min_darg0(double x, double y) {
512+
// CHECK-NEXT: double _d_x = 1;
513+
// CHECK-NEXT: double _d_y = 0;
514+
// CHECK-NEXT: ValueAndPushforward<const double &, const double &> _t0 = clad::custom_derivatives::std::min_pushforward(x, y, std::greater<double>(), _d_x, _d_y, std::greater<double>());
515+
// CHECK-NEXT: return _t0.pushforward;
516+
// CHECK-NEXT: }
517+
502518
int main () { //expected-no-diagnostics
503519
float f_result[2];
504520
double d_result[2];
@@ -673,5 +689,11 @@ int main () { //expected-no-diagnostics
673689
auto d_pow_zero = clad::differentiate(f_pow_zero, 0);
674690
printf("Result is = %.6f\n", d_pow_zero.execute(0, 0)); // CHECK-EXEC: Result is = 0.000000
675691

692+
auto d_custom_max = clad::differentiate(f_custom_max, 0);
693+
printf("Result is = %.6f\n", d_custom_max.execute(1, 2)); // CHECK-EXEC: Result is = 1.000000
694+
695+
auto d_custom_min = clad::differentiate(f_custom_min, 0);
696+
printf("Result is = %.6f\n", d_custom_min.execute(2, 3)); // CHECK-EXEC: Result is = 0.000000
697+
676698
return 0;
677699
}

test/Gradient/FunctionCalls.C

Lines changed: 28 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1254,6 +1254,31 @@ double fn34(double x, double y) {
12541254
// CHECK-NEXT: }
12551255
// CHECK-NEXT: }
12561256

1257+
double fn35(double x, double y) {
1258+
double val_max = std::max(x, y, std::greater<double>()); // min(x, y)
1259+
double val_min = std::min(x, y, std::greater<double>()); // max(x, y)
1260+
return 2 * val_max - 3 * val_min;
1261+
}
1262+
1263+
// CHECK: void fn35_grad(double x, double y, double *_d_x, double *_d_y) {
1264+
// CHECK-NEXT: double _d_val_max = 0.;
1265+
// CHECK-NEXT: double val_max = std::max(x, y, std::greater<double>());
1266+
// CHECK-NEXT: double _d_val_min = 0.;
1267+
// CHECK-NEXT: double val_min = std::min(x, y, std::greater<double>());
1268+
// CHECK-NEXT: {
1269+
// CHECK-NEXT: _d_val_max += 2 * 1;
1270+
// CHECK-NEXT: _d_val_min += 3 * -1;
1271+
// CHECK-NEXT: }
1272+
// CHECK-NEXT: {
1273+
// CHECK-NEXT: std::greater<double> _r1 = {};
1274+
// CHECK-NEXT: clad::custom_derivatives::std::min_pullback(x, y, std::greater<double>(), _d_val_min, _d_x, _d_y, &_r1);
1275+
// CHECK-NEXT: }
1276+
// CHECK-NEXT: {
1277+
// CHECK-NEXT: std::greater<double> _r0 = {};
1278+
// CHECK-NEXT: clad::custom_derivatives::std::max_pullback(x, y, std::greater<double>(), _d_val_max, _d_x, _d_y, &_r0);
1279+
// CHECK-NEXT: }
1280+
// CHECK-NEXT: }
1281+
12571282
template<typename T>
12581283
void reset(T* arr, int n) {
12591284
for (int i=0; i<n; ++i)
@@ -1414,6 +1439,9 @@ int main() {
14141439

14151440
INIT(fn34);
14161441
TEST2(fn34, 2, 1); // CHECK-EXEC: {1.00, 1.00}
1442+
1443+
INIT(fn35);
1444+
TEST2(fn35, 10, 1); // CHECK-EXEC: {-3.00, 2.00}
14171445
}
14181446

14191447
double sq_defined_later(double x) {

0 commit comments

Comments
 (0)