Use this to learn the idea, then write your own version.
123 45class Solution(object):6 def subtreeInversionSum(self, edges, nums, k):7 """8 :type edges: List[List[int]]9 :type nums: List[int]10 :type k: int11 :rtype: int12 """13 def iter_dfs():14 result = []15 stk = [(1, (0, -1, result))]16 while stk:17 step, args = stk.pop()18 if step == 1:19 u, p, ret = args20 ret[:] = [[nums[u]]*k, [nums[u]]*k]21 stk.append((4, (u, p, ret)))22 stk.append((2, (u, p, 0, ret)))23 elif step == 2:24 u, p, i, ret = args25 if i == len(adj[u]):26 continue27 v = adj[u][i]28 stk.append((2, (u, p, i+1, ret)))29 if v == p:30 continue31 new_ret = []32 stk.append((3, (new_ret, ret)))33 stk.append((1, (v, u, new_ret)))34 elif step == 3:35 new_ret, ret = args36 new_dp1, new_dp2 = new_ret37 dp1, dp2 = ret38 for i in xrange(k2):39 dp1[i] = max(dp1[i]+new_dp1[(k-2)-i], dp1[(k-2)-i]+new_dp1[i])40 dp2[i] = min(dp2[i]+new_dp2[(k-2)-i], dp2[(k-2)-i]+new_dp2[i])41 for i in xrange(k2, k):42 dp1[i] += new_dp1[i]43 dp2[i] += new_dp2[i]44 for i in reversed(xrange(k-1)):45 dp1[i] = max(dp1[i], dp1[i+1])46 dp2[i] = min(dp2[i], dp2[i+1])47 elif step == 4:48 u, p, ret = args49 dp1, dp2 = ret50 dp1.insert(0, max(dp1[0], -dp2[-1]))51 dp2.insert(0, min(dp2[0], -dp1[-1]))52 dp1.pop()53 dp2.pop()54 return result[0][0]55 56 adj = [[] for _ in xrange(len(nums))]57 for u, v in edges:58 adj[u].append(v)59 adj[v].append(u)60 return iter_dfs()61 62 63646566class Solution2(object):67 def subtreeInversionSum(self, edges, nums, k):68 """69 :type edges: List[List[int]]70 :type nums: List[int]71 :type k: int72 :rtype: int73 """74 def dfs(u, p):75 dp1, dp2 = [nums[u]]*k, [nums[u]]*k76 for v in adj[u]:77 if v == p:78 continue79 new_dp1, new_dp2 = dfs(v, u)80 for i in xrange(k2):81 dp1[i] = max(dp1[i]+new_dp1[(k-2)-i], dp1[(k-2)-i]+new_dp1[i])82 dp2[i] = min(dp2[i]+new_dp2[(k-2)-i], dp2[(k-2)-i]+new_dp2[i])83 for i in xrange(k2, k):84 dp1[i] += new_dp1[i]85 dp2[i] += new_dp2[i]86 for i in reversed(xrange(k-1)):87 dp1[i] = max(dp1[i], dp1[i+1])88 dp2[i] = min(dp2[i], dp2[i+1])89 dp1.insert(0, max(dp1[0], -dp2[-1]))90 dp2.insert(0, min(dp2[0], -dp1[-1]))91 dp1.pop()92 dp2.pop()93 return dp1, dp294 95 adj = [[] for _ in xrange(len(nums))]96 for u, v in edges:97 adj[u].append(v)98 adj[v].append(u)99 dp1, _ = dfs(0, -1)100 return dp1[0]101