跳转链接
https://www.acwing.com/problem/content/1238/
题目描述
给定三个整数数组 A=[A1,A2,…AN], B=[B1,B2,…BN], C=[C1,C2,…CN], 请你统计有多少个三元组 (i,j,k) 满足: 1≤i,j,k≤N Ai<Bj<Ck
输入格式 第一行包含一个整数 N。 第二行包含 N 个整数 A1,A2,…AN。 第三行包含 N 个整数 B1,B2,…BN。 第四行包含 N 个整数 C1,C2,…CN。 输出格式 一个整数表示答案。 数据范围 1≤N≤105 , 0≤Ai,Bi,Ci≤10^5^
输入样例 3 1 1 1 2 2 2 3 3 3 输出样例 27
题解思路
最暴力的做法是三重循环枚举,但是会超时 我们可以发现一个性质,a[i]几乎只和b[i]有关联,c[i]几乎只和b[i]有关联 所以a[i]和c[i]是具有互斥性质的,而b[i]同时和a[i],c[i]有关联 因此我们只需要找出比b[i]小的总个数和比b[i]大的总个数,两者通过乘法原理相乘即可得出三元组个数!
而对于本题有两种做法 做法一(前缀和 复杂度O(N)) 1.cnt[]可以用来表示在A/B中i这个值出现多少次 2.通过s[]记录当前数组从1到i中,所有数出现的次数的前缀和 3.as[]记录在A中由多少个数小于B[i],as[i] = s[b[i] - 1](b[i] - 1是比b[i]小的数字中最大的) 4.cs[]记录在C中由多少个数大于B[i], cs[i] = s[N - 1] - s[b[i]](通过前缀和公式,计算区间{b[i] + 1,N - 1},即比b[i]大的数) 5.由于a[]和c[]互斥,通过乘法原理可得res += as[i] * cs[i]
做法二(二分 + sort 复杂度O(nlogn)) 1.对3个数组分别进行从小到大排序 2.找到a[]数组中最小的大于等于bi的元素位置la,便可知小于bi的个数为num1 3.找到c[]数组中最大的小于等于bi的元素位置lb,便可知大于bi的个数为num2 4.由于a[]和c[]互斥,通过乘法原理可知符合条件的个数为num1 * num2 Warning:当la == 0 或者 lb == n + 1时,则表示不符合条件
代码
做法一(前缀和)
#include <bits/stdc++.h>
using namespace std;
const int N = 100010;
int n;
int a[N], b[N], c[N], cnt[N];
int s[N], as[N], cs[N];
int main()
{
cin >> n;
for(int i = 1; i <= n; i ++ ) cin >> a[i], a[i] ++ ; //给每个值都加1,避免算前缀和时加不到cnt[0]和s[b[i] - 1] = s[-1]的越界情况
for(int i = 1; i <= n; i ++ ) cin >> b[i], b[i] ++ ;
for(int i = 1; i <= n; i ++ ) cin >> c[i], c[i] ++ ;
for(int i = 1; i <= n; i ++ ) cnt[a[i]] ++ ;
for(int i = 1; i <= N; i ++ ) s[i] = s[i - 1] + cnt[i]; 求从1到i中,所有数出现的次数的前缀和
for(int i = 1; i <= n; i ++ ) as[i] = s[b[i] - 1]; //计算a数组中比b[i]小的数一共有多少
memset(cnt, 0, sizeof cnt);
memset(s, 0, sizeof s);
for(int i = 1; i <= n; i ++ ) cnt[c[i]] ++ ;
for(int i = 1; i <= N; i ++ ) s[i] = s[i - 1] + cnt[i]; 求从1到i中,所有数出现的次数的前缀和
for(int i = 1; i <= n; i ++ ) cs[i] = s[N - 1] - s[b[i]]; 计算c数组中比b[i]大的数一共有多少
long long res = 0; //累加每个b[i]可以构成的三元组数量
for(int i = 1; i <= n; i ++ ) res += (long long)as[i] * cs[i];
cout << res;
return 0;
}做法二(二分 + sort)
#include <bits/stdc++.h>
using namespace std;
const int N = 100010;
int n;
int a[N], b[N], c[N];
int main()
{
cin >> n;
for(int i = 1; i <= n; i ++ ) cin >> a[i], a[i] ++ ;
for(int i = 1; i <= n; i ++ ) cin >> b[i], b[i] ++ ;
for(int i = 1; i <= n; i ++ ) cin >> c[i], c[i] ++ ;
sort(a + 1, a + 1 + n);
sort(b + 1, b + 1 + n);
sort(c + 1, c + 1 + n);
long long res = 0;
for(int i = 1; i <= n; i ++ )
{
int la = 0, ra = n + 1;
while(la < ra) //找到a[]数组中最小的大于等于bi的元素位置
{
int mid = la + ra >> 1;
if(a[mid] >= b[i]) ra = mid;
else la = mid + 1;
}
int lc = 0, rc = n + 1;
while(lc < rc) //找到c[]数组中最大的小于等于bi的元素位置
{
int mid = lc + rc + 1 >> 1;
if(c[mid] <= b[i]) lc = mid;
else rc = mid - 1;
}
if(la == 0 || lc == n + 1) continue;
res += (long long) (la - 1) * (n - lc);
}
cout << res;
return 0;
}
