我们发现会有 M - N 个数没有被选,所以我们加上这些数,那么递推式就变成了 D_n = (M - N) \times D_{n - 1} + (n-1) \times (D_{n-1} + D_{n-2})。
最终答案就是把上面两个乘起来。
code
#include<bits/stdc++.h>
using namespace std;
constexpr int maxn=5e5+10,p=1e9+7;
int read() {int x=0,f=1;char ch=getchar();while (ch<'0' || ch>'9'){if (ch=='-') f=-1;ch=getchar();}while (ch>='0' && ch<='9'){x=(x<<1)+(x<<3)+ch-'0';ch=getchar();}return x*f;}
int n,m;
long long f[maxn],g[maxn];
long long ksm(long long x,long long y,long long mo)
{
long long cnt=1;
while(y)
{
if (y&1) cnt=(cnt*x)%mo;
x=(x*x)%mo;
y>>=1;
}
return cnt;
}
long long NY(long long a,long long p) {return ksm(a,p-2,p);}
void getfg(long long n)
{
f[0]=g[0]=1;
for (long long i=1;i<=n;i++) (f[i]=i*f[i-1])%=p,g[i]=NY(f[i],p);
}
long long getA(long long x,long long y) {return f[x]*g[x-y]%p;}
long long d[maxn];
void getD(long long n,long long m)
{
d[0]=1,d[1]=m-n;
for (long long i=2;i<=n;i++) (d[i]=(m-n)*d[i-1]%p+(i-1)*(d[i-1]+d[i-2])%p+p)%=p;
}
int main()
{
n=read(),m=read();
getfg(m);getD(n,m);
printf("%lld\n",(d[n]*getA(m,n))%p);
return 0;
}