Approach
Breadth-first search
For Total Sum of Interaction Cost in Tree Groups, the implementation explores reachable states in layers, which is the standard shape for unweighted shortest paths and minimum-step transitions.
- Model each valid configuration as a state and each legal move as an edge.
- Seed the queue with the starting state and mark it immediately.
- Expand each state once, recording distance or reachability for unseen neighbours.
Code notes
- 152 lines of C++ from the credited upstream file total-sum-of-interaction-cost-in-tree-groups.cpp.
- The implementation visibly relies on sequence storage, hash lookup.
- 20 loop blocks detected, together with recursive traversal.
Complexity
Verify that each state and transition is processed only a bounded number of times; that determines the traversal cost.
Check the problem constraints before deciding whether this complexity will pass.
Use this to learn the idea, then write your own version.
123 45class Solution {6public:7 long long interactionCosts(int n, vector<vector<int>>& edges, vector<int>& group) {8 vector<vector<int>> adj(n);9 for (const auto& e : edges) {10 adj[e[0]].emplace_back(e[1]);11 adj[e[1]].emplace_back(e[0]);12 }13 const auto& mx = ranges::max(group);14 vector<int64_t> total(mx);15 for (const auto& x : group) {16 ++total[x - 1];17 }18 int64_t result = 0;19 const auto& bfs = [&]() {20 vector<int> order = {0};21 vector<int> parent(n, -1);22 for (int i = 0; i < size(adj); ++i) {23 const auto u = order[i];24 for (const auto& v : adj[u]) {25 if (v == parent[u]) {26 continue;27 }28 parent[v] = u;29 order.emplace_back(v);30 }31 }32 return pair(order, parent);33 };34 35 const auto& [order, parent] = bfs();36 vector<vector<int64_t>> cnt(n, vector<int64_t>(mx));37 for (int i = size(order) - 1; i >= 0; --i) {38 const auto& u = order[i];39 ++cnt[u][group[u] - 1];40 for (const auto& v : adj[u]) {41 if (u != parent[v]) {42 continue;43 }44 for (int k = 0; k < size(cnt[v]); ++k) {45 result += cnt[v][k] * (total[k] - cnt[v][k]);46 cnt[u][k] += cnt[v][k];47 }48 }49 }50 return result;51 }52};53 54555657class Solution2 {58public:59 long long interactionCosts(int n, vector<vector<int>>& edges, vector<int>& group) {60 vector<vector<int>> adj(n);61 for (const auto& e : edges) {62 adj[e[0]].emplace_back(e[1]);63 adj[e[1]].emplace_back(e[0]);64 } 65 unordered_map<int, int64_t> total;66 for (const auto& x : group) {67 ++total[x];68 }69 int64_t result = 0;70 const auto& bfs = [&]() {71 vector<int> order = {0};72 vector<int> parent(n, -1);73 for (int i = 0; i < size(adj); ++i) {74 const auto u = order[i];75 for (const auto& v : adj[u]) {76 if (v == parent[u]) {77 continue;78 }79 parent[v] = u;80 order.emplace_back(v);81 }82 }83 return pair(order, parent);84 };85 86 const auto& [order, parent] = bfs();87 vector<unordered_map<int, int64_t>> cnt(n);88 for (int i = size(order) - 1; i >= 0; --i) {89 const auto& u = order[i];90 ++cnt[u][group[u]];91 for (const auto& v : adj[u]) {92 if (u != parent[v]) {93 continue;94 }95 for (const auto& [k, c] : cnt[v]) {96 result += c * (total[k] - c);97 }98 if (size(cnt[v]) > size(cnt[u])) {99 swap(cnt[u], cnt[v]);100 }101 for (const auto& [k, c] : cnt[v]) {102 cnt[u][k] += c;103 }104 cnt[v].clear();105 }106 }107 return result;108 }109};110 111112113114class Solution3 {115public:116 long long interactionCosts(int n, vector<vector<int>>& edges, vector<int>& group) {117 vector<vector<int>> adj(n);118 for (const auto& e : edges) {119 adj[e[0]].emplace_back(e[1]);120 adj[e[1]].emplace_back(e[0]);121 } 122 unordered_map<int, int64_t> total;123 for (const auto& x : group) {124 ++total[x];125 }126 int64_t result = 0;127 const auto dfs = [&](this auto&& dfs, int u, int p) -> unordered_map<int, int64_t> {128 unordered_map<int, int64_t> cnt;129 ++cnt[group[u]];130 for (const auto& v : adj[u]) {131 if (v == p) {132 continue;133 }134 auto new_cnt = dfs(v, u);135 for (const auto& [k, c] : new_cnt) {136 result += c * (total[k] - c);137 }138 if (size(new_cnt) > size(cnt)) {139 swap(cnt, new_cnt);140 }141 for (const auto& [k, c] : new_cnt) {142 cnt[k] += c;143 }144 }145 return cnt;146 };147 148 dfs(0, -1);149 return result;150 }151};152