【模板】线段树 2

GitHub跳转原题关系图返回列表

把区间乘和区间加统一成仿射懒标记,维护模意义下的区间和。

OJ: luogu

题目 ID: P3373

难度:普及/提高-

标签:线段树懒标记区间乘区间加取模python

日期: 2026-07-16 23:59

题意

支持区间乘、区间加和区间求和,所有结果对给定模数取模。

思路

把每个元素的待执行操作写成 v -> v * mul + add。当前节点整段应用 (mul, add) 时,区间和变为 sum * mul + length * add;旧标记 (old_mul, old_add) 与新标记复合为 mul = old_mul * muladd = old_add * mul + add

Python 知识

  • 在更新处及时 % modulus,避免大整数继续膨胀。
  • list(map(int, input().split())) 同时兼容三种不同长度的操作行。
  • 用平行列表保存 summultiplyaddition,避免大量节点对象。

代码

python
import sys


sys.setrecursionlimit(1_000_000)
input = sys.stdin.buffer.readline
n, operations, modulus = map(int, input().split())
values = list(map(int, input().split()))
tree = [0] * (4 * n)
multiply = [1] * (4 * n)
addition = [0] * (4 * n)


def build(node, left, right):
    if left == right:
        tree[node] = values[left - 1] % modulus
        return
    middle = (left + right) // 2
    build(node * 2, left, middle)
    build(node * 2 + 1, middle + 1, right)
    tree[node] = (tree[node * 2] + tree[node * 2 + 1]) % modulus


def apply(node, length, mul, add):
    tree[node] = (tree[node] * mul + length * add) % modulus
    multiply[node] = multiply[node] * mul % modulus
    addition[node] = (addition[node] * mul + add) % modulus


def push(node, left, right):
    if left == right or (multiply[node] == 1 and addition[node] == 0):
        return
    middle = (left + right) // 2
    apply(node * 2, middle - left + 1, multiply[node], addition[node])
    apply(node * 2 + 1, right - middle, multiply[node], addition[node])
    multiply[node], addition[node] = 1, 0


def update(node, left, right, query_left, query_right, mul, add):
    if query_left <= left and right <= query_right:
        apply(node, right - left + 1, mul, add)
        return
    push(node, left, right)
    middle = (left + right) // 2
    if query_left <= middle:
        update(node * 2, left, middle, query_left, query_right, mul, add)
    if middle < query_right:
        update(node * 2 + 1, middle + 1, right, query_left, query_right, mul, add)
    tree[node] = (tree[node * 2] + tree[node * 2 + 1]) % modulus


def query(node, left, right, query_left, query_right):
    if query_left <= left and right <= query_right:
        return tree[node]
    push(node, left, right)
    middle = (left + right) // 2
    answer = 0
    if query_left <= middle:
        answer += query(node * 2, left, middle, query_left, query_right)
    if middle < query_right:
        answer += query(node * 2 + 1, middle + 1, right, query_left, query_right)
    return answer % modulus


build(1, 1, n)
answers = []
for _ in range(operations):
    operation = list(map(int, input().split()))
    if operation[0] == 1:
        update(1, 1, n, operation[1], operation[2], operation[3] % modulus, 0)
    elif operation[0] == 2:
        update(1, 1, n, operation[1], operation[2], 1, operation[3] % modulus)
    else:
        answers.append(str(query(1, 1, n, operation[1], operation[2])))
print("\n".join(answers))

原有 C++ 版本仍保留:

cpp
#include <bits/stdc++.h>
using namespace std;

const int MAXN = 100000 + 5;

int n, q;
long long mod_value;
long long a[MAXN];
long long seg_sum[MAXN << 2];
long long lazy_mul[MAXN << 2];
long long lazy_add[MAXN << 2];

void push_up(int u) {
    seg_sum[u] = (seg_sum[u << 1] + seg_sum[u << 1 | 1]) % mod_value;
}

void apply_mul(int u, int l, int r, long long val) {
    seg_sum[u] = seg_sum[u] * val % mod_value;
    lazy_mul[u] = lazy_mul[u] * val % mod_value;
    lazy_add[u] = lazy_add[u] * val % mod_value;
}

void apply_add(int u, int l, int r, long long val) {
    seg_sum[u] = (seg_sum[u] + (r - l + 1) * val) % mod_value;
    lazy_add[u] = (lazy_add[u] + val) % mod_value;
}

void push_down(int u, int l, int r) {
    if (l == r) {
        lazy_mul[u] = 1;
        lazy_add[u] = 0;
        return;
    }

    int mid = (l + r) >> 1;
    if (lazy_mul[u] != 1) {
        apply_mul(u << 1, l, mid, lazy_mul[u]);
        apply_mul(u << 1 | 1, mid + 1, r, lazy_mul[u]);
        lazy_mul[u] = 1;
    }
    if (lazy_add[u] != 0) {
        apply_add(u << 1, l, mid, lazy_add[u]);
        apply_add(u << 1 | 1, mid + 1, r, lazy_add[u]);
        lazy_add[u] = 0;
    }
}

void build(int u, int l, int r) {
    lazy_mul[u] = 1;
    lazy_add[u] = 0;

    if (l == r) {
        seg_sum[u] = a[l] % mod_value;
        return;
    }

    int mid = (l + r) >> 1;
    build(u << 1, l, mid);
    build(u << 1 | 1, mid + 1, r);
    push_up(u);
}

void range_mul(int u, int l, int r, int ql, int qr, long long val) {
    if (ql <= l && r <= qr) {
        apply_mul(u, l, r, val);
        return;
    }

    push_down(u, l, r);
    int mid = (l + r) >> 1;
    if (ql <= mid) {
        range_mul(u << 1, l, mid, ql, qr, val);
    }
    if (qr > mid) {
        range_mul(u << 1 | 1, mid + 1, r, ql, qr, val);
    }
    push_up(u);
}

void range_add(int u, int l, int r, int ql, int qr, long long val) {
    if (ql <= l && r <= qr) {
        apply_add(u, l, r, val);
        return;
    }

    push_down(u, l, r);
    int mid = (l + r) >> 1;
    if (ql <= mid) {
        range_add(u << 1, l, mid, ql, qr, val);
    }
    if (qr > mid) {
        range_add(u << 1 | 1, mid + 1, r, ql, qr, val);
    }
    push_up(u);
}

long long query_sum(int u, int l, int r, int ql, int qr) {
    if (ql <= l && r <= qr) {
        return seg_sum[u];
    }

    push_down(u, l, r);
    int mid = (l + r) >> 1;
    long long ans = 0;
    if (ql <= mid) {
        ans = (ans + query_sum(u << 1, l, mid, ql, qr)) % mod_value;
    }
    if (qr > mid) {
        ans = (ans + query_sum(u << 1 | 1, mid + 1, r, ql, qr)) % mod_value;
    }
    return ans;
}

int main() {
    ios::sync_with_stdio(false);
    cin.tie(nullptr);

    cin >> n >> q >> mod_value;
    for (int i = 1; i <= n; i++) {
        cin >> a[i];
    }

    build(1, 1, n);

    while (q--) {
        int op;
        cin >> op;
        if (op == 1) {
            int l, r;
            long long x;
            cin >> l >> r >> x;
            range_mul(1, 1, n, l, r, x % mod_value);
        } else if (op == 2) {
            int l, r;
            long long x;
            cin >> l >> r >> x;
            range_add(1, 1, n, l, r, x % mod_value);
        } else {
            int l, r;
            cin >> l >> r;
            cout << query_sum(1, 1, n, l, r) % mod_value << '\n';
        }
    }

    return 0;
}

复杂度

建树 O(n),每次操作 O(log n),空间 O(n)

总结

区间乘加的懒标记本质是函数复合;先写出代数公式,再实现线段树会更可靠。