n,m = map(int,input().split())
w = list(map(int,input().split()))
t = 0
def jiancha(a):
for i in range(len(a)):
if 0 in a:
a.remove(0)
elif 0 not in a:
break
flag = 0
while len(w)>m:
t += 1
for i in range(m):
if w[i] < 1:
continue
else:
w[i]-=1
jiancha(w)
if len(w)<=m:
t += max(w)
print(t)