1236. 递增三元组 题解

跳转链接

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时,则表示不符合条件

代码

cpp
做法一(前缀和)
#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;
}
cpp
做法二(二分 + 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;
}
1204. 错误票据 题解
1210. 连号区间数 题解
Valaxy v1.0.0-rc.3 驱动|主题-Yunv1.0.0-rc.3