#!/usr/bin/python3
import sys,threading,random
sys.setrecursionlimit(200200)
threading.stack_size(2**28)

timer = 0
def dfs(node,child,val):
    global timer
    if len(child[node]) == 0:
        val[node] = timer
        timer += 1

    child[node] = sorted(child[node])

    for nxt in child[node]:
        dfs(nxt,child,val)


def solve():
    n,m = map(int, input().split())
    child = [[] for i in range(n)]
    for i in range(n-1):
        a,b = map(int,input().split())
        a -= 1
        b -= 1
        child[a].append(b)

    ds = [int(input())-1 for __ in range(m)]

    val = {}
    timer = 0
    dfs(0,child,val)

    rs = [-1 for i in range(n)]
    for i in range(n-1,-1,-1):
        if len(child[i]) == 0:
            rs[i] = (val[i], val[i])
        else:
            rs[i] = (rs[child[i][0]][0], rs[child[i][-1]][1])

    ptr = 0
    for idx,d in enumerate(ds):
        ptr = max(ptr, rs[d][0])
        if ptr > rs[d][1]:
            print(idx)
            break
    else:
        print(m)

threading.Thread(target=solve).start()
