Skip to content

Commit 773992b

Browse files
authored
Merge pull request #136 from gjkennedy/gpu-fix
modified expressions to use the A2D-specific implementations of unary…
2 parents c7a0a6f + 536f883 commit 773992b

5 files changed

Lines changed: 265 additions & 247 deletions

File tree

include/ad/a2dbinary.h

Lines changed: 24 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -182,7 +182,7 @@ A2D_1ST_BINARY_BASIC(MultExpr, operator*, a.value() * b.value(),
182182
// AREVBODY, BREVBODY)
183183
A2D_1ST_BINARY(Divide, operator/, a.value() / b.value(), T(1.0) / b.value(),
184184
tmp*(a.bvalue() - tmp * a.value() * b.bvalue()), tmp* bval,
185-
-tmp* tmp* a.value() * bval)
185+
-tmp * tmp * a.value() * bval)
186186
A2D_1ST_BINARY(Max, max2,
187187
(RealPart(a.value()) > RealPart(b.value()) ? a.value()
188188
: b.value()),
@@ -196,7 +196,7 @@ A2D_1ST_BINARY(Min, min2,
196196
tmp* a.bvalue() + (1.0 - tmp) * b.value(), tmp* bval,
197197
(1.0 - tmp) * bval)
198198

199-
A2D_1ST_BINARY(Atan2, atan2, atan2(a.value(), b.value()),
199+
A2D_1ST_BINARY(Atan2, atan2, A2D::atan2(a.value(), b.value()),
200200
T(1.0) / (a.value() * a.value() + b.value() * b.value()),
201201
tmp*(b.value() * a.bvalue() - a.value() * b.bvalue()),
202202
b.value() * tmp * bval, -a.value() * tmp * bval)
@@ -404,11 +404,12 @@ A2D_2ND_BINARY_BASIC(MultExpr2, operator*, a.value() * b.value(),
404404
// A2D_2ND_BINARY(OBJNAME, OPERNAME, FUNCBODY, TEMPBODY, AREVBODY,
405405
// BREVBODY, HFORWARDBODY, HAREVBODY, HBREVBODY)
406406
A2D_2ND_BINARY(Divide2, operator/, a.value() / b.value(), T(1.0) / b.value(),
407-
tmp* bval, -tmp* tmp* a.value() * bval,
407+
tmp* bval, -tmp * tmp * a.value() * bval,
408408
tmp*(a.pvalue() - tmp * a.value() * b.pvalue()),
409409
tmp*(hval - tmp * bval * b.pvalue()),
410-
tmp* tmp * (2.0 * tmp * a.value() * bval * b.pvalue() -
411-
a.value() * hval - bval * a.pvalue()))
410+
tmp * tmp *
411+
(2.0 * tmp * a.value() * bval * b.pvalue() -
412+
a.value() * hval - bval * a.pvalue()))
412413
A2D_2ND_BINARY(Max2, max2,
413414
(RealPart(a.value()) > RealPart(b.value()) ? a.value()
414415
: b.value()),
@@ -424,7 +425,7 @@ A2D_2ND_BINARY(Min2, min2,
424425
tmp* a.bvalue() + (1.0 - tmp) * b.value(), tmp* hval,
425426
(1.0 - tmp) * hval)
426427
// atan2(y, x): a=y, b=x, tmp = 1/(x^2 + y^2)
427-
A2D_2ND_BINARY(Atan22, atan2, atan2(a.value(), b.value()),
428+
A2D_2ND_BINARY(Atan22, atan2, A2D::atan2(a.value(), b.value()),
428429
T(1.0) / (a.value() * a.value() + b.value() * b.value()),
429430
b.value() * tmp * bval, -a.value() * tmp * bval,
430431
tmp*(b.value() * a.pvalue() - a.value() * b.pvalue()),
@@ -501,10 +502,12 @@ A2D_1ST_BINARY_LEFT_BASIC(LMultExpr, operator*, a.value() * b, a.bvalue() * b,
501502
b* bval)
502503
A2D_1ST_BINARY_LEFT_BASIC(LDivide, operator/, a.value() / b, a.bvalue() / b,
503504
bval / b)
504-
A2D_1ST_BINARY_LEFT_BASIC(PowExpr, pow, pow(a.value(), b),
505-
a.bvalue() * b * pow(a.value(), b - 1.0),
506-
b* pow(a.value(), b - 1.0) * bval)
507-
505+
// A2D_1ST_BINARY_LEFT_BASIC(PowExpr, pow, pow(a.value(), b),
506+
// a.bvalue() * b * pow(a.value(), b - 1.0),
507+
// b* pow(a.value(), b - 1.0) * bval)
508+
A2D_1ST_BINARY_LEFT_BASIC(PowExpr, pow, A2D::pow(a.value(), b),
509+
a.bvalue() * b * A2D::pow(a.value(), b - T(1.0)),
510+
b* A2D::pow(a.value(), b - T(1.0)) * bval)
508511
/*
509512
Definitions for memory-less forward and reverse-mode first-order AD
510513
@@ -587,12 +590,18 @@ A2D_2ND_BINARY_LEFT_BASIC(LMultExpr2, operator*, a.value() * b, b* bval,
587590
b* a.pvalue(), b* hval)
588591
A2D_2ND_BINARY_LEFT_BASIC(LDivide2, operator/, a.value() / b, bval / b,
589592
a.pvalue() / b, hval / b)
590-
A2D_2ND_BINARY_LEFT_BASIC(PowExpr2, pow, pow(a.value(), b),
591-
bval* b* pow(a.value(), b - 1.0),
592-
a.pvalue() * b * pow(a.value(), b - 1.0),
593-
hval* b* pow(a.value(), b - 1.0) +
593+
// A2D_2ND_BINARY_LEFT_BASIC(PowExpr2, pow, pow(a.value(), b),
594+
// bval * b * pow(a.value(), b - 1.0),
595+
// a.pvalue() * b * pow(a.value(), b - 1.0),
596+
// hval * b * pow(a.value(), b - 1.0) +
597+
// bval * a.pvalue() * b * (b - 1.0) *
598+
// pow(a.value(), b - 2.0));
599+
A2D_2ND_BINARY_LEFT_BASIC(PowExpr2, pow, A2D::pow(a.value(), b),
600+
bval * b * A2D::pow(a.value(), b - 1.0),
601+
a.pvalue() * b * A2D::pow(a.value(), b - 1.0),
602+
hval * b * A2D::pow(a.value(), b - 1.0) +
594603
bval * a.pvalue() * b * (b - 1.0) *
595-
pow(a.value(), b - 2.0));
604+
A2D::pow(a.value(), b - 2.0));
596605

597606
/*
598607
Definitions for memory-less forward and reverse-mode first-order AD

include/ad/a2dscalarops.h

Lines changed: 0 additions & 77 deletions
Original file line numberDiff line numberDiff line change
@@ -172,83 +172,6 @@ A2D_FUNCTION auto Eval(Expr&& expr, A2DObj<T&> out) {
172172
return EvalExprRef2<Expr, T>(a2d_forward<Expr>(expr), out);
173173
}
174174

175-
namespace Test {
176-
template <typename T>
177-
class ScalarTest : public A2DTest<T, T, T, T> {
178-
public:
179-
using Input = VarTuple<T, T, T>;
180-
using Output = VarTuple<T, T>;
181-
182-
// Assemble a string to describe the test
183-
std::string name() { return std::string("ScalarOperations"); }
184-
185-
// Evaluate the matrix-matrix product
186-
Output eval(const Input& x) {
187-
T a, b, f;
188-
x.get_values(a, b);
189-
190-
f = log(a * a * sqrt(exp(a * sin(a) + 3.0 * a)) + 2.0 * a * a * a * a) +
191-
max2(a, min2(a * b, b * b)) - 4.0 * a / b + pow(5.0 / (b * b), 2.0) +
192-
acos(a * 0.1);
193-
194-
return MakeVarTuple<T>(f);
195-
}
196-
197-
// Compute the derivative
198-
void deriv(const Output& seed, const Input& x, Input& g) {
199-
T a0, ab, b0, bb;
200-
ADObj<T&> a(a0, ab), b(b0, bb);
201-
ADObj<T> f;
202-
x.get_values(a.value(), b.value());
203-
204-
auto stack = MakeStack(Eval(
205-
log(a * a * sqrt(exp(a * sin(a) + 3.0 * a)) + 2.0 * a * a * a * a) +
206-
max2(a, min2(a * b, b * b)) - 4.0 * a / b +
207-
pow(5.0 / (b * b), 2.0) + acos(a * 0.1),
208-
f));
209-
210-
seed.get_values(f.bvalue());
211-
stack.reverse();
212-
g.set_values(a.bvalue(), b.bvalue());
213-
}
214-
215-
// Compute the second-derivative
216-
void hprod(const Output& seed, const Output& hval, const Input& x,
217-
const Input& p, Input& h) {
218-
T a0, ab, ap, ah, b0, bb, bp, bh;
219-
A2DObj<T&> a(a0, ab, ap, ah), b(b0, bb, bp, bh);
220-
A2DObj<T> f;
221-
x.get_values(a.value(), b.value());
222-
p.get_values(a.pvalue(), b.pvalue());
223-
224-
auto stack = MakeStack(Eval(
225-
log(a * a * sqrt(exp(a * sin(a) + 3.0 * a)) + 2.0 * a * a * a * a) +
226-
max2(a, min2(a * b, b * b)) - 4.0 * a / b +
227-
pow(5.0 / (b * b), 2.0) + acos(a * 0.1),
228-
f));
229-
230-
seed.get_values(f.bvalue());
231-
hval.get_values(f.hvalue());
232-
stack.hproduct();
233-
h.set_values(a.hvalue(), b.hvalue());
234-
}
235-
};
236-
237-
inline bool ScalarTestAll(bool component, bool write_output) {
238-
using Tc = A2D_complex_t<double>;
239-
240-
bool passed = true;
241-
ScalarTest<Tc> test1;
242-
test1.set_step_size(1e-8); // inverse trigonometric functions may suffer from
243-
// subtraction cancellation even for complex step
244-
// with certain underlying implementation
245-
passed = passed && Run(test1, component, write_output);
246-
247-
return passed;
248-
}
249-
250-
} // namespace Test
251-
252175
} // namespace A2D
253176

254177
#endif // A2D_SCALAR_OPS_H

0 commit comments

Comments
 (0)