| 96 | } |
| 97 | |
| 98 | inline std::vector<euler::common::IDWeightPair> Node::__SampleNeighbor( |
| 99 | const std::vector<int32_t>& edge_types, |
| 100 | int32_t count, |
| 101 | const NeighborInfo& ni) const { |
| 102 | std::vector<euler::common::IDWeightPair> err_vec; |
| 103 | std::vector<euler::common::IDWeightPair> empty_vec; |
| 104 | std::vector<euler::common::IDWeightPair> vec(count); |
| 105 | euler::common::CompactWeightedCollection<int32_t> sub_edge_group_collection_; |
| 106 | if (edge_types.size() > 1 && |
| 107 | edge_types.size() < ni.edge_group_collection.GetSize()) { |
| 108 | std::vector<std::pair<int32_t, float>> edge_type_weight(edge_types.size()); |
| 109 | // rebuild weighted collection |
| 110 | for (size_t i = 0; i < edge_types.size(); ++i) { |
| 111 | int32_t edge_type = edge_types[i]; |
| 112 | if (edge_type >= 0 && edge_type < |
| 113 | static_cast<int32_t>(ni.edge_group_collection.GetSize())) { |
| 114 | edge_type_weight[i] = ni.edge_group_collection.Get(edge_type); |
| 115 | } else { |
| 116 | EULER_LOG(ERROR) << "input edge types vec error:" << edge_type; |
| 117 | return err_vec; |
| 118 | } |
| 119 | } |
| 120 | sub_edge_group_collection_.Init(edge_type_weight); |
| 121 | } |
| 122 | |
| 123 | for (int32_t i = 0; i < count; ++i) { |
| 124 | int32_t edge_type = 0; |
| 125 | if (edge_types.size() == 1) { |
| 126 | edge_type = edge_types[0]; |
| 127 | if (edge_type < 0 || edge_type >= |
| 128 | static_cast<int32_t>(ni.edge_group_collection.GetSize())) { |
| 129 | return err_vec; |
| 130 | } |
| 131 | int32_t pre_idx = edge_type == 0 ? 0 : |
| 132 | ni.neighbor_groups_idx[edge_type - 1]; |
| 133 | int32_t cur_idx = ni.neighbor_groups_idx[edge_type] - 1; |
| 134 | if (cur_idx < pre_idx) { |
| 135 | return empty_vec; |
| 136 | } |
| 137 | } else if (edge_types.size() > 1 && |
| 138 | edge_types.size() < ni.edge_group_collection.GetSize()) { |
| 139 | if (sub_edge_group_collection_.GetSumWeight() == 0) { |
| 140 | return empty_vec; |
| 141 | } |
| 142 | edge_type = sub_edge_group_collection_.Sample().first; |
| 143 | } else { // sampling in all edge groups |
| 144 | if (ni.edge_group_collection.GetSumWeight() == 0) { |
| 145 | return empty_vec; |
| 146 | } |
| 147 | edge_type = ni.edge_group_collection.Sample().first; |
| 148 | } |
| 149 | // sample neighbor |
| 150 | int32_t interval_idx_begin = edge_type == 0 ? 0 : |
| 151 | ni.neighbor_groups_idx[edge_type - 1]; |
| 152 | int32_t interval_idx_end = ni.neighbor_groups_idx[edge_type] - 1; |
| 153 | size_t mid = euler::common::RandomSelect<euler::common::NodeID>( |
| 154 | ni.neighbors_weight, interval_idx_begin, interval_idx_end); |
| 155 | float pre_sum_weight = mid <= 0 ? 0 : ni.neighbors_weight[mid - 1]; |