[TJOI2015] 旅游

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

HLD 加区间加线段树维护有向路径上的最大买卖差值。

OJ: luogu

题目 ID: P3976

难度:提高

标签:重链剖分线段树最大子段差python

日期: 2026-07-17 02:00

题意

沿 a -> b 的路径选择先买后卖的两个城市,输出最大利润,然后让路径上所有价格增加 v

思路

线段树节点维护最小值、最大值、正向最大差和反向最大差。正向合并的跨段候选是 right.max-left.min,反向则是 left.max-right.min。HLD 拆段时,a 侧区间反向、b 侧区间正向,按旅行顺序合并;查询后再做路径区间加。

Python 知识

  • 元组 (min, max, forward, backward) 让方向信息随区间一起返回。
  • reverse_info 只交换两个方向的差值,最值不变。
  • None 作为空合并结果,避免构造无效哨兵。

代码

python
import sys


sys.setrecursionlimit(1_000_000)
input = sys.stdin.buffer.readline
n = int(input())
prices = [0] + list(map(int, input().split()))
graph = [[] for _ in range(n + 1)]
for _ in range(n - 1):
    u, v = map(int, input().split())
    graph[u].append(v)
    graph[v].append(u)

parent = [0] * (n + 1)
depth = [0] * (n + 1)
order = [1]
for node in order:
    for neighbor in graph[node]:
        if neighbor != parent[node]:
            parent[neighbor] = node
            depth[neighbor] = depth[node] + 1
            order.append(neighbor)
subtree = [1] * (n + 1)
heavy = [0] * (n + 1)
for node in reversed(order[1:]):
    subtree[parent[node]] += subtree[node]
    if subtree[node] > subtree[heavy[parent[node]]]:
        heavy[parent[node]] = node
top = [0] * (n + 1)
dfn = [0] * (n + 1)
timer = 0
chains = [(1, 1)]
while chains:
    node, chain_top = chains.pop()
    while node:
        top[node] = chain_top
        timer += 1
        dfn[node] = timer
        for neighbor in graph[node]:
            if neighbor != parent[node] and neighbor != heavy[node]:
                chains.append((neighbor, neighbor))
        node = heavy[node]

minimum = [0] * (4 * n)
maximum = [0] * (4 * n)
forward = [0] * (4 * n)
backward = [0] * (4 * n)
lazy = [0] * (4 * n)
base = [0] * (n + 1)
for node in range(1, n + 1):
    base[dfn[node]] = prices[node]


def pull(node):
    left, right = node * 2, node * 2 + 1
    minimum[node] = min(minimum[left], minimum[right])
    maximum[node] = max(maximum[left], maximum[right])
    forward[node] = max(forward[left], forward[right], maximum[right] - minimum[left])
    backward[node] = max(backward[left], backward[right], maximum[left] - minimum[right])


def build(node, left, right):
    if left == right:
        minimum[node] = maximum[node] = base[left]
        return
    middle = (left + right) // 2
    build(node * 2, left, middle)
    build(node * 2 + 1, middle + 1, right)
    pull(node)


def apply(node, value):
    minimum[node] += value
    maximum[node] += value
    lazy[node] += value


def push(node):
    if lazy[node]:
        apply(node * 2, lazy[node])
        apply(node * 2 + 1, lazy[node])
        lazy[node] = 0


def update(node, left, right, query_left, query_right, value):
    if query_left <= left and right <= query_right:
        apply(node, value)
        return
    push(node)
    middle = (left + right) // 2
    if query_left <= middle:
        update(node * 2, left, middle, query_left, query_right, value)
    if middle < query_right:
        update(node * 2 + 1, middle + 1, right, query_left, query_right, value)
    pull(node)


def merge(first, second):
    if first is None:
        return second
    if second is None:
        return first
    first_min, first_max, first_best, first_reverse = first
    second_min, second_max, second_best, second_reverse = second
    return (min(first_min, second_min), max(first_max, second_max),
            max(first_best, second_best, second_max - first_min),
            max(first_reverse, second_reverse, first_max - second_min))


def query(node, left, right, query_left, query_right):
    if query_left <= left and right <= query_right:
        return minimum[node], maximum[node], forward[node], backward[node]
    push(node)
    middle = (left + right) // 2
    result = None
    if query_left <= middle:
        result = query(node * 2, left, middle, query_left, query_right)
    if middle < query_right:
        result = merge(result, query(node * 2 + 1, middle + 1, right, query_left, query_right))
    return result


def reverse_info(info):
    return info[0], info[1], info[3], info[2]


build(1, 1, n)


def path_info(x, y):
    left_parts = []
    right_parts = []
    while top[x] != top[y]:
        if depth[top[x]] >= depth[top[y]]:
            left_parts.append(reverse_info(query(1, 1, n, dfn[top[x]], dfn[x])))
            x = parent[top[x]]
        else:
            right_parts.append(query(1, 1, n, dfn[top[y]], dfn[y]))
            y = parent[top[y]]
    if depth[x] >= depth[y]:
        left_parts.append(reverse_info(query(1, 1, n, dfn[y], dfn[x])))
    else:
        right_parts.append(query(1, 1, n, dfn[x], dfn[y]))
    result = None
    for info in left_parts:
        result = merge(result, info)
    for info in reversed(right_parts):
        result = merge(result, info)
    return result


query_count = int(input())
answers = []
for _ in range(query_count):
    x, y, increase = map(int, input().split())
    info = path_info(x, y)
    answers.append(str(max(0, info[2])))
    while top[x] != top[y]:
        if depth[top[x]] < depth[top[y]]:
            x, y = y, x
        update(1, 1, n, dfn[top[x]], dfn[x], increase)
        x = parent[top[x]]
    if depth[x] > depth[y]:
        x, y = y, x
    update(1, 1, n, dfn[x], dfn[y], increase)
print("\n".join(answers))

原有 C++ 版本仍保留:

cpp
/**
 * Author by Rainboy blog: https://rainboylv.com github: https://github.com/rainboylvx
 * rbook: -> https://rbook.roj.ac.cn  https://rbook2.roj.ac.cn
 * rainboy的学习导航网站: https://idx.roj.ac.cn
 * create_at: 2026-07-17 01:04
 * update_at: 2026-07-17 01:04
 */
#include <bits/stdc++.h>
using namespace std;

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

    return 0;
}

复杂度

每次查询和更新 O(log^2 n),空间 O(n)

总结

路径有方向时,区间统计量必须同时保存正向和反向两套信息。