[R50D]三元组2


题目

给定三个数组 a,b,ca,b,c(长度分别为 na,nb,ncn_a,n_b,n_c)和非负整数 dd。求满足 aibjcka_i\le b_j\le c_kckaidc_k-a_i\le d 的三元组 (i,j,k)(i,j,k) 的个数。

对于 100%100\% 的数据,1na,nb,nc2×1051\le n_a,n_b,n_c\le 2\times 10^50d1090\le d\le 10^91ai,bj,ck1091\le a_i,b_j,c_k\le 10^9

思路

先把 a,b,ca,b,c 分别升序排序,元素的相对顺序不影响计数。

两个条件可以改写为:固定一对 (j,k)(j,k) 满足 bjckb_j\le c_k,则合法的 aia_i 需要同时满足

aibj,ckaid    aickd.a_i\le b_j,\qquad c_k-a_i\le d\iff a_i\ge c_k-d.

ai[ckd,bj]a_i\in[c_k-d,\,b_j]。这要求 ckdbjc_k-d\le b_j,也就是 bjckdb_j\ge c_k-d;再加上 bjckb_j\le c_k,得到 bj[ckd,ck]b_j\in[c_k-d,\,c_k]

于是答案可以按 kk 拆分:

ans=kj:ckdbjck(#{aibj}#{ai<ckd}).\text{ans}=\sum_{k}\sum_{\substack{j:\,c_k-d\le b_j\le c_k}}\bigl(\#\{a_i\le b_j\}-\#\{a_i<c_k-d\}\bigr).

对固定的 kk,记 =ckd\ell=c_k-dbase=#{ai<}\text{base}=\#\{a_i<\ell\}(这是关于 kk 的常量),并把 hj=#{aibj}h_j=\#\{a_i\le b_j\} 预处理出来。则内层求和为

j[jl,jr](hjbase)=(j[jl,jr]hj)base(jrjl+1),\sum_{j\in[j_l,j_r]}(h_j-\text{base}) =\left(\sum_{j\in[j_l,j_r]}h_j\right)-\text{base}\cdot(j_r-j_l+1),

其中 jl=lower_bound(b,)j_l=\text{lower\_bound}(b,\ell)jr=upper_bound(b,ck)1j_r=\text{upper\_bound}(b,c_k)-1。对 hh 做前缀和 HH 后,区间和 hj=H[jr+1]H[jl]\sum h_j=H[j_r+1]-H[j_l]O(1)O(1) 取得。

hj=upper_bound(a,bj)h_j=\text{upper\_bound}(a,b_j),对每个 jj 二分一次即可。

复杂度

  • 时间:排序 O(nlogn)O(n\log n),预处理 h,Hh,HO(nblogna)O(n_b\log n_a)O(nb)O(n_b),枚举 kk 每次 O(log)O(\log),总计 O((na+nb+nc)logn)O((n_a+n_b+n_c)\log n)
  • 空间:O(n)O(n)
  • 答案最大可达 (2×105)3=8×1015(2\times 10^5)^3=8\times 10^{15},需用 64 位整数。

仓颉实现

import std.env.*
import std.convert.*
import std.sort.*

// lower_bound: 第一个 >= key 的下标,不存在返回 n
func lowerBound(a: Array<Int64>, key: Int64): Int64 {
    var lo: Int64 = 0
    var hi: Int64 = Int64(a.size)
    while (lo < hi) {
        let mid = (lo + hi) / 2
        if (a[mid] < key) {
            lo = mid + 1
        } else {
            hi = mid
        }
    }
    return lo
}

// upper_bound: 第一个 > key 的下标,不存在返回 n
func upperBound(a: Array<Int64>, key: Int64): Int64 {
    var lo: Int64 = 0
    var hi: Int64 = Int64(a.size)
    while (lo < hi) {
        let mid = (lo + hi) / 2
        if (a[mid] <= key) {
            lo = mid + 1
        } else {
            hi = mid
        }
    }
    return lo
}

main(): Int64 {
    let reader = getStdIn()
    let firstLine = reader.readln().getOrThrow().split(" ", removeEmpty: true)
    let na = Int64.parse(firstLine[0])
    let nb = Int64.parse(firstLine[1])
    let nc = Int64.parse(firstLine[2])
    let d = Int64.parse(firstLine[3])
    let a = reader.readln().getOrThrow().split(" ", removeEmpty: true).map({ s: String => Int64.parse(s) })
    let b = reader.readln().getOrThrow().split(" ", removeEmpty: true).map({ s: String => Int64.parse(s) })
    let c = reader.readln().getOrThrow().split(" ", removeEmpty: true).map({ s: String => Int64.parse(s) })

    sort(a)
    sort(b)
    sort(c)

    // h[j] = countA(<= b[j])
    let nbInt = nb
    let h = Array<Int64>(nbInt, { _ => 0 })
    for (j in 0..nbInt) {
        h[j] = upperBound(a, b[j])
    }
    // H[i] = h[0]+...+h[i-1]
    let H = Array<Int64>(nbInt + 1, { _ => 0 })
    for (j in 0..nbInt) {
        H[j + 1] = H[j] + h[j]
    }

    var ans: Int64 = 0
    let ncInt = nc
    for (k in 0..ncInt) {
        let ck = c[k]
        let lo = ck - d
        let jl = lowerBound(b, lo)
        let jr = upperBound(b, ck) - 1
        if (jl > jr) {
            continue
        }
        let base = lowerBound(a, lo)
        let cnt = jr - jl + 1
        ans += (H[jr + 1] - H[jl]) - base * cnt
    }

    println(ans.toString())
    return 0
}