#include <bits/stdc++.h>
using namespace std;
long long m,n,a[100005],b[100005],sum;
long long find(int x){
int l=1,r=m,mid;
while(l<=r){
mid=(l+r)/2;
if(x<a[mid]){
r=mid-1;
}
if(x>a[mid]){
l=mid+1;
}
if(x==a[mid]){
return 0;
}
}
int c[4];
c[1]=abs(x-a[l]),c[2]=abs(x-a[r]),c[3]=abs(x-a[mid]);
for(int i=1;i<=3;i++){
for(int j=1;j<=3;j++){
if(c[i]<c[j]){
swap(c[i],c[j]);
}
}
}
return c[1];
}
int main()
{
cin>>m>>n;
for(int i=1;i<=m;i++){
cin>>a[i];
}
sort(a+1,a+m+1);
for(int i=1;i<=n;i++){
cin>>b[i];
if(m==1)
sum+=abs(b[i]-a[1]);
else
sum+=find(b[i]);
}
cout<<sum;
return 0;
}