n,k=map(int,input().split()) M=998244353 d=n//k-1 A=list(map(int,input().split())) B=list(map(int,input().split())) ans=0 b=sum(B) for i in range(n): ans+=d*pow(n-1,-1,M)*(b-B[i])*A[i] ans%=M print(ans)