Hướng dẫn cho Bài 4: Thay Đổi Dãy (THT B Hà Tĩnh 2026)


Chỉ sử dụng khi thực sự cần thiết như một cách tôn trọng tác giả và người viết hướng dẫn này.

Chép code từ bài hướng dẫn để nộp bài là hành vi có thể dẫn đến khóa tài khoản.

Tóm tắt đề bài

Cho mảng \(a\) gồm \(n\) phần tử. Có \(q\) truy vấn, mỗi truy vấn gồm ba số nguyên \(l, r, x\). Với mỗi truy vấn, ta thay tất cả các phần tử \(a_i\) có giá trị thỏa mãn \(l \le a_i \le r\) thành giá trị \(x\). Nhiệm vụ của bạn là in ra tổng của mảng ban đầu và tổng của mảng sau mỗi lần thực hiện truy vấn.

Phân tích

  • Điều kiện: \(n, q \le 2 \cdot 10^5\), các giá trị \(a_i, l, r, x\) có thể lên đến \(10^9\).
  • Nhận xét:
    • Các truy vấn không thay đổi giá trị theo vị trí mà thay đổi theo giá trị.
    • Số lượng giá trị khác nhau trong mảng ban đầu tối đa là \(n\). Sau mỗi truy vấn, ta gộp nhiều giá trị cũ thành một giá trị mới \(x\).
    • Tổng của mảng có thể vượt quá giới hạn của kiểu số nguyên 32-bit, nên cần dùng kiểu dữ liệu long long (C++) hoặc mặc định (Python).

Cách làm đơn giản (Brute Force)

Ý tưởng

Với mỗi truy vấn \((l, r, x)\), ta duyệt qua toàn bộ mảng \(a\). Nếu \(l \le a_i \le r\), ta cập nhật \(a_i = x\). Sau đó tính tổng toàn bộ mảng.

Độ phức tạp

  • Thời gian: \(O(q \cdot n)\)
  • Đánh giá: Với \(n, q = 2 \cdot 10^5\), độ phức tạp \(O(n \cdot q)\) sẽ lên tới \(4 \cdot 10^{10}\) phép tính, không thể vượt qua giới hạn thời gian (thường là 1s). Cách này chỉ phù hợp cho Subtask 2 (\(n, q \le 2000\)).

Code Brute Force

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

int main() {
    int n, q;
    cin >> n >> q;
    vector<long long> a(n);
    long long current_sum = 0;
    for (int i = 0; i < n; i++) {
        cin >> a[i];
        current_sum += a[i];
    }
    cout << current_sum << endl;
    while (q--) {
        long long l, r, x;
        cin >> l >> r >> x;
        current_sum = 0;
        for (int i = 0; i < n; i++) {
            if (a[i] >= l && a[i] <= r) {
                a[i] = x;
            }
            current_sum += a[i];
        }
        cout << current_sum << endl;
    }
    return 0;
}
Python
Python
import sys

def solve():
    input = sys.stdin.read().split()
    n = int(input[0])
    q = int(input[1])
    a = list(map(int, input[2:2+n]))

    current_sum = sum(a)
    print(current_sum)

    ptr = 2 + n
    for _ in range(q):
        l = int(input[ptr])
        r = int(input[ptr+1])
        x = int(input[ptr+2])
        ptr += 3

        current_sum = 0
        for i in range(n):
            if l <= a[i] <= r:
                a[i] = x
            current_sum += a[i]
        print(current_sum)

solve()

Hướng giải quyết (Tối ưu)

Nhận xét quan trọng

Thay vì quản lý từng phần tử ở từng vị trí, ta quản lý số lượng xuất hiện của mỗi giá trị.

  • Sử dụng một cấu trúc dữ liệu để lưu trữ: (giá trị: số lượng). Trong C++ ta dùng std::map, trong Python ta dùng dict hoặc collections.Counter.
  • Khi có truy vấn \((l, r, x)\):
    1. Tìm tất cả các giá trị \(v\) trong cấu trúc dữ liệu sao cho \(l \le v \le r\).
    2. Với mỗi giá trị \(v\) tìm được, giả sử nó xuất hiện \(count\) lần:
      • Trừ khỏi tổng hiện tại một lượng: \(v \cdot count\).
      • Lưu lại tổng số lượng các phần tử bị thay đổi: \(total\_count = \sum count\).
      • Xóa giá trị \(v\) khỏi cấu trúc dữ liệu.
    3. Thêm giá trị \(x\) vào cấu trúc dữ liệu với số lượng là \(total\_count\).
    4. Cộng thêm vào tổng hiện tại một lượng: \(x \cdot total\_count\).

Tại sao cách này nhanh?

  • Mỗi giá trị ban đầu chỉ bị "xóa" một lần duy nhất.
  • Mỗi truy vấn chỉ thêm vào tối đa một giá trị mới \(x\).
  • Số lần xóa giá trị trong suốt quá trình chạy tối đa là \(n + q\).
  • Trong C++, std::map cho phép tìm các giá trị trong đoạn \([l, r]\) hiệu quả bằng hàm lower_bound.

Độ phức tạp

  • Thời gian: \(O((n + q) \log n)\) do mỗi phần tử được thêm và xóa khỏi map tối đa một lần.
  • Bộ nhớ: \(O(n + q)\) để lưu trữ map.

Code tham khảo

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

int main() {
    // Tối ưu tốc độ nhập xuất
    ios::sync_with_stdio(false);
    cin.tie(nullptr);

    int n, q;
    cin >> n >> q;

    map<long long, long long> countByValue;
    long long totalSum = 0;

    // Đếm số lần xuất hiện của từng giá trị ban đầu
    for (int i = 0; i < n; i++) {
        long long value;
        cin >> value;
        countByValue[value]++;
        totalSum += value;
    }

    // In tổng ban đầu
    cout << totalSum << '\n';

    for (int i = 0; i < q; i++) {
        long long l, r, x;
        cin >> l >> r >> x;

        // Tìm phần tử đầu tiên có giá trị >= l
        auto it = countByValue.lower_bound(l);
        long long movedCount = 0;
        vector<long long> valuesToErase;

        // Duyệt qua các phần tử có giá trị <= r
        while (it != countByValue.end() && it->first <= r) {
            long long val = it->first;
            long long cnt = it->second;

            movedCount += cnt;
            totalSum -= val * cnt; // Trừ giá trị cũ khỏi tổng
            valuesToErase.push_back(val);
            it++;
        }

        // Xóa các giá trị đã xử lý khỏi map
        for (long long v : valuesToErase) {
            countByValue.erase(v);
        }

        // Thêm giá trị mới x vào map
        if (movedCount > 0) {
            countByValue[x] += movedCount;
            totalSum += x * movedCount; // Cộng giá trị mới vào tổng
        }

        cout << totalSum << '\n';
    }

    return 0;
}
Python
Python
import sys
from bisect import bisect_left, bisect_right

def solve():
    # Đọc toàn bộ dữ liệu để xử lý nhanh hơn
    input_data = sys.stdin.read().split()
    if not input_data:
        return

    n = int(input_data[0])
    q = int(input_data[1])

    count_by_value = {}
    total_sum = 0

    # Đếm số lần xuất hiện ban đầu
    for i in range(n):
        val = int(input_data[2 + i])
        count_by_value[val] = count_by_value.get(val, 0) + 1
        total_sum += val

    print(total_sum)

    # Sắp xếp các key để tìm kiếm nhị phân
    sorted_keys = sorted(count_by_value.keys())

    ptr = 2 + n
    for _ in range(q):
        l = int(input_data[ptr])
        r = int(input_data[ptr+1])
        x = int(input_data[ptr+2])
        ptr += 3

        # Tìm các giá trị trong đoạn [l, r]
        idx_l = bisect_left(sorted_keys, l)
        idx_r = bisect_right(sorted_keys, r)

        moved_count = 0
        # Duyệt các giá trị cần xóa
        for i in range(idx_l, idx_r):
            val = sorted_keys[i]
            cnt = count_by_value.pop(val)
            moved_count += cnt
            total_sum -= val * cnt

        # Cập nhật danh sách key và map
        if moved_count > 0:
            total_sum += x * moved_count
            count_by_value[x] = count_by_value.get(x, 0) + moved_count

        # Chỉ sắp xếp lại khi có sự thay đổi thực sự để đảm bảo hiệu năng
        if idx_l < idx_r:
            sorted_keys = sorted(count_by_value.keys())

        print(total_sum)

if __name__ == "__main__":
    solve()

Bình luận

Mới nhất
Tải bình luận...

Không có bình luận nào.