Skip to content

Commit e4f1925

Browse files
leetcodezvgvassilev
authored andcommitted
Fix TBRAnalyzer visiting the true branch twice. It ignored the false branch completely, causing silent wrong derivatives. Fixed the typo, added a regression test, and updated the assignments test.
1 parent aa1bb82 commit e4f1925

3 files changed

Lines changed: 38 additions & 6 deletions

File tree

lib/Differentiator/TBRAnalyzer.cpp

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -247,7 +247,7 @@ bool TBRAnalyzer::TraverseConditionalOperator(clang::ConditionalOperator* CO) {
247247

248248
auto thenBranch = std::move(m_BlockData[m_CurBlockID]);
249249
m_BlockData[m_CurBlockID] = std::move(elseBranch);
250-
TraverseStmt(CO->getTrueExpr());
250+
TraverseStmt(CO->getFalseExpr());
251251

252252
merge(m_BlockData[m_CurBlockID].get(), thenBranch.get());
253253
return false;

test/Gradient/Assignments.C

Lines changed: 9 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -416,26 +416,30 @@ double f12(double x, double y) {
416416

417417
//CHECK: void f12_grad(double x, double y, double *_d_x, double *_d_y) {
418418
//CHECK-NEXT: double _t0;
419+
//CHECK-NEXT: double _t1;
419420
//CHECK-NEXT: double _d_t = 0.;
420421
//CHECK-NEXT: double t;
421422
//CHECK-NEXT: bool _cond0 = x > y;
422423
//CHECK-NEXT: if (_cond0)
423424
//CHECK-NEXT: _t0 = t;
424-
//CHECK-NEXT: double *_t1 = &(_cond0 ? (t = x) : (t = y));
425-
//CHECK-NEXT: double _t2 = *_t1;
426-
//CHECK-NEXT: *_t1 *= y;
425+
//CHECK-NEXT: else
426+
//CHECK-NEXT: _t1 = t;
427+
//CHECK-NEXT: double *_t2 = &(_cond0 ? (t = x) : (t = y));
428+
//CHECK-NEXT: double _t3 = *_t2;
429+
//CHECK-NEXT: *_t2 *= y;
427430
//CHECK-NEXT: _d_t += 1;
428431
//CHECK-NEXT: {
429-
//CHECK-NEXT: *_t1 = _t2;
432+
//CHECK-NEXT: *_t2 = _t3;
430433
//CHECK-NEXT: double _r_d0 = (_cond0 ? _d_t : _d_t);
431434
//CHECK-NEXT: (_cond0 ? _d_t : _d_t) = 0.;
432435
//CHECK-NEXT: (_cond0 ? _d_t : _d_t) += _r_d0 * y;
433-
//CHECK-NEXT: *_d_y += *_t1 * _r_d0;
436+
//CHECK-NEXT: *_d_y += *_t2 * _r_d0;
434437
//CHECK-NEXT: if (_cond0) {
435438
//CHECK-NEXT: t = _t0;
436439
//CHECK-NEXT: *_d_x += _d_t;
437440
//CHECK-NEXT: _d_t = 0.;
438441
//CHECK-NEXT: } else {
442+
//CHECK-NEXT: t = _t1;
439443
//CHECK-NEXT: *_d_y += _d_t;
440444
//CHECK-NEXT: _d_t = 0.;
441445
//CHECK-NEXT: }

test/Gradient/TBRTernary.C

Lines changed: 28 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,28 @@
1+
// RUN: %cladclang %s -I%S/../../include -oReverseMode.out 2>&1 | %filecheck %s
2+
// RUN: ./ReverseMode.out | %filecheck_exec %s
3+
4+
#include "clad/Differentiator/Differentiator.h"
5+
#include <iostream>
6+
7+
double f(bool cond, double a, double b) {
8+
double x = cond ? (a * a) : (b * b);
9+
a = 0;
10+
b = 0;
11+
return x;
12+
}
13+
14+
int main() {
15+
auto df = clad::gradient(f, "a,b");
16+
double da = 0.0, db = 0.0;
17+
18+
df.execute(false, 2.0, 3.0, &da, &db);
19+
std::cout << "da: " << da << ", db: " << db << std::endl;
20+
// CHECK-EXEC: da: 0, db: 6
21+
22+
da = 0.0; db = 0.0;
23+
df.execute(true, 2.0, 3.0, &da, &db);
24+
std::cout << "da: " << da << ", db: " << db << std::endl;
25+
// CHECK-EXEC: da: 4, db: 0
26+
27+
return 0;
28+
}

0 commit comments

Comments
 (0)