Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
60 changes: 28 additions & 32 deletions src/substruct/recursive_preprocessor.cu
Original file line number Diff line number Diff line change
Expand Up @@ -13,9 +13,12 @@
// See the License for the specific language governing permissions and
// limitations under the License.

#include <GraphMol/SmilesParse/SmartsWrite.h>

#include <algorithm>
#include <cstdio>
#include <string>
#include <unordered_map>
#include <utility>

#include "src/substruct/molecules_device.cuh"
Expand Down Expand Up @@ -43,13 +46,17 @@ void LeafSubpatterns::buildAllPatterns(const MoleculesHost& queriesHost) {
continue;
}

std::unordered_map<std::string, int> uniquePatternIndices;
for (const auto& entry : recursiveInfo.patterns) {
if (entry.queryMol == nullptr) {
continue;
}

LeafSubpatternKey key{queryIdx, entry.patternId};
if (patternIndexMap.find(key) != patternIndexMap.end()) {
const std::string patternKey = std::to_string(entry.depth) + ':' + RDKit::MolToSmarts(*entry.queryMol);
const auto uniqueIt = uniquePatternIndices.find(patternKey);
if (uniqueIt != uniquePatternIndices.end()) {
patternIndexMap[key] = uniqueIt->second;
continue;
}

Expand Down Expand Up @@ -86,6 +93,7 @@ void LeafSubpatterns::buildAllPatterns(const MoleculesHost& queriesHost) {
}

patternIndexMap[key] = molIdx;
uniquePatternIndices.emplace(patternKey, molIdx);
}
}

Expand All @@ -112,6 +120,7 @@ void LeafSubpatterns::buildAllPatterns(const MoleculesHost& queriesHost) {

perQueryMaxDepth[queryIdx] = recursiveInfo.maxDepth;

std::array<std::unordered_map<int, size_t>, kMaxSmartsNestingDepth + 1> entryIndicesByPattern;
for (const auto& entry : recursiveInfo.patterns) {
if (entry.queryMol == nullptr) {
continue;
Expand All @@ -126,14 +135,23 @@ void LeafSubpatterns::buildAllPatterns(const MoleculesHost& queriesHost) {
continue;
}

auto& entries = perQueryPatterns[queryIdx][entry.depth];
auto& entryIndices = entryIndicesByPattern[entry.depth];
const auto existing = entryIndices.find(patternMolIdx);
if (existing != entryIndices.end()) {
entries[existing->second].patternMask |= 1u << entry.patternId;
continue;
}

BatchedPatternEntry batchEntry;
batchEntry.mainQueryIdx = queryIdx;
batchEntry.patternId = entry.patternId;
batchEntry.patternMask = 1u << entry.patternId;
batchEntry.patternMolIdx = patternMolIdx;
batchEntry.depth = entry.depth;
batchEntry.localIdInParent = entry.localIdInParent;

perQueryPatterns[queryIdx][entry.depth].push_back(batchEntry);
entryIndices.emplace(patternMolIdx, entries.size());
entries.push_back(batchEntry);
}
}

Expand Down Expand Up @@ -317,7 +335,6 @@ void RecursivePatternPreprocessor::preprocessMiniBatch(

void preprocessRecursiveSmarts(SubstructTemplateConfig templateConfig,
const MoleculesDevice& targetsDevice,
const MoleculesHost& queriesHost,
const LeafSubpatterns& leafSubpatterns,
MiniBatchResultsDevice& miniBatchResults,
const int numQueries,
Expand All @@ -343,40 +360,19 @@ void preprocessRecursiveSmarts(SubstructTemplateConfig templateConfig,

const int firstQueryInMiniBatch = miniBatchPairOffset % numQueries;
const int numUniqueQueries = std::min(miniBatchSize, numQueries);
const int recursivePatternsSize = static_cast<int>(queriesHost.recursivePatterns.size());

int maxDepth = 0;
int maxDepth = 0;
for (int i = 0; i < numUniqueQueries; ++i) {
const int queryIdx = (firstQueryInMiniBatch + i) % numQueries;

if (queryIdx >= recursivePatternsSize) {
continue;
}

const auto& recursiveInfo = queriesHost.recursivePatterns[queryIdx];
if (recursiveInfo.empty()) {
if (queryIdx >= static_cast<int>(leafSubpatterns.perQueryPatterns.size())) {
continue;
}

maxDepth = std::max(maxDepth, recursiveInfo.maxDepth);

for (const auto& entry : recursiveInfo.patterns) {
if (entry.queryMol == nullptr) {
continue;
}

const int patternMolIdx = leafSubpatterns.getPatternIndex(queryIdx, entry.patternId);
if (patternMolIdx < 0) {
throw std::runtime_error("Pattern not found in pre-built LeafSubpatterns: queryIdx=" +
std::to_string(queryIdx) + ", patternId=" + std::to_string(entry.patternId));
}

BatchedPatternEntry& batchEntry = patternEntriesHost.emplace_back();
batchEntry.mainQueryIdx = queryIdx;
batchEntry.patternId = entry.patternId;
batchEntry.patternMolIdx = patternMolIdx;
batchEntry.depth = entry.depth;
batchEntry.localIdInParent = entry.localIdInParent;
const int queryMaxDepth = leafSubpatterns.perQueryMaxDepth[queryIdx];
maxDepth = std::max(maxDepth, queryMaxDepth);
for (int depth = 0; depth <= std::min(queryMaxDepth, kMaxSmartsNestingDepth); ++depth) {
const auto& entries = leafSubpatterns.perQueryPatterns[queryIdx][depth];
patternEntriesHost.insert(patternEntriesHost.end(), entries.begin(), entries.end());
}
}

Expand Down
2 changes: 0 additions & 2 deletions src/substruct/recursive_preprocessor.h
Original file line number Diff line number Diff line change
Expand Up @@ -298,7 +298,6 @@ class RecursivePatternPreprocessor {
* depth level for pipeline synchronization.
*
* @param targetsDevice Device-resident target molecules
* @param queriesHost Host-side query data (contains recursivePatterns per query)
* @param leafSubpatterns Pre-built leaf subpattern molecules (device-resident)
* @param miniBatchResults The mini-batch results buffer where recursiveMatchBits will be written
* @param numQueries Total number of queries (for computing pair indices)
Expand All @@ -313,7 +312,6 @@ class RecursivePatternPreprocessor {
*/
void preprocessRecursiveSmarts(SubstructTemplateConfig templateConfig,
const MoleculesDevice& targetsDevice,
const MoleculesHost& queriesHost,
const LeafSubpatterns& leafSubpatterns,
MiniBatchResultsDevice& miniBatchResults,
int numQueries,
Expand Down
11 changes: 5 additions & 6 deletions src/substruct/substruct_algos.cuh
Original file line number Diff line number Diff line change
Expand Up @@ -52,7 +52,7 @@ enum class SubstructOutputMode {
*/
struct PaintModeParams {
uint32_t* recursiveBits; ///< Buffer to paint bits into [maxTargetAtoms per pair]
int patternId; ///< Bit position (0-31) to set
uint32_t patternMask; ///< Bits to set for all equivalent pattern occurrences
int maxTargetAtoms; ///< Stride for indexing recursiveBits
int outputPairIdx; ///< Which pair's buffer to write to
};
Expand Down Expand Up @@ -387,14 +387,13 @@ __device__ void gsiBFSSearchGPU(const TargetMoleculeView&
} else {
// Paint mode: set bit for this target atom
atomicOr(&paintParams.recursiveBits[paintParams.outputPairIdx * paintParams.maxTargetAtoms + t],
1u << paintParams.patternId);
paintParams.patternMask);
atomicAdd(reportedCount, 1);
if constexpr (kDebugGSI) {
printf("[GSI Paint] pairIdx=%d, targetAtom=%d, patternId=%d, bit=0x%x\n",
printf("[GSI Paint] pairIdx=%d, targetAtom=%d, patternMask=0x%x\n",
paintParams.outputPairIdx,
t,
paintParams.patternId,
1u << paintParams.patternId);
paintParams.patternMask);
}
}
} else {
Expand Down Expand Up @@ -561,7 +560,7 @@ __device__ void gsiBFSSearchGPU(const TargetMoleculeView&
const int firstTargetAtom = partial[0];
atomicOr(
&paintParams.recursiveBits[paintParams.outputPairIdx * paintParams.maxTargetAtoms + firstTargetAtom],
1u << paintParams.patternId);
paintParams.patternMask);
atomicAdd(reportedCount, 1);
}
} else {
Expand Down
4 changes: 2 additions & 2 deletions src/substruct/substruct_dfs.cuh
Original file line number Diff line number Diff line change
Expand Up @@ -426,7 +426,7 @@ __device__ void dfsSearchPair(const TargetMoleculeView& ta
const int root = roots.lowest();
roots.clearLowest();
atomicOr(&out.paint.recursiveBits[out.paint.outputPairIdx * out.paint.maxTargetAtoms + root],
1u << out.paint.patternId);
out.paint.patternMask);
}
} else {
unsigned char mapping[1];
Expand Down Expand Up @@ -473,7 +473,7 @@ __device__ void dfsSearchPair(const TargetMoleculeView& ta
// first hit settles the root and skips the rest of its subtree.
if (!terminals.empty()) {
atomicOr(&out.paint.recursiveBits[out.paint.outputPairIdx * out.paint.maxTargetAtoms + mapping[0]],
1u << out.paint.patternId);
out.paint.patternMask);
verdict.rootDone = true;
}
} else {
Expand Down
16 changes: 8 additions & 8 deletions src/substruct/substruct_kernels.cu
Original file line number Diff line number Diff line change
Expand Up @@ -534,9 +534,9 @@ __global__ void substructPaintKernelT(TargetMoleculesDeviceView targets,
return;
}

const int mainQueryIdx = patternEntries ? patternEntries[localPatternIdx].mainQueryIdx : defaultMainQueryIdx;
const int patternId = patternEntries ? patternEntries[localPatternIdx].patternId : defaultPatternId;
const int patternMolIdx = patternEntries ? patternEntries[localPatternIdx].patternMolIdx : localPatternIdx;
const int mainQueryIdx = patternEntries ? patternEntries[localPatternIdx].mainQueryIdx : defaultMainQueryIdx;
const uint32_t patternMask = patternEntries ? patternEntries[localPatternIdx].patternMask : (1u << defaultPatternId);
const int patternMolIdx = patternEntries ? patternEntries[localPatternIdx].patternMolIdx : localPatternIdx;

const int globalPairIdx = targetIdx * outputNumQueries + mainQueryIdx;

Expand Down Expand Up @@ -601,7 +601,7 @@ __global__ void substructPaintKernelT(TargetMoleculesDeviceView targets,

PaintModeParams paintParams;
paintParams.recursiveBits = outputRecursiveBits;
paintParams.patternId = patternId;
paintParams.patternMask = patternMask;
paintParams.maxTargetAtoms = maxTargetAtoms;
paintParams.outputPairIdx = batchLocalPairIdx;

Expand Down Expand Up @@ -749,9 +749,9 @@ __launch_bounds__(dfs::kBlockSize, dfs::minBlocksPerSM<MaxTargetAtoms, MaxQueryA
return;
}

const int mainQueryIdx = patternEntries ? patternEntries[localPatternIdx].mainQueryIdx : defaultMainQueryIdx;
const int patternId = patternEntries ? patternEntries[localPatternIdx].patternId : defaultPatternId;
const int patternMolIdx = patternEntries ? patternEntries[localPatternIdx].patternMolIdx : localPatternIdx;
const int mainQueryIdx = patternEntries ? patternEntries[localPatternIdx].mainQueryIdx : defaultMainQueryIdx;
const uint32_t patternMask = patternEntries ? patternEntries[localPatternIdx].patternMask : (1u << defaultPatternId);
const int patternMolIdx = patternEntries ? patternEntries[localPatternIdx].patternMolIdx : localPatternIdx;

const int globalPairIdx = targetIdx * outputNumQueries + mainQueryIdx;
if (globalPairIdx < miniBatchPairOffset || globalPairIdx >= miniBatchPairOffset + miniBatchSize) {
Expand All @@ -763,7 +763,7 @@ __launch_bounds__(dfs::kBlockSize, dfs::minBlocksPerSM<MaxTargetAtoms, MaxQueryA

dfs::DfsPairOutput out;
out.paint.recursiveBits = outputRecursiveBits;
out.paint.patternId = patternId;
out.paint.patternMask = patternMask;
out.paint.maxTargetAtoms = maxTargetAtoms;
out.paint.outputPairIdx = globalPairIdx - miniBatchPairOffset;

Expand Down
12 changes: 6 additions & 6 deletions src/substruct/substruct_types.h
Original file line number Diff line number Diff line change
Expand Up @@ -41,15 +41,15 @@ struct ZeroBuffersSpec {
* @brief Per-pattern metadata for batched recursive preprocessing kernel.
*
* Each entry describes one recursive pattern in the combined batch:
* which main query it belongs to, what bit to paint, and where the
* which main query it belongs to, what bits to paint, and where the
* pattern data starts in the combined pattern batch.
*/
struct BatchedPatternEntry {
int mainQueryIdx; ///< Index of the main query this pattern belongs to
int patternId; ///< Bit position (0-31) to paint for this pattern
int patternMolIdx; ///< Index into the combined patterns MoleculesDevice
int depth; ///< Nesting depth (0=leaf, higher=parent of children)
int localIdInParent; ///< Bit position in parent's input (for nested patterns)
int mainQueryIdx; ///< Index of the main query this pattern belongs to
uint32_t patternMask; ///< Bits to paint for all equivalent occurrences
int patternMolIdx; ///< Index into the combined patterns MoleculesDevice
int depth; ///< Nesting depth (0=leaf, higher=parent of children)
int localIdInParent; ///< Bit position in parent's input (for nested patterns)
};

// =============================================================================
Expand Down
1 change: 0 additions & 1 deletion src/testutils/substruct_validation.cu
Original file line number Diff line number Diff line change
Expand Up @@ -344,7 +344,6 @@ std::vector<std::vector<uint8_t>> computeGpuLabelMatrix(const RDKit::ROMol& targ
std::vector<BatchedPatternEntry> scratchPatternEntries;
preprocessRecursiveSmarts(SubstructTemplateConfig::Config_T128_Q64_B8,
targetDevice,
queryHost,
leafSubpatterns,
miniBatchResults,
1,
Expand Down
1 change: 0 additions & 1 deletion tests/test_graph_labeler_recursive.cu
Original file line number Diff line number Diff line change
Expand Up @@ -311,7 +311,6 @@ class RecursivePaintTest : public ::testing::Test {
std::vector<BatchedPatternEntry> scratchPatternEntries;
preprocessRecursiveSmarts(SubstructTemplateConfig::Config_T128_Q64_B8,
targetDevice,
queryHost,
leafSubpatterns,
*results_,
numQueries_,
Expand Down
65 changes: 65 additions & 0 deletions tests/test_recursive_preprocessor.cu
Original file line number Diff line number Diff line change
Expand Up @@ -178,6 +178,71 @@ TEST(RecursivePreprocessorTest, PaintsBitsForSimpleRecursivePattern) {
EXPECT_FALSE(hasRecursiveBit(1, 1));
}

TEST(RecursivePreprocessorTest, DeduplicatesRepeatedPatterns) {
auto query = makeMolFromSmarts("[$(*-N),$(*-O),$(*-N)]");
ASSERT_NE(query, nullptr);

std::vector<const RDKit::ROMol*> queries = {query.get()};
std::vector<int> emptySortOrder;
MoleculesHost queriesHost = nvMolKit::buildQueryBatchParallel(queries, emptySortOrder, 1);

ASSERT_EQ(queriesHost.recursivePatterns.size(), 1);
EXPECT_EQ(queriesHost.recursivePatterns[0].size(), 3);

RecursivePatternPreprocessor preprocessor;
preprocessor.buildPatterns(queriesHost);

const LeafSubpatterns& leafSubpatterns = preprocessor.leafSubpatterns();
EXPECT_EQ(leafSubpatterns.patternsHost.numMolecules(), 2);
ASSERT_EQ(leafSubpatterns.perQueryPatterns.size(), 1);
ASSERT_EQ(leafSubpatterns.perQueryPatterns[0][0].size(), 2);

std::vector<uint32_t> masks;
for (const auto& entry : leafSubpatterns.perQueryPatterns[0][0]) {
masks.push_back(entry.patternMask);
}
std::sort(masks.begin(), masks.end());
EXPECT_EQ(masks, (std::vector<uint32_t>{0x2u, 0x5u}));

EXPECT_EQ(leafSubpatterns.getPatternIndex(0, 0), leafSubpatterns.getPatternIndex(0, 2));
EXPECT_NE(leafSubpatterns.getPatternIndex(0, 0), leafSubpatterns.getPatternIndex(0, 1));
}

TEST(RecursivePreprocessorTest, DeduplicatesRepeatedNestedPatternsByDepth) {
auto query = makeMolFromSmarts("[$([C;$(*-N)]),$([C;$(*-O)]),$([C;$(*-N)])]");
ASSERT_NE(query, nullptr);

std::vector<const RDKit::ROMol*> queries = {query.get()};
std::vector<int> emptySortOrder;
MoleculesHost queriesHost = nvMolKit::buildQueryBatchParallel(queries, emptySortOrder, 1);

ASSERT_EQ(queriesHost.recursivePatterns.size(), 1);
EXPECT_EQ(queriesHost.recursivePatterns[0].size(), 6);
EXPECT_EQ(queriesHost.recursivePatterns[0].maxDepth, 2);

RecursivePatternPreprocessor preprocessor;
preprocessor.buildPatterns(queriesHost);

const LeafSubpatterns& leafSubpatterns = preprocessor.leafSubpatterns();
EXPECT_EQ(leafSubpatterns.patternsHost.numMolecules(), 4);
ASSERT_EQ(leafSubpatterns.perQueryPatterns.size(), 1);
ASSERT_EQ(leafSubpatterns.perQueryPatterns[0][0].size(), 2);
ASSERT_EQ(leafSubpatterns.perQueryPatterns[0][1].size(), 2);

std::vector<uint32_t> leafMasks;
std::vector<uint32_t> parentMasks;
for (const auto& entry : leafSubpatterns.perQueryPatterns[0][0]) {
leafMasks.push_back(entry.patternMask);
}
for (const auto& entry : leafSubpatterns.perQueryPatterns[0][1]) {
parentMasks.push_back(entry.patternMask);
}
std::sort(leafMasks.begin(), leafMasks.end());
std::sort(parentMasks.begin(), parentMasks.end());
EXPECT_EQ(leafMasks, (std::vector<uint32_t>{0x10u, 0x28u}));
EXPECT_EQ(parentMasks, (std::vector<uint32_t>{0x2u, 0x5u}));
}

/**
* @brief Leaf subpattern with more atoms than the caller's MaxQueryAtoms
* template tier should not overflow the shared memory label matrix.
Expand Down
2 changes: 1 addition & 1 deletion tests/test_substruct_algos.cu
Original file line number Diff line number Diff line change
Expand Up @@ -312,7 +312,7 @@ __global__ void testGSIPaintKernel(
int maxTargetAtoms,
int outputPairIdx) {
BitMatrix2DView<MaxTargetAtoms, MaxQueryAtoms> labelMatrix(labelMatrixStorage);
nvMolKit::PaintModeParams paintParams{recursiveBits, patternId, maxTargetAtoms, outputPairIdx};
nvMolKit::PaintModeParams paintParams{recursiveBits, 1u << patternId, maxTargetAtoms, outputPairIdx};

gsiBFSSearchGPU<MaxTargetAtoms, MaxQueryAtoms, MaxBondsPerAtom, SubstructOutputMode::PaintBits>(target,
query,
Expand Down
30 changes: 30 additions & 0 deletions tests/test_substruct_search.cu
Original file line number Diff line number Diff line change
Expand Up @@ -1499,6 +1499,36 @@ TEST_P(RecursiveSubstructureSearchTest, NestedRecursiveBatchProcessing) {
}
}

TEST_P(RecursiveSubstructureSearchTest, RepeatedRecursivePatternsMatchRDKit) {
std::vector<std::unique_ptr<RDKit::ROMol>> targetMols;
std::vector<std::unique_ptr<RDKit::ROMol>> queryMols;

const std::vector<std::string> targets = {"NCCCN", "OCCCO", "NCCCO", "NCCCC", "CCCCC", "NCCCNCCCN"};
parseMolecules(targets, {"[$(*-N)]C[$(*-N)]"}, targetMols, queryMols);

SubstructSearchResults results;
getSubstructMatches(getRawPtrs(targetMols), getRawPtrs(queryMols), results, algorithm(), stream_.stream());

for (size_t t = 0; t < targets.size(); ++t) {
expectMatchesRDKit(results, *targetMols[t], *queryMols[0], static_cast<int>(t), 0, targets[t]);
}
}

TEST_P(RecursiveSubstructureSearchTest, RepeatedNestedRecursivePatternsMatchRDKit) {
std::vector<std::unique_ptr<RDKit::ROMol>> targetMols;
std::vector<std::unique_ptr<RDKit::ROMol>> queryMols;

const std::vector<std::string> targets = {"NCCCN", "OCCCO", "NCCCO", "NCCCC", "CCCCC", "NCCCNCCCN"};
parseMolecules(targets, {"[$([C;$(*-N)])]C[$([C;$(*-N)])]"}, targetMols, queryMols);

SubstructSearchResults results;
getSubstructMatches(getRawPtrs(targetMols), getRawPtrs(queryMols), results, algorithm(), stream_.stream());

for (size_t t = 0; t < targets.size(); ++t) {
expectMatchesRDKit(results, *targetMols[t], *queryMols[0], static_cast<int>(t), 0, targets[t]);
}
}

TEST_P(RecursiveSubstructureSearchTest, AmideRecursiveQueryMatchesSmallTargets) {
std::vector<std::unique_ptr<RDKit::ROMol>> targetMols;
std::vector<std::unique_ptr<RDKit::ROMol>> queryMols;
Expand Down
Loading