对于线性代数的一点想法
好像会了一个把 GF(2) 线性方程反复求解可以到 \(\Theta(n^3/w + Q\cdot n^2)\) 的方法(固定矩阵 \(A\),many RHS 的时候很爽),传统的时间复杂度到\(\Theta(Q\cdot n^3/w)\)。不过 OI 里大概率也没啥用就是了。

UPD:这里只写 \(p=2\)。一般 \(p\) 也能做,但我懒得把常数和工程细节写到能打的程度。
Part1
固定矩阵 \(A\),反复解 \(A y \equiv c \pmod 2\),每次 RHS \(c\) 不一样。
把 \(A_2=A\bmod 2\) 当成 \(\mathbb F_2\) 上的线性映射之后,这事就很线代。
可解性只看 \(c\) 在不在像空间里:\(c\in\mathrm{Im}(A_2)\) 才能解。
能解的话解集长这样:
\(y = T_2 c + N_2 z\)
特解 + 齐次解。\(T_2\) 就是“固定挑一个特解”的线性算子,\(N_2\) 是核空间基,\(z\) 自由变量。
想做的事也简单:\(A\) 固定时,高斯消元只做一次。
把行变换记录下来,以后每个 RHS 来了就照着变换 RHS,然后回代。
复杂度账也很直白。
传统做法:每个 RHS 都消元一次。bitset/XOR 的视角下,一次消元量级大概是 \(O(mn\cdot \frac{n}{w})\)(\(w=64\)),有 \(Q\) 个 RHS 就乘 \(Q\)。
这套做法:预处理一次(同量级),之后每个 RHS 只需要过一遍脚本(脚本长度通常是 \(O(mn)\) 级别)再回代 \(O(rn)\)。大头不再乘 \(Q\)。
所以只解一两个 RHS 没必要折腾,多 RHS 才值。模 \(2^k\) 的逐次提升也是同理:每一轮都在解同一个 \(A_2\),不复用的话纯浪费。
Part2
代码(只做 many RHS 的模 2):
#include<bits/stdc++.h>
#define L(i, j, k) for(int i = (j); i <= (k); ++i)
#define R(i, j, k) for(int i = (j); i >= (k); --i)
#define ll long long
#define ull unsigned long long
#define vi vector<int>
#define sz(a) ((int)(a).size())
using namespace std;
struct GF2 {
int m, n, W, rk;
vector<vector<ull>> a;
vector<int> perm;
vector<int> pivc, pivr;
enum OP { SW, XO };
struct op { int t, x, y; };
vector<op> ops;
GF2(int _m = 0, int _n = 0) { init(_m, _n); }
void init(int _m, int _n) {
m = _m, n = _n;
W = (n + 63) >> 6;
a.assign(m, vector<ull>(W, 0));
perm.resize(n);
iota(perm.begin(), perm.end(), 0);
ops.clear();
rk = 0;
}
inline void setb(int r, int c, int v) {
if(v) a[r][c >> 6] |= 1ULL << (c & 63);
else a[r][c >> 6] &= ~(1ULL << (c & 63));
}
inline int getb(int r, int cpos) const {
return (a[r][cpos >> 6] >> (cpos & 63)) & 1ULL;
}
inline void swrow(int x, int y) {
if(x == y) return;
swap(a[x], a[y]);
ops.push_back({SW, x, y});
}
inline void xorow(int x, int y) {
if(x == y) return;
L(i, 0, W - 1) a[x][i] ^= a[y][i];
ops.push_back({XO, x, y});
}
inline void swcol(int c1, int c2) {
if(c1 == c2) return;
L(r, 0, m - 1) {
int b1 = getb(r, c1), b2 = getb(r, c2);
if(b1 ^ b2) {
a[r][c1 >> 6] ^= 1ULL << (c1 & 63);
a[r][c2 >> 6] ^= 1ULL << (c2 & 63);
}
}
swap(perm[c1], perm[c2]);
}
void preprocess() {
pivc.assign(m, -1);
pivr.assign(n, -1);
int r = 0;
L(c, 0, n - 1) if(r < m) {
int bc = -1, br = -1;
L(cc, c, n - 1) if(bc == -1) {
L(rr, r, m - 1) if(getb(rr, cc)) { bc = cc, br = rr; break; }
}
if(bc == -1) break;
if(bc != c) swcol(bc, c);
if(br != r) swrow(br, r);
pivc[r] = c;
pivr[c] = r;
L(rr, r + 1, m - 1) if(getb(rr, c)) xorow(rr, r);
++r;
}
rk = 0;
L(i, 0, m - 1) if(pivc[i] != -1) ++rk;
}
void preprocess_reduced() {
preprocess();
R(r, rk - 1, 0) {
int pc = pivc[r];
L(rr, 0, r - 1) if(getb(rr, pc)) xorow(rr, r);
}
}
inline void apply_ops(vector<unsigned char> &rhs) const {
for(auto &o : ops) {
if(o.t == SW) swap(rhs[o.x], rhs[o.y]);
else rhs[o.x] ^= rhs[o.y];
}
}
bool solvable(const vector<unsigned char> &c) const {
vector<unsigned char> rhs = c;
apply_ops(rhs);
L(r, rk, m - 1) {
bool z = 1;
L(i, 0, W - 1) if(a[r][i]) { z = 0; break; }
if(z && rhs[r]) return 0;
}
return 1;
}
vector<unsigned char> solve1(const vector<unsigned char> &c) const {
vector<unsigned char> rhs = c;
apply_ops(rhs);
vector<unsigned char> ypos(n, 0);
R(r, rk - 1, 0) {
int pc = pivc[r];
unsigned char s = rhs[r];
L(col, pc + 1, n - 1) if(getb(r, col)) s ^= ypos[col];
ypos[pc] = s;
}
vector<unsigned char> y(n, 0);
L(pos, 0, n - 1) y[perm[pos]] = ypos[pos];
return y;
}
};
int main() {
ios::sync_with_stdio(false);
cin.tie(nullptr);
int m = 3, n = 4;
GF2 S(m, n);
int Araw[3][4] = {
{1,0,1,1},
{0,1,1,0},
{1,1,0,1}
};
L(i, 0, m - 1) L(j, 0, n - 1) S.setb(i, j, Araw[i][j]);
S.preprocess_reduced();
vector<unsigned char> c = {1,0,1};
if(S.solvable(c)) {
auto y = S.solve1(c);
cout << "Solvable, one y = ";
L(i, 0, n - 1) cout << int(y[i]) << " \n"[i == n - 1];
} else cout << "No solution\n";
return 0;
}
Part3
模 \(2^k\) 的提升。每一轮都要解一次
\(A_2 y_t \equiv c_t \pmod 2\)
\(A_2\) 固定,所以 Part2 的 solve1 直接拿来用。
更新:
\(x_{t+1} = x_t + 2^t y_t\)
残差:
\(r_t \equiv b - A x_t \pmod{2^{t+1}}\)
取 RHS:
\(c_t \equiv (r_t / 2^t) \pmod 2\)
每轮就是:算 \(A x_t\) 的低 \(t+1\) 位,取残差第 \(t\) 位,solve1,加回去。中间某一轮 solvable 过不去就无解。
代码(ull 演示,\(k\le 62\)):
#include<bits/stdc++.h>
#define L(i, j, k) for(int i = (j); i <= (k); ++i)
#define R(i, j, k) for(int i = (j); i >= (k); --i)
#define ll long long
#define ull unsigned long long
#define vi vector<int>
#define sz(a) ((int)(a).size())
using namespace std;
struct GF2 {
int m, n, W, rk;
vector<vector<ull>> a;
vector<int> perm;
vector<int> pivc, pivr;
enum OP { SW, XO };
struct op { int t, x, y; };
vector<op> ops;
GF2(int _m = 0, int _n = 0) { init(_m, _n); }
void init(int _m, int _n) {
m = _m, n = _n;
W = (n + 63) >> 6;
a.assign(m, vector<ull>(W, 0));
perm.resize(n);
iota(perm.begin(), perm.end(), 0);
ops.clear();
rk = 0;
}
inline void setb(int r, int c, int v) {
if(v) a[r][c >> 6] |= 1ULL << (c & 63);
else a[r][c >> 6] &= ~(1ULL << (c & 63));
}
inline int getb(int r, int cpos) const {
return (a[r][cpos >> 6] >> (cpos & 63)) & 1ULL;
}
inline void swrow(int x, int y) {
if(x == y) return;
swap(a[x], a[y]);
ops.push_back({SW, x, y});
}
inline void xorow(int x, int y) {
if(x == y) return;
L(i, 0, W - 1) a[x][i] ^= a[y][i];
ops.push_back({XO, x, y});
}
inline void swcol(int c1, int c2) {
if(c1 == c2) return;
L(r, 0, m - 1) {
int b1 = getb(r, c1), b2 = getb(r, c2);
if(b1 ^ b2) {
a[r][c1 >> 6] ^= 1ULL << (c1 & 63);
a[r][c2 >> 6] ^= 1ULL << (c2 & 63);
}
}
swap(perm[c1], perm[c2]);
}
void preprocess() {
pivc.assign(m, -1);
pivr.assign(n, -1);
int r = 0;
L(c, 0, n - 1) if(r < m) {
int bc = -1, br = -1;
L(cc, c, n - 1) if(bc == -1) {
L(rr, r, m - 1) if(getb(rr, cc)) { bc = cc, br = rr; break; }
}
if(bc == -1) break;
if(bc != c) swcol(bc, c);
if(br != r) swrow(br, r);
pivc[r] = c;
pivr[c] = r;
L(rr, r + 1, m - 1) if(getb(rr, c)) xorow(rr, r);
++r;
}
rk = 0;
L(i, 0, m - 1) if(pivc[i] != -1) ++rk;
}
void preprocess_reduced() {
preprocess();
R(r, rk - 1, 0) {
int pc = pivc[r];
L(rr, 0, r - 1) if(getb(rr, pc)) xorow(rr, r);
}
}
inline void apply_ops(vector<unsigned char> &rhs) const {
for(auto &o : ops) {
if(o.t == SW) swap(rhs[o.x], rhs[o.y]);
else rhs[o.x] ^= rhs[o.y];
}
}
bool solvable(const vector<unsigned char> &c) const {
vector<unsigned char> rhs = c;
apply_ops(rhs);
L(r, rk, m - 1) {
bool z = 1;
L(i, 0, W - 1) if(a[r][i]) { z = 0; break; }
if(z && rhs[r]) return 0;
}
return 1;
}
vector<unsigned char> solve1(const vector<unsigned char> &c) const {
vector<unsigned char> rhs = c;
apply_ops(rhs);
vector<unsigned char> ypos(n, 0);
R(r, rk - 1, 0) {
int pc = pivc[r];
unsigned char s = rhs[r];
L(col, pc + 1, n - 1) if(getb(r, col)) s ^= ypos[col];
ypos[pc] = s;
}
vector<unsigned char> y(n, 0);
L(pos, 0, n - 1) y[perm[pos]] = ypos[pos];
return y;
}
};
struct Lift2k {
int m, n;
GF2 S;
vector<vector<ull>> A;
Lift2k(int _m=0,int _n=0):m(_m),n(_n),S(_m,_n){
A.assign(m, vector<ull>(n, 0));
}
void setA(int i,int j, ull v){
A[i][j]=v;
S.setb(i,j, (int)(v & 1ULL));
}
void preprocess(){
S.preprocess_reduced();
}
vector<ull> Ax_mod(const vector<ull>& x, int e) const{
ull MOD = 1ULL << e;
ull mask = MOD - 1;
vector<ull> y(m, 0);
L(i,0,m-1){
__uint128_t s = 0;
L(j,0,n-1){
s += (__uint128_t)(A[i][j] & mask) * (__uint128_t)(x[j] & mask);
}
y[i] = (ull)(s & mask);
}
return y;
}
pair<bool, vector<ull>> solve(const vector<ull>& b, int k) const{
vector<ull> x(n, 0);
ull maskk = (1ULL<<k) - 1;
L(t, 0, k-1){
int e = t + 1;
ull MOD = 1ULL << e;
ull mask = MOD - 1;
auto Ax = Ax_mod(x, e);
vector<unsigned char> ct(m, 0);
L(i,0,m-1){
ull r = ( (b[i] & mask) + MOD - (Ax[i] & mask) ) & mask;
ct[i] = (unsigned char)((r >> t) & 1ULL);
}
if(!S.solvable(ct)) return {false, {}};
auto y = S.solve1(ct);
L(j,0,n-1) if(y[j]) x[j] = (x[j] + (1ULL<<t)) & maskk;
}
return {true, x};
}
};
int main() {
ios::sync_with_stdio(false);
cin.tie(nullptr);
int m = 3, n = 4;
Lift2k L2(m, n);
ull Araw[3][4] = {
{1,0,1,1},
{0,1,1,0},
{1,1,0,1}
};
L(i,0,m-1) L(j,0,n-1) L2.setA(i, j, Araw[i][j]);
L2.preprocess();
int k = 10;
vector<ull> b = {5,7,9};
auto res = L2.solve(b, k);
if(!res.first) {
cout << "No solution mod 2^" << k << "\n";
return 0;
}
auto x = res.second;
cout << "One solution x mod 2^" << k << ":\n";
L(i,0,n-1) cout << x[i] << " \n"[i==n-1];
ull maskk = (1ULL<<k) - 1;
auto Ax = L2.Ax_mod(x, k);
bool ok = 1;
L(i,0,m-1) if(((Ax[i] - b[i]) & maskk) != 0) ok = 0;
cout << "check = " << ok << "\n";
return 0;
}

浙公网安备 33010602011771号