@@ -20,18 +20,20 @@ struct ParallelPageRankBindData : public TableFunctionData {
2020 int64_t iterations = 100 ;
2121 double tolerance = 1e-6 ;
2222 bool directed = false ;
23+ bool weighted = false ;
2324};
2425struct ParallelPageRankGlobalState : public GlobalTableFunctionState {
2526 std::mutex input_mutex;
2627 std::vector<int64_t > src_nodes, dst_nodes, result_nodes;
27- std::vector<double > result_ranks;
28+ std::vector<double > weights, result_ranks;
2829 idx_t output_idx = 0 ; bool computed = false ;
2930 idx_t MaxThreads () const override { return 1 ; }
3031};
3132
3233static unique_ptr<FunctionData> ParallelPageRankBind (ClientContext &ctx, TableFunctionBindInput &input, vector<LogicalType> &rt, vector<string> &nm) {
3334 auto bd = make_uniq<ParallelPageRankBindData>();
3435 CheckInt64Input (input, " onager_par_pagerank" );
36+ bd->weighted = input.input_table_types .size () >= 3 && input.input_table_types [2 ] == LogicalType::DOUBLE ;
3537 for (auto &kv : input.named_parameters ) {
3638 if (kv.first == " damping" ) bd->damping = GetRequiredParam<double >(" onager_par_pagerank" , " damping" , kv.second );
3739 if (kv.first == " iterations" ) bd->iterations = GetNonNegativeParam (" onager_par_pagerank" , " iterations" , kv.second );
@@ -44,22 +46,26 @@ static unique_ptr<FunctionData> ParallelPageRankBind(ClientContext &ctx, TableFu
4446}
4547static unique_ptr<GlobalTableFunctionState> ParallelPageRankInitGlobal (ClientContext &ctx, TableFunctionInitInput &input) { return make_uniq<ParallelPageRankGlobalState>(); }
4648static OperatorResultType ParallelPageRankInOut (ExecutionContext &ctx, TableFunctionInput &data, DataChunk &input, DataChunk &output) {
49+ auto &bd = data.bind_data ->Cast <ParallelPageRankBindData>();
4750 auto &gs = data.global_state ->Cast <ParallelPageRankGlobalState>();
4851 std::lock_guard<std::mutex> lock (gs.input_mutex );
49- AppendInt64Edges (input, gs.src_nodes , gs.dst_nodes , " onager_par_pagerank" );
52+ if (bd.weighted ) {
53+ AppendWeightedEdges (input, gs.src_nodes , gs.dst_nodes , gs.weights , " onager_par_pagerank" );
54+ } else {
55+ AppendInt64Edges (input, gs.src_nodes , gs.dst_nodes , " onager_par_pagerank" );
56+ }
5057 ONAGER_SET_CARDINALITY (output, 0 ); return OperatorResultType::NEED_MORE_INPUT ;
5158}
5259static OperatorFinalizeResultType ParallelPageRankFinal (ExecutionContext &ctx, TableFunctionInput &data, DataChunk &output) {
5360 auto &bd = data.bind_data ->Cast <ParallelPageRankBindData>(); auto &gs = data.global_state ->Cast <ParallelPageRankGlobalState>();
5461 std::lock_guard<std::mutex> lock (gs.input_mutex );
5562 if (!gs.computed ) {
5663 if (gs.src_nodes .empty ()) { gs.computed = true ; ONAGER_SET_CARDINALITY (output, 0 ); return OperatorFinalizeResultType::FINISHED ; }
57- // graphina's pagerank_parallel ignores edge weights, so no weights are passed here.
58- // Weighted PageRank is available through onager_ctr_pagerank.
59- int64_t nc = ::onager::onager_compute_pagerank_parallel (gs.src_nodes .data (), gs.dst_nodes .data (), gs.src_nodes .size (), nullptr , 0 , bd.damping , bd.iterations , bd.tolerance , bd.directed , nullptr , nullptr );
64+ const double *w = gs.weights .empty () ? nullptr : gs.weights .data ();
65+ int64_t nc = ::onager::onager_compute_pagerank_parallel (gs.src_nodes .data (), gs.dst_nodes .data (), gs.src_nodes .size (), w, gs.weights .size (), bd.damping , bd.iterations , bd.tolerance , bd.directed , nullptr , nullptr );
6066 if (nc < 0 ) throw InvalidInputException (" Parallel PageRank failed: " + GetOnagerError ());
6167 gs.result_nodes .resize (nc); gs.result_ranks .resize (nc);
62- int64_t rc = ::onager::onager_compute_pagerank_parallel (gs.src_nodes .data (), gs.dst_nodes .data (), gs.src_nodes .size (), nullptr , 0 , bd.damping , bd.iterations , bd.tolerance , bd.directed , gs.result_nodes .data (), gs.result_ranks .data ());
68+ int64_t rc = ::onager::onager_compute_pagerank_parallel (gs.src_nodes .data (), gs.dst_nodes .data (), gs.src_nodes .size (), w, gs. weights . size () , bd.damping , bd.iterations , bd.tolerance , bd.directed , gs.result_nodes .data (), gs.result_ranks .data ());
6369 if (rc != nc) throw InvalidInputException (" Parallel PageRank failed: " + GetOnagerError ());
6470 gs.computed = true ;
6571 }
0 commit comments