Skip to content

Commit 83f17f2

Browse files
Shubham Shuklavgvassilev
authored andcommitted
Add support for omp critical directive in forward and reverse mode
1 parent 861bf36 commit 83f17f2

6 files changed

Lines changed: 171 additions & 0 deletions

File tree

include/clad/Differentiator/BaseForwardModeVisitor.h

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -131,6 +131,7 @@ class BaseForwardModeVisitor
131131
StmtDiff VisitOMPParallelDirective(const clang::OMPParallelDirective* D);
132132
StmtDiff
133133
VisitOMPParallelForDirective(const clang::OMPParallelForDirective* D);
134+
StmtDiff VisitOMPCriticalDirective(const clang::OMPCriticalDirective* D);
134135
clang::OMPClause* VisitOMPPrivateClause(const clang::OMPPrivateClause* C);
135136
clang::OMPClause*
136137
VisitOMPFirstprivateClause(const clang::OMPFirstprivateClause* C);

include/clad/Differentiator/ReverseModeVisitor.h

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -480,6 +480,7 @@ namespace clad {
480480
VisitOMPExecutableDirective(const clang::OMPExecutableDirective* D);
481481
StmtDiff
482482
VisitOMPParallelForDirective(const clang::OMPParallelForDirective* D);
483+
StmtDiff VisitOMPCriticalDirective(const clang::OMPCriticalDirective* D);
483484

484485
/// Helper function that builds `T* _this = malloc(sifeof(T));`
485486
/// and `free(_this)`.

lib/Differentiator/BaseForwardModeVisitorOpenMP.cpp

Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -152,4 +152,17 @@ StmtDiff BaseForwardModeVisitor::VisitOMPParallelForDirective(
152152
CLAD_COMPAT_CLANG19_SemaOpenMP(m_Sema).EndOpenMPDSABlock(Res.getStmt());
153153
return Res;
154154
}
155+
156+
StmtDiff BaseForwardModeVisitor::VisitOMPCriticalDirective(
157+
const OMPCriticalDirective* D) {
158+
StmtDiff BodyDiff = Visit(D->getAssociatedStmt());
159+
DeclarationNameInfo DirName = D->getDirectiveName();
160+
OpenMPDirectiveKind CancelRegion = OMPD_unknown;
161+
llvm::SmallVector<OMPClause*, 0> Clauses;
162+
return CLAD_COMPAT_CLANG19_SemaOpenMP(m_Sema)
163+
.ActOnOpenMPExecutableDirective(OMPD_critical, DirName, CancelRegion,
164+
Clauses, BodyDiff.getStmt(),
165+
D->getBeginLoc(), D->getEndLoc())
166+
.get();
167+
}
155168
} // namespace clad

lib/Differentiator/ReverseModeVisitorOpenMP.cpp

Lines changed: 24 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -569,4 +569,28 @@ StmtDiff ReverseModeVisitor::VisitOMPParallelForDirective(
569569
CLAD_COMPAT_CLANG19_SemaOpenMP(m_Sema).EndOpenMPDSABlock(SDiff.getStmt());
570570
return SDiff;
571571
}
572+
573+
StmtDiff
574+
ReverseModeVisitor::VisitOMPCriticalDirective(const OMPCriticalDirective* D) {
575+
StmtDiff BodyDiff = Visit(D->getAssociatedStmt());
576+
DeclarationNameInfo DirName = D->getDirectiveName();
577+
OpenMPDirectiveKind CancelRegion = OMPD_unknown;
578+
llvm::SmallVector<OMPClause*, 0> Clauses;
579+
580+
Stmt* ForwardCritical =
581+
CLAD_COMPAT_CLANG19_SemaOpenMP(m_Sema)
582+
.ActOnOpenMPExecutableDirective(OMPD_critical, DirName, CancelRegion,
583+
Clauses, BodyDiff.getStmt(),
584+
D->getBeginLoc(), D->getEndLoc())
585+
.get();
586+
587+
Stmt* ReverseCritical =
588+
CLAD_COMPAT_CLANG19_SemaOpenMP(m_Sema)
589+
.ActOnOpenMPExecutableDirective(OMPD_critical, DirName, CancelRegion,
590+
Clauses, BodyDiff.getStmt_dx(),
591+
D->getBeginLoc(), D->getEndLoc())
592+
.get();
593+
594+
return {ForwardCritical, ReverseCritical};
595+
}
572596
} // namespace clad

test/ForwardMode/OpenMP.C

Lines changed: 30 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -80,6 +80,33 @@ double accumulate_scale_parallel(double scale, int n) {
8080
// CHECK-NEXT: return _d_sum;
8181
// CHECK-NEXT: }
8282

83+
double accumulate_critical(const double *x, int n) {
84+
double result = 0;
85+
#pragma omp parallel for
86+
for (int i = 0; i < n; i++) {
87+
#pragma omp critical
88+
{
89+
result += x[0] * x[0];
90+
}
91+
}
92+
return result;
93+
}
94+
95+
// CHECK: double accumulate_critical_darg0_0(const double *x, int n) {
96+
// CHECK-NEXT: int _d_n = 0;
97+
// CHECK-NEXT: double _d_result = 0;
98+
// CHECK-NEXT: double result = 0;
99+
// CHECK-NEXT: #pragma omp parallel for
100+
// CHECK-NEXT: for (int i = 0; i < n; i++) {
101+
// CHECK-NEXT: #pragma omp critical
102+
// CHECK-NEXT: {
103+
// CHECK-NEXT: _d_result += 1. * x[0] + x[0] * 1.;
104+
// CHECK-NEXT: result += x[0] * x[0];
105+
// CHECK-NEXT: }
106+
// CHECK-NEXT: }
107+
// CHECK-NEXT: return _d_result;
108+
// CHECK-NEXT: }
109+
83110
int main() {
84111
double x[5] = {1, 2, 3, 4, 5};
85112
int n = 5;
@@ -90,5 +117,8 @@ int main() {
90117

91118
auto d_accum_wrt_scale = clad::differentiate(accumulate_scale_parallel, "scale");
92119
double d3 = d_accum_wrt_scale.execute(3.5, 8);
120+
121+
auto d_critical = clad::differentiate(accumulate_critical, "x[0]");
122+
double d4 = d_critical.execute(x, n);
93123
return 0;
94124
}

test/Gradient/OpenMP.C

Lines changed: 102 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -816,6 +816,98 @@ double fn20(const double* x, int n) {
816816
// CHECK-NEXT: }
817817
// CHECK-NEXT: }
818818

819+
double fn21(const double *x, int n) {
820+
double result = 0;
821+
#pragma omp parallel for
822+
for (int i = 0; i < n; i++) {
823+
#pragma omp critical
824+
{
825+
result += x[0] * x[0];
826+
}
827+
}
828+
return result;
829+
}
830+
831+
// CHECK: void fn21_grad(const double *x, int n, double *_d_x, int *_d_n) {
832+
// CHECK-NEXT: double _d_result = 0.;
833+
// CHECK-NEXT: double result = 0;
834+
// CHECK-NEXT: #pragma omp parallel
835+
// CHECK-NEXT: {
836+
// CHECK-NEXT: int _t_chunklo0 = 0;
837+
// CHECK-NEXT: int _t_chunkhi0 = 0;
838+
// CHECK-NEXT: clad::GetStaticSchedule(0, n - 1, 1, &_t_chunklo0, &_t_chunkhi0);
839+
// CHECK-NEXT: for (int i = _t_chunklo0; i <= _t_chunkhi0; i += 1) {
840+
// CHECK-NEXT: #pragma omp critical
841+
// CHECK-NEXT: {
842+
// CHECK-NEXT: result += x[0] * x[0];
843+
// CHECK-NEXT: }
844+
// CHECK-NEXT: }
845+
// CHECK-NEXT: }
846+
// CHECK-NEXT: _d_result += 1;
847+
// CHECK-NEXT: #pragma omp parallel
848+
// CHECK-NEXT: {
849+
// CHECK-NEXT: int _t_chunklo1 = 0;
850+
// CHECK-NEXT: int _t_chunkhi1 = 0;
851+
// CHECK-NEXT: clad::GetStaticSchedule(0, n - 1, 1, &_t_chunklo1, &_t_chunkhi1);
852+
// CHECK-NEXT: for (int i = _t_chunkhi1; i >= _t_chunklo1; i -= 1) {
853+
// CHECK-NEXT: #pragma omp critical
854+
// CHECK-NEXT: {
855+
// CHECK-NEXT: {
856+
// CHECK-NEXT: double _r_d0 = _d_result;
857+
// CHECK-NEXT: _d_x[0] += _r_d0 * x[0];
858+
// CHECK-NEXT: _d_x[0] += x[0] * _r_d0;
859+
// CHECK-NEXT: }
860+
// CHECK-NEXT: }
861+
// CHECK-NEXT: }
862+
// CHECK-NEXT: }
863+
// CHECK-NEXT: }
864+
865+
double fn22(const double *x, int n) {
866+
double result = 0;
867+
#pragma omp parallel for
868+
for (int i = 0; i < n; i++) {
869+
#pragma omp critical
870+
{
871+
result += x[0] * x[1];
872+
}
873+
}
874+
return result;
875+
}
876+
877+
// CHECK: void fn22_grad(const double *x, int n, double *_d_x, int *_d_n) {
878+
// CHECK-NEXT: double _d_result = 0.;
879+
// CHECK-NEXT: double result = 0;
880+
// CHECK-NEXT: #pragma omp parallel
881+
// CHECK-NEXT: {
882+
// CHECK-NEXT: int _t_chunklo0 = 0;
883+
// CHECK-NEXT: int _t_chunkhi0 = 0;
884+
// CHECK-NEXT: clad::GetStaticSchedule(0, n - 1, 1, &_t_chunklo0, &_t_chunkhi0);
885+
// CHECK-NEXT: for (int i = _t_chunklo0; i <= _t_chunkhi0; i += 1) {
886+
// CHECK-NEXT: #pragma omp critical
887+
// CHECK-NEXT: {
888+
// CHECK-NEXT: result += x[0] * x[1];
889+
// CHECK-NEXT: }
890+
// CHECK-NEXT: }
891+
// CHECK-NEXT: }
892+
// CHECK-NEXT: _d_result += 1;
893+
// CHECK-NEXT: #pragma omp parallel
894+
// CHECK-NEXT: {
895+
// CHECK-NEXT: int _t_chunklo1 = 0;
896+
// CHECK-NEXT: int _t_chunkhi1 = 0;
897+
// CHECK-NEXT: clad::GetStaticSchedule(0, n - 1, 1, &_t_chunklo1, &_t_chunkhi1);
898+
// CHECK-NEXT: for (int i = _t_chunkhi1; i >= _t_chunklo1; i -= 1) {
899+
// CHECK-NEXT: #pragma omp critical
900+
// CHECK-NEXT: {
901+
// CHECK-NEXT: {
902+
// CHECK-NEXT: double _r_d0 = _d_result;
903+
// CHECK-NEXT: _d_x[0] += _r_d0 * x[1];
904+
// CHECK-NEXT: _d_x[1] += x[0] * _r_d0;
905+
// CHECK-NEXT: }
906+
// CHECK-NEXT: }
907+
// CHECK-NEXT: }
908+
// CHECK-NEXT: }
909+
// CHECK-NEXT: }
910+
819911
template <size_t N>
820912
void reset(double (&arr)[N], double val = 0) {
821913
for (size_t i = 0; i < N; ++i)
@@ -924,5 +1016,15 @@ int main() {
9241016
auto fn20_grad = clad::gradient(fn20);
9251017
fn20_grad.execute(x, 4, dx, &dn);
9261018
printf("{%.2f, %.2f, %.2f, %.2f}\n", dx[0], dx[1], dx[2], dx[3]); // CHECK-EXEC: {4.00, 6.00, 8.00, 10.00}
1019+
1020+
reset(dx);
1021+
auto fn21_grad = clad::gradient(fn21);
1022+
fn21_grad.execute(x, 4, dx, &dn);
1023+
printf("{%.2f, %.2f, %.2f, %.2f}\n", dx[0], dx[1], dx[2], dx[3]); // CHECK-EXEC: {16.00, 0.00, 0.00, 0.00}
1024+
1025+
reset(dx);
1026+
auto fn22_grad = clad::gradient(fn22);
1027+
fn22_grad.execute(x, 4, dx, &dn);
1028+
printf("{%.2f, %.2f, %.2f, %.2f}\n", dx[0], dx[1], dx[2], dx[3]); // CHECK-EXEC: {12.00, 8.00, 0.00, 0.00}
9271029
return 0;
9281030
}

0 commit comments

Comments
 (0)