P9377 百合 题解

· · 题解

P9377 百合

推销我的博客园。

题意

给定三个整数 k,m,s,现有 2^k 个点和 m 条带权边,点的下标从 0 开始。

同时,还有 k 个整数 a_1\sim a_k,你可以花费 a_t 的代价从 x 走到 y,其中 txy 在二进制表示下不同的比特数。

问从 s 开始到每个点的最短路长度。

数据范围

思路

一道比较复杂的最短路问题。

首先能想到把 a_i 进行差分,然后在 2^k\times k 个点之间跑最短路……

但有两个问题,没有保证 a 具有单调性,即差分后可能存在负权边,也就是不能使用 dij,二是有负权边则可能会出现反复横跳的情况,即从 (x,i)(y,i+1) 再到 (x,i+2),这种情况是不合法的,但并不好判断。

于是考虑建更多的点,令 (x,i,j) 表示考虑前 i 位是否进行反转,一共反转了 j 位,最终数值为 x。但如果用差分数组来表示边权的话则仍会有负权边的问题,于是考虑最后再算 a 的值,具体而言:

这样就能够避免反复横跳和负权边的问题,对这张图跑 dij,复杂度 O((2^kk^2+m)\log 2^kk^2),难以通过。

注意到 (x,i,j) 部分有大量权为 0 的边,无需优先队列,所以考虑用 bfs 特殊处理这一部分,这样就可以把复杂度优化至 O(2^kk^2+(2^kk+m)\log 2^kk)

同时由于空间紧张,bfs 部分需要开 bool 数组。

复杂度

Code

#include <iostream>
#include <vector>
#include <queue>
#include <tuple>
#define _1 (__int128)1

using namespace std;
using ll = long long;
using pii = pair<int, int>;
using tp = tuple<int, int, int>;

void FileIO (const string s) {
  freopen((s + ".in").c_str(), "r", stdin);
  freopen((s + ".out").c_str(), "w", stdout);
}

const int T = (1 << 17) + 5;

int p, m, st, dis[T], mx, a[20];
bool vis[T][20][20];
vector<pii> g[T];
priority_queue<pii, vector<pii>, greater<pii>> pq;
queue<tp> q;

void Record (int x, int lv) {
  if (lv > mx + 10 || (dis[x] && lv >= dis[x])) return ;
  pq.push({lv, x}), dis[x] = lv;
}

void Record_ (int x, int y, int z) {
  if (vis[x][y][z]) return ;
  vis[x][y][z] = 1, q.push({x, y, z});
}

void bfs (int x, int lv) {
  Record_(x, 0, 0);
  while (q.size()) {
    auto [y, i, j] = q.front();
    q.pop(), Record(y, lv + a[j]);
    if (i < p) Record_(y, i + 1, j), Record_(y ^ (1 << i), i + 1, j + 1);
  }
}

void dij () {
  Record(st, 1);
  while (pq.size()) {
    auto [lv, x] = pq.top();
    pq.pop();
    if (dis[x] != lv) continue;
    bfs(x, lv);
    for (auto [i, j] : g[x])
      Record(i, lv + j);
  }
}

signed main () {
  ios::sync_with_stdio(0), cin.tie(0);
  // FileIO("");
  cin >> p >> m >> st;
  for (int i = 1; i <= p; i++)
    cin >> a[i], mx = max(mx, a[i]);
  for (int i = 1, x, y, z; i <= m; i++)
    cin >> x >> y >> z, g[x].push_back({y, z}), g[y].push_back({x, z});
  dij();
  for (int i = 0; i < (1 << p); i++) 
    cout << dis[i] - 1 << ' ';
  return 0;
}