Given the root of a binary tree, return its maximum depth as an integer. The maximum depth is the number of nodes along the longest path from the root node down to any leaf node. If the tree is empty, return 0.
def max_depth(root=None, val=None, left=None, right=None):