Leetcode每日一题 —— 1339. 分裂二叉树的最大乘积

1339. 分裂二叉树的最大乘积

思路
因为是树结构,所以容易想到遍历每条边,求子树的值之和与剩余值的积即可。

代码

class Solution {
    long ans;
    int total;
    public int maxProduct(TreeNode root) {
        ans = 0;
        total = traversal(root);
        dfs(root.left);
        dfs(root.right);
        return (int) (ans % 1000000007);
    }

    private long dfs(TreeNode node) {
        if (node == null) return 0;
        long sum = node.val + dfs(node.left) + dfs(node.right);
        ans = Math.max(ans, sum * (total - sum));
        return sum;
    }

    private int traversal(TreeNode node) {
        if (node == null) return 0;
        return node.val + traversal(node.left) + traversal(node.right);
    }
}
7 个赞

佬加油呀

2 个赞

一起加油 :partying_face:

今日刷题

2 个赞

每日一题

1 个赞

先求总和,然后枚举分离每一棵子树来求结果,不断更新最大值即可。

注意求的是最大值取模,而不是取模后的最大值。

我这写出来不知道为什么这么慢,想着求出子树和后临时存入 val,但其实有点多余了,等下再优化一下。

class Solution {
public:
    int maxProduct(TreeNode* root) {
        // 首先节点值都是正数
        // 因此子树中节点越多,和自然就越大
        // 拆掉边后,两棵子树的节点数应该相近或者相同

        // 其实可以枚举拆掉每条边,子树的和可以算出来,拆掉边就临时扣去某棵子树的和
        int modulus = (int)(1e9 + 7);
        // 先把所有节点的和算出来
        // 假设所有节点都是 10000,最大总和是 500000000,int 类型放得下
        function<int(TreeNode*)> dfsSum;
        dfsSum = [&](TreeNode* node) {
            if (node == nullptr) {
                return 0;
            }
            int sum = dfsSum(node->left) + dfsSum(node->right) + node->val;
            // 把节点值替换为以这个节点为根的子树节点值总和
            node->val = sum;
            return sum;
        };
        int totalSum = dfsSum(root);
        // cout << "TOTAL: " << totalSum << endl;
        // 枚举移除掉每条边,看看哪个得到的结果最大
        long long res = 0;
        function<void(TreeNode*)> dfsFind;
        dfsFind = [&](TreeNode* node) {
            if (node == nullptr) {
                return;
            }
            // 看看把当前节点为根的子树拆出来,结果如何
            // 注意找的应该是 **真实最大值** 取模后的结果,而不是取模后的最大值
            res = max(res,
                      (long long)node->val * (long long)(totalSum - node->val));
            dfsFind(node->left);
            dfsFind(node->right);
        };
        dfsFind(root);
        return (int)(res % modulus);
    }
};

优化了一下,快多了,原来是 C++ 的一些小技巧,我这 std::function + lambda 的递归写法开销比较大(有间接调用,即运行时才知道调用哪个函数,有点阻止编译器优化了)。参考了前面 0ms 的代码才发现还有直接转发引用 lambda 的写法… (auto&& 自动推导是左值引用还是右值引用) 学到新东西了。

class Solution {
public:
    int maxProduct(TreeNode* root) {
        // 首先节点值都是正数
        // 因此子树中节点越多,和自然就越大
        // 拆掉边后,两棵子树的节点数应该相近或者相同

        // 其实可以枚举拆掉每条边,子树的和可以算出来,拆掉边就临时扣去某棵子树的和
        int modulus = (int)(1e9 + 7);
        // 先把所有节点的和算出来
        // 假设所有节点都是 10000,最大总和是 500000000,int 类型放得下
        auto dfsSum = [&](auto&& self, TreeNode* node) {
            if (node == nullptr) {
                return 0;
            }
            int sum =
                self(self, node->left) + self(self, node->right) + node->val;
            return sum;
        };
        int totalSum = dfsSum(dfsSum, root);
        // cout << "TOTAL: " << totalSum << endl;
        // 枚举移除掉每条边,看看哪个得到的结果最大
        long long res = 0;
        auto dfsFind = [&](auto&& self, TreeNode* node) {
            if (node == nullptr) {
                return 0;
            }
            // 看看把当前节点为根的子树拆出来,结果如何
            // 注意找的应该是 **真实最大值** 取模后的结果,而不是取模后的最大值
            int sum =
                self(self, node->left) + self(self, node->right) + node->val;
            res = max(res, (long long)sum * (long long)(totalSum - sum));
            return sum;
        };
        dfsFind(dfsFind, root);
        return (int)(res % modulus);
    }
};

用其他语言写可能就没有这么神奇的优化空间了

3 个赞
# Definition for a binary tree node.
# class TreeNode:
#     def __init__(self, val=0, left=None, right=None):
#         self.val = val
#         self.left = left
#         self.right = right
class Solution:
    @cache
    def maxProduct(self, root: Optional[TreeNode]) -> int:
        md = 1_000_000_007
        def dfs(root):
            if not root:
                return 0
            cur = root.val + dfs(root.left) + dfs(root.right)
            st.append(cur)
            return cur


        # def dfs1(root):
        #     if not root:
        #         return 0
        #     cur = root.val + dfs1(root.left) + dfs1(root.right)
        #     nonlocal ans
        #     ans = max(ans, cur * (sm - cur))
        #     return cur

        st = []
        sm = dfs(root)
        ans = max([x * (sm - x) for x in st])
        return ans % md


1 个赞

初版,用DFS+层序遍历,会TLE:

# Definition for a binary tree node.
# class TreeNode:
#     def __init__(self, val=0, left=None, right=None):
#         self.val = val
#         self.left = left
#         self.right = right
class Solution:
    @cache
    def maxProduct(self, root: Optional[TreeNode]) -> int:
        mod = 1e9 + 7
        ans = []
        def dfs(root):
            if not root:
                return 0
            return root.val + dfs(root.left) + dfs(root.right)
        
        total = dfs(root)
        while len(q):
            node = q.popleft()
            if node.left:
                q.append(node.left)
            if node.right:
                q.append(node.right)
            cur_sum = dfs(node)
            ans.append(cur_sum * (total - cur_sum))
        return int(max(ans) % mod)

优化版,在遍历过程中就存储中间结果,少遍历一次:

# Definition for a binary tree node.
# class TreeNode:
#     def __init__(self, val=0, left=None, right=None):
#         self.val = val
#         self.left = left
#         self.right = right
class Solution:
    @cache
    def maxProduct(self, root: Optional[TreeNode]) -> int:
        mod = 1e9 + 7
        node_sum = []
        def dfs(root):
            if not root:
                return 0
            cur_sum = root.val + dfs(root.left) + dfs(root.right)
            node_sum.append(cur_sum)
            return cur_sum
        
        total = dfs(root)
        ans = max([x * (total - x) for x in node_sum])
        return int(ans % mod)
2 个赞

auto&& self self(self, node->left)

哇哦,真是神奇,学废了学废了!

3 个赞

优秀,这样只遍历一次就行了。

哦不对,还是两次,只是把第二次遍历从树改队列了,时间复杂度还是一样的。

1 个赞

那个 :face_with_peeking_eye:第一段代码是不是少了点什么?
看起来像是想像昨天那样用队列来存储层?不太懂python,不过看起来这样dfs子树会重复行数次的,而且队列感觉也会有时间的损耗 :joy:

3 个赞

后面那个是用队列对整棵二叉树进行层序遍历,记录每个节点作为根节点时的子树和 owo

每个节点相当于经历了多次DFS,有很多次都是没必要的
(根节点的DFS其实包含了这些子节点的DFS,属于重复遍历,也就TLE了 owo

3 个赞

每日一题打卡 :distorted_face:

/**
 * @param {TreeNode} root
 * @return {number}
 */
var maxProduct = function (root) {
  const subtreeSums = [];

  const calcSubtreeSum = (node) => {
    if (!node) return 0;
    const sum =
      node.val + calcSubtreeSum(node.left) + calcSubtreeSum(node.right);
    subtreeSums.push(sum);
    return sum;
  };

  const totalSum = calcSubtreeSum(root);
  let maxProduct = 0;

  for (const subSum of subtreeSums) {
    const product = subSum * (totalSum - subSum);
    if (product > maxProduct) maxProduct = product;
  }

  return maxProduct % (10 ** 9 + 7);
};
    def maxProduct(self, root: Optional[TreeNode]) -> int:
        sums = []
        def dfs(node):
            if not node:
                return 0
            sums.append((t := node.val + dfs(node.left) + dfs(node.right)))
            return t

        total = dfs(root)

        return max([s * (total-s) for s in sums]) % int(1e9+7)