线段树模板

 单点修线段树:

#include "bits/stdc++.h"
using namespace std;

// 1. 定义节点信息
struct Info {
  long long sum = 0; // 区间求和建议开 long long 防止溢出
};

// 2. 重载 + 运算符用于合并节点信息
Info operator+(const Info &a, const Info &b) {
  return {a.sum + b.sum};
}

struct SegmentTree {
  vector<Info> info;
  int n;

  // 默认构造(建空树)
  SegmentTree(int size) {
    n = size;
    info.assign(4 * (n + 5), Info{});
    build(1, 0, n);
  }

  // 用初始数组建树
  SegmentTree(const vector<long long>& arr) {
    n = arr.size();
    info.assign(4 * (n + 5), Info{});
    build(1, 0, n, arr);
  }

  void pull(int p) {
    info[p] = info[2 * p] + info[2 * p + 1];
  }

  // 辅助函数:空建树
  void build(int p, int l, int r) {
    if (r - l == 1) {
      info[p] = Info();
      return;
    }
    int m = l + (r - l) / 2;
    build(2 * p, l, m);
    build(2 * p + 1, m, r);
    pull(p);
  }

  // 辅助函数:带初始数组建树
  void build(int p, int l, int r, const vector<long long> &arr) {
    if (r - l == 1) {
      info[p] = {arr[l]};
      return;
    }
    int m = l + (r - l) / 2;
    build(2 * p, l, m, arr);
    build(2 * p + 1, m, r, arr);
    pull(p);
  }

  // 单点修改:将 pos 位置的值修改为 v 
  // (如果是 "在原基础上增加 v",把 info[p].sum = v 改为 += 即可)
  void modify(int p, int l, int r, int pos, long long v) {
    if (r - l == 1) {
      info[p].sum = v; 
      return;
    }
    int m = l + (r - l) / 2;
    if (pos < m) {
      modify(2 * p, l, m, pos, v);
    } else {
      modify(2 * p + 1, m, r, pos, v);
    }
    pull(p);
  }

  // 区间查询:返回 Info 结构体
  Info rangeQuery(int p, int l, int r, int x, int y) {
    // 越界返回单位元(sum 默认为 0)
    if (l >= y || r <= x) {
      return Info(); 
    }
    // 完全包含直接返回
    if (l >= x && r <= y) {
      return info[p];
    }
    int m = l + (r - l) / 2;
    return rangeQuery(2 * p, l, m, x, y) + rangeQuery(2 * p + 1, m, r, x, y);
  }

  // --- 外部调用接口 ---
  
  // 单点修改
  void modify(int pos, long long v) { 
      modify(1, 0, n, pos, v); 
  }
  
  // 区间查询(返回区间 [l, r) 的总和)
  long long rangeQuery(int l, int r) { 
      return rangeQuery(1, 0, n, l, r).sum; 
  }
};

#include "bits/stdc++.h"
using namespace std;
struct Info {
  int x = 0;
  int sum = 1;
};

struct SegmentTree {
  vector<Info> info;
  int n;
  SegmentTree(int size) {
    n = size;
    info.assign(4 * (n + 5), Info{});
  
    build(1, 0, n);
  }
  void pull(int p) {
    if (info[p].x == 0)
      info[p].sum = 1;
    else
      info[p].sum = info[2 * p].sum + info[2 * p + 1].sum;
  }

  void build(int p, int l, int r) {
    if (r - l == 1) {
      info[p] = Info();
      return;
    }
    int m = (l + r + 1) / 2;
    build(2 * p, l, m);
    build(2 * p + 1, m, r);
    pull(p);
  }

  void modify(int p, int l, int r, int x, int y) {
    if (l >= y || r <= x) {
      return;
    }
    if (l >= x && r <= y) {
      info[p].x = 1;
      pull(p);
      return;
    }
    int m = (l + r + 1) / 2;
    modify(2 * p, l, m, x, y);
    modify(2 * p + 1, m, r, x, y);
    pull(p);
  }

  int rangeQuery(int p, int l, int r, int x, int y) {
    if (l >= y || r <= x) {
      return 0;
    }
    if (l >= x && r <= y) {
      return info[p].sum;
    }
    int m = (l + r + 1) / 2;
    return rangeQuery(2 * p, l, m, x, y) + rangeQuery(2 * p + 1, m, r, x, y) +
           !info[p].x;
  }

  void modify(int x, int y) { modify(1, 0, n, x, y); }
  int rangeQuery(int l, int r) { return rangeQuery(1, 0, n, l, r); }
};

懒标记线段树:

#include "bits/stdc++.h"
using namespace  std;

struct Tag {
  int tag = 0;
  void apply(const Tag &t) { tag += t.tag; }
};

struct Info {
  int x = -2e9;
  void apply(const Tag &t) { x += t.tag; }
};

Info operator+(const Info &a, const Info &b) { return {max(a.x, b.x)}; };


struct SegmentTree {
  vector<Info> info;
  vector<Tag> tag;
  int n;
  // 选一个写
  SegmentTree(int size) {
    n = size;
    info.assign(4 * (n + 5), Info{});
    tag.assign(4 * (n + 5), Tag{});
    build(1, 0, n);
  }
  void build(int p, int l, int r) {
    if (r - l == 1) {
      info[p] = Info();
      return;
    }
    int m = (l + r) / 2;
    build(2 * p, l, m);
    build(2 * p + 1, m, r);
    pull(p);
  }
  SegmentTree(const vector<Info> &arr) {
    n = arr.size();
    info.assign(4 * (n + 5), Info{});
    tag.assign(4 * (n + 5), Tag{});
    build(1, 0, n, arr);
  }
  void build(int p, int l, int r, const vector<Info> &arr) {
    if (r - l == 1) {
      info[p] = arr[l];
      return;
    }
    int m = (l + r) / 2;
    build(2 * p, l, m, arr);
    build(2 * p + 1, m, r, arr);
    pull(p);
  }
  void pull(int p) { info[p] = info[2 * p] + info[2 * p + 1]; }

  void applyNode(int p, const Tag &v) {
    info[p].apply(v);
    tag[p].apply(v);
  }

  void push(int p) {
    applyNode(2 * p, tag[p]);
    applyNode(2 * p + 1, tag[p]);
    tag[p] = Tag();
  }

  void modify(int p, int l, int r, int x, const Info &v) {
    if (r - l == 1) {
      info[p] = v;
      return;
    }
    int m = (l + r) / 2;
    push(p);
    if (x < m) {
      modify(2 * p, l, m, x, v);
    } else {
      modify(2 * p + 1, m, r, x, v);
    }
    pull(p);
  }

  Info rangeQuery(int p, int l, int r, int x, int y) {
    if (l >= y || r <= x) {
      return Info();
    }
    if (l >= x && r <= y) {
      return info[p];
    }
    int m = (l + r) / 2;
    push(p);
    return rangeQuery(2 * p, l, m, x, y) + rangeQuery(2 * p + 1, m, r, x, y);
  }

  void rangeApply(int p, int l, int r, int x, int y, const Tag &v) {
    if (l >= y || r <= x) {
      return;
    }
    if (l >= x && r <= y) {
      applyNode(p, v);
      return;
    }
    int m = (l + r) / 2;
    push(p);
    rangeApply(2 * p, l, m, x, y, v);
    rangeApply(2 * p + 1, m, r, x, y, v);
    pull(p);
  }
  template <class F>
  int findFirst(int p, int l, int r, int x, int y, F &&pred) {
    if (l >= y || r <= x) {
      return -1;
    }
    if (l >= x && r <= y && !pred(info[p])) {
      return -1;
    }
    if (r - l == 1) {
      return l;
    }
    int m = (l + r) / 2;
    push(p);
    int res = findFirst(2 * p, l, m, x, y, pred);
    if (res == -1) {
      res = findFirst(2 * p + 1, m, r, x, y, pred);
    }
    return res;
  }

  template <class F> int findLast(int p, int l, int r, int x, int y, F &&pred) {
    if (l >= y || r <= x) {
      return -1;
    }
    if (l >= x && r <= y && !pred(info[p])) {
      return -1;
    }
    if (r - l == 1) {
      return l;
    }
    int m = (l + r) / 2;
    push(p);
    int res = findLast(2 * p + 1, m, r, x, y, pred);
    if (res == -1) {
      res = findLast(2 * p, l, m, x, y, pred);
    }
    return res;
  }
  template <class F> int findLast(int l, int r, F &&pred) {
    return findLast(1, 0, n, l, r, pred);
  }
  template <class F> int findFirst(int l, int r, F &&pred) {
    return findFirst(1, 0, n, l, r, pred);
  }
  void modify(int x, const Info &v) { modify(1, 0, n, x, v); }
  Info rangeQuery(int x, int y) { return rangeQuery(1, 0, n, x, y); }
  void rangeApply(int x, int y, const Tag &v) { rangeApply(1, 0, n, x, y, v); }
};

int main(){

}

这里是极简版的线段树二分(findFirst / findLast)使用总结:

核心原理: 利用线段树节点维护的 Info 进行剪枝,时间内找到区间 [l, r) 中满足条件的点。

谓词函数 pred 怎么写:

传入 Lambda 表达式,判断当前区间有没有可能包含答案:

  • 返回 true:这区间可能有答案,放行向下搜。

  • 返回 false:这区间绝对没答案,直接剪枝(提速关键)。

极简代码模板(以寻找区间 [l, r) 第一个 >= k 的位置为例):

C++

 

// 1. 写条件:当前区间最大值 >= k,说明里面可能有我们要的数
auto pred = [&](const Info& info) {
    return info.x >= k; 
};

// 2. 调接口:找第一个用 findFirst,找最后一个用 findLast
int pos = seg.findFirst(l, r, pred);

// 3. 判无解:找不到一定会返回 -1,千万记得特判!
if (pos != -1) {
    // 找到了,pos 是下标
}

三大避坑指南:

  1. 前提条件:你的 Info 必须维护了能用于判断的属性(比如找 >= k,Info 里必须有区间最大值)。

  2. pred 的逻辑:是判断“区间是否有潜力”,而不是“区间是否全满足”。

  3. 防 RE:一定要处理返回的 -1,别直接拿去当数组下标。

开点线段树:

#include "bits/stdc++.h"
using namespace  std;

struct node{
	  int v;
      int l;
	  int c;
};

const int N = 5e5 + 10;
vector<node> g[N];

template<class Info>
struct DynamicSegmentTree {
    struct Node {
        int ls = 0; // 左孩子下标
        int rs = 0; // 右孩子下标
        Info info;
    };

    int n; 
     vector<Node> nodes; 

    DynamicSegmentTree(int n_) : n(n_) {
        nodes.push_back(Node{});
    }

    int newNode() {
        nodes.emplace_back();
        return (int)nodes.size() - 1;
    }

    // 上拉更新
    void pull(int p) {
   
        const Info &lhs = nodes[p].ls ? nodes[nodes[p].ls].info : Info();
        const Info &rhs = nodes[p].rs ? nodes[nodes[p].rs].info : Info();
        nodes[p].info = lhs + rhs;
    }

    //如果是新的树开一个点
    //如果是旧树不变
    int modify(int p, int l, int r, int x, const Info &v) {
        if (!p) p = newNode(); // 动态开点
        
        if (r - l == 1) {
            nodes[p].info = v;
            return p;
        }
        
        int m = l + (r - l) / 2;
        
        // 移除 push(p)
        
        if (x < m) {
         
            int new_ls = modify(nodes[p].ls, l, m, x, v);
            nodes[p].ls = new_ls;
        } else {
            int new_rs = modify(nodes[p].rs, m, r, x, v);
            nodes[p].rs = new_rs;
        }
        
        pull(p); // 更新当前节点信息
        return p;
    }

    // 外部接口:单点修改
    void modify(int &root, int x, const Info &v) {
        root = modify(root, 0, n, x, v);
    }

    // 区间查询
    Info rangeQuery(int p, int l, int r, int x, int y) {
        // 如果节点不存在,或者区间无交集,返回单位元
        if (!p || l >= y || r <= x) {
            return Info();
        }
        // 完全包含
        if (l >= x && r <= y) {
            return nodes[p].info;
        }
        
        int m = l + (r - l) / 2;

        
        return rangeQuery(nodes[p].ls, l, m, x, y) + rangeQuery(nodes[p].rs, m, r, x, y);
    }

    // 外部接口:区间查询
    Info rangeQuery(int root, int l, int r) {
        return rangeQuery(root, 0, n, l, r);
    }
    
    // 线段树二分:查找区间 [x, y) 内第一个满足 pred 的位置
    template<class F>
    int findFirst(int p, int l, int r, int x, int y, F &&pred) {
        if (!p || l >= y || r <= x) return -1;
        
        // 剪枝:如果当前区间的信息都不满足条件,直接返回
        if (l >= x && r <= y && !pred(nodes[p].info)) return -1;
        
        if (r - l == 1) return l;
        
        int m = l + (r - l) / 2;
        
        // 移除 push(p)
        
        int res = findFirst(nodes[p].ls, l, m, x, y, pred);
        if (res == -1) res = findFirst(nodes[p].rs, m, r, x, y, pred);
        return res;
    }

    template<class F>
    int findFirst(int root, int l, int r, F &&pred) {
        return findFirst(root, 0, n, l, r, pred);
    }
  // --- 你要求的 findLast ---
    template<class F>
    int findLast(int p, int l, int r, int x, int y, F &&pred) {
        if (!p || l >= y || r <= x) return -1;
        if (l >= x && r <= y && !pred(nodes[p].info)) return -1;
        if (r - l == 1) return l;
        
        int m = l + (r - l) / 2;
        
        // 优先找右边
        int res = findLast(nodes[p].rs, m, r, x, y, pred);
        if (res == -1) res = findLast(nodes[p].ls, l, m, x, y, pred);
        return res;
    }

    template<class F>
    int findLast(int root, int l, int r, F &&pred) {
        return findLast(root, 0, n, l, r, pred);
    }

};



struct Info {  
    int mx = -1; // 默认为 0 或 -INF,视具体题目而定
};

// 重载 + 运算符用于合并
Info operator+(const Info &a, const Info &b) {
    Info res;
    res.mx = max(a.mx, b.mx);
    return res;
}

开点懒标记线段树:

#include "bits/stdc++.h"
using namespace  std;

struct node{
	  int v;
      int l;
	  int c;
};

const int N = 5e5 + 10;
vector<node> g[N];

template<class Info>
struct DynamicSegmentTree {
    struct Node {
        int ls = 0; // 左孩子下标
        int rs = 0; // 右孩子下标
        Info info;
    };

    int n; 
     vector<Node> nodes; 

    DynamicSegmentTree(int n_) : n(n_) {
        nodes.push_back(Node{});
    }

    int newNode() {
        nodes.emplace_back();
        return (int)nodes.size() - 1;
    }

    // 上拉更新
    void pull(int p) {
   
        const Info &lhs = nodes[p].ls ? nodes[nodes[p].ls].info : Info();
        const Info &rhs = nodes[p].rs ? nodes[nodes[p].rs].info : Info();
        nodes[p].info = lhs + rhs;
    }

    //如果是新的树开一个点
    //如果是旧树不变
    int modify(int p, int l, int r, int x, const Info &v) {
        if (!p) p = newNode(); // 动态开点
        
        if (r - l == 1) {
            nodes[p].info = v;
            return p;
        }
        
        int m = l + (r - l) / 2;
        
        // 移除 push(p)
        
        if (x < m) {
         
            int new_ls = modify(nodes[p].ls, l, m, x, v);
            nodes[p].ls = new_ls;
        } else {
            int new_rs = modify(nodes[p].rs, m, r, x, v);
            nodes[p].rs = new_rs;
        }
        
        pull(p); // 更新当前节点信息
        return p;
    }

    // 外部接口:单点修改
    void modify(int &root, int x, const Info &v) {
        root = modify(root, 0, n, x, v);
    }

    // 区间查询
    Info rangeQuery(int p, int l, int r, int x, int y) {
        // 如果节点不存在,或者区间无交集,返回单位元
        if (!p || l >= y || r <= x) {
            return Info();
        }
        // 完全包含
        if (l >= x && r <= y) {
            return nodes[p].info;
        }
        
        int m = l + (r - l) / 2;

        
        return rangeQuery(nodes[p].ls, l, m, x, y) + rangeQuery(nodes[p].rs, m, r, x, y);
    }

    // 外部接口:区间查询
    Info rangeQuery(int root, int l, int r) {
        return rangeQuery(root, 0, n, l, r);
    }
    
    // 线段树二分:查找区间 [x, y) 内第一个满足 pred 的位置
    template<class F>
    int findFirst(int p, int l, int r, int x, int y, F &&pred) {
        if (!p || l >= y || r <= x) return -1;
        
        // 剪枝:如果当前区间的信息都不满足条件,直接返回
        if (l >= x && r <= y && !pred(nodes[p].info)) return -1;
        
        if (r - l == 1) return l;
        
        int m = l + (r - l) / 2;
        
        // 移除 push(p)
        
        int res = findFirst(nodes[p].ls, l, m, x, y, pred);
        if (res == -1) res = findFirst(nodes[p].rs, m, r, x, y, pred);
        return res;
    }

    template<class F>
    int findFirst(int root, int l, int r, F &&pred) {
        return findFirst(root, 0, n, l, r, pred);
    }
  // --- 你要求的 findLast ---
    template<class F>
    int findLast(int p, int l, int r, int x, int y, F &&pred) {
        if (!p || l >= y || r <= x) return -1;
        if (l >= x && r <= y && !pred(nodes[p].info)) return -1;
        if (r - l == 1) return l;
        
        int m = l + (r - l) / 2;
        
        // 优先找右边
        int res = findLast(nodes[p].rs, m, r, x, y, pred);
        if (res == -1) res = findLast(nodes[p].ls, l, m, x, y, pred);
        return res;
    }

    template<class F>
    int findLast(int root, int l, int r, F &&pred) {
        return findLast(root, 0, n, l, r, pred);
    }

};



struct Info {  
    int mx = -1; // 默认为 0 或 -INF,视具体题目而定
};

// 重载 + 运算符用于合并
Info operator+(const Info &a, const Info &b) {
    Info res;
    res.mx = max(a.mx, b.mx);
    return res;
}

 

posted @ 2026-05-15 17:04  r_123  阅读(10)  评论(0)    收藏  举报