同余最短路

什么是同余最短路

在处理一些问题是根据数的同余关系建图, 来达到优化空间复杂度的目的.

一些经典问题

诸如:

  1. 给定 $n$ 个整数, 求这 $n$ 个数能拼凑出的多少其他整数 ( $n$ 个数字可以重复取).
  2. 给定 $n$ 个整数, 求这 $n$ 个数不能拼凑出的 最大/最小 的整数.
  3. 至少需要拼几次才可以拼出模 $k$ 余 $b$ 的数.

详细解释在例题中给出.

做题是最好的学习 (不是)

例题

一. 简单不等式

给定长度为 $n$ 的正整数序列 $A$ , 请问有多少个整数属于 $[l, r]$, 可以表示为形如

其中 $P$ 为你构造的任意非负整数序列.

数据

input1

2 5 10
3 5

output1

5

input2

3 101 100000000
99 23333 66666

output2

99573689

数据范围与提示

$ 1 \leq n \leq 12 $, $1 \leq l \leq r \leq 10^{12} $, $1 \leq a_i \leq 2 \times 10^5 $

见到求区间答案的这种题首先考虑考虑是否可以差分, 用 $ solve(x) $ 表示 $1 \to x$ 的答案, $ l \to r $ 就是 $ solve(r) - solve(l - 1) $.

再看题目, 此题等价于 “给定 n 个整数, 求这 n 个数能拼凑出的多少其他整数 (n个数字可以重复取).”
我们记录 $ mn = min(a_i) $, 构造在 余 mn 意义下的最短路, 在 $i \in [0, mn) $ 与 $ (i + a[j]) \% mn $ 之间连单向边, 边权为 $a[j]$ , 意思是从 $ i $ 加到 $ (i + a[j]) $ 在模 $ mn $ 意义下需要的代价为 $a[j]$. (一定要注意都是在模 $k$ 意义下!!!!感性理解理解) , 而我们从 0 跑到 $i \in [0, mn) $ 的最短路(所花费的代价)就是 $T$ , 而且 $T + k\cdot mn$ 也是答案之一.
因此对于 $ 0 \to i \in [0, mn) $ 的最短路 $dis[i]$ 对 [1, x] 的贡献为:

加一是因为 $dis[i]$ 本身就是一个答案.
分数在计算 $T + k\cdot mn$ 的个数.

附上代码, 也有详细解释.

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95

#include <bits/stdc++.h>
#define ll long long
#define int long long
#define rep(i, a, b) for(int i = (a); i <= (b); i ++)
#define inf 0x3f3f3f3f
#define maxn 1000001
#define love return
#define you 0
using namespace std;

int n, l, r;//三输入
int a[30], mn = inf;

int head[maxn], nxt[maxn], dat[maxn], val[maxn], cnt;
void add(int u, int v, int w)
{
dat[++ cnt] = v; nxt[cnt] = head[u]; head[u] = cnt; val[cnt] = w;
}

struct node
{
int u, val;
bool operator < (const node& a) const {return val > a.val; } //优先队列从小到大排序
};

int dis[maxn];
bool vis[maxn];

void tj(int x)//最短路
{
priority_queue < node > q;
q.push({0, 0});

while(!q.empty())
{
int u = q.top().u, k = q.top().val;
q.pop();
if(vis[u] && dis[u] != k)
continue;

vis[u] = 1;

for(int i = head[u]; i; i = nxt[i])
{
int t = dat[i], w = val[i];
if(dis[t] > dis[u] + w)
{
dis[t] = dis[u] + w;
q.push({t, dis[t]});
}
}
}
love ;
}

int value(int x)
{
int res = 0;
rep(i, 0, mn - 1)
{
if(dis[i] <= x)//注意dis[j] 超过 x 的话就不用计算了
res += (x - dis[i]) / mn + 1;
}
love res;
}

signed main()
{
freopen("a.in", "r", stdin);
freopen("a.out", "w", stdout);
cin >> n >> l >> r;

rep(i, 1, n)
{
cin >> a[i];
mn = min(a[i], mn);//找到最小值, 使复杂度最小
}

rep(i, 0, mn - 1)
{
rep(j, 1, n)
{
if(a[j] == mn) continue;
add(i, (i + a[j]) % mn, a[j]);//建边
}
dis[i] = inf;
}dis[0] = 0;
tj(0);

cout << value(r) - value(l - 1) << endl;

love you; //:D
}

luogu简化版本P3403.刷经验

附上代码

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
#include <bits/stdc++.h>
#define ll long long
#define int unsigned long long
#define rep(i, a, b) for(int i = (a); i <= (b); i ++)
#define inf 0x3f3f3f3f
#define maxn 500001
#define love return
#define you 0
using namespace std;

int n, l, r;//三输入
int mn = inf;
ll h;
int a, b, c;

int head[maxn], nxt[maxn], dat[maxn], val[maxn], cnt;
void add(int u, int v, int w)
{
dat[++ cnt] = v; nxt[cnt] = head[u]; head[u] = cnt; val[cnt] = w;
}

struct node
{
int u, val;
bool operator < (const node& a) const {return val > a.val; } //优先队列从小到大排序
};

long long dis[maxn];
bool vis[maxn];

void tj()//最短路
{
priority_queue < node > q;
q.push({1, 1});dis[1] = 1;

while(!q.empty())
{
int u = q.top().u, k = q.top().val;
q.pop();
if(dis[u] != k && vis[u])
continue;

vis[u] = 1;

for(int i = head[u]; i; i = nxt[i])
{
int t = dat[i], w = val[i];
if(dis[t] > dis[u] + w)
{
dis[t] = dis[u] + w;
q.push({t, dis[t]});
}
}
}
love ;
}

int value(int x)
{
int res = 0;
rep(i, 0, a - 1)
{
// cout << dis[i] << endl;
if(dis[i] <= x)//注意dis[j] 超过 x 的话就不用计算了
res += (x - dis[i]) / a + 1;
}
love res;
}

signed main()
{
freopen("a.in", "r", stdin);
freopen("a.out", "w", stdout);
cin >> h >> a >> b >> c;
if(a == 1 || b == 1 || c == 1)
{
cout << h << endl;
return 0;
}
for(int i = 0; i < a; i ++)
{
add(i, (i + b) % a, b);
add(i, (i + c) % a, c);
dis[i] = 18446744073709551610;
}
tj();
cout << value(h) << endl;
love you; //:D
}