文本分词

用带频率的相邻边集合和懒删除优先队列维护字节对合并,每次只更新合并位置附近的边。

OJ: shumeng

题目 ID: CSP202406C

难度:提高+/省选-

标签:模拟字符串链表优先队列

日期: 2026-07-31 16:21

形式化题目

给定 nn 个带频率的单词,初始把每个单词拆成单个字母。重复以下操作直到词表达到 mm 个词:

  1. 找出所有相邻词汇对中“加权出现次数”最大的一个(同一单词内的相邻对按该单词频率加权);
  2. 若有并列,依次按合并后总长度更短、前词更短、合并串字典序更小、左词汇编号更小选择;
  3. 把每个序列中该相邻对合并为一个新词,加入词表。

输出加入词表的顺序。

思路

这是 BPE(字节对编码)的模拟题,难点在于多轮合并时不能每轮全局重扫。

朴素做法:每轮重新扫描

先看直觉做法:每轮重新扫描全部序列,统计所有相邻词汇对的加权次数,选出最优对后替换所有出现位置。

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-31 16:21
 * update_at: 2026-08-17 22:39
 */
// brute.cpp:小数据暴力解,每轮重新扫描全部序列统计词汇对,只适合小数据验证合并规则。
#include <bits/stdc++.h>
using namespace std;

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

    int n, vocabulary_size;
    cin >> n >> vocabulary_size;
    vector<string> words(n);
    vector<long long> frequency(n);
    bool exists[26] = {};
    for (int i = 0; i < n; i++) {
        cin >> words[i] >> frequency[i];
        for (int j = 0; j < (int)words[i].size(); j++) exists[words[i][j] - 'a'] = true;
    }

    // 初始词表:出现过的字母各占一个编号
    vector<string> vocabulary;
    int letter_token[26];
    for (int i = 0; i < 26; i++) letter_token[i] = -1;
    for (int i = 0; i < 26; i++) {
        if (exists[i]) {
            letter_token[i] = (int)vocabulary.size();
            vocabulary.push_back(string(1, (char)('a' + i)));
        }
    }

    // 把每个单词转成词汇编号序列
    vector<vector<int> > sequences(n);
    for (int i = 0; i < n; i++) {
        for (int j = 0; j < (int)words[i].size(); j++) {
            sequences[i].push_back(letter_token[words[i][j] - 'a']);
        }
    }

    while ((int)vocabulary.size() < vocabulary_size) {
        // 每轮重新统计全部相邻词汇对的加权出现次数
        map<pair<int, int>, long long> count;
        for (int i = 0; i < n; i++) {
            for (int j = 0; j + 1 < (int)sequences[i].size(); j++) {
                count[make_pair(sequences[i][j], sequences[i][j + 1])] += frequency[i];
            }
        }
        if (count.empty()) break;

        // 选出最优词汇对:权重最大,同权重按题目规则比较拼接串
        pair<int, int> best_pair = count.begin()->first;
        long long best_weight = count.begin()->second;
        for (map<pair<int, int>, long long>::iterator iterator = count.begin(); iterator != count.end(); ++iterator) {
            pair<int, int> candidate = iterator->first;
            long long candidate_weight = iterator->second;
            bool better = false;
            if (candidate_weight != best_weight) {
                better = candidate_weight > best_weight;
            } else {
                size_t candidate_length = vocabulary[candidate.first].size() + vocabulary[candidate.second].size();
                size_t best_length = vocabulary[best_pair.first].size() + vocabulary[best_pair.second].size();
                if (candidate_length != best_length) {
                    better = candidate_length < best_length;
                } else if (vocabulary[candidate.first].size() != vocabulary[best_pair.first].size()) {
                    better = vocabulary[candidate.first].size() < vocabulary[best_pair.first].size();
                } else {
                    string candidate_text = vocabulary[candidate.first] + vocabulary[candidate.second];
                    string best_text = vocabulary[best_pair.first] + vocabulary[best_pair.second];
                    if (candidate_text != best_text) {
                        better = candidate_text < best_text;
                    } else {
                        better = candidate < best_pair;
                    }
                }
            }
            if (better) {
                best_pair = candidate;
                best_weight = candidate_weight;
            }
        }

        // 生成新词汇,并把所有序列中的该相邻对替换成新词汇
        int merged_token = (int)vocabulary.size();
        vocabulary.push_back(vocabulary[best_pair.first] + vocabulary[best_pair.second]);
        for (int i = 0; i < n; i++) {
            vector<int> next_sequence;
            for (int j = 0; j < (int)sequences[i].size();) {
                if (j + 1 < (int)sequences[i].size() && sequences[i][j] == best_pair.first
                        && sequences[i][j + 1] == best_pair.second) {
                    next_sequence.push_back(merged_token);
                    j += 2;
                } else {
                    next_sequence.push_back(sequences[i][j]);
                    j++;
                }
            }
            sequences[i].swap(next_sequence);
        }
    }

    int output_count = min(vocabulary_size, (int)vocabulary.size());
    for (int i = 0; i < output_count; i++) cout << vocabulary[i] << '\n';
    return 0;
}

做法完全贴合规则,但每轮 O(L)O(L) 扫描,总长度 LL 可达 2.5×1052.5 \times 10^5、词表目标 mm 可达 50005000,无法通过。

主解:链表 + 局部边维护

把每个单词的词汇序列建成双向链表,节点编号唯一。一个相邻边由“左节点编号”唯一确定,存入对应词汇对的 set<左节点>;该词汇对的加权出现次数就是它所有边位置的频率之和。

合并 (A,B)(A, B) 时,取出该对的所有边位置,逐个执行:把左节点改为新词、删除右节点,并重新连接链表。合并只会改变合并点两侧各一条边,因此只需删除和添加这几条边的计数即可,不需要全局重扫。

用惰性删除堆选最优

用优先队列维护所有词汇对的排序。每次边的增删都会改变某对词汇的权重或边集合,给它加版本号并压入新条目;取堆顶时,丢弃版本号或权重与当前记录不一致的过期条目,直到拿到有效的最优对。

特例:合并 A=AA=A 时(如 a a a),按位置从小到大依次合并,得到 aa a,符合题目从前到后的合并规则。

代码

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-31 16:21
 * update_at: 2026-08-17 22:39
 */
#include <bits/stdc++.h>
using namespace std;

// 词表:token_table[i].text 是编号 i 的词汇(初始为单字母)
struct Token {
    string text;
};

// 链表节点:每个节点是某个单词序列中的一个词汇,节点下标即节点编号
struct Occurrence {
    int token;     // 该节点当前的词汇编号
    int word;      // 该节点属于第几个单词(用于加权频率)
    int previous;  // 前驱节点,0 表示链表头
    int next;      // 后继节点,0 表示链表尾
    bool alive;    // 节点是否仍存活(被合并后置为 false)
};

// 一对相邻词汇 (左词汇, 右词汇) 的统计信息
struct PairInfo {
    long long weight;  // 按单词频率加权的出现次数
    int version;       // 版本号,用于惰性删除过期的堆条目
    set<int> positions; // 该词汇对出现的所有边位置(左节点编号)
    PairInfo() : weight(0), version(0) {}
};

// 优先队列中的条目
struct HeapEntry {
    unsigned long long key; // 编码了 (左词汇, 右词汇) 的 64 位键
    long long weight;
    int version;
};

vector<Token> token_table;      // 词表,下标即词汇编号
vector<Occurrence> occurrence;  // 全部单词的词汇序列链表节点,下标 0 为哑节点
vector<long long> word_frequency; // 每个单词的出现频率
unordered_map<unsigned long long, PairInfo> pair_info; // 词汇对 -> 统计信息

// 把 (left, right) 两个词汇编号打包成一个 64 位键
unsigned long long make_key(int left, int right) {
    return (unsigned long long)(unsigned int)left << 32 | (unsigned int)right;
}

// 从键中解出左词汇编号
int key_left(unsigned long long key) {
    return (int)(key >> 32);
}

// 从键中解出右词汇编号
int key_right(unsigned long long key) {
    return (int)(key & 0xffffffffu);
}

// 比较两个词汇拼接串 (a1+a2) 与 (b1+b2) 的字典序:返回 -1/0/1
int compare_concatenation(int first_left, int first_right, int second_left, int second_right) {
    const string &first_left_text = token_table[first_left].text;
    const string &first_right_text = token_table[first_right].text;
    const string &second_left_text = token_table[second_left].text;
    const string &second_right_text = token_table[second_right].text;
    size_t first_length = first_left_text.size() + first_right_text.size();
    size_t second_length = second_left_text.size() + second_right_text.size();
    size_t common_length = min(first_length, second_length);
    for (size_t i = 0; i < common_length; i++) {
        char first_char = i < first_left_text.size() ? first_left_text[i] : first_right_text[i - first_left_text.size()];
        char second_char = i < second_left_text.size() ? second_left_text[i] : second_right_text[i - second_left_text.size()];
        if (first_char < second_char) return -1;
        if (first_char > second_char) return 1;
    }
    if (first_length < second_length) return -1;
    if (first_length > second_length) return 1;
    return 0;
}

// 判断 first 是否比 second 更优(更大的权重更优,其次按题目规则比较拼接串)
bool better_entry(const HeapEntry &first, const HeapEntry &second) {
    if (first.weight != second.weight) return first.weight > second.weight;
    int first_left = key_left(first.key);
    int first_right = key_right(first.key);
    int second_left = key_left(second.key);
    int second_right = key_right(second.key);
    size_t first_length = token_table[first_left].text.size() + token_table[first_right].text.size();
    size_t second_length = token_table[second_left].text.size() + token_table[second_right].text.size();
    if (first_length != second_length) return first_length < second_length;
    if (token_table[first_left].text.size() != token_table[second_left].text.size()) {
        return token_table[first_left].text.size() < token_table[second_left].text.size();
    }
    int text_compare = compare_concatenation(first_left, first_right, second_left, second_right);
    if (text_compare != 0) return text_compare < 0;
    return first.key < second.key;
}

// 把 better_entry 转成优先队列的小于号(greater 优先弹出)
struct HeapCompare {
    bool operator()(const HeapEntry &first, const HeapEntry &second) const {
        return better_entry(second, first);
    }
};

priority_queue<HeapEntry, vector<HeapEntry>, HeapCompare> heap; // 惰性删除堆

// 若词汇对 key 还有存活边,就把当前版本压入堆
void push_current_pair(unsigned long long key) {
    PairInfo &info = pair_info[key];
    if (info.positions.empty()) return;
    heap.push({key, info.weight, info.version});
}

// 添加边:左节点 left 与其后继构成一个新的相邻词汇对
void add_edge(int left) {
    if (left == 0 || !occurrence[left].alive || occurrence[left].next == 0) return;
    int right = occurrence[left].next;
    unsigned long long key = make_key(occurrence[left].token, occurrence[right].token);
    PairInfo &info = pair_info[key];
    info.weight += word_frequency[occurrence[left].word];
    info.positions.insert(left);
    info.version++;
    heap.push({key, info.weight, info.version});
}

// 删除边:左节点 left 与其后继的相邻词汇对不再存在
void remove_edge(int left) {
    if (left == 0 || !occurrence[left].alive || occurrence[left].next == 0) return;
    int right = occurrence[left].next;
    unsigned long long key = make_key(occurrence[left].token, occurrence[right].token);
    unordered_map<unsigned long long, PairInfo>::iterator iterator = pair_info.find(key);
    if (iterator == pair_info.end()) return;
    PairInfo &info = iterator->second;
    info.weight -= word_frequency[occurrence[left].word];
    info.positions.erase(left);
    info.version++;
    push_current_pair(key);
}

// 在链表位置 left 处执行一次合并:把 (first_token, second_token) 替换成 merged_token
void merge_occurrence(int left, int first_token, int second_token, int merged_token) {
    int right = occurrence[left].next;
    if (right == 0 || !occurrence[left].alive || !occurrence[right].alive) return;
    if (occurrence[left].token != first_token || occurrence[right].token != second_token) return;

    int previous = occurrence[left].previous;
    int next = occurrence[right].next;
    // 合并会改变左右两侧的相邻关系,先把相关边删掉
    remove_edge(previous);
    remove_edge(left);
    remove_edge(right);

    // 左节点保留为合并后的词汇,右节点作废
    occurrence[left].token = merged_token;
    occurrence[left].next = next;
    if (next != 0) occurrence[next].previous = left;
    occurrence[right].alive = false;
    occurrence[right].previous = 0;
    occurrence[right].next = 0;

    // 重新建立合并点两侧的相邻边
    add_edge(previous);
    add_edge(left);
}

// 合并词汇对 (first_token, second_token):所有出现位置一起合并,生成新词汇
void merge_pair(int first_token, int second_token) {
    unsigned long long key = make_key(first_token, second_token);
    PairInfo &info = pair_info[key];
    vector<int> positions(info.positions.begin(), info.positions.end());
    string merged_text = token_table[first_token].text + token_table[second_token].text;
    int merged_token = (int)token_table.size();
    token_table.push_back({merged_text});
    for (int i = 0; i < (int)positions.size(); i++) {
        merge_occurrence(positions[i], first_token, second_token, merged_token);
    }
}

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

    int n, vocabulary_size;
    cin >> n >> vocabulary_size;
    vector<string> words(n);
    word_frequency.resize(n);
    bool exists[26] = {};
    for (int i = 0; i < n; i++) {
        cin >> words[i] >> word_frequency[i];
        for (int j = 0; j < (int)words[i].size(); j++) exists[words[i][j] - 'a'] = true;
    }

    // 初始词表:出现过的字母各占一个词汇编号
    int letter_token[26];
    for (int i = 0; i < 26; i++) letter_token[i] = -1;
    for (int i = 0; i < 26; i++) {
        if (exists[i]) {
            letter_token[i] = (int)token_table.size();
            token_table.push_back({string(1, (char)('a' + i))});
        }
    }

    // 把每个单词串成双向链表,节点下标 0 为链表头哨兵
    occurrence.push_back({0, 0, 0, 0, false});
    for (int i = 0; i < n; i++) {
        int previous = 0;
        for (int j = 0; j < (int)words[i].size(); j++) {
            int node = (int)occurrence.size();
            occurrence.push_back({letter_token[words[i][j] - 'a'], i, previous, 0, true});
            if (previous != 0) occurrence[previous].next = node;
            previous = node;
        }
    }
    for (int i = 1; i < (int)occurrence.size(); i++) add_edge(i);

    // 反复合并当前最优词汇对,直到词表达到目标大小
    while ((int)token_table.size() < vocabulary_size) {
        // 弹出所有过期的堆条目
        while (!heap.empty()) {
            HeapEntry top = heap.top();
            unordered_map<unsigned long long, PairInfo>::iterator iterator = pair_info.find(top.key);
            if (iterator != pair_info.end() && iterator->second.version == top.version
                    && iterator->second.weight == top.weight && !iterator->second.positions.empty()) break;
            heap.pop();
        }
        if (heap.empty()) break;
        HeapEntry top = heap.top();
        heap.pop();
        merge_pair(key_left(top.key), key_right(top.key));
    }

    int output_count = min(vocabulary_size, (int)token_table.size());
    for (int i = 0; i < output_count; i++) cout << token_table[i].text << '\n';
    return 0;
}

复杂度

设初始字母总长度为 LL,目标词表大小 mm

  • 时间:每次实际合并至少减少一个链表节点,所有被处理的位置总数 O(L)O(L);每次边集合与堆操作为对数级,总复杂度约 O(LlogL+mlogL)O(L \log L + m \log L),另计拼接串比较的开销。
  • 空间:链表节点、边集合与词表共 O(L+m)O(L + m)

总结

把“全局重新统计”改成“只维护发生变化的相邻边”,配合惰性删除堆处理动态最优选择,就把多轮合并的代价压缩到所有链表节点的摊销规模。理解“一次合并只影响两条边”是写出高效实现的关键。