From 930569fc813cf0b8d0248280a9e798d0570a0a26 Mon Sep 17 00:00:00 2001 From: mkolodner Date: Tue, 28 Jul 2026 09:40:12 +0000 Subject: [PATCH 1/2] Extract PPR top-k selection helper --- gigl-core/core/sampling/ppr_forward_push.cpp | 195 +++++++++++-------- 1 file changed, 117 insertions(+), 78 deletions(-) diff --git a/gigl-core/core/sampling/ppr_forward_push.cpp b/gigl-core/core/sampling/ppr_forward_push.cpp index a2ab255a8..ff7eebe43 100644 --- a/gigl-core/core/sampling/ppr_forward_push.cpp +++ b/gigl-core/core/sampling/ppr_forward_push.cpp @@ -266,11 +266,113 @@ void PPRForwardPush::pushResiduals( } } +// Helper function for selecting one seed/node-type's finalized PPR rows. +// +// Inputs: +// nodeTypeState: finalized PPR scores and residuals for one seed/node type. +// finalizedPPRNodeLimit: maximum finalized-PPR rows to select before top-up. +// +// Expected output: (node_id, raw_ppr_score) pairs selected by raw PPR score. +// The order is unspecified; callers that emit these directly should sort the +// returned vector before writing output tensors. +static std::vector> selectFinalizedPPRPairs(const SeedNodeTypeState& nodeTypeState, + int32_t finalizedPPRNodeLimit) { + const auto& scores = nodeTypeState.pprScores; + const auto higherScore = [](const auto& a, const auto& b) { return a.second > b.second; }; + + const int32_t pprTopK = std::min(finalizedPPRNodeLimit, static_cast(scores.size())); + std::vector> selectedPairs; + selectedPairs.reserve(static_cast(pprTopK)); + if (pprTopK > 0) { + std::vector> scorePairs(scores.begin(), scores.end()); + if (pprTopK < static_cast(scorePairs.size())) { + std::nth_element(scorePairs.begin(), scorePairs.begin() + pprTopK, scorePairs.end(), higherScore); + } + + for (int32_t rankIdx = 0; rankIdx < pprTopK; ++rankIdx) { + selectedPairs.emplace_back(scorePairs[rankIdx].first, scorePairs[rankIdx].second); + } + } + + return selectedPairs; +} + +// Helper function for extending one seed/node-type's selected PPR rows with +// residual top-up candidates. +// +// Inputs: +// nodeTypeState: finalized PPR scores and residuals for one seed/node type. +// selectedPairs: mutable finalized-PPR rows already selected for this seed. +// sequenceLength: maximum total rows after residual top-up. +// +// Expected output: selectedPairs has up to sequenceLength rows after appending +// highest-scoring residual candidates that are not already selected. This helper +// does not sort selectedPairs; callers sort only when their output needs it. +static void appendResidualTopUpPairs(const SeedNodeTypeState& nodeTypeState, + std::vector>& selectedPairs, + int32_t sequenceLength) { + const int32_t residualTopUpBudget = + std::max(0, sequenceLength - static_cast(selectedPairs.size())); + if (residualTopUpBudget > 0) { + const auto& scores = nodeTypeState.pprScores; + const auto higherScore = [](const auto& a, const auto& b) { return a.second > b.second; }; + std::unordered_set selectedPPRNodeIds; + selectedPPRNodeIds.reserve(selectedPairs.size()); + for (const auto& selectedPair : selectedPairs) { + selectedPPRNodeIds.insert(selectedPair.first); + } + + std::vector> residualPairs; + residualPairs.reserve(nodeTypeState.residuals.size()); + for (const auto& [nodeId, residual] : nodeTypeState.residuals) { + if (residual <= 0.0 || selectedPPRNodeIds.find(nodeId) != selectedPPRNodeIds.end()) { + continue; + } + + auto scoreIter = scores.find(nodeId); + double pprScore = (scoreIter != scores.end()) ? scoreIter->second : 0.0; + double outputScore = pprScore + residual; + residualPairs.emplace_back(nodeId, outputScore); + } + + const int32_t residualTopK = std::min(residualTopUpBudget, static_cast(residualPairs.size())); + if (residualTopK > 0) { + if (residualTopK < static_cast(residualPairs.size())) { + std::nth_element( + residualPairs.begin(), residualPairs.begin() + residualTopK, residualPairs.end(), higherScore); + } + + for (int32_t rankIdx = 0; rankIdx < residualTopK; ++rankIdx) { + selectedPairs.emplace_back(residualPairs[rankIdx].first, residualPairs[rankIdx].second); + } + } + } +} + +// Helper function for moving finalized-PPR rows onto the same score scale as +// residual top-up rows. +// +// Inputs: +// nodeTypeState: residual table for one seed/node type. +// selectedPairs: mutable finalized-PPR rows with raw PPR scores. +// +// Expected output: selectedPairs scores are updated in-place to +// ppr_score + residual(node) when residual mass exists for that node. +static void addResidualMassToPPRPairs(const SeedNodeTypeState& nodeTypeState, + std::vector>& selectedPairs) { + for (auto& [nodeId, score] : selectedPairs) { + auto residualIter = nodeTypeState.residuals.find(nodeId); + if (residualIter != nodeTypeState.residuals.end()) { + score += residualIter->second; + } + } +} + std::unordered_map> PPRForwardPush:: extractTopKWithResidualTopUp(int32_t maxPPRNodes, bool enableResidualTopUp) { TORCH_CHECK(maxPPRNodes >= 0, "maxPPRNodes must be non-negative, got ", maxPPRNodes, "."); - std::unordered_map> result; + std::unordered_map> extractedPPRByNodeTypeId; // Emit an entry for every node type, even if unreachable in this batch (empty tensors, // all-zero valid_counts). This keeps the output shape consistent across batches so // downstream model architectures see a fixed set of PPR edge types every iteration. @@ -281,83 +383,20 @@ std::unordered_map b.second; }; - - const int32_t topK = std::min(maxPPRNodes, static_cast(scores.size())); - const int32_t residualTopUpBudget = enableResidualTopUp ? maxPPRNodes - topK : 0; - std::vector> selectedPairs; - selectedPairs.reserve(static_cast(topK) + static_cast(residualTopUpBudget)); - std::unordered_set selectedPPRNodeIds; - if (topK > 0) { - if (residualTopUpBudget > 0) { - selectedPPRNodeIds.reserve(static_cast(topK)); - } - std::vector> scorePairs(scores.begin(), scores.end()); - // Selection is intentionally two-phase: finalized nodes are selected - // first by raw PPR score, and residual candidates only compete for - // the remaining budget. - if (enableResidualTopUp) { - // The final emitted order is sorted by ppr_score + residual - // after top-up candidates are selected, so this pass only - // needs to partition out the raw-PPR top K. - if (topK < static_cast(scorePairs.size())) { - std::nth_element(scorePairs.begin(), scorePairs.begin() + topK, scorePairs.end(), higherScore); - } - } else { - std::partial_sort(scorePairs.begin(), scorePairs.begin() + topK, scorePairs.end(), higherScore); - } - - for (int32_t rankIdx = 0; rankIdx < topK; ++rankIdx) { - int32_t nodeId = scorePairs[rankIdx].first; - double outputScore = scorePairs[rankIdx].second; - if (enableResidualTopUp) { - auto residualIter = nodeTypeState.residuals.find(nodeId); - if (residualIter != nodeTypeState.residuals.end()) { - outputScore += residualIter->second; - } - } - selectedPairs.emplace_back(nodeId, outputScore); - if (residualTopUpBudget > 0) { - selectedPPRNodeIds.insert(nodeId); - } - } + auto selectedPairs = selectFinalizedPPRPairs(nodeTypeState, maxPPRNodes); + if (enableResidualTopUp) { + addResidualMassToPPRPairs(nodeTypeState, selectedPairs); + appendResidualTopUpPairs(nodeTypeState, selectedPairs, maxPPRNodes); } - if (residualTopUpBudget > 0) { - std::vector> residualPairs; - residualPairs.reserve(nodeTypeState.residuals.size()); - for (const auto& [nodeId, residual] : nodeTypeState.residuals) { - if (residual <= 0.0 || selectedPPRNodeIds.find(nodeId) != selectedPPRNodeIds.end()) { - continue; - } - - auto scoreIter = scores.find(nodeId); - double pprScore = (scoreIter != scores.end()) ? scoreIter->second : 0.0; - double outputScore = pprScore + residual; - residualPairs.emplace_back(nodeId, outputScore); - } - - const int32_t residualTopK = std::min(residualTopUpBudget, static_cast(residualPairs.size())); - if (residualTopK > 0) { - // Residual candidates only need selection here; selected - // finalized and residual rows are sorted together below. - if (residualTopK < static_cast(residualPairs.size())) { - std::nth_element(residualPairs.begin(), - residualPairs.begin() + residualTopK, - residualPairs.end(), - higherScore); - } - - for (int32_t rankIdx = 0; rankIdx < residualTopK; ++rankIdx) { - selectedPairs.emplace_back(residualPairs[rankIdx].first, residualPairs[rankIdx].second); - } - } - } - - if (enableResidualTopUp && selectedPairs.size() > 1) { + // Empty and singleton outputs are already ordered. Multi-row outputs can + // be unordered because the selection helpers use nth_element, so sort once + // to rank emitted rows by the score returned to callers. + if (selectedPairs.size() > 1) { + const auto higherScore = [](const auto& a, const auto& b) { return a.second > b.second; }; std::sort(selectedPairs.begin(), selectedPairs.end(), higherScore); } + for (const auto& [nodeId, score] : selectedPairs) { flatIds.push_back(static_cast(nodeId)); flatWeights.push_back(score); @@ -365,11 +404,11 @@ std::unordered_map(selectedPairs.size())); } - result[nodeTypeId] = {torch::tensor(flatIds, torch::kLong), - torch::tensor(flatWeights, torch::kDouble), - torch::tensor(validCounts, torch::kLong)}; + extractedPPRByNodeTypeId[nodeTypeId] = {torch::tensor(flatIds, torch::kLong), + torch::tensor(flatWeights, torch::kDouble), + torch::tensor(validCounts, torch::kLong)}; } - return result; + return extractedPPRByNodeTypeId; } int32_t PPRForwardPush::getTotalDegree(int32_t nodeId, int32_t nodeTypeId) const { From d06fbf14a04b4170319a3cdae17599185eec0481 Mon Sep 17 00:00:00 2001 From: mkolodner Date: Wed, 29 Jul 2026 16:43:13 +0000 Subject: [PATCH 2/2] Address PPR helper review comments --- gigl-core/core/sampling/ppr_forward_push.cpp | 45 ++++++++++++-------- 1 file changed, 27 insertions(+), 18 deletions(-) diff --git a/gigl-core/core/sampling/ppr_forward_push.cpp b/gigl-core/core/sampling/ppr_forward_push.cpp index ff7eebe43..915df7c56 100644 --- a/gigl-core/core/sampling/ppr_forward_push.cpp +++ b/gigl-core/core/sampling/ppr_forward_push.cpp @@ -278,18 +278,20 @@ void PPRForwardPush::pushResiduals( static std::vector> selectFinalizedPPRPairs(const SeedNodeTypeState& nodeTypeState, int32_t finalizedPPRNodeLimit) { const auto& scores = nodeTypeState.pprScores; - const auto higherScore = [](const auto& a, const auto& b) { return a.second > b.second; }; - const int32_t pprTopK = std::min(finalizedPPRNodeLimit, static_cast(scores.size())); + const int32_t numReturnedPairs = std::min(finalizedPPRNodeLimit, static_cast(scores.size())); std::vector> selectedPairs; - selectedPairs.reserve(static_cast(pprTopK)); - if (pprTopK > 0) { + selectedPairs.reserve(static_cast(numReturnedPairs)); + if (numReturnedPairs > 0) { std::vector> scorePairs(scores.begin(), scores.end()); - if (pprTopK < static_cast(scorePairs.size())) { - std::nth_element(scorePairs.begin(), scorePairs.begin() + pprTopK, scorePairs.end(), higherScore); + if (numReturnedPairs < static_cast(scorePairs.size())) { + std::nth_element(scorePairs.begin(), + scorePairs.begin() + numReturnedPairs, + scorePairs.end(), + [](const auto& a, const auto& b) { return a.second > b.second; }); } - for (int32_t rankIdx = 0; rankIdx < pprTopK; ++rankIdx) { + for (int32_t rankIdx = 0; rankIdx < numReturnedPairs; ++rankIdx) { selectedPairs.emplace_back(scorePairs[rankIdx].first, scorePairs[rankIdx].second); } } @@ -314,8 +316,7 @@ static void appendResidualTopUpPairs(const SeedNodeTypeState& nodeTypeState, const int32_t residualTopUpBudget = std::max(0, sequenceLength - static_cast(selectedPairs.size())); if (residualTopUpBudget > 0) { - const auto& scores = nodeTypeState.pprScores; - const auto higherScore = [](const auto& a, const auto& b) { return a.second > b.second; }; + const std::unordered_map& pprScoresByNodeId = nodeTypeState.pprScores; std::unordered_set selectedPPRNodeIds; selectedPPRNodeIds.reserve(selectedPairs.size()); for (const auto& selectedPair : selectedPairs) { @@ -325,12 +326,15 @@ static void appendResidualTopUpPairs(const SeedNodeTypeState& nodeTypeState, std::vector> residualPairs; residualPairs.reserve(nodeTypeState.residuals.size()); for (const auto& [nodeId, residual] : nodeTypeState.residuals) { + // Forward push residuals are non-negative in normal operation. Pushed + // nodes remain in the map with zero residual, so skip drained entries + // and any unexpected non-positive values. if (residual <= 0.0 || selectedPPRNodeIds.find(nodeId) != selectedPPRNodeIds.end()) { continue; } - auto scoreIter = scores.find(nodeId); - double pprScore = (scoreIter != scores.end()) ? scoreIter->second : 0.0; + std::unordered_map::const_iterator pprScoreIter = pprScoresByNodeId.find(nodeId); + double pprScore = (pprScoreIter != pprScoresByNodeId.end()) ? pprScoreIter->second : 0.0; double outputScore = pprScore + residual; residualPairs.emplace_back(nodeId, outputScore); } @@ -338,8 +342,10 @@ static void appendResidualTopUpPairs(const SeedNodeTypeState& nodeTypeState, const int32_t residualTopK = std::min(residualTopUpBudget, static_cast(residualPairs.size())); if (residualTopK > 0) { if (residualTopK < static_cast(residualPairs.size())) { - std::nth_element( - residualPairs.begin(), residualPairs.begin() + residualTopK, residualPairs.end(), higherScore); + std::nth_element(residualPairs.begin(), + residualPairs.begin() + residualTopK, + residualPairs.end(), + [](const auto& a, const auto& b) { return a.second > b.second; }); } for (int32_t rankIdx = 0; rankIdx < residualTopK; ++rankIdx) { @@ -389,12 +395,15 @@ std::unordered_map 1) { - const auto higherScore = [](const auto& a, const auto& b) { return a.second > b.second; }; - std::sort(selectedPairs.begin(), selectedPairs.end(), higherScore); + std::sort(selectedPairs.begin(), selectedPairs.end(), [](const auto& a, const auto& b) { + return a.second > b.second; + }); } for (const auto& [nodeId, score] : selectedPairs) {