# your code goes here
import sys

def solve():
    input = sys.stdin.read
    data = input().split()
    if not data:
        return
    
    n = int(data[0])
    k = int(data[1])
    h = [int(x) for x in data[2:]]
    
    dp = [float('inf')] * n
    dp[0] = 0
    
    for i in range(n):
        for j in range(1, k + 1):
            if i + j < n:
                dp[i + j] = min(dp[i + j], dp[i] + abs(h[i + j] - h[i]))
    
    print(dp[n - 1])

if __name__ == '__main__':
    solve()