家人们,谁懂啊,为什么指针线段树最后一个点TLE了啊?!
查看原帖
家人们,谁懂啊,为什么指针线段树最后一个点TLE了啊?!
817044
cjwdyzxfblzs楼主2023/5/18 21:27

奉上我的指针版线段树,求调

#include <bits/stdc++.h>
using namespace std;
#define int long long
const int N = 500005;
const int M = 20000005;
int n, m;
int a[N];
int gcd(int a, int b)
{
    if (!b) 
        return a;
    else
        return gcd(b, a % b);
}
struct node
{
    node *ls, *rs;
    int l, r;
    int val;
    // int size;
    int tag;
    explicit node()
    {
        ls = rs = nullptr;
        // size = 0;
        val = 0;
        tag = 0;
        l = r = 0;
    }
};
node *root = new node();
void build(node *u, int l, int r)
{
    if (u == nullptr)
        return;
    u->ls = new node();
    u->rs = new node();
    u->l = l;
    u->r = r;
    if (l == r)
    {
        u->val = a[l];
        return;
    }
    int mid = (l + r) >> 1;
    build(u->ls, l, mid);
    build(u->rs, mid + 1, r);
    return;
}
void push_down(node *u)
{
    if (u == nullptr)
        return;
    if (u->tag > 0)
    {
        if (u->ls)
        {
            if (u->ls->l == u->ls->r)
                u->ls->val += u->tag;
            else
                u->ls->tag += u->tag;
        }
        if (u->rs)
        {
            if (u->rs->r == u->rs->l)
                u->rs->val += u->tag;
            else
                u->rs->tag += u->tag;
        }
        u->tag = 0;
        // if (u->ls)
        //     u->ls->tag += u->tag,
        //         u->ls->val += (u->ls->r - u->ls->l + 1) * u->tag;
        // if (u->rs)
        //     u->rs->tag += u->tag,
        //         u->rs->val += (u->rs->r - u->rs->l + 1) * u->tag;
        // if (u->ls && u->rs == nullptr)
        //     u->val += (u->ls->r - u->ls->l + 1) * u->tag;
        // if (u->rs && u->ls == nullptr)
        //     u->val += (u->rs->r - u->rs->l + 1) * u->tag;
        // if (u->ls && u->rs)
        //     u->val += (u->r - u->l + 1) * u->tag;
        // u->tag = 0;
    }
    return;
}
void add(node *u, int l, int r, int val)
{
    if (u == nullptr)
        return;
    if (l <= u->l && r >= u->r)
    {
        if (u->l == u->r)
            u->val += val;
        else
            u->tag += val;
        // u->val += val * (u->r - u->l + 1);
        // u->tag += val;
        return;
    }
    push_down(u);
    int mid = u->l + u->r >> 1;
    if (l <= mid)
        add(u->ls, l, r, val);
    if (r > mid)
        add(u->rs, l, r, val);
    return;
}
int query(int p, node *u)
{
    if (u == nullptr)
        return false;
    if (u->l == u->r)
        return u->val;
    push_down(u);
    int mid = u->l + u->r >> 1;
    if (p <= mid)
        return query(p, u->ls);
    return query(p, u->rs);
}
void output(node *u, int l, int r)
{
    if (u == nullptr)
        return;
    if (l == r)
    {
        cout << u->val << endl;
        return;
    }
    int mid = (l + r) >> 1;
    output(u->ls, l, mid);
    output(u->rs, mid + 1, r);
}
int vis[M];
int phi(int n)
{
    if (vis[n] != -1)
        return vis[n];
    int res = n, now = n;
    for (int i = 2; 1LL * i * i <= now; i++)
    {
        if (now % i == 0)
        {
            res = res / i * (i - 1);
            while (now % i == 0)
                now /= i;
        }
    }
    if (now > 1)
        res = res / now * (now - 1);
    return vis[n] = res;
}
pair<int, bool> pow(int a, int b, int mod)
{
    bool st = 0;
    int res = 1;
    if (res >= mod)
        st = true, res %= mod;
    if (a >= mod)
        st = true, a %= mod;
    while (b > 0)
    {
        if (b & 1)
        {
            res = res * a;
            if (res >= mod)
            {
                st = true;
                res = res % mod;
            }
        }
        a = a * a;
        if (a >= mod)
        {
            st = true;
            a = a % mod;
        }
        b = b >> 1;
    }
    return {res, st};
}
pair<int, bool> dfs(int l, int r, int mod)
{
    int val = query(l, root);
    if (l == r || mod == 1)
        return {val < mod ? val : val % mod + mod, val >= mod};
    int p = phi(mod);
    pair<int, bool> res = dfs(l + 1, r, p);
    if (res.second)
        res.first += p;
    return pow(val, res.first, mod);
}
signed main()
{
    // ios::sync_with_stdio(false);
    memset(vis, -1, sizeof(vis));
    cin >> n >> m;
    for (int i = 1; i <= n; i++)
        cin >> a[i];
    build(root, 1, n);
    // output(root, 1, n);
    // return 0;
    for (int i = 1; i <= m; i++)
    {
        int opt, l, r, x, p;
        cin >> opt >> l >> r;
        if (opt == 1)
        {
            cin >> x;
            add(root, l, r, x);
        }
        else
        {
            cin >> p;
            cout << dfs(l, r, p).first % p << "\n";
        }
    }
    delete root;
    return 0;
}
2023/5/18 21:27
加载中...