如下是我的代码,本机样例都过不去就鬼畜报错,一交100pts,十分莫名其妙,Orz_stO求大佬释疑。
#include<bits/stdc++.h>
using namespace std;
#define int long long
int n,m,T,mod;
int qp(int n,int p,int mod){
if(p==0)return 1;
if(p==1)return n;
if(p%2==0)return qp(n*n%mod,p/2,mod)%mod;
return n*qp(n*n%mod,p/2,mod)%mod;
}
int exgcd(int a,int b,int &x,int &y){
if(b==0){
x=1;y=0;return a;
}
int res=exgcd(b,a%b,y,x);
y-=a/b*x;
return res;
}
int inv(int n,int p,int mod){
int xx,yy;
exgcd(n,p,xx,yy);
xx=(xx%mod+mod)%mod;
return xx;
}
int cnt(int n,int p,int k){
if(n==0)return 1;
int ans=1,a2=1;
for(int i=1;i<k;i++){
if(i%p!=0)ans=ans*i%k;
}
ans=qp(ans,n/k,k)%k;
for(int i=k*(n/k);i<=n;i++){
if(i%p!=0)a2=a2*(i%k)%k;
}
return ans%k*a2%k*cnt(n/p,p,k)%k;
}
int C1(int n,int m,int p,int k){
if(n<m)return 0;
int f1=cnt(n,p,k),f2=cnt(m,p,k),f3=cnt(n-m,p,k),c=0;
for(int i=m;i>0;i/=p)c-=i/p;
for(int i=n;i>0;i/=p)c+=i/p;
for(int i=n-m;i>0;i/=p)c-=i/p;
return f1*inv(f2,k,k)%k*inv(f3,k,k)%k*qp(p,c,k)%k;
}
int CRT(int a[],int t[],int m[],int n,int mod){
int res=0;
for(int i=1;i<=n;i++)res=(res+a[i]%mod*t[i]%mod*m[i]%mod)%mod;
return res;
}
int C(int n,int m,int mod){
int a[100005],pk[100005],mi[100005];
int n2=mod,point=1;
for(int i=2;i<=sqrt(mod)&&n2;i++){
int tmp=1;
while(n2%i==0){
n2/=i;
tmp*=i;
}
if(tmp>1){
a[point]=C1(n,m,i,tmp);
pk[point]=tmp;
point++;
}
}
if(n2>1){
a[point]=C1(n,m,n2,n2);
pk[point]=n2;
point++;
}
for(int i=1;i<point;i++)mi[i]=inv((mod/pk[i])%pk[i],pk[i],pk[i])%pk[i];
for(int i=1;i<point;i++)pk[i]=mod/pk[i];
return CRT(a,pk,mi,point-1,mod);
}
int a[100005],ans=1;
signed main(){
cin>>mod;
cin>>n>>m;
int sum=0;
for(int i=1;i<=m;i++){
cin>>a[i];
if(n<a[i]){
cout<<"Impossible\n";
exit(0);
}
ans=(ans*C(n,a[i],mod))%mod;
n-=a[i];
}
cout<<ans;
}