[语言月赛 202407] speech 的题解


记住只在没有思路时使用题解,不要从它复制粘贴代码。请尊重题目和题解的作者。
在解题之前提交题解的代码会导致封禁。

作者: admin

概述

本题要求找出魅力值最大的语言编号,其中一套语言的魅力值定义为其语法数量与该语言使用人数的乘积,若有多套语言的魅力值并列最大则输出编号最小者。核心解法是用桶数组统计每套语言的使用人数,再按编号从小到大扫描一次,仅在严格大于当前最大值时更新答案。

分析
核心观察

一套语言的魅力值只由两个量决定:它自身的语法数量 a_i 与使用它的居民人数 c_i。而 c_i 并不需要额外存储整条居民序列,把每个居民使用的语言编号直接当作下标做桶计数,就能在遍历居民的过程中统计出全部 c_i。

思路

记 c_i 为使用语言 i 的居民人数,则语言 i 的魅力值为 s_i = a_i \times c_i,答案为使 s_i 取到最大值的、编号最小的 i。

统计 c_i 有暴力与桶计数两种做法。暴力做法对每套语言都完整扫描一遍居民序列,时间为 O(nm),在 n, m \le 10^3 时约为 10^6 次操作,虽然可以接受但并不必要。注意到第 j 个居民使用的语言编号 b_j 恰好就是要累加的下标,于是可以在读入 b_j 时直接令 c_{b_j} 自增 1,居民序列遍历完毕时全部 c_i 也就统计完毕,这一步只需 O(n) 时间。

得到 c_i 后,从 1 到 m 顺序枚举语言,用变量 bestValue 记录已扫描过的魅力值最大值、变量 best 记录取到该值的编号。更新条件写成严格大于 s_i > bestValue:当 s_i 与当前最大值相等时不更新,编号较小的语言因此得以保留,恰好满足并列取最小编号的要求;若写成大于等于 s_i \ge bestValue,并列时会不断替换成编号更大的语言,答案就错了。桶计数与一次扫描合起来,总时间为 O(n + m)。

具体示例

以样例 1 为例,n = 3、m = 2,两套语言的语法数量分别为 a_1 = 1、a_2 = 2,居民使用的语言编号依次为 1, 1, 2。

桶计数得到 c_1 = 2、c_2 = 1。语言 1 的魅力值为 1 \times 2 = 2,语言 2 的魅力值为 2 \times 1 = 2,二者并列最大;由于 1 < 2,输出编号 1。

作为对照,样例 2 只把 a_2 改为 3,此时语言 2 的魅力值为 3 \times 1 = 3,严格大于语言 1 的 2,输出变为 2。可见答案的走向完全由“严格大于”这一更新条件驱动。

算法步骤
  1. 读入 n 与 m。
  2. 读入 m 个整数存入数组 a,其中 a_i 表示语言 i 的语法数量。
  3. 依次读入 n 个居民使用的语言编号,每读入一个编号就把对应的桶 cnt 加 1,读完后 cnt_i 即为语言 i 的使用人数。
  4. 令 best 为 1,bestValue 为 -1。
  5. 从 1 到 m 枚举语言编号 i,计算 a_i \times cnt_i;若该值严格大于 bestValue,则用它更新 bestValue,并令 best 等于 i。
  6. 输出 best。
复杂度分析
  • 时间:O(n + m)。读入语法数量、统计居民语言、扫描全部语言各为一次线性遍历。
  • 空间:O(m)。只需要长度为 m + 1 的数组 a 与 cnt,居民的语言编号可以边读边统计,无需存储。
实现注意事项
  • 更新最大值必须使用严格大于;写成大于等于会在魅力值并列时留下编号较大的语言,与输出最小编号的要求相反。
  • bestValue 初始化为 -1。由于 a_i \ge 0、c_i \ge 0,所有魅力值均非负,而 -1 严格小于任何可能的魅力值,故第一次比较必然成功,编号 1 至少会被记录一次。
  • a_i 可以取 0,此时该语言的魅力值为 0,仍然要参与比较,不能提前跳过。
  • 魅力值的上界为 10^3 \times 10^3 = 10^6,int 足以容纳,改用 64 位整型相乘可彻底避免溢出隐患。
  • 数据规模较小,使用 cin 直接读入即可;若追求稳妥可关闭同步流再读入。
源代码
#include <iostream>
#include <vector>

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

    int n, m;
    std::cin >> n >> m;

    std::vector<int> a(m + 1, 0), cnt(m + 1, 0);
    for (int i = 1; i <= m; ++i) {
        std::cin >> a[i];
    }

    for (int j = 0; j < n; ++j) {
        int b;
        std::cin >> b;
        ++cnt[b];
    }

    int best = 1;
    long long bestValue = -1;
    for (int i = 1; i <= m; ++i) {
        long long value = 1LL * a[i] * cnt[i];
        if (value > bestValue) {
            bestValue = value;
            best = i;
        }
    }

    std::cout << best << '\n';
    return 0;
}
import sys


def main():
    data = sys.stdin.buffer.read().split()
    n = int(data[0])
    m = int(data[1])

    a = [0] * (m + 1)
    for i in range(1, m + 1):
        a[i] = int(data[1 + i])

    cnt = [0] * (m + 1)
    for j in range(n):
        b = int(data[2 + m + j])
        cnt[b] += 1

    best = 1
    best_value = -1
    for i in range(1, m + 1):
        value = a[i] * cnt[i]
        if value > best_value:
            best_value = value
            best = i

    sys.stdout.write(str(best) + "\n")


main()

评论

目前没有评论。