Skip to content

Commit 0f3a496

Browse files
author
ovdiiuv
committed
Don't create CUDA atomics for basic indices
1 parent b17af42 commit 0f3a496

2 files changed

Lines changed: 138 additions & 24 deletions

File tree

lib/Differentiator/ReverseModeVisitor.cpp

Lines changed: 134 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -56,7 +56,10 @@
5656
#include "clad/Differentiator/CladUtils.h"
5757
#include "clad/Differentiator/Compatibility.h"
5858

59-
using namespace clang;
59+
#include "clang/ASTMatchers/ASTMatchFinder.h"
60+
#include "clang/ASTMatchers/ASTMatchers.h"
61+
62+
using namespace clang::ast_matchers;
6063

6164
namespace clad {
6265

@@ -133,30 +136,141 @@ Expr* ReverseModeVisitor::getStdInitListSizeExpr(const Expr* E) {
133136
return CladTapeResult{*this, PushExpr, PopExpr, TapeRef};
134137
}
135138

139+
static bool isInjectiveE(const clang::Expr* E, clang::ASTContext& ctx) {
140+
StatementMatcher matcher =
141+
binaryOperator(
142+
hasOperatorName("+"),
143+
hasLHS(expr(
144+
has(callExpr(hasDescendant(memberExpr(
145+
hasObjectExpression(opaqueValueExpr(hasSourceExpression(
146+
declRefExpr(to(varDecl(hasName("threadIdx"))))))))))),
147+
unless(binaryOperator()))),
148+
hasRHS(binaryOperator(
149+
hasOperatorName("*"),
150+
has(expr(
151+
has(callExpr(hasDescendant(memberExpr(hasObjectExpression(
152+
opaqueValueExpr(hasSourceExpression(declRefExpr(
153+
to(varDecl(hasName("blockIdx"))))))))))),
154+
unless(binaryOperator()))),
155+
has(expr(
156+
has(callExpr(hasDescendant(memberExpr(hasObjectExpression(
157+
opaqueValueExpr(hasSourceExpression(declRefExpr(
158+
to(varDecl(hasName("blockDim"))))))))))),
159+
unless(binaryOperator()))))))
160+
.bind("targetVar");
161+
162+
class InjectivePatternMatchCallback : public MatchFinder::MatchCallback {
163+
const clang::Expr* targetExpr;
164+
165+
public:
166+
InjectivePatternMatchCallback(const clang::Expr* E) : targetExpr(E) {}
167+
bool matched = false;
168+
virtual void run(const MatchFinder::MatchResult& Result) override {
169+
if (const auto* expr =
170+
Result.Nodes.getNodeAs<clang::BinaryOperator>("targetVar")) {
171+
if (targetExpr == expr)
172+
matched = true;
173+
}
174+
}
175+
bool hasMatched() const { return matched; }
176+
};
177+
178+
MatchFinder Finder;
179+
InjectivePatternMatchCallback Callback(E);
180+
Finder.addMatcher(matcher, &Callback);
181+
Finder.matchAST(ctx);
182+
183+
return Callback.hasMatched();
184+
}
185+
186+
static bool isInjective(const clang::Expr* E, clang::ASTContext& ctx) {
187+
class InjectiveChecker
188+
: public clang::RecursiveASTVisitor<InjectiveChecker> {
189+
clang::ASTContext& m_Context;
190+
191+
public:
192+
InjectiveChecker(clang::ASTContext& Context) : m_Context(Context){};
193+
194+
bool isInjectiveIdx(const clang::Expr* E) {
195+
// NOLINTNEXTLINE(cppcoreguidelines-pro-type-const-cast)
196+
return TraverseStmt(const_cast<clang::Expr*>(E));
197+
}
198+
199+
bool TraverseBinaryOperator(clang::BinaryOperator* BinOp) {
200+
const auto opCode = BinOp->getOpcode();
201+
Expr* L = BinOp->getLHS();
202+
Expr* R = BinOp->getRHS();
203+
204+
if (opCode == BO_Add || opCode == BO_Mul) {
205+
Expr::EvalResult dummy;
206+
207+
bool isConstL =
208+
clad_compat::Expr_EvaluateAsConstantExpr(L, dummy, m_Context);
209+
bool isConstR =
210+
clad_compat::Expr_EvaluateAsConstantExpr(R, dummy, m_Context);
211+
212+
if (isConstL && isConstR)
213+
return false;
214+
215+
if (!isConstL && isConstR)
216+
return TraverseStmt(L);
217+
218+
if (isConstL && !isConstR)
219+
return TraverseStmt(R);
220+
221+
return isInjectiveE(BinOp, m_Context);
222+
}
223+
return false;
224+
}
225+
226+
bool TraverseDeclRefExpr(clang::DeclRefExpr* DRE) {
227+
if (auto* VD = dyn_cast<VarDecl>(DRE->getDecl())) {
228+
if (auto* init = VD->getInit())
229+
return isInjectiveE(init->IgnoreImpCasts(), m_Context);
230+
}
231+
return false;
232+
}
233+
234+
bool TraverseIntegerLiteral(clang::IntegerLiteral* IL) { return false; }
235+
236+
} checker(ctx);
237+
238+
return checker.isInjectiveIdx(E);
239+
}
240+
136241
bool ReverseModeVisitor::shouldUseCudaAtomicOps(const Expr* E) {
137242
if (!m_Context.getLangOpts().CUDA)
138243
return false;
139-
140-
if (!isa<DeclRefExpr>(E))
141-
return false;
142-
143-
const auto* DRE = cast<DeclRefExpr>(E);
144-
145-
if (const auto* PVD = dyn_cast<ParmVarDecl>(DRE->getDecl())) {
146-
if (m_DiffReq->hasAttr<clang::CUDAGlobalAttr>())
147-
// Check whether this param is in the global memory of the GPU
148-
return m_DiffReq.HasIndependentParameter(PVD);
149-
if (m_DiffReq->hasAttr<clang::CUDADeviceAttr>()) {
150-
for (auto index : m_DiffReq.CUDAGlobalArgsIndexes) {
151-
const auto* PVDOrig = m_DiffReq->getParamDecl(index);
152-
if ("_d_" + PVDOrig->getNameAsString() == PVD->getNameAsString() &&
153-
(utils::isArrayOrPointerType(PVDOrig->getType()) ||
154-
PVDOrig->getType()->isReferenceType()))
155-
return true;
244+
if (const auto* DRE = dyn_cast<DeclRefExpr>(E)) {
245+
if (const auto* PVD = dyn_cast<ParmVarDecl>(DRE->getDecl())) {
246+
if (m_DiffReq->hasAttr<clang::CUDAGlobalAttr>())
247+
// Check whether this param is in the global memory of the GPU
248+
return m_DiffReq.HasIndependentParameter(PVD);
249+
if (m_DiffReq->hasAttr<clang::CUDADeviceAttr>()) {
250+
for (auto index : m_DiffReq.CUDAGlobalArgsIndexes) {
251+
const auto* PVDOrig = m_DiffReq->getParamDecl(index);
252+
if ("_d_" + PVDOrig->getNameAsString() == PVD->getNameAsString() &&
253+
(utils::isArrayOrPointerType(PVDOrig->getType()) ||
254+
PVDOrig->getType()->isReferenceType()))
255+
return true;
256+
}
156257
}
157258
}
259+
} else if (const auto* ASE = dyn_cast<ArraySubscriptExpr>(E)) {
260+
if (m_DiffReq->hasAttr<clang::CUDAGlobalAttr>()) {
261+
auto* idx = ASE->getIdx();
262+
auto* base = dyn_cast<DeclRefExpr>(ASE->getBase()->IgnoreImpCasts());
263+
if (const auto* PVD = dyn_cast<ParmVarDecl>(base->getDecl())) {
264+
if (m_DiffReq->hasAttr<clang::CUDAGlobalAttr>()) {
265+
// Check whether this param is in the global memory of the GPU and
266+
// if index is injective.
267+
return m_DiffReq.HasIndependentParameter(PVD) &&
268+
!isInjective(idx, m_Context);
269+
}
270+
}
271+
return true;
272+
}
158273
}
159-
160274
return false;
161275
}
162276

@@ -1341,7 +1455,7 @@ Expr* ReverseModeVisitor::getStdInitListSizeExpr(const Expr* E) {
13411455
// Create the (target += dfdx) statement.
13421456
if (dfdx()) {
13431457
Expr* add_assign = nullptr;
1344-
if (shouldUseCudaAtomicOps(target))
1458+
if (shouldUseCudaAtomicOps(ASE))
13451459
add_assign = BuildCallToCudaAtomicAdd(result, dfdx());
13461460
else
13471461
add_assign = BuildOp(BO_AddAssign, result, dfdx());

test/CUDA/GradientKernels.cu

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -77,7 +77,7 @@ __global__ void add_kernel_3(int *out, int *in) {
7777
//CHECK-NEXT: {
7878
//CHECK-NEXT: out[index0] = _t0;
7979
//CHECK-NEXT: int _r_d0 = _d_out[index0];
80-
//CHECK-NEXT: atomicAdd(&_d_in[index0], _r_d0);
80+
//CHECK-NEXT: _d_in[index0] += _r_d0;
8181
//CHECK-NEXT: }
8282
//CHECK-NEXT:}
8383

@@ -347,7 +347,7 @@ __global__ void dup_kernel_with_device_call_2(double *out, const double *in, dou
347347
//CHECK-NEXT: int _d_index = 0;
348348
//CHECK-NEXT: int index0 = threadIdx.x + blockIdx.x * blockDim.x;
349349
//CHECK-NEXT: {
350-
//CHECK-NEXT: atomicAdd(&_d_in[index0], _d_y);
350+
//CHECK-NEXT: _d_in[index0] += _d_y;
351351
//CHECK-NEXT: *_d_val += _d_y;
352352
//CHECK-NEXT: }
353353
//CHECK-NEXT:}
@@ -382,7 +382,7 @@ __global__ void kernel_with_device_call_3(double *out, double *in, double *val)
382382
//CHECK-NEXT: int _d_index = 0;
383383
//CHECK-NEXT: int index0 = threadIdx.x + blockIdx.x * blockDim.x;
384384
//CHECK-NEXT: {
385-
//CHECK-NEXT: atomicAdd(&_d_in[index0], _d_y);
385+
//CHECK-NEXT: _d_in[index0] += _d_y;
386386
//CHECK-NEXT: atomicAdd(_d_val, _d_y);
387387
//CHECK-NEXT: }
388388
//CHECK-NEXT:}
@@ -418,7 +418,7 @@ __global__ void kernel_with_nested_device_call(double *out, double *in, double v
418418
//CHECK-NEXT: int _d_index = 0;
419419
//CHECK-NEXT: int index0 = threadIdx.x + blockIdx.x * blockDim.x;
420420
//CHECK-NEXT: {
421-
//CHECK-NEXT: atomicAdd(&_d_in[index0], _d_y);
421+
//CHECK-NEXT: _d_in[index0] += _d_y;
422422
//CHECK-NEXT: *_d_val += _d_y;
423423
//CHECK-NEXT: }
424424
//CHECK-NEXT:}

0 commit comments

Comments
 (0)