Skip to content

Commit 361245a

Browse files
author
ovdiiuv
committed
Rework isInjectiveE not to use ASTMatchers
1 parent 0f3a496 commit 361245a

2 files changed

Lines changed: 176 additions & 47 deletions

File tree

lib/Differentiator/ReverseModeVisitor.cpp

Lines changed: 124 additions & 47 deletions
Original file line numberDiff line numberDiff line change
@@ -51,16 +51,14 @@
5151
#include <algorithm>
5252
#include <cstddef>
5353
#include <iterator>
54+
#include <memory>
5455
#include <numeric>
56+
#include <string>
57+
#include <vector>
5558

5659
#include "clad/Differentiator/CladUtils.h"
5760
#include "clad/Differentiator/Compatibility.h"
5861

59-
#include "clang/ASTMatchers/ASTMatchFinder.h"
60-
#include "clang/ASTMatchers/ASTMatchers.h"
61-
62-
using namespace clang::ast_matchers;
63-
6462
namespace clad {
6563

6664
Expr* getArraySizeExpr(const ArrayType* AT, ASTContext& context,
@@ -136,51 +134,130 @@ Expr* ReverseModeVisitor::getStdInitListSizeExpr(const Expr* E) {
136134
return CladTapeResult{*this, PushExpr, PopExpr, TapeRef};
137135
}
138136

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;
137+
static bool isInjectiveE(const clang::Expr* E) {
138+
class InjectiveCheckerExpr
139+
: public clang::RecursiveASTVisitor<InjectiveCheckerExpr> {
140+
struct IdxNode {
141+
struct ExprOrBinOp {
142+
const clang::DeclRefExpr* m_E = nullptr;
143+
clang::BinaryOperatorKind m_Opcode;
144+
enum class Kind { Expr, Opcode } m_Kind;
145+
146+
ExprOrBinOp(const clang::DeclRefExpr* E)
147+
: m_E(E), m_Kind(Kind::Expr) {}
148+
ExprOrBinOp(clang::BinaryOperatorKind op)
149+
: m_Opcode(op), m_Kind(Kind::Opcode) {}
150+
151+
[[nodiscard]] bool isExpr() const { return m_Kind == Kind::Expr; }
152+
[[nodiscard]] bool isOpcode() const { return m_Kind == Kind::Opcode; }
153+
};
154+
155+
ExprOrBinOp Node;
156+
std::unique_ptr<IdxNode> left;
157+
std::unique_ptr<IdxNode> right;
158+
159+
IdxNode(ExprOrBinOp N) : Node(N) {}
160+
};
161+
std::unique_ptr<IdxNode> m_Root;
162+
IdxNode* m_ParentNode{};
163+
164+
enum class side { left, right } m_Side;
164165

165166
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;
167+
InjectiveCheckerExpr() = default;
168+
169+
bool comparePatternToTree(const IdxNode* current,
170+
const std::vector<std::string>& patternIdx,
171+
size_t i = 1) {
172+
173+
if (!current && (i >= patternIdx.size() || patternIdx[i].empty()))
174+
return true;
175+
176+
if (!current || i >= patternIdx.size() || patternIdx[i].empty())
177+
return false;
178+
179+
std::string expected = patternIdx[i];
180+
181+
if (current->Node.isOpcode()) {
182+
std::string actualOp =
183+
clang::BinaryOperator::getOpcodeStr(current->Node.m_Opcode).str();
184+
if (actualOp != expected)
185+
return false;
186+
} else if (current->Node.isExpr()) {
187+
std::string actualName =
188+
current->Node.m_E->getNameInfo().getAsString();
189+
if (actualName != expected)
190+
return false;
173191
}
192+
193+
bool sameOrder =
194+
comparePatternToTree(current->left.get(), patternIdx, 2 * i) &&
195+
comparePatternToTree(current->right.get(), patternIdx, 2 * i + 1);
196+
197+
if (sameOrder)
198+
return true;
199+
200+
return comparePatternToTree(current->left.get(), patternIdx,
201+
2 * i + 1) &&
202+
comparePatternToTree(current->right.get(), patternIdx, 2 * i);
174203
}
175-
bool hasMatched() const { return matched; }
176-
};
177204

178-
MatchFinder Finder;
179-
InjectivePatternMatchCallback Callback(E);
180-
Finder.addMatcher(matcher, &Callback);
181-
Finder.matchAST(ctx);
205+
bool isInjectiveIdx(const clang::Expr* E) {
206+
// NOLINTNEXTLINE(cppcoreguidelines-pro-type-const-cast)
207+
TraverseStmt(const_cast<clang::Expr*>(E));
208+
std::vector<std::string> pattern = {"", "+", "threadIdx", "*",
209+
"", "", "blockIdx", "blockDim"};
210+
return comparePatternToTree(m_Root.get(), pattern);
211+
}
182212

183-
return Callback.hasMatched();
213+
bool TraverseBinaryOperator(clang::BinaryOperator* BinOp) {
214+
const auto opCode = BinOp->getOpcode();
215+
if (opCode == BO_Add || opCode == BO_Mul) {
216+
std::unique_ptr<IdxNode> curr = std::make_unique<IdxNode>(opCode);
217+
IdxNode* currPtr;
218+
219+
if (!m_Root) {
220+
m_Root = std::move(curr);
221+
currPtr = m_Root.get();
222+
} else {
223+
currPtr = curr.get();
224+
225+
if (m_Side == side::left)
226+
m_ParentNode->left = std::move(curr);
227+
228+
if (m_Side == side::right)
229+
m_ParentNode->right = std::move(curr);
230+
}
231+
232+
Expr* L = BinOp->getLHS();
233+
Expr* R = BinOp->getRHS();
234+
235+
m_ParentNode = currPtr;
236+
m_Side = side::left;
237+
TraverseStmt(L);
238+
m_ParentNode = currPtr;
239+
240+
m_Side = side::right;
241+
TraverseStmt(R);
242+
}
243+
return true;
244+
}
245+
246+
bool TraverseDeclRefExpr(clang::DeclRefExpr* DRE) {
247+
std::unique_ptr<IdxNode> curr = std::make_unique<IdxNode>(DRE);
248+
249+
if (m_ParentNode) {
250+
if (m_Side == side::left)
251+
m_ParentNode->left = std::move(curr);
252+
253+
if (m_Side == side::right)
254+
m_ParentNode->right = std::move(curr);
255+
}
256+
return true;
257+
}
258+
259+
} checker;
260+
return checker.isInjectiveIdx(E);
184261
}
185262

186263
static bool isInjective(const clang::Expr* E, clang::ASTContext& ctx) {
@@ -218,15 +295,15 @@ Expr* ReverseModeVisitor::getStdInitListSizeExpr(const Expr* E) {
218295
if (isConstL && !isConstR)
219296
return TraverseStmt(R);
220297

221-
return isInjectiveE(BinOp, m_Context);
298+
return isInjectiveE(BinOp);
222299
}
223300
return false;
224301
}
225302

226303
bool TraverseDeclRefExpr(clang::DeclRefExpr* DRE) {
227304
if (auto* VD = dyn_cast<VarDecl>(DRE->getDecl())) {
228305
if (auto* init = VD->getInit())
229-
return isInjectiveE(init->IgnoreImpCasts(), m_Context);
306+
return isInjectiveE(init->IgnoreImpCasts());
230307
}
231308
return false;
232309
}
@@ -258,7 +335,7 @@ Expr* ReverseModeVisitor::getStdInitListSizeExpr(const Expr* E) {
258335
}
259336
} else if (const auto* ASE = dyn_cast<ArraySubscriptExpr>(E)) {
260337
if (m_DiffReq->hasAttr<clang::CUDAGlobalAttr>()) {
261-
auto* idx = ASE->getIdx();
338+
const auto* idx = ASE->getIdx();
262339
auto* base = dyn_cast<DeclRefExpr>(ASE->getBase()->IgnoreImpCasts());
263340
if (const auto* PVD = dyn_cast<ParmVarDecl>(base->getDecl())) {
264341
if (m_DiffReq->hasAttr<clang::CUDAGlobalAttr>()) {

test/CUDA/GradientKernels.cu

Lines changed: 52 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -689,6 +689,56 @@ void launch_add_kernel_4(int *out, int *in, const int N) {
689689
//CHECK-NEXT: cudaFree(_d_out_dev);
690690
//CHECK-NEXT:}
691691

692+
__global__ void indices_perm(int *out, int *in) {
693+
int index1 = threadIdx.x + blockIdx.x * blockDim.x;
694+
int index2 = threadIdx.x + blockDim.x * blockIdx.x;
695+
int index3 = blockIdx.x * blockDim.x + threadIdx.x;
696+
int index4 = blockDim.x * blockIdx.x + threadIdx.x;
697+
out[index1] += in[index1];
698+
out[index2] += in[index2];
699+
out[index3] += in[index3];
700+
out[index4] += in[index4];
701+
}
702+
703+
// CHECK: void indices_perm_grad(int *out, int *in, int *_d_out, int *_d_in) {
704+
// CHECK-NEXT: int _d_index1 = 0;
705+
// CHECK-NEXT: int index1 = threadIdx.x + blockIdx.x * blockDim.x;
706+
// CHECK-NEXT: int _d_index2 = 0;
707+
// CHECK-NEXT: int index2 = threadIdx.x + blockDim.x * blockIdx.x;
708+
// CHECK-NEXT: int _d_index3 = 0;
709+
// CHECK-NEXT: int index3 = blockIdx.x * blockDim.x + threadIdx.x;
710+
// CHECK-NEXT: int _d_index4 = 0;
711+
// CHECK-NEXT: int index4 = blockDim.x * blockIdx.x + threadIdx.x;
712+
// CHECK-NEXT: int _t0 = out[index1];
713+
// CHECK-NEXT: out[index1] += in[index1];
714+
// CHECK-NEXT: int _t1 = out[index2];
715+
// CHECK-NEXT: out[index2] += in[index2];
716+
// CHECK-NEXT: int _t2 = out[index3];
717+
// CHECK-NEXT: out[index3] += in[index3];
718+
// CHECK-NEXT: int _t3 = out[index4];
719+
// CHECK-NEXT: out[index4] += in[index4];
720+
// CHECK-NEXT: {
721+
// CHECK-NEXT: out[index4] = _t3;
722+
// CHECK-NEXT: int _r_d3 = _d_out[index4];
723+
// CHECK-NEXT: _d_in[index4] += _r_d3;
724+
// CHECK-NEXT: }
725+
// CHECK-NEXT: {
726+
// CHECK-NEXT: out[index3] = _t2;
727+
// CHECK-NEXT: int _r_d2 = _d_out[index3];
728+
// CHECK-NEXT: _d_in[index3] += _r_d2;
729+
// CHECK-NEXT: }
730+
// CHECK-NEXT: {
731+
// CHECK-NEXT: out[index2] = _t1;
732+
// CHECK-NEXT: int _r_d1 = _d_out[index2];
733+
// CHECK-NEXT: _d_in[index2] += _r_d1;
734+
// CHECK-NEXT: }
735+
// CHECK-NEXT: {
736+
// CHECK-NEXT: out[index1] = _t0;
737+
// CHECK-NEXT: int _r_d0 = _d_out[index1];
738+
// CHECK-NEXT: _d_in[index1] += _r_d0;
739+
// CHECK-NEXT: }
740+
// CHECK-NEXT: }
741+
692742
#define TEST(F, grid, block, shared_mem, use_stream, x, dx, N) \
693743
{ \
694744
int *fives = (int*)malloc(N * sizeof(int)); \
@@ -963,6 +1013,8 @@ int main(void) {
9631013
launch_kernel_4_test.execute(zeros_int, fives_int, 10, out_res, in_res);
9641014
printf("%d, %d, %d\n", in_res[0], in_res[1], in_res[2]); // CHECK-EXEC: 5, 5, 5
9651015

1016+
TEST_2(indices_perm, dim3(1), dim3(5, 1, 1), 0, false, "out, in", dummy_out, dummy_in, d_out, d_in, 5); // CHECK-EXEC: 20, 20, 20, 20, 20
1017+
9661018
free(res);
9671019
free(fives);
9681020
free(zeros);

0 commit comments

Comments
 (0)