不开O2会tle,被卡常求助
查看原帖
不开O2会tle,被卡常求助
349344
BotBw楼主2022/8/26 21:32

板子很长,但是应该没有地方写的太烂,有什么优化常数的建议吗. 正确性应该没有问题,不知道是不是modint或者segtree写的太烂导致常数过大。

#include <bits/stdc++.h>
using namespace std;
#ifdef LOCAL
#include "debug.h"
#include "leet.h"
#include "mySTL.h"
#else
#define debug(...)
#endif
#define FOR(i, a, b) for(int i = (a); i <= (int)(b); ++i)
#define _FOR(i, a, b) for(int i = (a); i >= (int)(b); --i)
#define INT_INF 0x3f3f3f3f
#define LLONG_INF 0x3f3f3f3f3f3f3f3f
typedef long long ll;
typedef pair<int, int> pii;
typedef pair<ll, ll> pll;
#ifndef SEGTREE_H
#define SEGTREE_H

#include <vector>

#define lc(x) ((x) * 2)
#define rc(x) ((x) * 2 + 1)
#define len(x) (t[(x)].r - t[(x)].l + 1)

template <typename tag, void(*push_up)(tag&, const tag&, const tag&), void(*push_down)(tag&, tag&, int, tag&, int), tag(*ie)()>
class segtree {
  struct node {
    int l, r;
    tag val;
  };

  int lo, hi;
  std::vector<node> t;

  void build(int x, int l, int r, const vector<tag> &init) {
    t[x] = {l, r, ie()};
    if(l == r) {
      t[x] = {l, r, init[l]};
      return;
    }
    int mi = (l + r) / 2;
    build(lc(x), l, mi, init);
    build(rc(x), mi + 1, r, init);
    push_up(t[x].val, t[lc(x)].val, t[rc(x)].val);
  }

  void build(int x, int l, int r, const tag* init) {
    t[x] = {l, r, ie()};
    if(l == r) {
      t[x] = {l, r, init[l]};
      return;
    }
    int mi = (l + r) / 2;
    build(lc(x), l, mi, init);
    build(rc(x), mi + 1, r, init);
    push_up(t[x].val, t[lc(x)].val, t[rc(x)].val);
  }
  
  template <typename modifier>
  void _update(int x, int l, int r, modifier m) {
    if(l <= t[x].l && t[x].r <= r) {
      m(t[x].l, t[x].r, t[x].val);
      return;
    }
    push_down(t[x].val, t[lc(x)].val, len(lc(x)), t[rc(x)].val, len(rc(x)));
    int mi = (t[x].l + t[x].r) / 2;
    if(l <= mi) _update(lc(x), l, r, m);
    if(r >= mi + 1) _update(rc(x), l, r, m);
    push_up(t[x].val, t[lc(x)].val, t[rc(x)].val);
  }

  tag _query(int x, int l, int r) {
    if(l <= t[x].l && t[x].r <= r) return t[x].val;
    push_down(t[x].val, t[lc(x)].val, len(lc(x)), t[rc(x)].val, len(rc(x)));
    int mi = (t[x].l + t[x].r) / 2;
    if(l <= mi && r >= mi + 1) {
      tag L = _query(lc(x), l, r);
      tag R = _query(rc(x), l, r);
      tag ret;
      push_up(ret, L, R);
      return ret;
    } else if(l <= mi) return _query(lc(x), l, r);
    else return _query(rc(x), l, r);
  }


 public:
  segtree(int _lo, int _hi): lo(_lo), hi(_hi) {
    int n = hi - lo + 1;
    vector<tag> init(n + 1, ie());
    t = vector<node>(4*n);
    build(1, lo, hi, init);
  }

  segtree(int _lo, int _hi, const vector<tag> &init): lo(_lo), hi(_hi) {
    int n = hi - lo + 1;
    t = vector<node>(4*n);
    build(1, lo, hi, init);
  }

  segtree(int _lo, int _hi, const tag *init): lo(_lo), hi(_hi) {
    int n = hi - lo + 1;
    t = vector<node>(4*n);
    build(1, lo, hi, init);
  }

  template <typename modifier>
  void update(int x, modifier m) {
    assert(lo <= x && x <= hi);
    _update(1, x, x, m);
  }

  template <typename modifier>
  void update(int l, int r, modifier m) {
    assert(l <= r && lo <= l && r <= hi);
    _update(1, l, r, m);
  }

  tag query(int x) {
    assert(lo <= x && x <= hi);
    return _query(1, x, x);
  }

  tag query(int l, int r) {
    assert(l <= r && lo <= l && r <= hi);
    return _query(1, l, r);
  }
};

#undef lc
#undef rc
#undef len

#endif
#ifndef MODINT_H
#define MODINT_H

#include <iostream>
#include <vector>
using namespace std;

// modint
struct modint {
  typedef long long ll;
  static int MOD;
  ll val;

  modint(ll v = 0) : val(v % MOD) {
    if (val < 0) val += MOD;
  }

  inline static void setmod(int _MOD) { MOD = _MOD; }
  inline int getmod() const { return MOD; }
  inline modint operator-() const { return val ? MOD - val : 0; }
  inline modint operator+(const modint& r) const {
    return modint(*this) += r;
  }
  inline modint operator-(const modint& r) const {
    return modint(*this) -= r;
  }
  inline modint operator*(const modint& r) const {
    return modint(*this) *= r;
  }
  inline modint operator/(const modint& r) const {
    return modint(*this) /= r;
  }
  inline modint& operator+=(const modint& r) {
    val += r.val;
    if (val >= MOD) val -= MOD;
    return *this;
  }
  inline modint& operator-=(const modint& r) noexcept {
    val -= r.val;
    if (val < 0) val += MOD;
    return *this;
  }
  inline modint& operator*=(const modint& r) noexcept {
    val = (val % MOD) * (r.val % MOD) % MOD;
    return *this;
  }
  inline modint& operator/=(const modint& r) noexcept {
    ll a = r.val, b = MOD, u = 1, v = 0;
    while (b) {
      ll t = a / b;
      a -= t * b, swap(a, b);
      u -= t * v, swap(u, v);
    }
    val = val * u % MOD;
    if (val < 0) val += MOD;
    return *this;
  }
  bool operator==(const modint& r) const noexcept { return this->val == r.val; }
  bool operator!=(const modint& r) const noexcept { return this->val != r.val; }
  friend istream& operator>>(istream& is, modint& x) noexcept {
    is >> x.val;
    x.val %= MOD;
    if (x.val < 0) x.val += MOD;
    return is;
  }
  friend ostream& operator<<(ostream& os, const modint& x) noexcept {
    return os << x.val;
  }
  friend modint modpow(const modint& r, int n) noexcept {
    if (n == 0) return 1;
    if (n < 0) return modpow(modinv(r), -n);
    auto t = modpow(r, n / 2);
    t = t * t;
    if (n & 1) t = t * r;
    return t;
  }
  friend modint modinv(const modint& r) noexcept {
    ll a = r.val, b = MOD, u = 1, v = 0;
    while (b) {
      ll t = a / b;
      a -= t * b, swap(a, b);
      u -= t * v, swap(u, v);
    }
    return modint(u);
  }
};

int modint::MOD = 1000000007;

#endif

struct tag {
  modint sum, add, mul;
};

void push_up(tag& fa, const tag &l, const tag &r) {
  fa.sum = l.sum + r.sum;
}

void push_down(tag &fa, tag &l, int ll, tag &r, int lr) {
  modint add = fa.add, mul = fa.mul;
  fa.add = 0;
  fa.mul = 1;

  l.sum *= mul;
  l.sum += modint(ll) * add;
  l.mul *= mul;
  l.add *= mul;
  l.add += add;

  r.sum *= mul;
  r.sum += modint(lr) * add;
  r.mul *= mul;
  r.add *= mul;
  r.add += add;
}

tag ie() {
  return {0, 0, 1};
}

int main() {
  ios::sync_with_stdio(0); cin.tie(0); cout.tie(0);
  int n, m, p;
  cin >> n >> m >> p;
  modint::setmod(p);
  vector<tag> a(n + 1, ie());
  FOR(i, 1, n) cin >> a[i].sum;
  segtree<tag, push_up, push_down, ie> seg(1, n, a);
  FOR(i, 1, m) {
    int op;
    cin >> op;
    if(op == 1) {
      int x, y, k;
      cin >> x >> y >> k;
      if(x > y) swap(x, y);
      seg.update(x, y, [&](int l, int r, tag &x) {
        x.add *= k;
        x.mul *= k;
        x.sum *= k;
      });
    } else if (op == 2) {
      int x, y, k;
      cin >> x >> y >> k;
      if(x > y) swap(x, y);
      seg.update(x, y, [&](int l, int r, tag &x) {
        x.add += k;
        x.sum += (r - l + 1) * k;
      });
    } else {
      int x, y;
      cin >> x >> y;
      cout << seg.query(x, y).sum << '\n';
    }
  }
  return 0;
}
2022/8/26 21:32
加载中...