Loading

对于线性代数的一点想法

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

output

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;
}
posted @ 2026-01-09 22:59  MagnusSM2  阅读(32)  评论(0)    收藏  举报