diff --git a/src/substruct/recursive_preprocessor.cu b/src/substruct/recursive_preprocessor.cu index 5fcff86c..e4ed1b54 100644 --- a/src/substruct/recursive_preprocessor.cu +++ b/src/substruct/recursive_preprocessor.cu @@ -13,9 +13,12 @@ // See the License for the specific language governing permissions and // limitations under the License. +#include + #include #include #include +#include #include #include "src/substruct/molecules_device.cuh" @@ -43,13 +46,17 @@ void LeafSubpatterns::buildAllPatterns(const MoleculesHost& queriesHost) { continue; } + std::unordered_map 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; } @@ -86,6 +93,7 @@ void LeafSubpatterns::buildAllPatterns(const MoleculesHost& queriesHost) { } patternIndexMap[key] = molIdx; + uniquePatternIndices.emplace(patternKey, molIdx); } } @@ -112,6 +120,7 @@ void LeafSubpatterns::buildAllPatterns(const MoleculesHost& queriesHost) { perQueryMaxDepth[queryIdx] = recursiveInfo.maxDepth; + std::array, kMaxSmartsNestingDepth + 1> entryIndicesByPattern; for (const auto& entry : recursiveInfo.patterns) { if (entry.queryMol == nullptr) { continue; @@ -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); } } @@ -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, @@ -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(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(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()); } } diff --git a/src/substruct/recursive_preprocessor.h b/src/substruct/recursive_preprocessor.h index 5a5e37c1..0789b048 100644 --- a/src/substruct/recursive_preprocessor.h +++ b/src/substruct/recursive_preprocessor.h @@ -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) @@ -313,7 +312,6 @@ class RecursivePatternPreprocessor { */ void preprocessRecursiveSmarts(SubstructTemplateConfig templateConfig, const MoleculesDevice& targetsDevice, - const MoleculesHost& queriesHost, const LeafSubpatterns& leafSubpatterns, MiniBatchResultsDevice& miniBatchResults, int numQueries, diff --git a/src/substruct/substruct_algos.cuh b/src/substruct/substruct_algos.cuh index a6bc6480..eae1cb98 100644 --- a/src/substruct/substruct_algos.cuh +++ b/src/substruct/substruct_algos.cuh @@ -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 }; @@ -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 { @@ -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 { diff --git a/src/substruct/substruct_dfs.cuh b/src/substruct/substruct_dfs.cuh index b17f84c2..84d26be9 100644 --- a/src/substruct/substruct_dfs.cuh +++ b/src/substruct/substruct_dfs.cuh @@ -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]; @@ -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 { diff --git a/src/substruct/substruct_kernels.cu b/src/substruct/substruct_kernels.cu index 7ccd2106..a865da41 100644 --- a/src/substruct/substruct_kernels.cu +++ b/src/substruct/substruct_kernels.cu @@ -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; @@ -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; @@ -749,9 +749,9 @@ __launch_bounds__(dfs::kBlockSize, dfs::minBlocksPerSM= miniBatchPairOffset + miniBatchSize) { @@ -763,7 +763,7 @@ __launch_bounds__(dfs::kBlockSize, dfs::minBlocksPerSM> computeGpuLabelMatrix(const RDKit::ROMol& targ std::vector scratchPatternEntries; preprocessRecursiveSmarts(SubstructTemplateConfig::Config_T128_Q64_B8, targetDevice, - queryHost, leafSubpatterns, miniBatchResults, 1, diff --git a/tests/test_graph_labeler_recursive.cu b/tests/test_graph_labeler_recursive.cu index bf4ee121..6c4b6700 100644 --- a/tests/test_graph_labeler_recursive.cu +++ b/tests/test_graph_labeler_recursive.cu @@ -311,7 +311,6 @@ class RecursivePaintTest : public ::testing::Test { std::vector scratchPatternEntries; preprocessRecursiveSmarts(SubstructTemplateConfig::Config_T128_Q64_B8, targetDevice, - queryHost, leafSubpatterns, *results_, numQueries_, diff --git a/tests/test_recursive_preprocessor.cu b/tests/test_recursive_preprocessor.cu index 01c099f2..9b001cac 100644 --- a/tests/test_recursive_preprocessor.cu +++ b/tests/test_recursive_preprocessor.cu @@ -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 queries = {query.get()}; + std::vector 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 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{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 queries = {query.get()}; + std::vector 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 leafMasks; + std::vector 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{0x10u, 0x28u})); + EXPECT_EQ(parentMasks, (std::vector{0x2u, 0x5u})); +} + /** * @brief Leaf subpattern with more atoms than the caller's MaxQueryAtoms * template tier should not overflow the shared memory label matrix. diff --git a/tests/test_substruct_algos.cu b/tests/test_substruct_algos.cu index 16e8e0da..0dc5b06f 100644 --- a/tests/test_substruct_algos.cu +++ b/tests/test_substruct_algos.cu @@ -312,7 +312,7 @@ __global__ void testGSIPaintKernel( int maxTargetAtoms, int outputPairIdx) { BitMatrix2DView labelMatrix(labelMatrixStorage); - nvMolKit::PaintModeParams paintParams{recursiveBits, patternId, maxTargetAtoms, outputPairIdx}; + nvMolKit::PaintModeParams paintParams{recursiveBits, 1u << patternId, maxTargetAtoms, outputPairIdx}; gsiBFSSearchGPU(target, query, diff --git a/tests/test_substruct_search.cu b/tests/test_substruct_search.cu index 67387291..f42d3924 100644 --- a/tests/test_substruct_search.cu +++ b/tests/test_substruct_search.cu @@ -1499,6 +1499,36 @@ TEST_P(RecursiveSubstructureSearchTest, NestedRecursiveBatchProcessing) { } } +TEST_P(RecursiveSubstructureSearchTest, RepeatedRecursivePatternsMatchRDKit) { + std::vector> targetMols; + std::vector> queryMols; + + const std::vector 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(t), 0, targets[t]); + } +} + +TEST_P(RecursiveSubstructureSearchTest, RepeatedNestedRecursivePatternsMatchRDKit) { + std::vector> targetMols; + std::vector> queryMols; + + const std::vector 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(t), 0, targets[t]); + } +} + TEST_P(RecursiveSubstructureSearchTest, AmideRecursiveQueryMatchesSmallTargets) { std::vector> targetMols; std::vector> queryMols;