Use this to learn the idea, then write your own version.
123 4class UnionFind {5public:6 UnionFind(int n)7 : set_(n)8 , rank_(n) { 9 iota(begin(set_), end(set_), 0);10 }11 12 int find_set(int x) {13 vector<int> stk;14 while (set_[x] != x) { 15 stk.emplace_back(x);16 x = set_[x];17 }18 while (!empty(stk)) {19 const int y = stk.back(); stk.pop_back();20 set_[y] = x;21 }22 return x;23 }24 25 bool union_set(int x, int y) {26 x = find_set(x), y = find_set(y);27 if (x == y) {28 return false;29 }30 if (rank_[x] > rank_[y]) {31 swap(x, y);32 }33 set_[x] = y; 34 if (rank_[x] == rank_[y]) {35 ++rank_[y];36 }37 return true;38 }39 40private:41 vector<int> set_;42 vector<int> rank_;43};44 45int binary_search(int left, int right, const auto& check) {46 while (left <= right) {47 const auto& mid = left + (right - left) / 2;48 if (check(mid)) {49 right = mid - 1;50 } else {51 left = mid + 1;52 }53 }54 return left;55};56 5758class Solution {59public:60 vector<int> findMedian(int n, vector<vector<int>>& edges, vector<vector<int>>& queries) {61 vector<vector<pair<int, int>>> adj(size(edges) + 1);62 for (const auto& e : edges) {63 adj[e[0]].emplace_back(e[1], e[2]);64 adj[e[1]].emplace_back(e[0], e[2]);65 }66 const auto& iter_dfs = [&]() {67 vector<bool> lookup(size(adj));68 vector<vector<int>> lookup2(size(adj));69 for (int i = 0; i < size(queries); ++i) {70 for (const auto& x : queries[i]) {71 lookup2[x].emplace_back(i);72 }73 }74 UnionFind uf(size(adj));75 vector<int> ancestor(size(adj));76 iota(begin(ancestor), end(ancestor), 0);77 vector<int64_t> dist(size(adj));78 vector<int> depth(size(adj));79 vector<int> lca(size(queries));80 vector<int64_t> result(size(queries));81 vector<tuple<int, int, int, int>> stk = {{1, 0, -1, -1}};82 while (!empty(stk)) {83 const auto [step, u, p, i] = stk.back(); stk.pop_back();84 if (step == 1) {85 for (const auto& i : lookup2[u]) {86 if (queries[i][0] == queries[i][1]) {87 lca[i] = u;88 continue;89 }90 result[i] += dist[u];91 for (const auto& x : queries[i]) {92 if (lookup[x]) {93 lca[i] = ancestor[uf.find_set(x)];94 result[i] -= 2 * dist[lca[i]];95 }96 }97 }98 lookup[u] = true;99 stk.emplace_back(2, u, -1, 0);100 } else if (step == 2) {101 if (i == size(adj[u])) {102 continue;103 }104 const auto& [v, w] = adj[u][i];105 stk.emplace_back(2, u, -1, i + 1);106 if (lookup[v]) {107 continue;108 }109 dist[v] = dist[u] + w;110 depth[v] = depth[u] + 1;111 stk.emplace_back(3, v, u, -1);112 stk.emplace_back(1, v, -1, -1);113 } else if (step == 3) {114 uf.union_set(u, p);115 ancestor[uf.find_set(p)] = p;116 }117 }118 return tuple(result, lca, dist, depth);119 };120 121 const auto& [result, lca, dist, depth] = iter_dfs();122 const auto& iter_dfs2 = [&]() {123 vector<vector<pair<int, int>>> lookup3(size(adj));124 for (int i = 0; i < size(queries); ++i) {125 const int u = queries[i][0], v = queries[i][1];126 if (2 * (dist[u] - dist[lca[i]]) >= result[i]) {127 lookup3[u].emplace_back(i, 0);128 } else {129 lookup3[v].emplace_back(i, 1);130 }131 }132 vector<int> result2(size(queries));133 vector<int> path;134 vector<tuple<int, int, int>> stk = {{1, 0, -1}};135 while (!empty(stk)) {136 const auto [step, u, i] = stk.back(); stk.pop_back();137 if (step == 1) {138 path.emplace_back(u);139 for (const auto& [i, t] : lookup3[u]) {140 const auto& d = depth[u] - depth[lca[i]];141 if (t == 0) {142 const auto& j = binary_search(0, d, [&](const auto& x) {143 return 2 * (dist[u] - dist[path[size(path) - (x + 1)]]) >= result[i];144 });145 result2[i] = path[size(path) - (j + 1)];146 } else {147 const auto& l = dist[queries[i][0]] - dist[lca[i]];148 const auto& j = binary_search(0, d - 1, [&](const auto& x) {149 return 2 * (l + (dist[path[size(path) - ((d - 1) + 1) + x]] - dist[lca[i]])) >= result[i];150 });151 result2[i] = path[size(path) - ((d - 1) + 1) + j];152 }153 }154 stk.emplace_back(3, u, -1);155 stk.emplace_back(2, u, 0);156 } else if (step == 2) {157 if (i == size(adj[u])) {158 continue;159 }160 const auto& [v, w] = adj[u][i];161 stk.emplace_back(2, u, i + 1);162 if (size(path) >= 2 && path[size(path) - 2] == v) {163 continue;164 }165 stk.emplace_back(1, v, -1);166 } else if (step == 3) {167 path.pop_back();168 }169 }170 return result2;171 };172 173 return iter_dfs2();174 }175};176 177178179180class Solution2 {181public:182 vector<int> findMedian(int n, vector<vector<int>>& edges, vector<vector<int>>& queries) {183 vector<vector<pair<int, int>>> adj(size(edges) + 1);184 for (const auto& e : edges) {185 adj[e[0]].emplace_back(e[1], e[2]);186 adj[e[1]].emplace_back(e[0], e[2]);187 }188 vector<bool> lookup(size(adj));189 vector<vector<int>> lookup2(size(adj));190 for (int i = 0; i < size(queries); ++i) {191 for (const auto& x : queries[i]) {192 lookup2[x].emplace_back(i);193 }194 }195 UnionFind uf(size(adj));196 vector<int> ancestor(size(adj));197 iota(begin(ancestor), end(ancestor), 0);198 vector<int64_t> dist(size(adj));199 vector<int> depth(size(adj));200 vector<int> lca(size(queries));201 vector<int64_t> result(size(queries));202 const function<void (int)> dfs = [&](int u) {203 for (const auto& i : lookup2[u]) {204 if (queries[i][0] == queries[i][1]) {205 lca[i] = u;206 continue;207 }208 result[i] += dist[u];209 for (const auto& x : queries[i]) {210 if (lookup[x]) {211 lca[i] = ancestor[uf.find_set(x)];212 result[i] -= 2 * dist[lca[i]];213 }214 }215 }216 lookup[u] = true;217 for (const auto& [v, w] : adj[u]) {218 if (lookup[v]) {219 continue;220 }221 dist[v] = dist[u] + w;222 depth[v] = depth[u] + 1;223 dfs(v);224 uf.union_set(v, u);225 ancestor[uf.find_set(u)] = u;226 }227 };228 229 dfs(0);230 vector<int> result2(size(queries));231 vector<vector<pair<int, int>>> lookup3(size(adj));232 for (int i = 0; i < size(queries); ++i) {233 const int u = queries[i][0], v = queries[i][1];234 if (2 * (dist[u] - dist[lca[i]]) >= result[i]) {235 lookup3[u].emplace_back(i, 0);236 } else {237 lookup3[v].emplace_back(i, 1);238 }239 }240 vector<int> path;241 const function<void (int)> dfs2 = [&](int u) {242 path.emplace_back(u);243 for (const auto& [i, t] : lookup3[u]) {244 const auto& d = depth[u] - depth[lca[i]];245 if (t == 0) {246 const auto& j = binary_search(0, d, [&](const auto& x) {247 return 2 * (dist[u] - dist[path[size(path) - (x + 1)]]) >= result[i];248 });249 result2[i] = path[size(path) - (j + 1)];250 } else {251 const auto& l = dist[queries[i][0]] - dist[lca[i]];252 const auto& j = binary_search(0, d - 1, [&](const auto& x) {253 return 2 * (l + (dist[path[size(path) - ((d - 1) + 1) + x]] - dist[lca[i]])) >= result[i];254 });255 result2[i] = path[size(path) - ((d - 1) + 1) + j];256 }257 }258 for (const auto& [v, w] : adj[u]) {259 if (size(path) >= 2 && path[size(path) - 2] == v) {260 continue;261 }262 dfs2(v);263 }264 path.pop_back();265 };266 267 dfs2(0);268 return result2;269 }270};271