Keyboard shortcuts

Press or to navigate between chapters

Press S or / to search in the book

Press ? to show this help

Press Esc to hide this help

底层实现:文本处理

语言模型不会直接读取文字,它看到的是一串数字。在文字与数字之间,文本处理程序会作出许多决定:哪些字节属于同一个字符,哪些字符应当被归一化,文本从哪里切开,哪些片段进入词表,以及词表之外的内容如何回退。这些决定一旦在训练语料中固定下来,就会影响模型能够学习什么;训练与推理稍有不一致,同一句话甚至会变成两组不同的输入。

本书沿着这条转换路径逐层实现。Unicode 与 UTF-8 解决字符如何落到字节;正则表达式引擎负责识别和切分文本模式;Trie 为词典前缀查询提供紧凑的数据结构;中文分词把连续汉字段变成可处理的基本片段;SentencePiece 与 BytePiece 再从语料中学习子词,并完成 Token ID 的编码与解码。最后的 W2V 和 LDA 作为番外,展示切分后的词语如何进入经典的表示学习与主题模型。

中文分词不是 Tokenizer 的必要前置步骤,但它可以成为 PreTokenize 的一种实现。空格语言已经有天然的候选边界,中文等非空格语言则往往以整段文本进入子词学习。先在中文内部建立词边界,可以让两类语言通过同一个 PreTokenize 接口进入后续流程,也能缩短单次搜索的片段,减少跨词候选。是否启用这一步取决于任务;本书选择实现它,是为了把这种工程取舍讲清楚。

书中的算法都有对应的 C++ 项目。正文关注的是实现一项技术所需的核心结构、推导和约束,而不是某个库的接口手册。代码可以替换,训练数据也会变化,但字符边界、概率路径、词表覆盖和训练—推理一致性这些问题不会消失。

Unicode 与 UTF-8

Unicode

Unicode 为世界各地文字、符号和控制字符分配统一的码点,例如:

  • 英文'a' → 97, 'b' → 98, 'c' → 99, …
  • 中文'你' → 20320, '我' → 25105, '他' → 20182, …
  • Emoji'😀' → 128512, '🌍' → 127757, …

Unicode 码点通常用 uint32_t 保存,但有效范围只有 U+0000U+10FFFF,其中代理项区间 U+D800U+DFFF 也不能表示 Unicode 标量值。因此,32 位整数只是方便的内存表示,并不意味着全部数值都合法。

虽然有了 Unicode,但如果把所有文字都用 uint32_t(4 字节)存储会比较低效:

  1. 存储效率问题:ASCII 字符只需要 7 位,用 4 字节存储会浪费大量空间
  2. 向后兼容:需要兼容 Unicode 之前的 ASCII 标准
  3. 网络传输:更少的字节意味着更快的传输速度

UTF-8(8-bit Unicode Transformation Format)使用 1 到 4 个字节编码一个 Unicode 标量值。ASCII 码点仍只占一个字节,并且编码结果与原有 ASCII 字节完全一致;更大的码点使用更多字节。

UTF-8

UTF-8 使用 1 到 4 个字节表示一个 Unicode 标量值。首字节决定序列长度,后续字节统一以 10 开头:

1字节:0xxxxxxx
2字节:110xxxxx 10xxxxxx
3字节:1110xxxx 10xxxxxx 10xxxxxx
4字节:11110xxx 10xxxxxx 10xxxxxx 10xxxxxx

解码时只需从 UTF-8 字节序列中取出所有 x 位,并按顺序拼接为 Unicode 码点:

1字节:0xxxxxxx → xxxxxxx
2字节:110xxxxx 10xxxxxx → xxxxx xxxxxx
3字节:1110xxxx 10xxxxxx 10xxxxxx → xxxx xxxxxx xxxxxx  
4字节:11110xxx 10xxxxxx 10xxxxxx 10xxxxxx → xxx xxxxxx xxxxxx xxxxxx

对应的合法码点范围为:

  • 1 字节U+0000U+007F
  • 2 字节U+0080U+07FF
  • 3 字节U+0800U+FFFF,排除代理项
  • 4 字节U+10000U+10FFFF

C++ 实现

函数 1:IsTrailByte

续字节的最高两位必须是 10

bool IsTrailByte(uint8_t x) {
    return (x & 0xC0) == 0x80;
}

实现原理

  • 0xC0(11000000)作为掩码,提取字节的最高 2 位
  • 0x80(10000000)是续字节的标准模式
  • 通过位与操作 & 提取最高 2 位,然后与 0x80 比较

它可以识别 0x800xBF 范围内的 64 个续字节。

函数 2:DecodeOneUTF8

解码不仅要检查字节前缀,还要拒绝过长编码、代理项和超过 U+10FFFF 的结果。非法序列返回替换字符 U+FFFD,并消费一个字节,使调用方可以继续处理后续输入:

uint32_t DecodeOneUTF8(const char* begin, const char* end, size_t* bytes) {
    constexpr uint32_t kError = 0xFFFD;
    const size_t len = end - begin;
    if (len == 0) {
        *bytes = 0;
        return kError;
    }

    const auto* data = reinterpret_cast<const uint8_t*>(begin);
    uint32_t cp = 0;
    size_t width = 0;
    uint32_t minimum = 0;

    if (data[0] < 0x80) {
        *bytes = 1;
        return data[0];
    } else if ((data[0] & 0xE0) == 0xC0) {
        width = 2; minimum = 0x80; cp = data[0] & 0x1F;
    } else if ((data[0] & 0xF0) == 0xE0) {
        width = 3; minimum = 0x800; cp = data[0] & 0x0F;
    } else if ((data[0] & 0xF8) == 0xF0) {
        width = 4; minimum = 0x10000; cp = data[0] & 0x07;
    }

    if (width != 0 && len >= width) {
        for (size_t i = 1; i < width; ++i) {
            if (!IsTrailByte(data[i])) {
                width = 0;
                break;
            }
            cp = (cp << 6) | (data[i] & 0x3F);
        }
        const bool surrogate = cp >= 0xD800 && cp <= 0xDFFF;
        if (width != 0 && cp >= minimum && cp <= 0x10FFFF && !surrogate) {
            *bytes = width;
            return cp;
        }
    }

    *bytes = 1;
    return kError;
}

minimum 是当前字节数可以表示的最小码点,用来排除本可用更短序列表示的过长编码。循环则逐个验证续字节,并通过左移 6 位还原码点。

函数 3:DecodeUTF8

基于 DecodeOneUTF8,可以实现整个 UTF-8 字符串的解码:

std::vector<uint32_t> DecodeUTF8(const std::string& str) {
    std::vector<uint32_t> codepoints;
    
    size_t pos = 0;
    while (pos < str.size()) {
        size_t bytes_consumed;
        
        uint32_t codepoint = DecodeOneUTF8(
            str.data() + pos, str.data() + str.size(), &bytes_consumed);
        
        codepoints.push_back(codepoint);
        pos += bytes_consumed;
    }
    
    return codepoints;
}

这个函数循环调用 DecodeOneUTF8,最终返回 Unicode 码点数组。非法输入会以替换字符保留在结果中,而不会截断整个字符串。

函数 4:EncodeOneUTF8

编码是解码的镜像操作,将 Unicode 码点转换为 UTF-8 字节序列:

size_t EncodeOneUTF8(uint32_t c, char* output) {
    constexpr uint32_t kError = 0xFFFD;
    if (c > 0x10FFFF || (c >= 0xD800 && c <= 0xDFFF)) {
        c = kError;
    }

    if (c <= 0x7F) {  // 0x7F: 01111111
        // 1字节:0xxxxxxx
        *output = static_cast<char>(c);
        return 1;
    }
    if (c <= 0x7FF) {  // 0x7FF: 011111111111
        // 2字节:110xxxxx 10xxxxxx
        output[1] = 0x80 | (c & 0x3F);         // 0x80: 10000000, 0x3F: 00111111 - 低6位 + 续字节前缀
        c >>= 6;
        output[0] = 0xC0 | c;                  // 0xC0: 11000000 - 高5位 + 首字节前缀
        return 2;
    }
    if (c <= 0xFFFF) {  // 0xFFFF: 1111111111111111
        // 3字节:1110xxxx 10xxxxxx 10xxxxxx
        output[2] = 0x80 | (c & 0x3F);         // 0x80: 10000000, 0x3F: 00111111 - 最低6位
        c >>= 6;
        output[1] = 0x80 | (c & 0x3F);         // 0x80: 10000000, 0x3F: 00111111 - 中间6位
        c >>= 6;
        output[0] = 0xE0 | c;                  // 0xE0: 11100000 - 最高4位 + 首字节前缀
        return 3;
    }
    // 4字节:11110xxx 10xxxxxx 10xxxxxx 10xxxxxx
    output[3] = 0x80 | (c & 0x3F);             // 0x80: 10000000, 0x3F: 00111111 - 最低6位
    c >>= 6;
    output[2] = 0x80 | (c & 0x3F);             // 0x80: 10000000, 0x3F: 00111111 - 次低6位
    c >>= 6;
    output[1] = 0x80 | (c & 0x3F);             // 0x80: 10000000, 0x3F: 00111111 - 次高6位
    c >>= 6;
    output[0] = 0xF0 | c;                      // 0xF0: 11110000 - 最高3位 + 首字节前缀
    return 4;
}

实现原理

编码采用“从低位到高位,倒序构建”的策略:

  1. 数据分解:用 c & 0x3F 提取低 6 位,然后右移 6 位处理下一组
  2. 格式添加:用 0x80 | 数据 给续字节添加 10 前缀
  3. 倒序填充:从最后一个字节开始向前填充,自然处理变长编码

函数 5:EncodeUTF8

最后将 Unicode 码点数组编码为 UTF-8 字符串:

std::string EncodeUTF8(const std::vector<uint32_t>& codepoints) {
    std::string result;
    
    for (uint32_t cp : codepoints) {
        char buffer[4];  // UTF-8 最多需要 4 个字节
        size_t bytes = EncodeOneUTF8(cp, buffer);
        
        result.append(buffer, bytes);
    }
    
    return result;
}

这个函数遍历码点数组,逐个编码后拼接成完整的 UTF-8 字符串。

完成编码与解码之后,后续算法就可以在码点层面处理字符,而把字节边界留在输入输出层。下一篇将以此为基础,实现能够正确处理 UTF-8 文本的正则表达式引擎。

配套实现:Ismantic/Ustr

正则表达式引擎:基础篇

Regex 介绍

正则表达式(Regex) —— 它是一种声明式的语言定义方法,通过简洁的语法描述复杂的字符串集合。

  • 语言 = 字符串的集合
  • 语法 = 生成规则(正则表达式就是一种语法)
  • 识别 = 判断给定字符串是否属于这个语言

举例:regex = 5?3* 其对应的集合 {ε, 5, 3, 53, 33, 533, 333, 5333, ...}。(其中 ε 表示空串)

更深层次,触及到了计算理论的核心:

  • 正则语言:正则表达式定义的语言类别,是Chomsky层次结构中最简单的一类
  • 有限状态机:每个正则表达式都等价于一个有限状态自动机
  • 语言的运算:连接/并集/闭包等操作对应正则表达式的语法结构。

正则语言的局限来自有限状态机只有有限记忆:它不能保存无界计数或任意深度的栈,因而无法识别任意层数的括号嵌套。这是语言表达能力的限制,与具体匹配程序是否回溯无关。

接下来介绍两种实现。Matcher 对应 src/regex-0.cc,用递归回溯直接匹配;Compiler 对应 src/regex-1.cc,将正则表达式编译成有限状态机,是本篇的重点。

Matcher

当对匹配效率要求不高以及语法比较简单的时候,可以先从匹配这个层面入手,快速实现一个正则引擎。 实现机制是同时跟踪模式与文本的当前位置,根据普通字符和量词决定消耗哪一边。当模式耗尽时,本次前缀匹配成功,文本可以仍有剩余字符;只有模式以 $ 结尾时,才要求剩余文本为空。

以下给出一个示例,来自开源项目(Wapiti),实现思路很简单:用三个函数分别处理不同层面的匹配逻辑

  • 字符层面的匹配,对特殊字符的处理(比如匹配数字字符)
  • 模式层面的匹配,对量词(*?)字符的处理
  • 双指针移动,处理锚点语法(^$)且在字符串中定位匹配位置

字符层面匹配

bool MatchCharacter(const std::string& pattern, size_t pos, char c) {
    if (c == '\0') return false;
    if (pattern[pos] == '.') return true; 
    if (pattern[pos] == '\\' && pos + 1 < pattern.length()) {
        switch (pattern[pos + 1]) {
            case 'a': return std::isalpha(c);
            case 'd': return std::isdigit(c);
            case 'l': return std::islower(c);
            case 'p': return std::ispunct(c);
            case 's': return std::isspace(c);
            case 'u': return std::isupper(c);
            case 'w': return std::isalnum(c);
            case 'A': return !std::isalpha(c);
            case 'D': return !std::isdigit(c);
            case 'L': return !std::islower(c);
            case 'P': return !std::ispunct(c);
            case 'S': return !std::isspace(c);
            case 'U': return !std::isupper(c);
            case 'W': return !std::isalnum(c);
            default: return pattern[pos + 1] == c;
        }
    }
    return pattern[pos] == c;
}

该函数判断给定字符 c 是否与 pattern[pos] 匹配:

  • . 是通配符,匹配任意字符,总是返回 true
  • \ 是转移字符,其后的 pattern[pos+1] 代表一类字符:
    • 小写字母(如\d 代表数字)表示匹配该字符
    • 大小字母(如\D 代表非数字)表示匹配该类字符的取反

模式层面匹配

bool MatchPattern(const std::string& re, const std::string& str, uint32_t& n) {
    if (re.empty()) return true;
    if (re[0] == '$' && re.length() == 1) return str.empty();

    size_t cn = (re[0] == '\\') ? 2 : 1;
    std::string next = re.substr(cn);

    if (!next.empty() && next[0] == '*') {
        next = next.substr(1);
        size_t pos = 0;
        do {
            uint32_t save = n;
            if (MatchPattern(next, str.substr(pos), n)) return true;
            n = save + 1;
            pos++;
        } while (pos <= str.length() && MatchCharacter(re, 0, str[pos-1]));
        return false;
    }

    if (!next.empty() && next[0] == '?') {
        next = next.substr(1);
        if (!str.empty() && MatchCharacter(re, 0, str[0])) {
            ++n;
            if (MatchPattern(next, str.substr(1), n)) return true;
            --n;
        }
        return MatchPattern(next, str, n);
    }

    ++n;
    return !str.empty() && MatchCharacter(re, 0, str[0]) &&
           MatchPattern(next, str.substr(1), n);
}

MatchPattern 的作用是:从 str 的当前位置开始,判断整个模式 re 能否匹配,并通过 n 记录匹配长度。

该函数实现了三种核心的模式层面匹配:

1. 量词 * (零次或多次匹配) * 量词的实现是最复杂的。它按照重复 0 次、1 次、2 次……的顺序尝试,通过回溯寻找第一个能让后续模式成功的分割点:

if (!next.empty() && next[0] == '*') {
    next = next.substr(1);        // 跳过 '*',获取后续模式
    size_t pos = 0;               // 当前尝试匹配的位置
    do {
        uint32_t save = n;        // 保存当前匹配计数
        // 尝试在当前位置匹配剩余模式
        if (MatchPattern(next, str.substr(pos), n)) return true;
        n = save + 1;             // 恢复计数并增加
        pos++;                    // 尝试下一个位置
    } while (pos <= str.length() && MatchCharacter(re, 0, str[pos-1]));
    // 只要当前匹配*仍然成立,就不会跳出循环
    return false;
}

执行逻辑

  1. 从零开始尝试:先尝试匹配 0 次(pos=0),直接匹配后续模式
  2. 逐步增加匹配:如果失败,则尝试匹配 1 次、2 次…直到不能匹配为止
  3. 尝试顺序:每次都先尝试匹配后续模式,而不是先消耗尽可匹配字符
  4. 回溯机制:如果某个匹配长度失败,则消耗一个字符继续尝试

示例:模式 a*b 匹配字符串 "aaab"

  • 尝试 0 个 a:匹配 "aaab"b → 失败
  • 尝试 1 个 a:匹配 "aab"b → 失败
  • 尝试 2 个 a:匹配 "ab"b → 失败
  • 尝试 3 个 a:匹配 "b"b → 成功!

这个实现主要是避免过度匹配导致的失败 经典问题:如果 * 贪婪地先消耗所有匹配的字符,可能会导致后续模式无法匹配 示例:模式 “.*b” 匹配字符串 “aabb”“

  • 贪婪实现:.* 先匹配整个“aabb““,然后b无字符可匹配,导致整体匹配失败
  • 当前实现:.* 从0开始尝试,逐步增加,直到找到合适的分割点

2. 量词 ?(零次或一次匹配)

? 量词相对简单,但也需要处理两种情况:

if (!next.empty() && next[0] == '?') {
    next = next.substr(1);        // 跳过 '?',获取后续模式
    // 情况1:尝试匹配一次
    if (!str.empty() && MatchCharacter(re, 0, str[0])) {
        ++n;                      // 增加匹配计数
        if (MatchPattern(next, str.substr(1), n)) return true;
        --n;                      // 匹配失败,恢复计数
    }
    // 情况2:匹配零次(跳过当前字符)
    return MatchPattern(next, str, n);
}

执行逻辑

  1. 优先匹配一次:如果当前字符能匹配,先尝试消耗一个字符
  2. 递归验证:检查剩余模式是否能匹配剩余字符串
  3. 回退到零次:如果匹配一次失败,则尝试匹配零次(不消耗字符)
  4. 贪婪特性:优先选择匹配一次而不是零次

示例:模式 a?b 匹配字符串 "ab"

  • 尝试匹配 1 个 a:消耗 "a",剩余 "b" 匹配模式 b → 成功!

示例:模式 a?b 匹配字符串 "b"

  • 尝试匹配 1 个 a:当前字符是 b,不匹配 a
  • 尝试匹配 0 个 a:直接用 "b" 匹配模式 b → 成功!

3. 顺序匹配(精确一次匹配)

当模式中没有量词时,执行顺序匹配:

++n;                              // 增加匹配计数
return !str.empty() &&           // 确保字符串非空
       MatchCharacter(re, 0, str[0]) &&  // 当前字符必须匹配
       MatchPattern(next, str.substr(1), n);  // 递归匹配剩余部分

执行逻辑

  1. 字符串检查:首先确保目标字符串不为空
  2. 字符匹配:使用 MatchCharacter 验证当前字符是否匹配模式
  3. 递归处理:消耗一个字符,继续匹配剩余的模式和字符串
  4. 计数更新:成功匹配时增加匹配字符计数

示例:模式 abc 匹配字符串 "abc"

  • 匹配 astr[0]='a' 匹配 re[0]='a' → 成功
  • 递归匹配 bc vs "bc"
    • 匹配 bstr[0]='b' 匹配 re[0]='b' → 成功
    • 递归匹配 c vs "c"
      • 匹配 cstr[0]='c' 匹配 re[0]='c' → 成功
      • 递归匹配 `` vs "":空模式匹配空串 → 成功

失败情况:模式 abc 匹配字符串 "axc"

  • 匹配 astr[0]='a' 匹配 re[0]='a' → 成功
  • 递归匹配 bc vs "xc"
    • 匹配 bstr[0]='x' 不匹配 re[0]='b' → 失败

双指针移动

int32_t MatchRegex(const std::string re, const std::string& str, uint32_t& n) {
    if (re[0] == '^') {
        n = 0;
        if (MatchPattern(re.substr(1), str, n)) return 0;
        return -1;
    }    

    for (size_t pos = 0; pos <= str.length(); ++pos) {
        n = 0;
        if (MatchPattern(re, str.substr(pos), n)) return pos;
    }
    return -1;
}

该函数是整个正则引擎的入口,负责处理锚点和搜索逻辑,若str不匹配re返回-1,否则返回匹配开始的位置:

锚点处理

开头锚点 ‘^’

if (re[0] == '^') {
    n = 0;
    if (MatchPattern(re.substr(1), str, n)) return 0;
    return -1;
}
  • 如果模式以 ^ 开头,则只在字符串开始位置尝试匹配
  • 去掉 ^ 后调用 MatchPattern 匹配剩余模式
  • 成功返回位置 0,失败返回 -1- 如果模式以 ‘^’ 开头,则只在字符串开始位置尝试匹配
  • 去掉 ‘^’ 后调用 ’MatchPattern

结尾锚点 $: 在 MatchPattern 中处理:

if (re[0] == '$' && re.length() == 1) return str.empty();
  • 只有当 $ 是整个模式时才作为结尾锚点
  • 要求剩余字符串为空才匹配成功

逐字搜索

for (size_t pos = 0; pos <= str.length(); ++pos) {
    n = 0;
    if (MatchPattern(re, str.substr(pos), n)) return pos;
}
return -1;

执行逻辑

  1. 遍历所有位置:从字符串的每个位置开始尝试匹配
  2. 重置计数器:每次尝试前将匹配计数 n 重置为 0
  3. 子串匹配:对从当前位置开始的子串进行模式匹配
  4. 返回首次匹配位置:找到匹配则立即返回位置,否则继续搜索
  5. 全部失败:所有位置都匹配失败则返回 -1

示例:模式 abc 在字符串 "xyzabc" 中搜索

  • pos=0:MatchPattern("abc", "xyzabc") → 失败
  • pos=1:MatchPattern("abc", "yzabc") → 失败
  • pos=2:MatchPattern("abc", "zabc") → 失败
  • pos=3:MatchPattern("abc", "abc") → 成功!返回 3

性能分析

Wapiti项目中用这个Regex实现来做特征抽取,场景比较固定,性能要求不高, 不过要是实现一个通用的正则引起,这个方案就不行了。除去代码中涉及到递归函数, 更关键的问题是对量词的处理上。

当前方案在失败较晚、多个量词相互组合时可能产生组合爆炸,最坏情况呈指数级增长。外层的逐位置搜索以及频繁的 substr() 复制还会带来额外开销。O(n^m) 可以作为理解多个量词组合数量的直观近似,但不是所有模式的统一精确上界。

经典案例分析

考虑模式 a*a*b 匹配字符串 "aaaaac"(最后是 c 不是 b,必然失败):

字符串: a a a a a c
模式:   a * a * b

对应的关键代码:

do {
    uint32_t save = n;
    if (MatchPattern(next, str.substr(pos), n)) return true;  // 分支
    n = save + 1;
    pos++;
} while (pos <= str.length() && MatchCharacter(re, 0, str[pos-1]));

数学分析

对于 n 个 a 字符,两个 a* 的分配方案数:

  • 第一个 a* 匹配 i 个,第二个匹配 (n-i) 个
  • i 可以从 0 到 n,共 (n+1) 种组合

实际上,回溯会尝试所有可能的组合

n=4 时的尝试次数

  • 第一个 a* 匹配 0 个:第二个 a* 尝试 0,1,2,3,4 → 5次
  • 第一个 a* 匹配 1 个:第二个 a* 尝试 0,1,2,3 → 4次
  • 第一个 a* 匹配 2 个:第二个 a* 尝试 0,1,2 → 3次
  • 第一个 a* 匹配 3 个:第二个 a* 尝试 0,1 → 2次
  • 第一个 a* 匹配 4 个:第二个 a* 尝试 0 → 1次

总计:5+4+3+2+1 = 15 = O(n²)

一般化公式

对于 m 个量词和 n 个字符的情况,复杂度为:

  • 2 个量词:O(n²)
  • 3 个量词:O(n³)
  • m 个量词:O(n^m)

Compiler

前面实现的匹配方法虽然简单直观,支持的语法也较少,更关键的是存在指数级时间复杂度的根本问题。高性能的正则表达式引擎,需要支持更多的语法,需要用更系统化的编译器方法:

  1. 语法分析:递归下降 Parser 直接读取正则字符,构建抽象语法树 AST;本实现没有独立的 Lexer
  2. NFA 构建:通过 Visitor 遍历 AST,使用 Thompson 构造将其转换为 NFA
  3. DFA 构建:通过子集构造将 NFA 转换为 DFA
  4. 匹配:逐个读取 Unicode code point,沿 DFA 的唯一转移前进

这样的编译器实现除了能把正则表达式的匹配问题转换为有限自动机的状态转换问题,实现线性时间复杂度的字符串匹配,还能让扩展正则表达式的功能也更容易些,这会是一种教科书级别的实现方案。

BNF 语法

BNF (Backus-Naur Form) 是一种用于描述上下文无关文法的标准表示法,其使用以下符号:

  • ::= 表示“定义为“或“产生“
  • | 表示“或者“,用于分隔不同的选择
  • <> 包围非终结符
  • 不在 <> 中的符合是终结符(具体的字符或Token)

正则表达式的BNF定义

以下实现的正则表达式语法支持以下操作:

<Pattern>    ::= <Sequence> ('|' <Sequence>)*
<Sequence>   ::= <Element>*
<Element>    ::= <Atom> <Quantifier>?
<Atom>       ::= <Literal> | '.' | '(' <Pattern> ')'
<Quantifier> ::= '*' | '+' | '?'
<Literal>    ::= UTF-8字符(除了特殊字符)

语法特性

优先级和结合性

该语法具有以下优先级(从高到低):

  1. 原子:字面量、点、括号表达式
  2. 量词*+?(后缀,右结合)
  3. 连接:序列中的元素连接(左结合)
  4. 选择| 操作符(左结合)

示例分析

  • ab*a(b*) 而不是 (ab)*
  • a|bca|(bc) 而不是 (a|b)c
  • a|b|c((a|b)|c) 左结合

递归下降解析器

递归下降解析是一种自顶向下的语法分析方法。每个非终结符对应一个解析函数,函数调用关系直接反映文法层次;Atom 遇到括号时再次调用 ParsePattern,因而能够处理嵌套表达式。本实现的文法不需要回溯。

实现详解

RegexParser类包含以下核心组件:

class RegexParser {
private:
    std::string pattern;    // 待解析的正则表达式
    size_t pos = 0;        // 当前解析位置
    int cnt = 0;           // 调试用的缩进计数

public:
    explicit RegexParser(std::string p) : pattern(std::move(p)) {}
    std::unique_ptr<Ast> Parse();  // 主解析入口

private:
    // 对应BNF中的每个非终结符
    std::unique_ptr<Ast> ParsePattern();
    std::unique_ptr<Ast> ParseSequence();
    std::unique_ptr<Ast> ParseElement();
    std::unique_ptr<Ast> ParseAtom();
};

ParsePattern函数 - 处理选择操作

std::unique_ptr<Ast> ParsePattern() {
    PrintEnter("ParsePattern", "<sequence> ('|' <sequence>)*");

    if (pos >= pattern.length()) {
        PrintExit("ParsePattern", "Empty");
        return std::make_unique<EmptyAst>();
    }

    auto n = std::make_unique<AlternativeAst>();
    n->InsertBranch(ParseSequence());  // 解析第一个序列

    while (pos < pattern.length() && pattern[pos] == '|') {
        std::cout << std::string(cnt*2, ' ') << "Got '|', Parse next branch\n";
        ++pos;  // 消耗'|'字符
        n->InsertBranch(ParseSequence());  // 解析下一个序列
    }

    PrintExit("ParsePattern", "AlternativeNode");
    return n;
}

关键设计点

  1. 空模式处理:如果到达字符串末尾,返回EmptyAst
  2. 至少一个分支:总是解析一个序列,确保Alternative节点至少有一个分支
  3. 循环处理多个分支:使用while循环处理所有|分隔的序列
  4. 位置管理:每次遇到|时递增pos以消耗该字符

ParseSequence函数 - 处理连接操作

std::unique_ptr<Ast> ParseSequence() {
    PrintEnter("ParseSequence", "<element>*");

    auto s = std::make_unique<SequenceAst>();

    while (pos < pattern.length() && 
           pattern[pos] != '|' && 
           pattern[pos] != ')') {
        std::cout << std::string(cnt*2, ' ') << "Parse next element\n";
        s->InsertElement(ParseElement());
    }

    PrintExit("ParseSequence", "SequenceNode");
    return s;
}

终止条件分析

  1. 到达字符串末尾pos >= pattern.length()
  2. 遇到选择分隔符pattern[pos] == '|'
  3. 遇到分组结束pattern[pos] == ')'

这些条件确保了序列解析在适当的边界停止,不会越界处理属于上层语法结构的字符。

ParseElement函数 - 处理量词

std::unique_ptr<Ast> ParseElement() {
    PrintEnter("ParseElement", "<atom> <quantifier>?");

    if (pos >= pattern.length()) {
        PrintExit("ParseElement", "Empty");
        return std::make_unique<EmptyAst>();
    }

    auto atom = ParseAtom();  // 先解析原子

    if (pos < pattern.length()) {
        char quantifier = pattern[pos];
        switch (quantifier) {
            case '*':
                std::cout << std::string(cnt * 2, ' ') << "Got Quantifier '*'\n";
                ++pos;
                PrintExit("ParseElement", "StarNode");
                return std::make_unique<StarAst>(std::move(atom));
            case '+':
                std::cout << std::string(cnt * 2, ' ') << "Got Quantifier '+'\n";
                ++pos;
                PrintExit("ParseElement", "PlusNode");
                return std::make_unique<PlusAst>(std::move(atom));
            case '?':
                std::cout << std::string(cnt * 2, ' ') << "Got Quantifier '?'\n";
                ++pos;
                PrintExit("ParseElement", "OptionalNode");
                return std::make_unique<OptionalAst>(std::move(atom));
        }
    }

    PrintExit("ParseElement", "AtomNode");
    return atom;
}

处理流程

  1. 原子优先:总是先解析原子部分
  2. 量词检测:检查原子后是否有量词字符
  3. 包装创建:根据量词类型创建相应的包装节点
  4. 位置递进:消耗量词字符并更新位置

ParseAtom函数 - 处理基本单元

std::unique_ptr<Ast> ParseAtom() {
    PrintEnter("ParseAtom", "<literal> | '.' | '(' <pattern> ')'");

    if (pos >= pattern.length()) {
        throw std::runtime_error("ParseAtom: Unexpected EOF");
    }

    char c = pattern[pos];
    std::unique_ptr<Ast> element;

    if (c == '(') {
        // 处理分组 '(' <pattern> ')'
        std::cout << std::string(cnt*2, ' ') << "Got '(', Parse down pattern\n";
        ++pos;  // 消耗'('
        element = ParsePattern();  // 递归解析子模式

        if (pos >= pattern.length() || pattern[pos] != ')') {
            throw std::runtime_error("ParseAtom: No matching ')'");
        }
        std::cout << std::string(cnt*2, ' ') << "Got ')', Parse down pattern done\n";
        ++pos;  // 消耗')'
        PrintExit("ParseAtom", "() Expression");
        return element;
        
    } else if (c == '.') {
        // 处理通配符
        std::cout << std::string(cnt*2, ' ') << "Got '.'\n";
        element = std::make_unique<DotAst>();
        ++pos;
        PrintExit("ParseAtom", "DotNode");
        return element;
        
    } else if (c != '|' && c != ')' && c != '*' && c != '+' && c != '?') {
        // 处理UTF-8字面量
        size_t bytes;
        uint32_t codepoint = DecodeUTF8At(pattern, pos, &bytes);

        if (bytes > 0) {
            std::string char_str = pattern.substr(pos, bytes);
            std::cout << std::string(cnt*2, ' ') << "Got UTF-8 '" << char_str 
                      << "' (U+" << std::hex << codepoint << std::dec << ")\n";
            element = std::make_unique<LiteralAst>(codepoint);
            pos += bytes;  // UTF-8字符可能占多个字节
            PrintExit("ParseAtom", "LiteralNode");
            return element;
        } else {
            throw std::runtime_error("ParseAtom: Invalid UTF-8 at " + std::to_string(pos));
        }
    } else {
        throw std::runtime_error("ParseAtom: Unexpected '" + std::string(1, c) + "' at pos " + std::to_string(pos));
    }
}

关键特性

  1. 分组处理:括号表达式通过递归调用ParsePattern处理
  2. UTF-8支持:使用专门的UTF-8解码函数处理多字节字符
  3. 错误检测:对不匹配的括号和无效字符进行错误处理
  4. 特殊字符过滤:确保特殊语法字符不被当作字面量处理

完整示例

通过一个完整的例子来观察解析过程:

输入"(a|b)*c"

ParsePattern() - pos=0
├── ParseSequence() - pos=0
    ├── ParseElement() - pos=0
    │   ├── ParseAtom() - pos=0
    │   │   ├── 发现'(' - pos=1
    │   │   └── ParsePattern() - pos=1 (递归)
    │   │       ├── ParseSequence() - pos=1
    │   │       │   └── ParseElement() - pos=1
    │   │       │       └── ParseAtom() - pos=1
    │   │       │           └── 解析'a' - pos=2
    │   │       ├── 发现'|' - pos=3
    │   │       └── ParseSequence() - pos=3
    │   │           └── ParseElement() - pos=3
    │   │               └── ParseAtom() - pos=3
    │   │                   └── 解析'b' - pos=4
    │   ├── 发现')' - pos=5
    │   └── 发现'*' - pos=6 (创建StarAst)
    └── ParseElement() - pos=6
        └── ParseAtom() - pos=6
            └── 解析'c' - pos=7

最终AST结构

SequenceAst
├── StarAst
│   └── AlternativeAst
│       ├── LiteralAst('a')
│       └── LiteralAst('b')
└── LiteralAst('c')

抽象语法树

**抽象语法树(Abstract Syntax Tree, AST)**是源代码语法结构的树状表示,它具有以下特点:

  1. 抽象化:去除了具体语法中的冗余信息(如括号,分隔符)
  2. 结构和:保留了语义相关的层次结构关系
  3. 类型化:每个节点都有明确的类型和语义
  4. 可遍历:支持各种遍历和转化操作

AST节点类型

正则表达式AST包含以下节点类型:

enum class AstType {
    Empty,        // 空表达式 ε
    Literal,      // 字面量字符
    Dot,          // 通配符 .
    Sequence,     // 序列连接
    Alternative,  // 选择 |
    Star,         // 零或多次 *
    Plus,         // 一或多次 +
    Optional      // 零或一次 ?
};

基础类

class Ast {
public:
    virtual ~Ast() = default;
    virtual AstType GetType() const = 0;
    virtual void Accept(AstVisitor* visitor) const = 0;  // 访问者模式接口
};

叶子节点类

空节点 - 表示空字符串:

class EmptyAst : public Ast {
public:
    AstType GetType() const override { return AstType::Empty; }
    void Accept(AstVisitor* v) const override { v->Visit(this); }
};

字面量节点 - 表示具体字符:

class LiteralAst : public Ast {
private:
    uint32_t point;  // Unicode码点

public:
    explicit LiteralAst(uint32_t p) : point(p) {}
    
    AstType GetType() const override { return AstType::Literal; }
    uint32_t GetPoint() const { return point; }
    void Accept(AstVisitor* v) const override { v->Visit(this); }
};

通配符节点 - 表示点操作符:

class DotAst : public Ast {
public:
    AstType GetType() const override { return AstType::Dot; }
    void Accept(AstVisitor* v) const override { v->Visit(this); }
};

复合节点类

序列节点 - 表示元素连接:

class SequenceAst : public Ast {
private:
    std::vector<std::unique_ptr<Ast>> elements;

public:
    void InsertElement(std::unique_ptr<Ast> e) {
        elements.push_back(std::move(e));
    }

    const std::vector<std::unique_ptr<Ast>>& GetElements() const {
        return elements;
    }

    AstType GetType() const override { return AstType::Sequence; }
    void Accept(AstVisitor* v) const override { v->Visit(this); }
};

选择节点 - 表示或操作:

class AlternativeAst : public Ast {
private:
    std::vector<std::unique_ptr<Ast>> branches;  // 注意:原代码中变量名有拼写错误

public:
    void InsertBranch(std::unique_ptr<Ast> branch) {
        branches.push_back(std::move(branch));
    }

    const std::vector<std::unique_ptr<Ast>>& GetBranches() const {
        return branches;
    }

    AstType GetType() const override { return AstType::Alternative; }
    void Accept(AstVisitor* v) const override { v->Visit(this); }
};

量词节点 - 表示重复操作:

// 星号量词 (0或多次)
class StarAst : public Ast {
private:
    std::unique_ptr<Ast> element;

public:
    explicit StarAst(std::unique_ptr<Ast> e) : element(std::move(e)) {}
    
    const Ast* GetElement() const { return element.get(); }
    AstType GetType() const override { return AstType::Star; }
    void Accept(AstVisitor* v) const override { v->Visit(this); }
};

// 加号量词 (1或多次) 和 问号量词 (0或1次) 结构类似...

访问者模式

**访问者模式(Visitor Pattern)**是一种行为设计模式,它允许在不修改现有类结构的情况下,对类层次结构添加新的操作。

核心组件

1. 访问者接口

class AstVisitor {
public:
    virtual ~AstVisitor() = default;
    virtual void Visit(const EmptyAst* node) = 0;
    virtual void Visit(const LiteralAst* node) = 0;
    virtual void Visit(const DotAst* node) = 0;
    virtual void Visit(const SequenceAst* node) = 0;
    virtual void Visit(const AlternativeAst* node) = 0;
    virtual void Visit(const StarAst* node) = 0;
    virtual void Visit(const PlusAst* node) = 0;
    virtual void Visit(const OptionalAst* node) = 0;
};

2. 可访问接口

// 在每个AST节点中
virtual void Accept(AstVisitor* visitor) const = 0;

// 具体实现(双分派机制)
void Accept(AstVisitor* v) const override { 
    v->Visit(this);  // this的类型决定了调用哪个Visit重载
}

访问者把算法与节点数据分开:打印器可以输出树形结构,NFABuilder 则使用同一套 Accept/Visit 接口把各类节点转换成 NFA 片段。后续增加新的 AST 节点时,需要同时为相关访问者补充对应的 Visit 方法。

NFA

非确定有限自动机(Nondeterministic Finite Automation, NFA)是一种理论计算模型,具有以下特征:

  1. 状态集合:有限个状态的集合
  2. 输入字母表:可接受的输入符号集合
  3. 转换函数:从状态和输入符号到状态集合的映射(非确定性)
  4. 初始状态:自动机的起始状态
  5. 接受状态集合:表示匹配成功的状态集合

NFA的“非确定性“体现在:

  • ε转换:不消耗输入字符的状态转换
  • 多重转换:从一个状态在同一输入上可以转换到多个状态
  • 并行执行:可以同时处于多个状态

下面通过代码详解。

状态设计

NFA状态类包含以下组件:

class NFAState {
public:
    int i;                                              // 状态编号
    bool end = false;                                   // 是否为接受状态
    std::map<uint32_t, std::vector<NFAState*>> transitions;  // 字符转换
    std::vector<NFAState*> e_transitions;              // ε转换
    
    static constexpr uint32_t DOT_CHAR = 0xFFFFFFFF;   // 通配符的特殊标记
    static int next;                                    // 全局状态计数器

    NFAState() : i(next++) {}

    void InsertTransition(uint32_t p, NFAState* t) {
        transitions[p].push_back(t);
    }

    void InsertEpsilonTransition(NFAState* t) {
        e_transitions.push_back(t);
    }

    void InsertDotTransition(NFAState* t) {
        transitions[DOT_CHAR].push_back(t);
    }
};

设计要点

  1. 状态编号:每个状态有唯一的编号,便于调试和可视化
  2. 接受标记end字段标识接受状态
  3. 字符转换表map<uint32_t, vector<NFAState*>>支持一对多的转换
  4. ε转换列表:专门处理不消耗字符的转换
  5. 通配符处理:使用特殊值0xFFFFFFFF表示通配符转换

片段结构

struct NFA {
    NFAState* start;  // 起始状态
    NFAState* end;    // 结束状态

    NFA(NFAState* s, NFAState* e) : start(s), end(e) {}
};

每个NFA片段都有明确的入口和出口,这种设计便于组合和连接。

Thompson构造法

Thompson构造法将正则表达式AST转换为NFA。这种方法为每种正则表达式操作定义了标准的NFA模板。

基本构造模板

1. 空表达式 (ε)

状态图:
[S] --ε--> [E]

代码实现:
void Visit(const EmptyAst* node) override {
    auto start = NewState();
    auto end = NewState();
    end->end = true;
    start->InsertEpsilonTransition(end);
    stack.push(NFA(start, end));
}

2. 字面量字符 (a)

状态图:
[S] --a--> [E]

代码实现:
void Visit(const LiteralAst* node) override {
    uint32_t p = node->GetPoint();
    auto start = NewState();
    auto end = NewState();
    end->end = true;
    start->InsertTransition(p, end);
    stack.push(NFA(start, end));
}

3. 通配符 (.)

状态图:
[S] --.--> [E]

代码实现:
void Visit(const DotAst* node) override {
    auto start = NewState();
    auto end = NewState();
    end->end = true;
    start->InsertDotTransition(end);
    stack.push(NFA(start, end));
}

复合构造模板

4. 序列连接 (AB)

对于序列A B C,构造过程:

A: [S1] --> [E1]
B: [S2] --> [E2]  
C: [S3] --> [E3]

连接后:
[S1] --> [E1] --ε--> [S2] --> [E2] --ε--> [S3] --> [E3]
void Visit(const SequenceAst* node) override {
    const auto& elements = node->GetElements();
    if (elements.empty()) {
        Visit(static_cast<const EmptyAst*>(nullptr));  // 创建空NFA
        return;
    }

    elements[0]->Accept(this);  // 构造第一个元素的NFA

    for (size_t i = 1; i < elements.size(); ++i) {
        auto left = stack.top(); stack.pop();    // 左侧NFA
        elements[i]->Accept(this);               // 构造右侧NFA
        auto right = stack.top(); stack.pop();  // 右侧NFA

        // 连接:左端点 --ε--> 右起点
        left.end->end = false;                   // 左端点不再是接受状态
        left.end->InsertEpsilonTransition(right.start);
        
        stack.push(NFA(left.start, right.end)); // 新NFA:左起点到右端点
    }
}

5. 选择操作 (A|B)

对于选择A | B,构造模板:

原始:
A: [S1] --> [E1]
B: [S2] --> [E2]

构造后:
        --ε--> [S1] --> [E1] --ε--
       /                           \
[新S]                               --> [新E]
       \                           /
        --ε--> [S2] --> [E2] --ε--
void Visit(const AlternativeAst* node) override {
    const auto& branches = node->GetBranches();
    if (branches.empty()) {
        Visit(static_cast<const EmptyAst*>(nullptr));
        return;
    }

    branches[0]->Accept(this);  // 构造第一个分支

    for (size_t i = 1; i < branches.size(); ++i) {
        auto left = stack.top(); stack.pop();    // 左分支NFA
        branches[i]->Accept(this);               // 构造右分支NFA
        auto right = stack.top(); stack.pop();  // 右分支NFA
        
        auto start = NewState();  // 新的起始状态
        auto end = NewState();    // 新的结束状态
        end->end = true;
        
        // 起始状态分别连接两个分支的起点
        start->InsertEpsilonTransition(left.start);
        start->InsertEpsilonTransition(right.start);
        
        // 两个分支的终点都连接到新的结束状态
        left.end->end = false;
        right.end->end = false;
        left.end->InsertEpsilonTransition(end);
        right.end->InsertEpsilonTransition(end);
        
        stack.push(NFA(start, end));
    }
}

6. 星号量词 (A)*

对于A*,构造模板:

       --ε--> [S1] --> [E1] --ε--
      /          ^           \   |
[新S]            |            v  |
      \          ε            [新E]
       \         |           /
        -------ε----------->
void Visit(const StarAst* node) override {
    node->GetElement()->Accept(this);  // 构造内部表达式A的NFA
    auto inner = stack.top(); stack.pop();
    
    auto start = NewState();
    auto end = NewState();
    end->end = true;
    
    // 可以直接跳过A(匹配0次)
    start->InsertEpsilonTransition(inner.start);
    start->InsertEpsilonTransition(end);
    
    // A的结束可以回到A的开始(匹配多次)或者结束
    inner.end->end = false;
    inner.end->InsertEpsilonTransition(inner.start);  // 循环
    inner.end->InsertEpsilonTransition(end);          // 结束
    
    stack.push(NFA(start, end));
}

7. 加号量词 (A+)

对于A+,构造模板(与A*类似,但必须至少匹配一次):

[新S] --ε--> [S1] --> [E1] --ε--> [新E]
                ^           |
                |           |
                ----ε-------
void Visit(const PlusAst* node) override {
    node->GetElement()->Accept(this);
    auto inner = stack.top(); stack.pop();
    
    auto start = NewState();
    auto end = NewState();
    end->end = true;
    
    // 必须至少匹配一次A
    start->InsertEpsilonTransition(inner.start);
    
    // A的结束可以回到A的开始(匹配多次)或者结束
    inner.end->end = false;
    inner.end->InsertEpsilonTransition(inner.start);  // 循环
    inner.end->InsertEpsilonTransition(end);          // 结束
    
    stack.push(NFA(start, end));
}

8. 问号量词 (A?)

对于A?,构造模板:

       --ε--> [S1] --> [E1] --ε--
      /                        \
[新S]                          --> [新E]
      \                        /
       --------ε -------->----
void Visit(const OptionalAst* node) override {
    node->GetElement()->Accept(this);
    auto inner = stack.top(); stack.pop();
    
    auto start = NewState();
    auto end = NewState();
    end->end = true;
    
    // 可以匹配A或者跳过A
    start->InsertEpsilonTransition(inner.start);  // 匹配A
    start->InsertEpsilonTransition(end);          // 跳过A
    
    // A匹配完成后到达结束状态
    inner.end->end = false;
    inner.end->InsertEpsilonTransition(end);
    
    stack.push(NFA(start, end));
}

NFA构建示例

通过构建"a*b"的NFA来演示完整过程:

步骤1:解析AST

SequenceAst
├── StarAst
│   └── LiteralAst('a')
└── LiteralAst('b')

步骤2:构建子NFA

构建LiteralAst(‘a’)

状态0 --'a'--> 状态1(接受)

构建StarAst

状态2 --ε--> 状态0 --'a'--> 状态1 --ε--> 状态3(接受)
  |                            |
  |                            v
   ----ε----> 状态3 <----ε----

构建LiteralAst(‘b’)

状态4 --'b'--> 状态5(接受)

步骤3:连接序列

a*b连接:

状态2 --ε--> 状态0 --'a'--> 状态1 --ε--> 状态3 --ε--> 状态4 --'b'--> 状态5(接受)
  |                            |
  |                            v
   ----ε----> 状态3 <----ε----

简化后的最终NFA:

状态2 --ε--> 状态0 --'a'--> 状态1 --ε--> 状态4 --'b'--> 状态5(接受)
  |                            |           ^
  |                            v           |
   ----ε----> 状态4 <----ε----            |
              |                            |
               --------'b'---------------->状态5

DFA

确定有限自动机(Deterministic Finite Automation, DFA) 是NFA的确定性版本,具有以下特征:

  1. 确定性转换:每个状态在给定输入下最多只能转换到一个状态
  2. 无ε转换:所有转换都必须消耗输入字符
  3. 高效匹配:匹配时间复杂度为O(n),其中n是输入字符串长度
  4. 空间换时间:可能产生指数级的状态数量

DFA vs NFA对比

特性NFADFA
状态转换非确定性(一对多)确定性(一对一)
ε转换支持不支持
并行状态可以同时处于多个状态任意时刻只在一个状态
构造复杂度简单,状态数较少复杂,状态数可能指数增长
匹配效率O(nm),需要跟踪多个状态O(n),直接状态转换

状态设计

class DFAState {
public:
    int i;                                    // 状态编号
    bool end = false;                         // 是否为接受状态
    std::map<uint32_t, DFAState*> transitions; // 确定性转换表(一对一)

    static int next;

    DFAState() : i(next++) {}
};

注意DFA状态与NFA状态的区别:

  • 转换表类型map<uint32_t, DFAState*> vs map<uint32_t, vector<NFAState*>>
  • 无ε转换:DFA不需要ε转换列表
  • 确定性:每个输入字符最多对应一个目标状态

子集构造法

子集构造法将NFA转换为等价的DFA。基本思想是:DFA的每个状态对应NFA状态的一个子集。

核心算法组件

1. ε闭包计算(Epsilon Closure)

ε闭包是指从给定状态集合出发,通过任意数量的ε转换能够到达的所有状态集合。

void EpsilonClosure(std::set<NFAState*>& states) {
    std::stack<NFAState*> stack;

    // 将所有当前状态压入栈
    for (auto state : states) {
        stack.push(state);
    }

    while (!stack.empty()) {
        auto current = stack.top();
        stack.pop();

        // 处理当前状态的所有ε转换
        for (auto target : current->e_transitions) {
            if (states.find(target) == states.end()) {
                states.insert(target);  // 添加新发现的状态
                stack.push(target);     // 继续探索新状态的ε转换
            }
        }
    }
}

算法示例

输入状态集合: {状态0}
状态0的ε转换: [状态1, 状态3]

执行过程:
1. 初始: states = {0}, stack = [0]
2. 处理状态0: 发现状态1和3,states = {0,1,3}, stack = [1,3]
3. 处理状态1: 无ε转换,stack = [3]
4. 处理状态3: 无ε转换,stack = []
5. 结果: states = {0,1,3}

2. Move操作

Move操作计算状态集合在给定输入字符下能够转换到的所有状态。

std::set<NFAState*> Move(const std::set<NFAState*>& states, uint32_t input) {
    std::set<NFAState*> result;

    for (auto state : states) {
        // 处理精确字符匹配
        auto it = state->transitions.find(input);
        if (it != state->transitions.end()) {
            for (auto target : it->second) {
                result.insert(target);
            }
        }

        // 处理通配符匹配(除了换行符)
        if (input != 10 && input != 13) {  // 不是\n或\r
            auto dot_it = state->transitions.find(NFAState::DOT_CHAR);
            if (dot_it != state->transitions.end()) {
                for (auto target : dot_it->second) {
                    result.insert(target);
                }
            }
        }
    }

    return result;
}

3. 接受状态检测

bool ContainsEndState(const std::set<NFAState*>& states) {
    for (auto state : states) {
        if (state->end) return true;
    }
    return false;
}

主构造算法

DFAState* Build(const NFA& nfa) {
    std::map<std::set<NFAState*>, DFAState*> state_map;  // NFA状态集合到DFA状态的映射
    std::queue<std::set<NFAState*>> queue;               // 待处理的状态集合队列

    // 1. 创建初始DFA状态
    std::set<NFAState*> start_set = {nfa.start};
    EpsilonClosure(start_set);  // 计算初始状态的ε闭包

    auto start_dfa = NewState();
    if (ContainsEndState(start_set)) {
        start_dfa->end = true;  // 如果ε闭包包含接受状态,则DFA起始状态也是接受状态
    }

    state_map[start_set] = start_dfa;
    queue.push(start_set);

    // 2. 处理队列中的每个状态集合
    while (!queue.empty()) {
        auto current_set = queue.front();
        queue.pop();

        auto current_dfa = state_map[current_set];

        // 3. 收集所有可能的输入字符
        std::set<uint32_t> alphabet;
        for (auto state : current_set) {
            for (const auto& [input, targets] : state->transitions) {
                alphabet.insert(input);
            }
        }

        // 4. 为每个输入字符构造转换
        for (uint32_t input : alphabet) {
            auto next_set = Move(current_set, input);  // 计算转换后的状态集合
            if (next_set.empty()) continue;

            EpsilonClosure(next_set);  // 计算ε闭包

            DFAState* next_dfa;
            if (state_map.find(next_set) == state_map.end()) {
                // 发现新的状态集合,创建对应的DFA状态
                next_dfa = NewState();
                if (ContainsEndState(next_set)) {
                    next_dfa->end = true;
                }
                state_map[next_set] = next_dfa;
                queue.push(next_set);  // 添加到队列中继续处理
            } else {
                // 状态集合已存在,直接获取对应的DFA状态
                next_dfa = state_map[next_set];
            }

            current_dfa->transitions[input] = next_dfa;  // 添加转换
        }
    }

    return start_dfa;
}

构建示例

通过"a*"的例子来演示DFA构建过程:

输入NFA

State 0 (START):
  --ε--> State 1
  --ε--> State 3

State 1:
  --'a'--> State 2

State 2:
  --ε--> State 1
  --ε--> State 3

State 3 (END):

构建步骤

步骤1:初始状态集合

NFA状态集合: {0}
ε闭包: {0, 1, 3}  # 状态0可以通过ε转换到达状态1和3
DFA状态: State 0 (接受状态,因为包含NFA状态3)

步骤2:处理输入’a’

当前集合: {0, 1, 3}
输入'a'的Move结果: {2}  # 只有状态1在输入'a'时转换到状态2
ε闭包({2}): {1, 2, 3}   # 状态2可以通过ε转换到状态1和3
创建DFA状态: State 1 (接受状态,因为包含NFA状态3)
添加转换: State 0 --'a'--> State 1

步骤3:处理新状态集合{1, 2, 3}

当前集合: {1, 2, 3}
输入'a'的Move结果: {2}  # 状态1在输入'a'时转换到状态2
ε闭包({2}): {1, 2, 3}   # 与已存在的状态集合相同
添加转换: State 1 --'a'--> State 1 (自循环)

最终DFA

=== DFA Struct ===
State 0 (START, END):
  --'a'--> State 1

State 1 (END):
  --'a'--> State 1

这个DFA完美地表示了a*的语义:

  • State 0:初始状态,也是接受状态(可以匹配空字符串)
  • State 1:匹配了至少一个’a’后的状态,也是接受状态
  • 自循环:State 1的自循环表示可以匹配任意多个’a’

匹配算法

DFA的匹配算法非常简单高效:

bool Match(DFAState* start, const std::string& text) {
    auto current = start;
    size_t pos = 0;

    while (pos < text.size()) {
        size_t bytes;
        uint32_t codepoint = DecodeUTF8At(text, pos, &bytes);

        if (bytes == 0) {
            return false;  // 无效的UTF-8编码
        }

        // 优先使用精确转移,否则回退到通配符转移
        auto it = current->transitions.find(codepoint);
        if (it != current->transitions.end()) {
            current = it->second;
        } else if (codepoint != '\n' && codepoint != '\r') {
            auto dot = current->transitions.find(NFAState::DOT_CHAR);
            if (dot == current->transitions.end()) {
                return false;
            }
            current = dot->second;
        } else {
            return false;
        }

        pos += bytes;
    }

    return current->end;  // 检查最终状态是否为接受状态
}

算法特点

  1. 线性时间复杂度:O(n),其中n是输入字符串长度
  2. 确定性执行:每次只需要查找一个转换
  3. UTF-8支持:正确处理多字节字符
  4. 简单直观:状态转换逻辑清晰

引擎实现

Regex类

现在把全部组件整合到一个Regex类中:

class Regex {
private:
    std::unique_ptr<Ast> ast;        // 抽象语法树
    DFAState* dfa = nullptr;         // 编译后的DFA
    DFABuilder builder;              // DFA构建器

public:
    Regex(const std::string& pattern) {
        std::cout << "\n=== Compile Regex Pattern: \""
                  << pattern << "\" ===\n";
        
        // 1. 语法分析:构建AST
        RegexParser parser(pattern);
        ast = parser.Parse();
        std::cout << "\n✓ Recursive Descent Parse Done\n";

        // 2. AST可视化
        std::cout << "\n=== AST Structure ===\n";
        RegexPrinter printer;
        ast->Accept(&printer);

        // 3. NFA构建
        std::cout << "\n=== AST -> NFA Conversion ===\n";
        NFABuilder nfa_builder;
        ast->Accept(&nfa_builder);
        auto nfa = nfa_builder.GetNFA();
        std::cout << "✓ AST -> NFA Conversion Done\n";
        nfa_builder.PrintNFA(nfa);

        // 4. DFA构建
        dfa = builder.Build(nfa);
        builder.PrintDFA(dfa);
    }

    bool Match(const std::string& text) {
        return builder.Match(dfa, text);
    }
};

Pipeline

该正则表达式引擎使用了标准的编译器流水线:

正则表达式字符串
        ↓
    [词法分析] (隐含在解析器中)
        ↓
   [语法分析] (递归下降解析器)
        ↓
   抽象语法树 (AST)
        ↓
   [语义分析] (访问者模式遍历)
        ↓
      [代码生成] (Thompson构造法)
        ↓
    非确定有限自动机 (NFA)
        ↓
     [优化] (子集构造法)
        ↓
    确定有限自动机 (DFA)
        ↓
    [执行] (状态机匹配)
        ↓
      匹配结果

性能分析

编译时复杂度

  1. 解析:O(m),其中m是模式长度
  2. NFA构建:O(m),每个AST节点处理一次
  3. DFA构建:O(2^n),最坏情况下指数级状态数

匹配时复杂度

  1. NFA匹配:O(mn),需要跟踪多个状态
  2. DFA匹配:O(n),确定性状态转换

空间复杂度

  1. AST:O(m),与模式长度成正比
  2. NFA:O(m),Thompson构造法保证线性状态数
  3. DFA:O(2^m),最坏情况下指数级

当前实现的边界

Compiler 实现的 Match 是全字符串匹配,不是 Matcher 那样的子串搜索。它暂不支持字符类、Unicode 属性、重复次数和锚点,这些能力将在高级篇继续扩展。

. 在 NFA 中用 DOT_CHAR 表示。DFA 构造时,具体字符转移会合并同一状态集合中的通配边;匹配时,如果不存在精确转移,则回退到 DOT_CHAR。换行符 \n\r 不执行这一回退,因此不会被 . 匹配。

配套实现:Ismantic/Regex

正则表达式引擎:高级篇

引言

本篇在正则表达式引擎:基础篇的实现上继续扩展,目标是支持 Tokenizer 预分词所需的正则语法。整体仍采用递归下降解析、AST、Thompson NFA、DFA 和匹配器组成的流水线,对应实现为 src/regex-2.cc

高级篇最关键的变化是引入字符集合。基础篇的一条 NFA 边只匹配一个确定的 codepoint,例如 'a';字符类 [abc]、Unicode 属性 \p{H} 和补集 [^...] 则要求一条边匹配一组甚至大量字符。直接枚举 Unicode codepoint 不现实,因此本篇用谓词函数表示集合:

using CharPred = std::function<bool(uint32_t)>;

输入字符使谓词返回 true,就可以沿这条 NFA 边转移。NFA 到这里已经能够直接运行,但 DFA 还需要高效地缓存确定转移。若两个 codepoint 对所有谓词都产生相同的真假结果,它们在这个正则中的行为就完全相同,可以归入同一个等价类。DFA 因而不必为每个 Unicode codepoint 分别保存转移,只需按等价类保存,并在实际遇到输入时按需构建状态。

整个设计因果链可以概括为:

字符集合
    ↓ 用 CharPred 表示集合
谓词化的 NFA 边
    ↓ 按所有谓词的真假结果分类
字符等价类
    ↓ 按等价类缓存并按需构建转移
Lazy DFA

其中,CharPred 是字符集合的程序表示;等价类则是 DFA 的构建与缓存优化,并不是 NFA 使用谓词边的必要条件。

  1. 字符类 [abc][^abc] — 匹配一组字符或其补集
  2. Unicode 属性 \p{A}\p{H}\p{N} — 匹配字母、汉字和数字
  3. 转义序列 \r\n\s — 匹配换行和空白
  4. 文本分段 Segment — 从左到右产生不重叠的最长匹配片段
  5. Lazy DFA — 按需构建 DFA 状态 + 等价类优化

这些特性服务于一个明确的 Tokenizer 预分词 Pattern,而不是追求兼容通用正则语法:

[^\r\n\p{A}\p{H}\p{N}]?\p{A}+
|\p{H}+
|\p{N}+
| ?[^\s\p{A}\p{H}\p{N}]+[\r\n]*
|\s*[\r\n]
|\s

架构概览

整体流程和 Regex 完全一致,是经典的编译流水线:

正则表达式字符串
    ↓ RegexParser(递归下降)
   AST(抽象语法树)
    ↓ NFABuilder(Thompson 构造,Visitor 模式)
   NFA(非确定有限自动机)
    ↓ LazyDFA(按需子集构造 + 等价类)
   DFA(确定有限自动机)
    ↓ Match / Segment
   匹配结果

与 Regex 的对应关系

组件RegexRegex 2变化
AST 节点8 种9 种+CharClassAst
Parser基础语法+[...] \p{A} \p{H} \p{N} \s面向目标 Pattern 扩展
NFA 边uint32_t 字符CharPred 谓词函数核心变化
DFA完整预构建Lazy 按需构建新方案
匹配Match(全文)Match + Segment+文本分段

Parser

自定义属性

为 Tokenizer 场景定义了三个近似属性。它们借用了 \\p{...} 的语法,但不是 Unicode Character Database 中标准属性的完整实现:

属性语法含义包含的字符
Alpha\p{A}字母(非汉字非数字)拉丁、西里尔、阿拉伯、日韩假名、天城文等
Han\p{H}汉字/CJKCJK 统一汉字及扩展区 A-H
Digit\p{N}数字子集ASCII 0-9、全角 0-9

这种手写范围适合受控的预分词场景,但会遗漏部分文字和数字,也可能覆盖 Unicode 中尚未分配的码位。需要严格的 Unicode 分类时,应使用 Unicode 数据表生成范围,或接入 ICU 等 Unicode 实现。

谓词函数

每个属性对应一个 C++ 函数 bool(uint32_t),在运行时对 Unicode codepoint 求值:

static bool IsAlpha(uint32_t c) { return IsWordChar(c) && !IsDigit(c) && !IsHan(c); }
static bool IsHan(uint32_t c) {
    return (c>=0x3400&&c<=0x4DBF) || (c>=0x4E00&&c<=0x9FFF) ||
           (c>=0xF900&&c<=0xFAFF) || (c>=0x20000&&c<=0x323AF);
}

这些函数直接作为 NFA 边的匹配条件,不需要枚举所有 Unicode 码位。

CharPred — 谓词化的字符匹配

这是高级实现与基础实现最核心的架构差异。

Regex 的方案

Regex 中 NFA 的边用 uint32_t 表示要匹配的字符:

// Regex: NFA 边 = 一个具体的 Unicode codepoint
std::map<uint32_t, std::vector<NFAState*>> transitions;

这种方案对字面量字符很自然,但无法高效表示 \p{H}(覆盖大量 codepoint)或 [^abc](取反集合)。

高级实现的方案

std::function<bool(uint32_t)> 作为边的匹配条件:

using CharPred = std::function<bool(uint32_t)>;

struct Edge {
    CharPred pred;    // 匹配条件:任意 bool(uint32_t) 函数
    NFAState* to;     // 目标状态
};

这样任何匹配规则都可以统一表示:

// 字面量 'a'
[](uint32_t c) { return c == 'a'; }

// Unicode 属性 \p{H}
IsHan  // 直接传函数指针

// 取反 [^abc]
[pred](uint32_t c) { return !pred(c); }

// 组合 [\r\n]
[a, b](uint32_t c) { return a(c) || b(c); }

字符类的解析

[...] 内部的每个元素被解析为一个 CharPred,然后用 || 组合:

CharPred ParseCharClass() {
    bool negated = Match('^');
    CharPred pred = [](uint32_t) { return false; };  // 空集

    while (!AtEnd() && pattern_[pos_] != ']') {
        CharPred cp;
        if (pattern_[pos_] == '\\') {
            // 转义:\s, \p{A}, 等
            ++pos_;  // 跳过反斜杠
            auto [p, r] = ParseEscape();
            cp = std::move(p);
        } else {
            uint32_t c = NextChar();
            cp = [c](uint32_t x) { return x == c; };
        }
        // 用 || 组合到 pred 中
        pred = [a=std::move(pred), b=std::move(cp)](uint32_t x) {
            return a(x) || b(x);
        };
    }

    if (negated) pred = [p=std::move(pred)](uint32_t x) { return !p(x); };
    return pred;
}

[^\r\n\p{A}\p{H}\p{N}] 为例,解析过程:

  1. ^ → 标记取反
  2. \rpred = (c == '\r')
  3. \npred = (c == '\r') || (c == '\n')
  4. \p{A}pred = ... || IsAlpha(c)
  5. \p{H}pred = ... || IsHan(c)
  6. \p{N}pred = ... || IsDigit(c)
  7. 取反 → pred = !(上述)

最终的 pred 函数含义:“既不是 \r\n,也不是字母、汉字、数字”——即标点和符号。

新增语法

在 Regex 的 BNF 基础上扩展:

Pattern     = Sequence ('|' Sequence)*
Sequence    = Quantified*
Quantified  = Atom Quantifier?
Quantifier  = '*' | '+' | '?'
Atom        = '(' Pattern ')'          -- 分组
            | '[' '^'? ClassItem* ']'  -- 字符类        [NEW]
            | '.'                       -- 任意字符
            | '\\' Escape              -- 转义          [NEW]
            | Literal                   -- 字面量
Escape      = 'r' | 'n' | 's'
            | 'p' '{' ('A'|'H'|'N') '}' -- Unicode 属性 [NEW]
            | <any char>               -- 转义字面量
ClassItem   = Escape | Char

AST

新增节点:CharClassAst

表示字符类 [...]、转义序列 \s、Unicode 属性 \p{A} 等:

class CharClassAst : public Ast {
    CharPred pred_;     // 匹配谓词
    std::string repr_;  // 用于打印的文本表示
};

和 LiteralAst 的区别:

  • LiteralAst 匹配一个确定的 codepoint
  • CharClassAst 匹配满足谓词的任意 codepoint

在 NFA 构建时,两者生成相同结构的片段(一条边连接起始和终止状态),区别只在边的 CharPred 内容。

NFA

CharPred 边

与 Regex 最大的区别。Regex 中 NFA 的 Visit(LiteralAst*) 创建一条 uint32_t 边:

// Regex
start->InsertTransition(codepoint, end);

高级实现统一使用 CharPred:

// regex-2
void PushPred(CharPred pred) {
    auto *s = NewState(), *e = NewState();
    e->accept = true;
    s->edges.push_back({std::move(pred), e});
    stack_.push({s, e});
}

void Visit(const LiteralAst* n) override {
    uint32_t p = n->GetPoint();
    PushPred([p](uint32_t c) { return c == p; });
}

void Visit(const CharClassAst* n) override {
    PushPred(n->GetPred());  // 直接使用解析好的谓词
}

对于字符类 [^\r\n\p{A}],NFA 的边就是一个组合谓词函数。NFA 结构不需要任何改动——只是边的“标签“从具体字符变成了谓词。

DFA

这是和 Regex 在 DFA 层面的主要区别。

问题:Unicode 状态爆炸

Regex 的 DFA 用完整的子集构造:预先计算所有可达的 DFA 状态。当 NFA 边是具体字符时,字母表有限,没有问题。

但高级实现的 NFA 边是谓词,一个 \p{A} 会覆盖大量码点。如果为每个码点维护 DFA 转移表,空间会迅速膨胀。

方案:等价类 + 按需构建

等价类:两个 codepoint 如果对所有 NFA 边的谓词给出相同的 true/false 结果,它们在 DFA 中的行为完全一致,归为同一个等价类。

例如,所有 ASCII 小写字母 a-z 都满足 IsAlpha=true, IsDigit=false, IsHan=false, IsWhitespace=false,它们属于同一个等价类。

等价类的本质是自动机无法区分的一组输入字符。从数学上看,它是一个 codepoint 集合:

C = {cp | Signature(cp) = [true, false, false, false]}

其中 Signature(cp) 表示依次用 NFA 中所有 CharPred 判断 cp 得到的真假序列。这个集合由当前正则中的谓词决定,并不是 Unicode 固有的字符类别;换一个 Pattern,使用的谓词发生变化,字符的等价类也可能随之变化。

实现中不会真的构造 std::set<uint32_t> 来枚举集合成员。Unicode 属性和补集可能包含大量 codepoint,完整保存既浪费空间,也没有必要。代码改用共同的谓词签名隐式表示这个集合:

数学上的集合:
{'a', 'b', 'c', ...}

代码中的表示:
[true, false, false, false] → class_id

凡是得到相同签名的 codepoint,都取得同一个 class_id,并在 DFA 转移表中共用同一条转移。sig_map_ 保存“谓词签名 → 等价类编号”,而 cp_cache_ 只缓存运行时已经遇到的“codepoint → 等价类编号”;两者都没有保存等价类的完整成员集合。

int ClassifyCP(uint32_t cp) {
    // 1. 查缓存
    auto it = cp_cache_.find(cp);
    if (it != cp_cache_.end()) return it->second;

    // 2. 计算签名:对每个 NFA 边的谓词求值
    std::vector<bool> sig;
    for (auto* p : all_preds_) sig.push_back((*p)(cp));

    // 3. 签名相同 → 同一个等价类 ID
    auto sit = sig_map_.find(sig);
    int cls;
    if (sit != sig_map_.end()) cls = sit->second;
    else { cls = next_cls_++; sig_map_[sig] = cls; }

    cp_cache_[cp] = cls;
    return cls;
}

在当前测试 Pattern 和已观察字符上,通常只会产生十余个等价类,因此 DFA 转移表比较紧凑。这不是一般上界:若有 \(P\) 个相互独立的谓词,理论上最多可能出现 \(2^P\) 种真假签名。

按需构建:DFA 状态不预先全部构建,而是在首次遇到某个 (DFA状态, 等价类) 组合时才计算:

int Step(int dfa_st, uint32_t cp) {
    int cls = ClassifyCP(cp);

    // 查 DFA 转移表缓存
    auto it = states_[dfa_st].trans.find(cls);
    if (it != states_[dfa_st].trans.end()) return it->second;

    // 首次遇到:执行 NFA 子集构造
    // 1. 找到当前 DFA 状态对应的 NFA 状态集
    // 2. 对集合中每个 NFA 状态的每条边,用 pred(cp) 测试
    // 3. 收集所有可达的 NFA 状态 → ε 闭包 → 新 DFA 状态
    // 4. 缓存结果
    ...
}

首次匹配时按需构建,后续匹配直接查表——兼顾了构建效率和匹配性能。

与 Regex DFA 的对比

Regex DFARegex 2 Lazy DFA
构建时机编译时一次性全部构建运行时按需构建
转移表 keyUnicode codepoint等价类 ID
空间可能很大取决于实际生成的状态和谓词签名
首次匹配快(已构建)略慢(需构建)
后续匹配同样快(已缓存)

Segment — 文本分段

基础实现只有 Match(全文匹配),高级实现增加了 Segment(从左到右产生不重叠的最长匹配片段),这是 Tokenizer 的核心需求。

算法:左到右、最长匹配

这里明确采用的是左端起点固定后的最长匹配策略。它不等同于常见回溯引擎的“左端优先、分支从左到右优先”:将所有分支合并成普通 DFA 后,只保留 accept 标记会丢失分支优先级。若目标是逐字节复现某个现有 Regex 引擎,还需要在接受状态中记录分支优先级,并定义长度与优先级的比较规则。

std::vector<std::string_view> Segment(std::string_view text) {
    std::vector<std::string_view> result;
    size_t pos = 0;

    while (pos < text.size()) {
        std::ptrdiff_t n = MatchAt(text.data() + pos, text.size() - pos);
        if (n > 0) {
            result.emplace_back(text.data() + pos, n);
            pos += n;      // 跳过已匹配的部分
        } else {
            // Tokenizer 不能静默丢弃未匹配输入
            size_t bytes = UTF8Len(text.data() + pos);
            result.emplace_back(text.data() + pos, bytes);
            pos += bytes;
        }
    }

    return result;
}

如果表达式能够匹配空串,MatchAt() 可能返回 0。上面的实现不输出空匹配,而是把当前 UTF-8 字符作为回退片段输出,以避免死循环和数据丢失;通用正则 API 则需要单独规定空匹配语义。

MatchAt 驱动 DFA 从位置 0 开始匹配,记录最后一个 accept 状态的位置(最长匹配):

std::ptrdiff_t MatchAt(const char* data, size_t len) {
    int cur = 0;                              // DFA 起始状态
    std::ptrdiff_t last = states_[0].accept ? 0 : -1;
    size_t pos = 0;

    while (pos < len) {
        uint32_t cp = DecodeUTF8(...);
        int next = Step(cur, cp);             // 查 DFA 转移
        if (next < 0) break;                  // 无转移,停止
        cur = next;
        pos += bytes;
        if (states_[cur].accept) last = pos;  // 记录最长匹配
    }

    return last;  // 返回最长匹配的字节数,-1 表示无匹配
}

示例

输入 "Hello, World!" 对本文的 Pattern 执行 Segment:

pos=0: MatchAt("Hello, World!") → 5 ("Hello")  → 输出 "Hello"
pos=5: MatchAt(", World!")      → 1 (",")       → 输出 ","
pos=6: MatchAt(" World!")       → 6 (" World")  → 输出 " World"
pos=12: MatchAt("!")            → 1 ("!")        → 输出 "!"

结果:['Hello', ',', ' World', '!']

Tokenizer Pattern

逐段拆解本实现使用的 pattern:

[^\r\n\p{A}\p{H}\p{N}]?\p{A}+
|\p{H}+
|\p{N}+
| ?[^\s\p{A}\p{H}\p{N}]+[\r\n]*
|\s*[\r\n]
|\s

分支 1:字母 run(可带一个前缀)

[^\r\n\p{A}\p{H}\p{N}]?\p{A}+

核心分支,处理所有字母文本。

  • [^\r\n\p{A}\p{H}\p{N}]? — 可选的一个前缀字符。排除了换行、字母、汉字、数字,剩下空格、标点、符号。最多一个。
  • \p{A}+ — 一个或多个字母字符

三种场景:

  • 无前缀HelloHello
  • 空格前缀 World World(空格粘到字母)
  • 标点前缀't't,world,world$hello$hello

常见英文缩写也能得到类似切分:don'tdon + 't。不过,这只是当前最长匹配策略下的结果;它不代表简化分支在所有输入上都与带大小写和分支优先级的原规则等价。

分支 2:汉字 run

\p{H}+

连续汉字,无前缀。空格遇到汉字时分支 1 不匹配(\p{A}+ 不含汉字),空格会被分支 6 独立输出。

分支 3:数字 run

\p{N}+

连续数字,无前缀,无长度限制。空格同样不会粘上来。

分支 4:标点 run(可带空格前缀)

 ?[^\s\p{A}\p{H}\p{N}]+[\r\n]*
  • ? — 可选的一个空格前缀(字面空格,不是 \s
  • [^\s\p{A}\p{H}\p{N}]+ — 一个或多个标点/符号
  • [\r\n]* — 可选的尾部换行

什么时候轮到分支 4?当标点后面不是字母时。如果后面跟字母,分支 1 匹配更长(,world → 分支 1 赢),分支 4 只在 , --- $(后跟数字或结尾)这类场景触发。

分支 5:换行

\s*[\r\n]

换行符(可带前导空白)。

分支 6:单个空白

\s

匹配单个空白字符。是最低优先级的兜底。

连续空格的处理:每个多余空格独立输出为一个 token,最后一个空格被下一个字母或标点分支的前缀吸收。连续标点则由分支 4 合并,例如 "___hello""___""hello"

运行与性能

下面给出简化 Pattern 的预期输出:


'Hello, World!'         -> ['Hello', ',', ' World', '!']
"don't"                 -> ["don", "'t"]
'$100'                  -> ['$', '100']
'24h'                   -> ['24', 'h']
'hello123world'         -> ['hello', '123', 'world']
' hello'                -> [' hello']
'  hello'               -> [' ', ' hello']
'   hello'              -> [' ', ' ', ' hello']
'hello  world'          -> ['hello', ' ', ' world']
'你好,世界!'            -> ['你好', ',', '世界', '!']
' 你好世界'              -> [' ', '你好世界']
'Hello, 你好! 123abc'   -> ['Hello', ',', ' ', '你好', '!', ' ', '123', 'abc']

性能分析

编译时复杂度

  • Parser:O(|pattern|)
  • NFA 构建:O(|pattern|)
  • DFA:按需构建,不在编译时消耗

匹配时复杂度

  • 等价类查询:首次 O(|preds|),后续 O(1)(缓存)
  • DFA 转移:首次 O(|NFA states|)(子集构造),后续 O(1)(缓存)
  • 对覆盖全部输入的 Tokenizer Pattern,Segment 通常按匹配长度向前推进,热缓存下接近 O(|text|)
  • 对一般 Pattern,某个起点可能扫描很远后失败,再从下一个字符重试,最坏可达到 O(|text|²)
  • Lazy DFA 的首次运行还要承担状态和转移的构建成本,不能与已经完整预构建的 DFA 简单视为相同常数开销
  • regex-2.ccMatchAt 在解码每个字符时会构造剩余文本的临时 std::string,因此实测结果还包含额外分配与复制成本。面向吞吐量优化时,应改为直接按指针和剩余长度解码

空间复杂度

  • NFA:O(|pattern|) 个状态
  • DFA:按需分配,最坏 O(2^|NFA|),实际远小于此
  • 等价类:最多受谓词真假签名数限制,理论上可达 O(2^|preds|),实际通常少得多

至此,正则引擎已经能将文本切分成适合后续统计的字符串片段。这些片段还不是最终的 token ID;后续 Tokenizer 章节将继续讲解归一化、预分词、频率统计、词表训练和编码过程。

配套实现:Ismantic/Regex

Double-Array Trie

概述

Double-Array Trie 是一种高效的 Trie 压缩表示。它通过数组布局和状态转移减少指针开销,兼顾紧凑存储与快速字符串检索,常用于词典匹配、形态分析和分词。

以存储单词 {“he”, “she”, “his”} 为例,朴素指针字典树结构如下:

    root(0)
    /     \
   h(1)    s(7)
  / | \    |
 e(2) i(3) h(8)
 |   |     |
$(4) s(5)  e(9)
     |     |
    $(6)  $(10)

节点编号及其含义:
0:root  1:h  2:he  3:hi  4:"he"  5:his  6:"his"  7:s  8:sh  9:she  10:"she"

该表达方式的问题

  • 节点需要存储多个子节点指针,内存开销大
  • 内存分散,缓存不友好
  • 指针访问增加间接寻址开销

Double-Array Trie 把树平铺到一维数组中:

  • state:保存计算子节点位置所需的状态值
  • character:存储字符信息
  • up:验证父子关系的正确性

数组表示:

struct Node {
    int state;
    int up;
    char character;
};
std::vector<Node> array;

状态转移:

// -> ct
state = array[pos].state
new_pos = pos ^ state ^ ct
// <- ct
array[new_pos].up = pos
array[new_pos].character = ct

状态转移把父节点位置、状态值和输入字节组合起来,由此可以把整棵 Trie 平铺到数组中。这里用到 XOR 的以下性质:

性质1:交换律

a ^ b = b ^ a

性质2:结合律

(a ^ b) ^ c = a ^ (b ^ c)

性质3:自消性

a ^ a = 0
a ^ 0 = a

性质4:可逆性

如果 c = a ^ b,那么 a = c ^ b,b = c ^ a

下面把 {he, his, she} 平铺到数组中。状态转移公式为:

uint32_t next_pos = pos ^ state ^ ct;

举例说明: 假设当前位置 pos = 5,状态值 state = array_[pos].state = 12,字符 ct = 'a'(97)

next_pos = 5 ^ 12 ^ 97
         = 104

验证机制

// 检查目标位置的字符是否匹配
if (array_[next_pos].character != ct) return false;

// 检查父子关系是否正确  
if (array_[next_pos].up != pos) return false;

state 由基址分配过程确定。GetFreeIndex 为当前节点的一组子节点寻找可用基址 index,再令 state = pos ^ index。查询时有:

pos ^ state ^ ct
= pos ^ pos ^ index ^ ct
= index ^ ct

完整示例: 对前面的 {he, his, she} 的例子,建好的数组如下表显示(eow表示是否词结束):

poscharacterupstateeow含义说明
0root-5false根状态分配index=5给子节点{h,s}
109h0103false’h’状态分配index=10给子节点{e,i}
118s098false’s’状态分配index=20给子节点{h}
111e109200true’he’状态叶子节点,state指向值节点200
99i109125false’hi’状态分配index=30给子节点{s}
124h11884false’sh’状态分配index=40给子节点{e}
93s99300true’his’状态叶子节点,state指向值节点300
77e124400true’she’状态叶子节点,state指向值节点400
200\000true“he“值节点存储值0
300\001true“his“值节点存储值1
400\002true“she“值节点存储值2

详细构建过程

  1. 根节点分配 (pos=0)
子节点: {h(104), s(115)}
GetFreeIndex找到: index = 5
子节点位置计算:
  h位置: 5 ^ 104 = 109
  s位置: 5 ^ 115 = 118
根节点state: 0 ^ 5 = 5
  1. h节点分配 (pos=109)
子节点: {e(101), i(105)}
GetFreeIndex找到: index = 10
子节点位置计算:
  e位置: 10 ^ 101 = 111
  i位置: 10 ^ 105 = 99
h节点state: 109 ^ 10 = 103
  1. s节点分配 (pos=118)
子节点: {h(104)}
GetFreeIndex找到: index = 20
子节点位置计算:
  h位置: 20 ^ 104 = 124
s节点state: 118 ^ 20 = 98
  1. hi节点分配 (pos=99)
子节点: {s(115)}
GetFreeIndex找到: index = 30
子节点位置计算:
  s位置: 30 ^ 115 = 93
hi节点state: 99 ^ 30 = 125
  1. sh节点分配 (pos=124)
子节点: {e(101)}
GetFreeIndex找到: index = 40
子节点位置计算:
  e位置: 40 ^ 101 = 77
sh节点state: 124 ^ 40 = 84

因此,构建时使用的 new_pos = index ^ ct 与查询时使用的 new_pos = pos ^ state ^ ct 完全等价。构建过程只需保存 state,查询过程就能恢复同一个基址。

至此,Double-Array Trie 的构建过程已经清楚:为一组子节点寻找可用的基址,将它们写入数组,再递归处理子节点,直到遍历完整棵 Trie。

实现

Double-Array Trie 的核心实现思路:

  1. 通过 XOR 实现状态转移(pos ^ state ^ ct
  2. 通过 up 验证父子关系
  3. 值节点复用 state 字段存储 value

构建过程

朴素 Trie → 递归平铺 → 双数组结构
  1. 初始化根节点(pos=0)
  2. 对每个节点的子节点集合:
    • 调用 GetFreeIndex 寻找合适基址
    • 计算所有子节点位置(index ^ char)
    • 设置父子关系(up 字段)
  3. 递归处理每个子节点

数据结构

Node 是数组中的基本单元。中间节点需要 state 完成状态转移,值节点不再需要它,因此两者通过 union 复用空间。插入 "abc" 时,构建过程会追加终止字节 '\0':字符 c 标记原词结束,随后创建满足 character == 0 && eow 的值节点,并在其中保存 value

struct Node {
    uint8_t character;  // 当前节点代表的字符
    bool eow;          // End of Word 标记
    uint32_t up;       // 父节点位置
    union { 
        uint32_t state;    // 状态值(用于中间节点)
        int32_t value;     // 存储值(用于叶子节点)
    };
    
    bool Empty() const {
        return character == 0 && !eow && up == 0 && state == 0;
    }
    
    bool Value() const {
        return character == 0 && eow;
    }
};

以及,构建过程中用到的数据结构:

struct TrieNode {
    std::map<uint8_t, std::unique_ptr<TrieNode>> down_nodes;
    int32_t value;
    bool eow;
};

std::unique_ptr<TrieNode> root_;
std::vector<Node> units_;
std::vector<bool> uses_;

uint32_t prev_pos_ = 0;

构建过程

阶段一:构建朴素字典树

void TrieInsert(const std::string& str, int32_t value) {
    TrieNode* current = root_.get();
    for (char c : str) {
        uint8_t t = static_cast<uint8_t>(c);
        if (current->down_nodes.find(t) == current->down_nodes.end()) {
            current->down_nodes[t] = std::make_unique<TrieNode>();
        }
        current = current->down_nodes[t].get();
    }
    current->eow = true;
    current->value = value;
}

阶段二:转换为双数组

核心的状态分配算法:

uint32_t SetupDownNodes(const std::vector<uint8_t>& es, uint32_t pos, TrieNode* node) {
    uint32_t index = GetFreeIndex(es);
    units_[pos].state = pos ^ index;  // 关键:利用 XOR 存储状态
    
    for (size_t i = 0; i < es.size(); ++i) {
        uint8_t e = es[i];
        uint32_t p = index ^ e;  // 利用 XOR 计算子节点位置
        
        uses_[p] = true;
        
        if (e == '\0') {
            units_[p].value = node->value;
            units_[p].eow = true;
            units_[pos].eow = true;
        }
        
        units_[p].character = e;
        units_[p].up = pos;
    }
    
    return index;
}

阶段三:递归转换

void NodeConvert(TrieNode* node, uint32_t pos) {
    std::vector<uint8_t> es;
    std::vector<TrieNode*> down_nodes;
    
    // 收集所有子节点
    for (const auto& n : node->down_nodes) {
        es.push_back(n.first);
        down_nodes.push_back(n.second.get());
    }
    
    // 如果是单词结尾,添加值节点
    if (node->eow) {
        es.push_back('\0');
        down_nodes.push_back(nullptr);
    }
    
    if (!es.empty()) {
        uint32_t index = SetupDownNodes(es, pos, node);
        
        // 递归处理子节点
        for (size_t i = 0; i < es.size(); ++i) {
            uint8_t e = es[i];
            if (e != '\0') {
                size_t down_pos = index ^ e;
                NodeConvert(down_nodes[i], down_pos);
            }
        }
    }
}

'\0' 节点没有子节点,因此递归构建不会继续深入;它在 SetupDownNodes 中被直接写成值节点。

具体示例

以 {“cat”, “car”} 为例,展示完整构建过程:

步骤1:构建传统树

root
 |
 c
 |
 a
/ \
t   r
$   $

步骤2:转换根节点

es = {'c'}
尝试 index = 1:
  'c'(99) 的位置 = 1 ^ 99 = 98
  检查 uses_[98] = false ✓

root: state = 0 ^ 1 = 1
'c' 节点放在位置 98

步骤3:转换 ‘c’ 节点

es = {'a'}  
尝试 index = 2:
  'a'(97) 的位置 = 2 ^ 97 = 99
  检查 uses_[99] = false ✓

'c': state = 98 ^ 2 = 100
'a' 节点放在位置 99

步骤4:转换 ‘a’ 节点

es = {'t', 'r'}
尝试 index = 10:
  't'(116) 的位置 = 10 ^ 116 = 126
  'r'(114) 的位置 = 10 ^ 114 = 124
  检查 uses_[126] = false, uses_[124] = false ✓

'a': state = 99 ^ 10 = 105
't' 节点放在位置 126,'r' 节点放在位置 124

检索方法

Piece GetPiece(const std::string& str) const {
    if (Empty()) return Piece();

    uint32_t pos = 0;
    const char* s = str.c_str();
    size_t n = str.size();

    for (size_t i = 0; i < n; ++i) {
        Node unit = array_[pos];
        uint32_t state = unit.state;
        uint8_t ct = static_cast<uint8_t>(s[i]);

        // 核心:XOR 状态转移
        uint32_t next_pos = pos ^ state ^ ct;

        if (next_pos >= size_) return Piece();

        pos = next_pos;
        unit = array_[pos];

        if (unit.character != ct) {
            return Piece();
        }
    }

    // 验证是否为完整单词
    Node unit = array_[pos];
    if (!unit.eow) return Piece();

    // 获取存储的值
    uint32_t value_pos = pos ^ unit.state;
    if (value_pos >= size_) return Piece();

    Node value_unit = array_[value_pos];
    if (!value_unit.Value()) 
        return Piece();
    
    return Piece(value_unit.value, n, str);
}

查找 “cat” 的完整过程

初始: pos = 0

第1步 ('c'):
  unit = array_[0], state = 1
  next_pos = 0 ^ 1 ^ 99 = 98
  检查 array_[98].character == 'c' ✓
  pos = 98

第2步 ('a'):  
  unit = array_[98], state = 100
  next_pos = 98 ^ 100 ^ 97 = 99
  检查 array_[99].character == 'a' ✓
  pos = 99

第3步 ('t'):
  unit = array_[99], state = 105  
  next_pos = 99 ^ 105 ^ 116 = 126
  检查 array_[126].character == 't' ✓
  pos = 126

验证单词结尾:
  array_[126].eow == true ✓
  value_pos = 126 ^ array_[126].state
  返回 array_[value_pos].value

前缀匹配

这个函数找出输入字符串中所有作为字典中单词前缀的部分:

std::vector<Piece> GetUpPieces(const std::string& str) const {
    std::vector<Piece> rs;
    if (Empty()) return rs;

    uint32_t pos = 0;
    const char* s = str.c_str();
    size_t n = str.size();

    for (size_t i = 0; i < n; ++i) {
        Node unit = array_[pos];
        uint32_t state = unit.state;
        uint8_t ct = static_cast<uint8_t>(s[i]);

        uint32_t next_pos = pos ^ state ^ ct;
        if (next_pos >= size_) break;

        pos = next_pos;
        Node next_unit = array_[pos];

        if (next_unit.character != ct) break;

        // 检查当前位置是否为单词结尾
        if (next_unit.eow) {
            size_t value_pos = pos ^ next_unit.state;
            if (value_pos < size_ && array_[value_pos].Value()) {
                Node value_node = array_[value_pos];
                rs.emplace_back(value_node.value, i+1, str.substr(i+1));
            }
        }
    }

    return rs;
}

示例:字典包含 {“car”, “card”, “care”},输入 “card”

输入: "card"

i=0, 字符='c': 
  找到 'c' 节点,不是单词结尾,继续

i=1, 字符='a': 
  找到 'a' 节点,不是单词结尾,继续

i=2, 字符='r': 
  找到 'r' 节点,是单词结尾 ✓
  添加 Piece(value=0, num=3, str="d") 到结果 ("car" 匹配)

i=3, 字符='d': 
  找到 'd' 节点,是单词结尾 ✓  
  添加 Piece(value=1, num=4, str="") 到结果 ("card" 匹配)

返回: [Piece("car", 3, "d"), Piece("card", 4, "")]

后缀匹配

这是最复杂的查找模式,找出所有以给定字符串为前缀的单词:

std::vector<Piece> GetDownPieces(const std::string& str) const {
    std::vector<Piece> rs;
    if (Empty()) return rs;

    // 先定位到前缀对应的节点
    uint32_t pos = 0;
    for (size_t i = 0; i < str.size(); ++i) {
        Node unit = array_[pos];
        pos ^= unit.state ^ static_cast<uint8_t>(str[i]);

        if (pos >= size_) return rs;

        unit = array_[pos];
        if (unit.character != static_cast<uint8_t>(str[i]))
            return rs;
    }

    // 从这个位置开始深度优先搜索
    std::string cw = str;
    CollectDownPieces(pos, cw, rs);
    return rs;
}

深度优先搜索实现

void CollectDownPieces(uint32_t pos, std::string& cw, 
                       std::vector<Piece>& rs) const {
    if (pos >= size_) return;

    Node unit = array_[pos];

    // 如果当前位置是单词结尾,添加到结果
    if (unit.eow) {
        uint32_t value_pos = pos ^ unit.state;
        if (value_pos < size_ && array_[value_pos].Value()) {
            Node value_node = array_[value_pos];
            rs.emplace_back(value_node.value, cw.size(), cw);
        }
    }

    // 遍历所有可能的子节点(利用 XOR 的遍历特性)
    uint32_t state = unit.state;
    for (int ct = 0; ct <= 255; ++ct) {
        uint32_t new_pos = pos ^ state ^ ct;

        if (new_pos >= size_) continue;
        if (new_pos == pos) continue;

        Node new_unit = array_[new_pos];

        // 验证是否为有效子节点
        if (new_unit.character == ct && new_unit.character != '\0'
            && new_unit.up == pos) {
            cw.push_back(static_cast<char>(ct));
            CollectDownPieces(new_pos, cw, rs);
            cw.pop_back();
        }
    }
}

示例:输入前缀 “ca”,字典包含 {“cat”, “car”, “card”, “can”}

1. 定位到 "ca" 对应的节点位置 pos=99

2. 从这个位置开始 DFS:
   遍历字符 0-255:
   
   ct=116('t'): 
     new_pos = 99 ^ state ^ 116 = 126
     验证 array_[126].character == 't' ✓
     验证 array_[126].up == 99 ✓
     递归搜索,发现 "cat" 是单词结尾
     
   ct=114('r'):
     new_pos = 99 ^ state ^ 114 = 124  
     验证 array_[124].character == 'r' ✓
     验证 array_[124].up == 99 ✓
     递归搜索,发现 "car" 是单词结尾
     继续从 124 搜索,发现 "card"
     
   ct=110('n'):
     类似地发现 "can"

返回: [Piece("cat"), Piece("car"), Piece("card"), Piece("can")]

基址分配

GetFreeIndex 为一组子节点分配基址 index,确保所有 index ^ byte 位置均未被占用。它的搜索效率会直接影响 Double-Array Trie 的构建速度和空间利用率。

基本要求

给定互不相同的字节集合 {c1, c2, ..., cn},需要找到一个 index,使得:

  • 所有位置 index ^ c1, index ^ c2, ..., index ^ cn 都未被占用
  • 尽量从上次分配位置附近找到结果,以改善空间局部性

对固定的 index,XOR 是一一映射,因此不同字节天然会得到不同位置,不需要额外检查相互冲突。

数学表达

对于字符集合 es = {e1, e2, ..., en}
找到一个 index,满足:
∀i ∈ [1,n] => uses[index ^ ei] == false

具体实现

涉及到的变量:

std::vector<bool> uses_;
uint32_t prev_pos_;
uint32_t GetFreeIndex(const std::vector<uint8_t>& es) {
    if (es.empty()) return 0;
    
    // 局部性优化:从上次分配位置附近开始搜索
    uint32_t start_index = (prev_pos_ > 256) ? prev_pos_ - 256 : 1;
    
    for (uint32_t i = start_index; ; ++i) {
        bool valid = true;
        
        // 检查所有子节点位置是否可用
        for (uint8_t e : es) {
            size_t p = i ^ e;
            EnsureSize(p + 1);
            
            if (uses_[p]) {
                valid = false;
                break;
            }
        }
        
        if (valid) {
            prev_pos_ = i;  // 记录这次分配位置,用于下次优化
            return i;
        }
    }
}

该函数采用线性探测,从上次分配位置附近开始,逐个检查整组子节点需要的位置。一个字节的取值范围是 0~255,因此 index ^ byte 始终位于与 index 相邻的 256 个位置范围内。这个简单策略通常能够获得可接受的局部性;更复杂的空闲区间索引可以进一步缩短构建时间。

输入es = {'a'(97), 'e'(101), 'i'(105), 'o'(111), 'u'(117)}

尝试 index=10:
  'a' 位置: 10 ^ 97 = 107
  'e' 位置: 10 ^ 101 = 111  
  'i' 位置: 10 ^ 105 = 99
  'o' 位置: 10 ^ 111 = 101
  'u' 位置: 10 ^ 117 = 127
  
检查所有位置 {107, 111, 99, 101, 127}:
  - 互不相同 ✓
  - 都未被占用 ✓
  → 返回 index=10

Double-Array Trie 以较高的构建成本换取紧凑、稳定的查询结构,适合词表构建完成后频繁执行前缀匹配。后面的中文分词会直接利用这种能力;下一篇则先介绍更适合动态更新的 Critbit Trie。

配套实现:Ismantic/Trie

Critbit Trie

Binary Trie

Binary Trie 将字符串展开为二进制位:0 走左分支,1 走右分支。一条从根节点到叶子节点的路径因此可以表示一个字符串。以 "a""b""c""d" 为例:

字符二进制表示:
'a' = 01100001 (97)
'b' = 01100010 (98)  
'c' = 01100011 (99)
'd' = 01100100 (100)

按照每层固定检查一个 bit(从 bit7 到 bit0)的规则,构建的完整二进制树如下:

                            root
                           /
                      bit7=0
                         |
                       node1
                         \
                      bit6=1
                         |
                       node2
                         \
                      bit5=1
                         |
                       node3
                        /
                   bit4=0
                      |
                    node4
                     /
                bit3=0
                   |
                 node5
                 /   \
            bit2=0   bit2=1
               |       |
             node6    'd'
             /   \
        bit1=0  bit1=1
           |       |
          'a'    node7
                 /   \
            bit0=0  bit0=1
               |       |
              'b'     'c'

这种结构保留了 Trie 的前缀检索能力,但每一位都需要一层节点。上例仅处理一个字节就需要 8 层,其中多数节点只有一个分支,空间利用率很低。

Critbit Trie

Critbit Trie 只保留真正产生分叉的二进制位,从而压缩连续的单分支节点。这里的 critbit(关键位),是两个字符串在公共二进制前缀之后遇到的第一个不同位。

Critbit 的定义

"a""b""c""d" 为例:

  • 二进制表示:01100001, 01100010, 01100011, 01100100
  • 相同前缀:01100(前5位相同)
  • 第一个不同位:第2位(从右数第2位)

这个 critbit 可以把四个字符分成两组:

  • bit2=0: {a, b, c}
  • bit2=1: {d}

继续递归分析:

  • {a, b, c} 的 critbit 是第 1 位,分成 {a} 和 {b, c}
  • {b, c} 的 critbit 是第 0 位,分成 {b} 和 {c}

最终得到压缩后的 Critbit Trie:

         (pos=0,bit=2)
         /           \
    bit2=0         bit2=1
       |             |
  (pos=0,bit=1)     'd'
   /         \
bit1=0     bit1=1
  |           |
 'a'    (pos=0,bit=0)
         /         \
    bit0=0       bit0=1
       |           |
      'b'         'c'

树中只剩下三个实际产生分叉的内部节点,不再为公共二进制前缀保存单分支节点。

变长字符串

字典序的理解

对于字符串比较,字典序规则是:

  1. 从左到右逐个字符比较
  2. 字符更小的排在前面
  3. 前缀相等时,短字符串排在前面(如 “a” < “an”)

逻辑补零

为了区分 "a""aa" 这类前缀关系,Critbit Trie 将字符串结束后的字节在逻辑上视为 \0。实现并不创建补齐后的副本,而是在节点位置超出字符串长度时直接取 0

例如,对于字符串集合 {"a", "b", "c", "d", "aa", "ab", "abc"},比较到三个字节时,可以采用下表中的等价表示:

原字符串逻辑表示说明
“a”“a\0\0”用2个\0填充
“b”“b\0\0”用2个\0填充
“c”“c\0\0”用2个\0填充
“d”“d\0\0”用2个\0填充
“aa”“aa\0”用1个\0填充
“ab”“ab\0”用1个\0填充
“abc”“abc”无需填充

二进制表示分析

bit位置标注: 76543210 (从右到左)

字符串pos=0pos=1pos=2
“a\0\0”a(97)=01100001\0(0)=00000000\0(0)=00000000
“b\0\0”b(98)=01100010\0(0)=00000000\0(0)=00000000
“c\0\0”c(99)=01100011\0(0)=00000000\0(0)=00000000
“d\0\0”d(100)=01100100\0(0)=00000000\0(0)=00000000
“aa\0”a(97)=01100001a(97)=01100001\0(0)=00000000
“ab\0”a(97)=01100001b(98)=01100010\0(0)=00000000
“abc”a(97)=01100001b(98)=01100010c(99)=01100011

关键 bit 分析

第 1 个关键 bit:分析整个集合

集合:{a\0\0, b\0\0, c\0\0, d\0\0, aa\0, ab\0, abc}
分析位置:pos=0

各字符串在pos=0的字符:

  • a\0\0, aa\0, ab\0, abc: ‘a’(97) = 01100001
  • b\0\0: ‘b’(98) = 01100010
  • c\0\0: ‘c’(99) = 01100011
  • d\0\0: ‘d’(100) = 01100100

bit位分组分析

  • bit7~bit3: 全部相同,无分组作用
  • bit2: 0组={a,b,c,aa,ab,abc}, 1组={d} ← 第一个有效分组!

结果:第1个关键bit = (pos=0, bit=2)

第 2 个关键 bit:分析左子树

集合:{a\0\0, b\0\0, c\0\0, aa\0, ab\0, abc} (bit2=0的组)
分析位置:pos=0

bit位分组分析

  • bit1: 0组={a,aa,ab,abc}, 1组={b,c} ← 第一个有效分组!

结果:第2个关键bit = (pos=0, bit=1)

第 3 个关键 bit:分析 {a, aa, ab, abc} 子集

集合:{a\0\0, aa\0, ab\0, abc} (pos=0,bit1=0的组)
分析位置:pos=1 (因为pos=0都是’a’,无法区分)

各字符串在pos=1的字符:

  • a\0\0: ‘\0’(0) = 00000000
  • aa\0: ‘a’(97) = 01100001
  • ab\0: ‘b’(98) = 01100010
  • abc: ‘b’(98) = 01100010

bit位分组分析

  • bit6: 0组={a}, 1组={aa,ab,abc} ← 第一个有效分组!

结果:第3个关键bit = (pos=1, bit=6)

继续分析其他子集

按照相同原理,可以得到:

  • 第4个关键bit = (pos=1, bit=1) - 区分{aa\0} vs {ab\0, abc}
  • 第5个关键bit = (pos=2, bit=6) - 区分{ab\0} vs {abc}
  • 第6个关键bit = (pos=0, bit=0) - 区分{b\0\0} vs {c\0\0}

最终树结构

root = Node(pos=0, bit=2)
├─ data[0] = Node(pos=0, bit=1)  ← bit2=0: {a,b,c,aa,ab,abc}
│  ├─ data[0] = Node(pos=1, bit=6)  ← bit1=0: {a,aa,ab,abc}
│  │  ├─ data[0] = Value("a")  ← bit6=0: {a}
│  │  └─ data[1] = Node(pos=1, bit=1)  ← bit6=1: {aa,ab,abc}
│  │     ├─ data[0] = Value("aa")  ← bit1=0: {aa}
│  │     └─ data[1] = Node(pos=2, bit=6)  ← bit1=1: {ab,abc}
│  │        ├─ data[0] = Value("ab")  ← bit6=0: {ab}
│  │        └─ data[1] = Value("abc")  ← bit6=1: {abc}
│  └─ data[1] = Node(pos=0, bit=0)  ← bit1=1: {b,c}
│     ├─ data[0] = Value("b")  ← bit0=0: {b}
│     └─ data[1] = Value("c")  ← bit0=1: {c}
└─ data[1] = Value("d")  ← bit2=1: {d}

Node(pos=1, bit=6) 为例,"a"pos=1 处按 \0 处理,因此走 data[0]"aa""ab""abc" 在该位都是 1,因此走 data[1]。逻辑补零由此把字符串结束位置纳入比较,无须额外保存终止节点。

节点裂变

核心性质

内部节点按照 (pos, bit) 的顺序排列。插入时,新节点会被放到第一个更靠后的检查位置之前,因此同一组字符串形成的判定位次序不依赖插入顺序。

节点裂变机制

插入新字符串的过程可以理解为节点裂变

  1. 找到候选叶子:按现有节点的关键位走到一个叶子
  2. 计算关键分叉点:比较叶子字符串与新字符串,找到 critbit
  3. 节点裂变:将1个节点裂变成3个节点的子结构

实例:从 {a, b, c, d} 到 {a, b, c, d, e}

原始树结构 {a,b,c,d}:

root = (pos=0, bit=2)
├─ data[0] = (pos=0, bit=1)  ← {a,b,c}
│  ├─ data[0] = 'a'           ← bit1=0
│  └─ data[1] = (pos=0, bit=0) ← bit1=1: {b,c}
│     ├─ data[0] = 'b'        ← bit0=0
│     └─ data[1] = 'c'        ← bit0=1
└─ data[1] = 'd'              ← bit2=1 (目标裂变节点)

字符分析

  • ‘d’ = 01100100
  • ‘e’ = 01100101
  • critbit(‘d’,‘e’) = bit0 (最低位不同)

裂变过程

裂变前(1个节点):

Value('d')

裂变后(3个节点):

    Node(pos=0,bit=0)
    /              \
 Value('d')     Value('e')

最终树结构 {a,b,c,d,e}:

root = (pos=0, bit=2)
├─ data[0] = (pos=0, bit=1)  ← {a,b,c}
│  ├─ data[0] = 'a'           ← bit1=0
│  └─ data[1] = (pos=0, bit=0) ← bit1=1: {b,c}
│     ├─ data[0] = 'b'        ← bit0=0
│     └─ data[1] = 'c'        ← bit0=1
└─ data[1] = (pos=0, bit=0)  ← bit2=1: {d,e} (裂变后)
   ├─ data[0] = 'd'           ← bit0=0
   └─ data[1] = 'e'           ← bit0=1

裂变机制的意义

  1. 局部更新:插入只新增一个内部节点和一个叶子
  2. 有序判定:沿路径检查的 (pos, bit) 单调向后
  3. 前缀兼容:逻辑补零可以区分字符串及其更长后缀
  4. 无需重建:新字符串加入后,不必重新生成整棵树

总结:Critbit Trie 与 Double-Array Trie 面向不同场景。前者通过关键差异位组织变长字符串,插入和查询都不需要预先构建完整字符转移表;后者更适合静态词典和高吞吐前缀检索。理解两者的差异,有助于根据数据更新方式与查询模式选择结构。

配套实现:Ismantic/Trie

中文分词:基础篇

基本原理

中文句子没有天然的词语分隔符。对于“南京市长江大桥”,至少存在两条表面上合理的切分路径:

南京 | 市 | 长江 | 大桥
南京 | 市长 | 江大桥

最长匹配只能做局部选择,无法比较两条完整路径。DictCut 因此把中文分词定义为一个最大概率切分问题。

设句子 x 的一种合法切分为:

y = (w₁, w₂, ..., wₙ)

其中所有词依次拼接后必须等于原句。Unigram 模型假设各词相互独立,因此这条切分路径的联合概率为:

P(y) = P(w₁) × P(w₂) × ... × P(wₙ)

分词的目标是在所有合法切分中选择联合概率最高的一条:

y* = argmax P(y)

这里是在固定词典概率下选择最可能的隐藏切分,并不是重新估计模型参数。

词典采用简单的文本格式:

南京    751
市      1003
长江    602
大桥    399
市长    150
江大桥  25

第二列是词频,不是已经归一化的概率。Cutter::Build 计算完整词典的总频次 sum,查询一个已登录词时再得到对数概率:

double GetTrieValue(const std::string& word) {
    auto result = da_.GetUnit(word);
    if (!result.found || result.value == 0) {
        return std::log(1.0 / sum_);
    }
    return std::log(
        static_cast<double>(result.value) / sum_);
}

为了避免连续相乘造成数值下溢,DictCut 在对数域中计算。对数函数保持大小关系,因此最大化概率乘积等价于最大化对数概率之和:

y* = argmax Σ log P(wᵢ)

假设完整词典归一化后得到以下权重:

南京   -2.59
市     -2.30
长江   -2.81
大桥   -3.22
市长   -4.20
江大桥 -5.99

两条主要路径的得分分别是:

南京 | 市 | 长江 | 大桥
-2.59 -2.30 -2.81 -3.22 = -10.92

南京 | 市长 | 江大桥
-2.59 -4.20 -5.99 = -12.78

因为 -10.92 > -12.78,第一条路径胜出。这个模型考虑的是整条切分路径,但每个词的概率仍与上下文无关,因此它不能替代真正的上下文语言模型。

DictCut 在进入中文分词前还会执行一次预切分:

他是英国人Tom,编号123
→ 他是英国人 | Tom | , | 编号 | 123

只有连续汉字段进入 Trie 和动态规划。字母、数字、空格与标点片段直接保留。空格会先被折叠并替换为可见符号 ;这里的 Normalize 只处理空格,不包含 NFKC 等完整 Unicode 规范化。

动态规划

DictCut 将一个连续汉字段表示成有向无环图。节点是 UTF-8 字节位置,边是词典中从当前位置开始的候选词。

南京市

0 ──南──> 3 ──京──> 6 ──市──> 9
└────南京────> 6
└──────南京市──────> 9

代码中的 DAG 使用 G[i] 保存从字节位置 i 出发的边,其元素是候选词最后一个字节的位置:

std::vector<std::set<int>> Cutter::DAG(
    const std::string& sentence) {
    int n = sentence.length();
    std::vector<std::set<int>> G(n);

    for (int i = 0; i < n;) {
        int charlen = ustr::CharLen(
            static_cast<uint8_t>(sentence[i]));

        for (const auto& match :
             da_.PrefixSearch(sentence.substr(i))) {
            size_t end = i + match.length;
            if (match.length > 0 && end <= n) {
                G[i].insert(end - 1);
            }
        }

        G[i].insert(i + charlen - 1);
        i += charlen;
    }
    return G;
}

Trie 负责找出词典候选,最后加入的单字符边负责回退。这样图中始终存在一条能够走到句末的路径。

后向动态规划

当前实现从句末向句首计算:

route[i].first  = 从位置 i 到句末的最高得分
route[i].second = 最优路径第一条边的结束位置

对于 G[i] 中的每条边 i → x,转移为:

候选得分
= log P(sentence[i : x + 1])
 + route[x + 1].first

对应源码是:

std::vector<float_i> Cutter::Compute(
    const std::string& sentence,
    const std::vector<std::set<int>>& G) {
    int n = sentence.length();
    const double inf = std::numeric_limits<double>::infinity();
    std::vector<float_i> route(n + 1, {-inf, -1});

    route[n] = {1.0, n};

    for (int i = n - 1; i >= 0; --i) {
        float_i best = {-inf, -1};
        for (int end : G[i]) {
            double score = GetTrieValue(
                sentence.substr(i, end - i + 1));
            score += route[end + 1].first;

            if (score > best.first) {
                best = {score, end};
            }
        }
        route[i] = best;
    }
    return route;
}

对数域中空后缀通常初始化为 0.0。DictCut 当前使用 1.0,相当于为所有完整路径加上同一个常数,不会改变最优路径,也不会改变剪枝时两个路径的得分差。

只有 UTF-8 字符边界上的 graph[i] 包含候选边。处于多字节字符内部的位置虽然也存在于数组中,但不会被有效路径访问。

恢复切分

动态规划完成后,从句首沿 route 向后移动即可恢复结果,不需要再反转:

std::vector<std::string> Cutter::CutSegment(
    const std::string& sentence) {
    auto G = DAG(sentence);
    auto route = Compute(sentence, G);

    std::vector<std::string> result;
    for (int start = 0; start < sentence.length();) {
        int end = route[start].second + 1;
        result.push_back(
            sentence.substr(start, end - start));
        start = end;
    }
    return result;
}

对于“南京市长江大桥”,route[0] 最终指向“南京”,随后依次指向“市”“长江”和“大桥”,得到整条最高分路径。

Trie 应用

动态规划需要在每个字符位置找到全部词典前缀。如果逐个遍历词典,代价会随词表规模增长。DictCut 使用 Double-Array Trie,将公共前缀压缩在同一条搜索路径上。

构建 Trie 前,词语和频次需要保持一一对应,并按照词语排序:

void Cutter::Build(const std::vector<std::string>& words,
                   const std::vector<int>& freqs) {
    sum_ = 0;
    for (int freq : freqs) {
        sum_ += freq;
    }

    std::vector<std::pair<std::string, int>> pairs;
    for (size_t i = 0; i < words.size(); ++i) {
        pairs.emplace_back(words[i], freqs[i]);
    }
    std::sort(pairs.begin(), pairs.end());

    std::vector<std::string> sorted_words;
    std::vector<int> sorted_freqs;
    for (auto& [word, freq] : pairs) {
        sorted_words.push_back(std::move(word));
        sorted_freqs.push_back(freq);
    }

    da_.Build(sorted_words, sorted_freqs);
}

Trie 的终止节点直接保存整数频次。切分时,PrefixSearch 返回当前位置能够匹配的全部词典前缀及其字节长度:

auto matches = da_.PrefixSearch(sentence.substr(i));
for (const auto& match : matches) {
    G[i].insert(i + match.length - 1);
}

例如词典中同时存在“南京”“南京市”和“南”,一次前缀搜索会把三条候选边全部加入 DAG。Trie 只负责发现候选;选择最长词还是多个短词,仍由整条路径的概率决定。

搜索成本主要取决于当前位置能够沿 Trie 继续匹配的最长字节数以及返回的候选数量,而不是整个词典的大小。

回退机制

词典不可能覆盖所有汉字。为了保证 DAG 始终存在一条从句首到句尾的完整路径,DictCut 会在每个 UTF-8 字符位置加入一条单字符边:

int charlen = ustr::CharLen(
    static_cast<uint8_t>(sentence[i]));
graph[i].insert(i + charlen - 1);

DAG 使用字节位置,而不是字符编号。普通汉字通常占 3 个字节,扩展区汉字可能占 4 个字节,因此不能简单使用 i + 1

文本:    未  知  𠀀
字节位置:0   3   6    10
回退边:  0 → 3 → 6 → 10

单字在词典中时使用它的词频;不在词典中时,使用 log(1 / sum_) 作为未知词惩罚。因此已登录词通常优先于未知字,但生僻字仍能参与完整路径。

Cutter::Cut 还规定:没有构建词典时,不执行概率计算,而是直接把连续汉字段拆成单个 UTF-8 字符。非汉字段无论是否有词典都保持预切分结果。

词频训练

前面的最大概率切分假设词典已经包含可靠的词频。训练阶段解决的是另一个问题:从候选词表和无标注语料中估计这些词频。

候选词表 + 无标注语料
        ↓
估计每个词的使用频次
        ↓
word<TAB>frequency

DictCut 首先使用正向最长匹配完成冷启动。Segmenter 从当前位置查询 Trie 中的全部前缀并选择最长候选;如果没有候选,就回退到一个 UTF-8 字符:

while (i < sentence.size()) {
    auto matches = da_.PrefixSearch(sentence.substr(i));
    size_t best_len = 0;
    for (const auto& match : matches) {
        best_len = std::max(best_len, match.length);
    }

    if (best_len > 0) {
        result.push_back(sentence.substr(i, best_len));
        i += best_len;
    } else {
        int len = ustr::CharLen(
            static_cast<uint8_t>(sentence[i]));
        result.push_back(sentence.substr(i, len));
        i += len;
    }
}

最长匹配只负责产生第一版切分。--count 对切分结果计数,并用候选词表作为白名单,得到第一份 word<TAB>frequency 词典。

接下来重复两个步骤:

当前词频
   ↓ 构建 Trie 与 DAG
Viterbi 最优切分
   ↓ 统计最优路径上的词
新的词频

这属于 Viterbi EM,也称 Hard EM:E 步只选择当前概率最高的一条切分路径,M 步统计这条路径上的词频。它没有计算所有可能路径的期望计数,因此比 Soft EM 更简单。

run_em.sh 会把语料切成多个分片,并行执行切分和计数,再合并各分片的词频。默认每轮剪枝前执行两次 Hard EM,词表达到目标大小后再执行两次最终重估;次数可以通过 SUB_ITERS 调整。

词表剪枝

初始候选词表可能包含大量低频词、错误组合和可以被其他词替代的冗余词。DictCut 在若干轮 Hard EM 后计算每个词对最优路径的贡献,再缩小词表。

对于当前最优路径中的一个多字词,CutWithLoss 暂时删除它对应的 DAG 边,然后重新运行动态规划:

loss(word)
= 原最优路径得分 - 删除该词后的最优路径得分
double best_score = route[0].first;

graph[word.start].erase(word.end);
double alternative = Compute(sentence, graph)[0].first;
loss[word.text] += best_score - alternative;
count[word.text]++;
graph[word.start].insert(word.end);

loss 越大,说明删除这个词造成的概率损失越大,它越难被其他切分替代。

单字边不能删除,否则可能破坏完整路径。对于单字,代码直接比较它作为已登录词和未知字时的得分:

double known = GetTrieValue(word);
double unknown = std::log(1.0 / sum_);
loss[word] += known - unknown;

训练脚本会合并所有语料分片上的累计 loss 与出现次数,过滤低于 MIN_COUNT 的多字词,再计算剪枝排序分数:

score = loss / sqrt(count) / character_count

当前脚本对非 ASCII 词使用 UTF-8 字节数 / 3 估计字符数。这适用于常见三字节汉字,但不是通用的 Unicode 字符计数;它在这里仅作为剪枝归一化因子。

全部单字始终保留。对于其余候选,每轮保留排序靠前的 75%,但不会让词表缩到目标大小以下:

new_size
= max(target_size - single_count,
      eligible_count × 75%)

完整训练循环为:

正向最长匹配冷启动
        ↓
多轮 Hard EM
        ↓
计算删除 loss 并剪枝
        ↓
词表仍然过大?── 是 ──→ 继续训练
        ↓ 否
最终 Hard EM
        ↓
word<TAB>frequency

最终词典既可以直接交给 DictCut 做中文分词,也可以由 PieceTokenizer 在 PreTokenize 阶段加载:DictCut 负责确定中文词边界,SentencePiece 再在每个词内部学习子词。

配套实现:Ismantic/DictCut

中文分词:高级篇

引入 CRF

前面章节的分词方法容易理解,也有很高的执行效率,但面对歧义、新词和复杂上下文时,仅靠局部匹配往往难以作出稳定判断。更进一步的做法,是把中文分词看成一个字序列标注问题:为句子中的每个字预测其在词语中的位置,再根据标签序列恢复词语边界。

条件随机场(Conditional Random Field,CRF)正是解决这类问题的经典模型,也是 Wapiti 使用的核心模型。给定字序列 \(x = (x_1, x_2, \ldots, x_n)\),CRF 不再逐字独立判断,而是为完整的标签序列 \(y = (y_1, y_2, \ldots, y_n)\) 计算条件概率,并从中选择整体得分最高的序列。这使模型既能利用当前字及其上下文,也能约束相邻标签之间的组合关系。

CRF 用于中文分词的关键能力

  • 条件建模:直接建模 \(P(y\mid x)\),关注已知句子时标签序列出现的概率
  • 全局解码:联合考虑整句话的标签,不以某个字的局部最优结果代替全局最优结果
  • 特征组合:可以同时使用字、上下文、字符类型和标签转移等特征

举例,“南京市长江大桥”只看局部可能会把“市长”识别成一个词,但结合完整上下文,更合理的切分是“南京市 / 长江大桥”。常用的 BMES 标签体系以 B、M、E、S 分别表示词首、词中、词尾和单字成词。这个切分可以表示为:

南  京  市  长  江  大  桥
B   M   E   B   M   M   E

这样,分词问题就变成了为字序列寻找最佳标签序列的问题。

目标函数

概率公式

CRF 的条件概率定义为:

$$ P( y|x ) = \frac{1}{Z(x)}\times\exp\left(\sum_{i}\sum_{k}\lambda_k f_k(y_{i-1},y_i,x,i)\right) $$

关键组成部分

  • \(Z(x)\):归一化因子,确保概率和为1
  • \(f_k(y_{i-1},y_i,x,i)\):特征函数
  • \(\lambda_k\):特征权重参数

特征函数类型

一元特征(Unigram Features) $$ f_1(y_i,x,i) = \begin{cases} 1 & \text{if } (y_i,x,i) \text{满足指定的一元特征模板} \\ 0 & \text{otherwise} \end{cases} $$

二元特征(Bigram Features) $$ f_2(y_{i-1},y_i,x,i) = \begin{cases} 1 & \text{if } (y_{i-1},y_i,x,i) \text{满足指定的二元特征模板} \\ 0 & \text{otherwise} \end{cases} $$

势函数 为简化表示,引入势函数: $$ \psi_t(y’,y,x) = \exp\left(\sum_k \lambda_k f_k(y’,y,x,t)\right) $$

则条件概率可重写为: $$ P( y| x) = \frac{1}{Z(x)} \times \prod_t \psi_t(y_{t-1},y_t,x) $$

归一化因子

归一化因子需要对全部可能的标签序列求和:

$$ Z(x) = \sum_y \prod_t \psi_t(y_{t-1},y_t,x) $$

问题:如果有 \(T\) 个位置,每个位置有 \(L\) 个可能标签,则需要计算 \(L^T\) 个序列!

方案:引入动态规划算法高效计算。

标注过程

维特比

目标:找到最优标签序列,使得 \(P(y|x)\) 最大。

等价于最大化: $$ \text{score}(y, x) = \sum_i \sum_k \lambda_k f_k(y_{i-1}, y_i, x, i) = \sum_i \log \psi_i(y_{i-1}, y_i, x) $$

动态规划

状态定义 $$ \delta_t(y) = \max_{y_1,…,y_{t-1}} \text{score}(y_1,…,y_{t-1},y, x_1,…,x_t) $$

\(\delta_t(y)\) 表示到位置\(t\)标签为\(y\)的最优路径得分。

递推公式 $$ \delta_1(y) = \log \psi_1(\text{START}, y, x) $$ $$ \delta_t(y) = \max_{y’} \left[\delta_{t-1}(y’) + \log \psi_t(y’, y, x)\right] $$

回溯指针 $$ \phi_t(y) = \arg\max_{y’} \left[\delta_{t-1}(y’) + \log \psi_t(y’, y, x)\right] $$

梯度推导

训练目标

目标是最大化对数似然函数:

$$ L(\lambda) = \sum_s \log P(y^{(s)}|x^{(s)}) - R(\lambda) $$

其中:

  • \(s\) 是训练样本索引
  • \(y^{(s)}\) 和 \(x^{(s)}\) 分别是第 \(s\) 个样本的真实标签序列和观察序列
  • \(R(\lambda)\) 是正则化项

概率公式

$$ P(y|x) = \frac{1}{Z(x)} \exp \left(\sum_i \sum_k \lambda_k f_k(y_{i-1},y_i,x,i)\right) $$

其中:

  • \(Z(x)\) 是归一化因子
  • \(f_k(y_{i-1},y_i,x,i)\) 是第\(k\) 个特征函数
  • \(\lambda_k\) 是对应的权重参数

推导过程

第一步:目标函数展开

把概率公式带入目标函数: $$ \log P(y|x) = \sum_i \sum_k \lambda_k f_k(y_{i-1},y_i,x,i) - \log Z(x) $$

目标函数变为: $$ L(\lambda) = \sum_s \left[\sum_i \sum_k \lambda_k f_k(y_{i-1}^{(s)},y_i^{(s)},x^{(s)},i) -\log Z(x^{(s)})\right] - R(\lambda) $$

第二步:对参数求偏导

对参数 \(\lambda_k\) 求偏导:

$$ \frac{\partial L}{\partial \lambda_k} = \sum_s \left[\sum_i f_k(y_{i-1}^{(s)},y_i^{(s)},x^{(s)},i) - \frac{\partial \log Z(x^{(s)})}{\partial \lambda_k}\right] - \frac{\partial R(\lambda)}{\partial \lambda_k} $$

分析各项:

  • 第一项: \(\sum_i f_k(y_{i-1}^{(s)},y_i^{(s)},x^{(s)},i)\) 是真实标签序列下的特征值总和(经验期望
  • 第二项: \(\frac{\partial \log Z(x^{(s)})}{\partial \lambda_k}\) 需要详细推导(模型期望
  • 第三项: 正则化项的导数,L1 会特殊些,要专门来处理

关键在于计算第二项!

第三步:计算 \(\frac{\partial \log Z(x)}{\partial \lambda_k}\)

使用链式法则:

$$ \frac{\partial \log Z(x)}{\partial \lambda_k} = \frac{1}{Z(x)} \cdot \frac{\partial Z(x)}{\partial \lambda_k} $$

第四步:计算 \(\frac{\partial Z(x)}{\partial \lambda_k}\)

归一化因子的定义:

$$ Z(x) = \sum_y \exp\left(\sum_i \sum_k \lambda_k f_k(y_{i-1},y_i,x,i)\right) $$

用势函数表示:

$$ Z(x) = \sum_y \prod_t \psi_t(y_{t-1},y_t,x) $$

其中势函数:

$$ \psi_t(y_{t-1},y_t,x) = \exp\left(\sum_k \lambda_k f_k(y_{t-1},y_t,x,t)\right) $$

对 \(\lambda_k\) 求偏导:

$$ \frac{\partial Z(x)}{\partial \lambda_k} = \sum_y \frac{\partial}{\partial \lambda_k} \prod_t \psi_t(y_{t-1},y_t,x) $$

第五步:势函数的偏导数

由于 \(\psi_t\) 是指数函数:

$$ \frac{\partial \psi_t}{\partial \lambda_k} = \psi_t \cdot f_k(y_{t-1},y_t,x,t) $$

使用乘积法则,对于乘积 \(\prod_t \psi_t\):

$$ \frac{\partial}{\partial \lambda_k} \prod_t \psi_t = \sum_i \left[\prod_{t \ne i} \psi_t \right] \cdot \frac{\partial \psi_i}{\partial \lambda_k} $$ $$ = \sum_i \left[\prod_{t \ne i} \psi_t \right] \cdot \psi_i \cdot f_k(y_{i-1},y_i,x,i) $$ $$ = \prod_t \psi_t \cdot \sum_i f_k(y_{i-1},y_i,x,i) $$

关键洞察:最后一步从各项中提取了公共因子 \(\prod_t \psi_t\),将表达式化为“连乘 × 连加”

第六步:代入求和

把结果代入 \(Z(x)\) 的偏导数:

$$ \frac {\partial Z(x)}{\partial \lambda_k} = \sum_y \prod_t \psi_t(y_{t-1},y_t,x) \cdot \sum_i f_k(y_{i-1},y_i,x,i) $$ $$ = \sum_y \left[\sum_i f_k(y_{i-1},y_i,x,i)\right] \cdot \prod_t \psi_t(y_{t-1},y_t,x) $$

第七步:转化为概率形式

关键转化:注意到

$$ \prod_t \psi_t(y_{t-1},y_t,x) = P(y|x) \cdot Z(x) $$

代入得:

$$ \frac{\partial Z(x)}{\partial \lambda_k} = \sum_y \left[\sum_i f_k(y_{i-1},y_i,x,i)\right] \cdot P(y|x) \cdot Z(x) $$ $$ = Z(x) \cdot \sum_y P(y|x) \cdot \left[\sum_i f_k(y_{i-1},y_i,x,i)\right] $$

第八步:最终结果

把结果代入链式法则:

$$ \frac{\partial \log Z(x)}{\partial \lambda_k} = \frac{1}{Z(x)} \cdot Z(x) \cdot \sum_y P(y|x) \cdot \left[\sum_i f_k(y_{i-1},y_i,x,i)\right] $$

这正是模型期望:

$$ E_{P(y|x)}\left[\sum_i f_k(y_{i-1},y_i,x,i)\right] $$

第九步:梯度的最终形式

完整梯度公式

$$ \frac{\partial L}{\partial \lambda_k} = \sum_s \left[\sum_i f_k(y_{i-1}^{(s)}, y_i^{(s)},x^{(s)},i)\right] {}- \sum_s E_{P(y|x^{(s)})}\left[\sum_i f_k(y_{i-1},y_i,x^{(s)},i)\right] {}- \frac{\partial R(\lambda)}{\partial \lambda_k} $$

或者表示成

$$ \frac{\partial L}{\partial \lambda_k} = \text{经验期望} - \text{模型期望} {}- \frac{\partial R(\lambda)}{\partial \lambda_k} $$

其中:

  • 经验期望:真实数据中特征 \(f_k\) 出现的次数
  • 模型期望: 当前模型认为特征 \(f_k\) 应该出现的次数

前向后向

核心问题 对于已经推导出来的梯度公式:

$$ \frac{\partial L}{\partial \lambda_k} = \text{经验期望} - \text{模型期望} {}- \frac{\partial R(\lambda)}{\partial \lambda_k} $$

其中模型期望是:

$$ E_{P(y|x)}[f_k] = \sum_y P(y|x) \cdot \left[\sum_i f_k(y_{i-1},y_i,x,i)\right] $$

要是直接计算需要遍历 \(L^T\) 个可能的标签序列,计算量太大了!

这里与维特比算法的区别只有一个关键运算:维特比在每个状态保留最大值,用于寻找最佳切分;前向后向算法对所有路径求和,用于计算边际概率和模型期望。

突破口

对于一元特征: $$ E[f_k^{(1)}] = \sum_y P(y|x) \sum_i f_k^{(1)}(y_i,x,i) $$

交换求和顺序 $$ = \sum_i \sum_y P(y|x) \cdot f_k^{(1)}(y_i,x,i) $$

进一步分解:只有当 \(y_i\) 取特定值时 \(f_k^{(1)}\) 才为1 $$ = \sum_i \sum_{y_i} f_k^{(1)}(y_i,x,i) \sum_{y_1,…,y_{i-1},y_{i+1},…,y_T}P(y_1,…,y_T|x) $$

得到边际概率: $$ = \sum_i \sum_{y_i} f_k^{(1)}(y_i,x,i) \cdot P(y_i|x) $$

类似地,对于二元特征: $$ E[f_k^{(2)}] = \sum_i \sum_{y_{i-1}} \sum_{y_i} f_k^{(2)}(y_{i-1},y_i,x,i) \cdot P(y_{i-1},y_i|x) $$

这样,就把求 \(P(y|x)\) 转变成了怎么去计算边际概率 \(P(y_i|x)\) 和 \(P(y_{i-1},y_i|x)\) 了,问题化简了不少。

思路

关键洞察:边际概率可以分解为到达当前位置的前向分数、当前转移的势函数和离开当前位置的后向分数。这里的前向量与后向量都是未归一化分数,不是条件概率。

$$ \psi_t(y’,y,x) = \exp\left(\sum_k \lambda_k f_k(y’,y,x,t)\right) $$

前向过程

分数定义

前向分数 \(\alpha_i(y)\) 是所有以标签 \(y\) 到达位置 \(i\) 的路径势函数之和:

$$ \alpha_i(y) = \sum_{y_1,\ldots,y_{i-1}} \prod_{t=1}^{i}\psi_t(y_{t-1},y_t,x) $$

初始化

$$ \alpha_1(y)=\psi_1(\mathrm{START},y,x) $$

递推公式

$$ \alpha_i(y) = \sum_{y’}\alpha_{i-1}(y’)\psi_i(y’,y,x) $$

最后一个位置的前向分数之和就是配分函数:

$$ Z(x)=\sum_y\alpha_T(y) $$

后向过程

后向分数 \(\beta_i(y)\) 是从位置 \(i\) 的标签 \(y\) 出发,到达序列末尾的所有后缀路径势函数之和:

$$ \beta_i(y) = \sum_{y_{i+1},\ldots,y_T} \prod_{t=i+1}^{T}\psi_t(y_{t-1},y_t,x) $$

初始化

$$ \beta_T(y)=1 $$

递推公式

$$ \beta_i(y) = \sum_{y’}\psi_{i+1}(y,y’,x)\beta_{i+1}(y’) $$

边际概率

一元边际概率

$$ P(y_i=y\mid x) = \frac{\alpha_i(y)\beta_i(y)}{Z(x)} $$

二元边际概率

$$ P(y_{i-1}=y’,y_i=y\mid x) = \frac{ \alpha_{i-1}(y’) \psi_i(y’,y,x) \beta_i(y) }{Z(x)} $$

这两个公式使用同一个全局配分函数 \(Z(x)\),不需要分别定义 \(z_i\)、\(Z_{\text{unigram}}\) 或 \(Z_{\text{bigram}}\)。

数值稳定性

直接连乘势函数容易上溢或下溢。实现时可以在对数域中保存前向与后向分数,并用 \(\operatorname{logsumexp}\) 完成求和。例如:

$$ \log\alpha_i(y) = \operatorname{logsumexp}_{y’} \left( \log\alpha_{i-1}(y’) {}+ \log\psi_i(y’,y,x) \right) $$

也可以在每个位置对向量进行缩放,但前向、后向和边际概率必须使用同一组缩放因子并保持一致的定义。

L1 正则化

CRF 会产生大量特征,其中许多权重接近零,却仍然占用模型空间并参与计算。L1 正则化通过惩罚权重绝对值产生稀疏参数。Wapic 实际最小化的目标同时支持 L1 和 L2:

$$ \begin{aligned} F(\lambda) &=-L(\lambda)+r_1\sum_k|\lambda_k| \\ &\quad+\frac{r_2}{2}\sum_k\lambda_k^2 \end{aligned} $$

其中 \(r_1\) 对应命令行参数 --rho1,控制稀疏程度;\(r_2\) 对应 --rho2,为光滑部分增加 L2 惩罚。GradientComputer::RunGradientComputation 返回这个目标值,并把 \(r_2\lambda_k\) 加入普通梯度;L1 项则由优化器单独处理。

零点不可导

L1 正则化项在非零点可导,但在零点不可导。绝对值函数的次微分为:

$$ \partial |\lambda_k| = \begin{cases} {1} & \text{if } \lambda_k > 0 \\ {-1} & \text{if } \lambda_k < 0 \\ [-1,1] & \text{if } \lambda_k = 0 \end{cases} $$

标准 L-BFGS 假设目标函数光滑,不能直接处理零点处的区间次梯度。后面的 OWL-QN 会用伪梯度和象限约束解决这个问题,并让一部分参数精确变为零。

L-BFGS

L-BFGS(Limited-memory Broyden-Fletcher-Goldfarb-Shanno)是一种拟牛顿方法。它适合优化不含 L1 项的光滑 CRF 目标,并用有限的历史向量近似二阶曲率。

Wapic 用同一个 LBFGSOptimizer 承担两种模式:当 \(r_1=0\) 时执行标准 L-BFGS;当 \(r_1\neq0\) 时启用伪梯度和象限投影,实际执行后文的 OWL-QN。命令行虽然统一写作 -a l-bfgs,默认 \(r_1=0.5\),因此默认训练会进入 OWL-QN 分支。

拟牛顿法

问题基础

CRF 的参数估计本质上是一个无约束优化问题:

$$ \min_{\lambda} F(\lambda) = -\ell(\lambda) + R(\lambda) $$

其中:

  • \(\ell(\lambda)\) 是对数似然函数
  • \(R(\lambda)\) 是正则化项
  • \(\lambda\) 是 CRF 的参数向量

梯度下降只使用一阶信息:

$$ \lambda^{(k+1)} = \lambda^{(k)} -\alpha_k \nabla F(\lambda^{(k)}) $$

牛顿法则使用 Hessian 修正更新方向:

$$ \lambda^{(k+1)} = \lambda^{(k)} - [\nabla^2 F(\lambda^{(k)})]^{-1} \nabla F(\lambda^{(k)}) $$

对于高维 CRF,显式构造和求解 Hessian 的成本过高。拟牛顿法因此不直接计算 Hessian,而是逐步近似它的逆矩阵。

拟牛顿法

用 \(H_k\) 表示 Hessian 逆矩阵的近似,搜索方向为:

$$ d_k = -H_k \nabla F(\lambda^{(k)}) $$

当 \(H_k\) 正定时,\(\nabla F(\lambda^{(k)})^T d_k<0\),因此它是下降方向。线搜索随后决定沿该方向前进多远。

BFGS

BFGS 使用相邻两次迭代的参数变化和梯度变化来更新曲率近似。定义:

$$ s_k=\lambda^{(k+1)}-\lambda^{(k)} $$

$$ y_k=\nabla F(\lambda^{(k+1)})-\nabla F(\lambda^{(k)}) $$

如果 \(H_k\) 表示 Hessian 逆矩阵的近似,那么它应满足割线条件:

$$ H_{k+1}y_k=s_k $$

标准 BFGS 的逆矩阵更新为:

$$ \begin{aligned} H_{k+1} &=(I-\rho_k s_k y_k^T) H_k (I-\rho_k y_k s_k^T) \\ &\quad+\rho_k s_k s_k^T \end{aligned} $$

其中:

$$ \rho_k=\frac{1}{y_k^T s_k} $$

曲率条件

如果 \(H_k\) 正定,并且满足:

$$ y_k^T s_k>0 $$

那么 BFGS 更新得到的 \(H_{k+1}\) 仍然正定,搜索方向:

$$ d_k=-H_k\nabla F(\lambda^{(k)}) $$

就是下降方向。Wapic 的光滑分支使用 Wolfe 条件进行线搜索,并在 UpdateHistory 中直接记录 \(s_k\)、\(y_k\) 和 \(\rho_k=1/(y_k^Ts_k)\)。

L-BFGS

标准 BFGS 需要存储 \(n \times n\) 的矩阵,对于大规模问题(\(n\) 很大)并不现实。L-BFGS 不显式存储 Hessian 逆矩阵,而是保存最近的 \(m\) 组向量对 \({s_k,y_k}\),用它们隐式表示曲率信息。

核心思想:

  • 内存需求: 从\(O(n^2)\)降低到\(O(mn)\)
  • 计算复杂度: 每次迭代从\(O(n^2)\)降低到\(O(mn)\)
  • Wapic 默认值: \(m=5\),可通过 --histsz 调整

两步循环算法

算法目标

两步循环算法的目标是计算搜索方向: $$ d_k = -H_k g_k $$

其中 \(H_k\) 是 Hessian 逆矩阵的 L-BFGS 近似,\(g_k\) 是当前梯度。

关键是:我们要直接计算出 \(H_k g_k\) 的结果,而不显式构造 \(H_k\) 矩阵。

完整算法

输入

  • 当前梯度:\(g_k\)
  • 历史信息:\({s_i, y_i}_{i=k-m}^{k-1}\)(最近\(m\)个向量对)
  • 初始 Hessian 逆矩阵近似:\(H_k^0\)(通常是标量乘以单位矩阵)

输出

  • 搜索方向:\(d_k = -H_k g_k\)

算法步骤

第一步:反向循环(Backward Loop)

初始化: \(q = g_k\)

反向遍历历史信息: \(\text{for } i = k-1, k-2, \ldots, k-m:\) \(\alpha_i = \frac{s_i^T q}{y_i^T s_i}\) \(q = q - \alpha_i y_i\) \(\text{存储 } \alpha_i \text{ 供第二步使用}\)

第二步:正向循环(Forward Loop)

应用初始 Hessian 逆矩阵近似: \(r = H_k^0 q\)

正向遍历历史信息: \(\text{for } i = k-m, k-m+1, \ldots, k-1:\) \(\beta = \frac{y_i^T r}{y_i^T s_i}\) \(r = r + s_i (\alpha_i - \beta)\)

返回结果: \(\text{return } -r \quad \text{// 注意负号,因为我们要的是 } -H_k g_k\)

数学原理

每一对 \((s_i,y_i)\) 都对应一次 BFGS 逆矩阵更新。连续使用最近 \(m\) 个向量对时,矩阵中出现的是这些更新变换的有序乘积,不能把它们简单写成若干秩一矩阵的和。两步循环按照与这些更新相同的顺序执行向量运算,因此能直接得到 \(H_kg_k\),而不需要显式构造 \(H_k\)。

计算复杂度

时间复杂度

  • 每个内积操作:\(O(n)\)
  • 每个循环有 \(m\) 次迭代,每次做常数个内积
  • 总复杂度:\(O(mn)\)

空间复杂度

  • 存储 \(m\) 个 \(s_i\) 向量:\(O(mn)\)
  • 存储 \(m\) 个 \(y_i\) 向量:\(O(mn)\)
  • 临时变量:\(O(n)\)
  • 总复杂度:\(O(mn)\)

相比标准 BFGS 的 \(O(n^2)\) 存储,两步循环只执行向量运算,因而适合高维 CRF 参数。Wapic 用 syp 保存循环历史;sy 使用 float 减少内存,内积仍以 float_t 累加。

初始 Hessian 逆矩阵的选择

Wapic 使用标量矩阵作为初始近似: \(\gamma = \frac{y_{k-1}^T s_{k-1}}{y_{k-1}^T y_{k-1}}\) \(H_k^0 = \gamma I\) \(r = \gamma q\)

OWL-QN

L1 正则化的挑战

令 \(f(\lambda)\) 表示包含负对数似然和 L2 惩罚的光滑部分,OWL-QN 处理的目标函数为: $$ F(\lambda) = f(\lambda) + C \sum_k |\lambda_k| $$

在 Wapic 中,\(C=r_1\)。

问题:

  • L1 正则化项在零点不可微
  • 标准 L-BFGS 要求目标函数处处可微
  • 需要特殊处理来保持L-BFGS的收敛性质

OWL-QN 算法

OWL-QN(Orthant-Wise Limited-memory Quasi-Newton)是L-BFGS在L1 正则化上的扩展。

伪梯度的定义

OWL-QN 引入**伪梯度(pseudo-gradient)**的概念来处理不可微性: $$ \widetilde{\nabla}F(\lambda)_k = \begin{cases} \nabla f(\lambda)_k + C & \text{if } \lambda_k > 0 \\ \nabla f(\lambda)_k - C & \text{if } \lambda_k < 0 \\ \nabla f(\lambda)_k + C & \text{if } \lambda_k = 0 \text{ and } \nabla f(\lambda)_k < -C \\ \nabla f(\lambda)_k - C & \text{if } \lambda_k = 0 \text{ and } \nabla f(\lambda)_k > C \\ 0 & \text{if } \lambda_k = 0 \text{ and } |\nabla f(\lambda)_k| \leq C \end{cases} $$

这对应 ComputePseudoGradient:普通梯度保存在 g,伪梯度保存在 pg

象限约束

OWL-QN 的关键思想是象限约束(orthant constraint)

搜索方向约束: 在计算搜索方向后,逐维保留与负伪梯度一致的分量: $$ \widetilde d_{k,i} = \begin{cases} d_{k,i} & \text{if } d_{k,i}(-\widetilde{\nabla}F(\lambda^{(k)})_i) > 0 \\ 0 & \text{otherwise} \end{cases} $$

这对应 ConstrainSearchDirection:如果某一维满足 d[f] * pg[f] >= 0,就把该方向分量清零。

象限投影

在线搜索过程中,需要将候选点投影回当前象限: $$ \lambda_{\text{projected}} = \text{project}(\lambda^{(k)} + \alpha d_k, \xi) $$

其中投影操作定义为: $$ \text{project}(\lambda, \xi)_i = \begin{cases} \lambda_i & \text{if } \lambda_i \cdot \xi_i > 0 \\ 0 & \text{otherwise} \end{cases} $$

参考象限不能一律取负伪梯度。对已经非零的参数,应保持它当前所在的象限;只有参数为零时,才使用负伪梯度决定准备进入的象限:

$$ \xi_i = \begin{cases} \operatorname{sign}(\lambda_i^{(k)}) & \text{if } \lambda_i^{(k)}\neq 0 \\ \operatorname{sign}(-\widetilde{\nabla}F(\lambda^{(k)})_i) & \text{if } \lambda_i^{(k)}=0 \end{cases} $$

ProjectToOrthant 使用上一次参数 xp 确定非零参数的象限;如果原参数为零,则使用 -pg 选择准备进入的象限。

回溯线搜索

Wapic 沿用原始 Wapiti 的 OWL-QN 回溯判定。候选点经过象限投影后,代码检查: $$ F(\lambda_{\text{projected}}) < F(\lambda^{(k)}) {}+ \gamma d_k^T (\lambda_{\text{projected}}-\lambda^{(k)}) $$

其中 \(\gamma=10^{-4}\)。CheckArmijoRule 直接使用搜索方向 d 与投影后实际位移的内积;这与标准 Armijo 条件使用目标函数方向导数的写法不同,因此这里把它称为 Wapiti 的回溯判定,而不把两者视为完全相同的公式。

算法流程

Wapic 的 LBFGSOptimizer::Optimize 在 \(r_1\neq0\) 时执行以下流程:

  1. 初始化: 设置初始参数 \(\lambda^{(0)}\),记忆深度 \(m\),正则化参数 \(C\)

  2. 主循环: 对于 \(k = 0, 1, 2, \ldots\)

    a. 计算伪梯度: \(g_k=\widetilde{\nabla}F(\lambda^{(k)})\)

    b. 检查收敛: 如果 \(|g_k|\) 足够小,则停止

    c. 计算搜索方向: 使用两步循环算法计算 \(d_k = -H_k g_k\)

    d. 象限约束: 删除不与负伪梯度同向的搜索方向分量

    e. 线搜索: 使用 Wapiti 的回溯判定寻找步长 \(\alpha_k\)

    f. 象限投影: \(\lambda^{(k+1)} = \text{project}(\lambda^{(k)} + \alpha_k d_k, \xi)\)

    g. 更新历史: 令 \(s_k=\lambda^{(k+1)}-\lambda^{(k)}\),并用光滑部分的梯度差 \(y_k=\nabla f(\lambda^{(k+1)})-\nabla f(\lambda^{(k)})\) 更新历史信息

实现参数

  • --rho1:L1 强度,默认 0.5;设为 0 时关闭 OWL-QN 分支
  • --rho2:L2 强度,默认 0.0001
  • --histsz:历史向量对数量,默认 5
  • --maxls:最大线搜索次数,默认 40
  • --maxiter:最大训练轮数,默认 100

CheckConvergence 同时检查伪梯度范数和一段目标函数历史中的相对改善量。线搜索失败、梯度足够小或目标函数改善不足时,训练停止。

三节的关系可以归结为:L1 正则化定义稀疏目标;L-BFGS 高效优化光滑目标;当训练既需要 L-BFGS 的曲率信息,又需要 L1 稀疏性时,使用 OWL-QN。

至此,中文分词可以被看作两类互补方案:词典方法通过词频寻找最大概率路径,CRF 则通过特征与标签转移进行全局序列预测。它们都可以作为预切分阶段,为后续子词模型提供更稳定的输入边界。

配套实现:Ismantic/Wapic

Tokenizer:SentencePiece

Tokenizer 不只是一个 BPE 合并算法。从原始文本到 token id,中间还要经过规范化、预切分、词表训练和推理编码。PieceTokenizer 中的 sentencepiece 方法将这条流程分成四个阶段:

原始文本
  ↓ Normalize
规范化文本
  ↓ PreTokenize
预切分片段
  ↓ Counter
SentencePiece 模型
  ↓ Tokenizer
token pieces / token ids

训练和推理必须遵守同一套文本处理规则。否则即使词表相同,相同文本也可能产生不同的 token 序列。

Normalize

Unicode 中可能存在外观相同、编码不同的字符。例如全角字母、兼容字符和组合字符如果不先统一,会被训练器当成不同符号,浪费词表空间。

Normalizer 使用预编译的 Unicode 映射表完成 NFKC 等规范化。映射表被编码为 Double-Array Trie,处理文本时对当前位置执行最长前缀匹配:

输入字节
  ↓ Trie 最长匹配
找到规则 ──→ 输出规范化结果
未找到   ──→ 保留当前 UTF-8 码点
非法字节 ──→ U+FFFD

空格也在这一阶段处理。默认情况下,连续空格会被合并,开头和结尾的空格被移除,其余空格替换成可见符号

"  Hello   world  "
        ↓
"Hello▁world"

这样空格就能像普通字符一样进入词表,同时仍可在解码时恢复。启用 reconstruct 后,Normalizer 会保留全部空格,不再执行合并和裁剪。

规范化规则、空格符号和 reconstruct 配置保存在 PreTokenizerSpec 中,并随模型一起保存。

PreTokenize

PreTokenizer 在 BPE 训练之前确定基本边界。PieceTokenizer 将它分成三个维度:

维度配置作用
Splitword / isolate决定空格、标点和英文缩写怎样切分
Digitkeep / split数字串整体保留或逐码点拆开
Cnnone / char / dict连续汉字不切、逐字切或按词典切

例如:

输入:Hello, world 你好123

word + keep + none
→ Hello | , | ▁world | ▁ | 你好 | 123

isolate + split + char
→ Hello | , | ▁ | world | ▁ | 你 | 好 | 1 | 2 | 3

这些边界非常重要。Counter 只在单个预切分片段内部学习 BPE,不会把两个片段直接合并成一个 piece。

中文词典模式同样发生在这里。CnCutter 将词典构造成 Double-Array Trie,并把词频转换成对数概率。切分时,它从每个 UTF-8 码点边界查询所有词典前缀,将候选词看成一条从起点到终点的边:

南京市

0 ──南──> 3 ──京──> 6 ──市──> 9
└────南京────> 6
└──────南京市──────> 9

Viterbi 动态规划为每个字节位置保存当前最高分及其前驱位置,最后从文本末尾沿前驱回溯,得到总分最高的切分路径。词典没有覆盖的位置会加入单个 Unicode 码点的 fallback 边,因此生僻字也不会使路径中断。

for (const auto& match : GetMatches(han_run)) {
    int start = match.end - match.length + 1;
    int end = match.end + 1;
    float_t score = scores[start] + match.weight;
    if (score > scores[end]) {
        scores[end] = score;
        routes[end] = start;
    }
}

中文切词完成后,每个词再独立交给 BPE 学习词内的子词:

南京市长江大桥
  ↓ 中文切词
南京市 | 长江大桥
  ↓ BPE
南京 | 市 | 长江 | 大桥

这样可以减少 BPE 仅凭局部共现频率学出跨词片段。这里的中文切词器是 PreTokenizer 内部的独立组件,不依赖后面的 BytePiece 或 SentencePiece 算法。dict=no 表示逐字模式;CLI 训练时还会将它组合成 isolate + digit split,使中文、数字和标点保持细粒度。

Split、Digit 等配置保存在模型中。外部词典内容不会写进模型,因此词典模式下的训练和推理必须使用同一份词典。

实现上,这三个维度最终收束到同一个入口:

std::vector<std::string> PreTokenizer::Split(
    std::string_view normalized) const {
    return ustr::SplitTextCn(
        normalized, space_, cn_cut_fn_, cut_, split_digits_);
}

std::vector<std::string> PreTokenizer::PreTokenize(
    std::string_view text) const {
    return Split(normalizer_.Normalize(text));
}

cn_cut_fn_ 为空时,SplitTextCn 仍会执行普通切分和数字拆分;配置中文模式时,它再对连续汉字调用逐字或词典切分。Counter 与 Tokenizer 都调用 PreTokenize,从代码层面保证边界一致。

Counter

SentencePieceCounter 负责从预切分语料学习 BPE 词表。它首先加入几类基础 piece:

  • <unk><s></s> 等控制符号;
  • 256 个 byte piece,覆盖所有可能的字节;
  • character_coverage 选出的必备字符。

字符覆盖率

Counter 先统计所有预切分片段及其频率,再按片段频率加权统计字符。例如 low 出现 5 次,其中的 low 都贡献 5 次。

字符按照频率从高到低加入必备字符集合,直到累计频率达到 character_coverage。覆盖率之外的稀有字符在训练语料中暂时替换成内部未知字符,避免大量罕见字符占满词表。

字符覆盖率解决的是“哪些字符值得作为普通 piece 保留”,byte fallback 解决的是“未保留的字符如何编码”。两者并不冲突:常见字符使用普通词表,罕见字符仍可退回 UTF-8 字节。

初始化 Symbol

BPE 从单个 Unicode 码点开始:

片段:lowest
初始:l | o | w | e | s | t

每个 Symbol 可以表示一个字符,也可以表示左右两个 Symbol 合并后的结果。它还记录频率以及该相邻对在语料中的位置:

struct Symbol {
    const Symbol* left;
    const Symbol* right;
    UnicodeText chars;
    size_t byte_size;
    uint64_t fp;
    uint64_t freq;
    std::vector<uint64_t> positions;
};

字符 Symbol 会被缓存,同一个字符不需要重复创建对象。两个相邻 Symbol 也通过 fingerprint 缓存成同一个候选 Symbol,positions 则记录这个 pair 出现在哪个片段、哪两个位置。

位置本身不保存完整对象,而是压缩成一个 uint64_t

高 32 位:片段编号 sid
中 16 位:左 Symbol 位置
低 16 位:右 Symbol 位置

语料中相同 pair 的所有位置因此可以集中到同一个候选上。计算频率时,再把对应片段的出现次数累加起来。

选择并合并 pair

考虑经过预切分后的简单语料:

low     5 次
lowest  2 次
newer   3 次

初始 pair 频率包括:

l+o = 7
o+w = 7
w+e = 5
n+e = 3
e+w = 3

频率相同时,Counter 再按照长度和字符串顺序得到确定结果。假设首先选择 l+o

l | o | w       → lo | w
l | o | w | ... → lo | w | ...

下一轮会出现新的候选 lo+w,频率为 7:

lo | w → low

训练循环就是不断重复“选择最高频 pair、合并全部有效位置、加入新邻居”,直到达到目标词表大小。

局部更新

一次合并只会改变它附近的 pair:

l | o | w | e | s | t
    ↓ 合并 l+o
lo | w | e | s | t
      ↓ 更新相邻候选
lo+w

右侧位置被置空,只有新 Symbol 左右两侧的候选需要重新加入。其他位置没有变化,不需要重新扫描。

旧候选的 positions 中可能仍然保存已经失效的位置。Counter 不会在每次合并后到所有候选中删除它们,而是在 ComputeFreq 时检查左右位置是否仍然指向原 Symbol,只累加有效位置,并顺便压缩位置列表。这是一种延迟清理策略。

void ComputeFreq(Symbol* symbol) {
    if (symbol->freq > 0) return;

    size_t write = 0;
    for (uint64_t encoded : symbol->positions) {
        Position pos = DecodePos(encoded);
        if (symbol->left == symbols_[pos.sid][pos.left] &&
            symbol->right == symbols_[pos.sid][pos.right]) {
            symbol->freq += freqs_[pos.sid];
            symbol->positions[write++] = encoded;
        }
    }
    symbol->positions.resize(write);
}

这里的 freqs_[pos.sid] 是预切分片段在语料中的出现次数。一个位置有效时累加的不是 1,而是该片段的完整权重。

选出最高频 pair 后,真正修改语料表示的代码只处理命中的位置及其两个邻居:

for (uint64_t encoded : best_symbol->positions) {
    Position pos = DecodePos(encoded);
    if (symbols_[pos.sid][pos.left] == nullptr) continue;

    int prev = GetPrevIndex(pos.sid, pos.left);
    int next = GetNextIndex(pos.sid, pos.right);

    ResetFreq(pos.sid, prev, pos.left, best_symbol);
    ResetFreq(pos.sid, pos.right, next, best_symbol);

    symbols_[pos.sid][pos.left] = best_symbol;
    symbols_[pos.sid][pos.right] = nullptr;

    AddNewPair(pos.sid, prev, pos.left);
    AddNewPair(pos.sid, pos.left, next);
}

右位置被置为 nullptr,左位置指向合并后的 Symbol。旧邻居候选的频率被标记为待重算,新产生的两个相邻 pair 则加入候选集合;整轮没有重新遍历语料。

控制候选规模

候选集合也不会始终保存所有 pair。Counter 定期计算候选频率,只保留较高频的一部分作为 active symbols;每隔一段合并再刷新一次。这使大规模语料上的训练仍能控制时间和内存。

pair 的长度还受 max_piece_size 限制,避免重复标点或噪声文本产生异常长的 piece。

生成模型

每次选中的 piece 按学习顺序获得分数:

第 1 个合并:score =  0
第 2 个合并:score = -1
第 3 个合并:score = -2

这里的 score 不是 piece 出现的概率,而是推理时重放 BPE 的优先级。越早学到的合并分数越高。

pieces_.emplace_back(
    best_symbol->ToString(),
    -static_cast<float>(pieces_.size()));

这行代码把训练顺序直接编码进模型。推理器不需要保存独立的 merge 表,只要读取 piece 的 score 就能恢复相同顺序。

达到目标词表大小后,Counter 补入必备字符,并保存:

CounterSpec
PreTokenizerSpec
Pieces: piece / score / type

模型采用可读文本格式。训练配置、预切分配置、普通 piece、byte piece 和控制 token 因而可以一起加载。

Tokenizer

SentencePieceTokenizer 读取模型后,建立 piece → id 哈希表。编码时先按照模型中的配置执行 Normalize 和 PreTokenize,再分别编码每个预切分片段。因此,普通切分、数字拆分和中文切分产生的边界都不会被 BPE 跨越,训练与推理遵守同一套文本处理规则。

EncodeResult SentencePieceTokenizer::Encode(
    std::string_view text) const {
    EncodeResult result;
    for (const auto& piece : pretokenizer_.PreTokenize(text)) {
        auto sub = EncodeSegment(piece);
        result.insert(result.end(),
                      std::make_move_iterator(sub.begin()),
                      std::make_move_iterator(sub.end()));
    }
    return result;
}

EncodeSegment 一次只接收一个预切分片段,所以后面的 BPE 合并天然无法跨越片段边界。

为什么不能只做最长匹配

BPE 词表不仅记录有哪些字符串,还隐含了它们的学习顺序。假设词表同时存在 abbcabcabc 必须先由 a+bb+c 形成,才能继续参与下一次合并。直接从左到右选择最长字符串,可能得到与训练过程不同的结果。

SentencePieceTokenizer 因而从单个 Unicode 码点开始,按照模型 score 重放 BPE,而不是直接使用 Trie 最长匹配。

初始化候选队列

单个片段先被表示成 Symbol 数组,并通过 prevnext 索引模拟双向链表:

l ⇄ o ⇄ w ⇄ e ⇄ s ⇄ t

Tokenizer 检查每一对相邻 Symbol。如果拼接结果存在于词表,就将它放入优先队列:

l | o | w | e | s | t
└ l+o → lo,score=0
    └ o+w → ow,score=-3

优先队列先弹出 score 较高的 lo,将 lo 合并:

lo ⇄ w ⇄ e ⇄ s ⇄ t

此时只需检查 lo 的左右邻居。如果 low 也在词表中,就把 lo+w 加入队列。其他位置的相邻关系没有改变。

候选只在拼接结果已经存在于模型时入队:

auto MaybeAddPair = [&](int left, int right) {
    if (left == -1 || right == -1) return;

    std::string_view piece(
        symbols[left].piece.data(),
        symbols[left].piece.size() + symbols[right].piece.size());
    auto it = pieces_.find(piece);
    if (it == pieces_.end()) return;

    SymbolPair* pair = symbol_pair_allocator.Allocate();
    pair->left = left;
    pair->right = right;
    pair->score = model_->GetPieces(it->second).GetScore();
    pair->size = piece.size();
    agenda.push(pair);
};

SymbolPair 由对象池统一分配,避免大量候选反复申请小块内存。

失效候选

合并 l+o 后,队列里原来的 o+w 已经失效,因为 o 已被合并。为了避免在优先队列中间执行昂贵的删除,Tokenizer 允许旧候选暂时保留,等它被弹出时再检查:

  • 左右 Symbol 是否已经被消费;
  • 两个 Symbol 的总长度是否仍与候选一致。

无效候选直接丢弃,有效候选才执行合并。这与 Counter 的延迟位置清理采用了相同思路:先保留可能过期的信息,使用时再验证。

while (!agenda.empty()) {
    SymbolPair* top = agenda.top();
    agenda.pop();

    if (symbols[top->left].piece.empty() ||
        symbols[top->right].piece.empty() ||
        symbols[top->left].piece.size() +
            symbols[top->right].piece.size() != top->size) {
        continue;
    }

    Symbol& left = symbols[top->left];
    Symbol& right = symbols[top->right];
    left.piece = std::string_view(
        left.piece.data(), left.piece.size() + right.piece.size());
    left.next = right.next;
    if (right.next >= 0) {
        symbols[right.next].prev = top->left;
    }
    right.piece = {};

    MaybeAddPair(symbols[top->left].prev, top->left);
    MaybeAddPair(top->left, symbols[top->left].next);
}

合并直接扩展左侧 string_view,再修改 nextprev 索引并清空右侧 Symbol。因为两个片段来自同一段连续文本,所以无需创建新的字符串。

完整编码示例

假设模型学到了 lolowest,以及形成 est 所需的中间 piece,并为最终结果分配了示例 id:

lo   → 300
low  → 301
est  → 417

编码 lowest 的过程可以表示为:

l | o | w | e | s | t
    ↓ l+o
lo | w | e | s | t
    ↓ lo+w
low | e | s | t
          ↓ 按模型中已有的中间合并形成 est
low | est
    ↓ 查询 id
301 | 417

实际结果完全由模型中的 piece 和 score 决定,示例 id 只用于说明过程。

Byte fallback

如果某个剩余片段不在普通词表中,Tokenizer 会将它转换成 UTF-8 字节,并逐字节输出对应的 byte piece:

未知字符
  ↓ UTF-8
E7 8C AB
  ↓ byte fallback
<0xE7> | <0x8C> | <0xAB>

因此任意输入都可以编码,不需要把整段文本替换成 <unk>

解码

解码执行相反过程:普通 piece 直接拼接,byte piece 还原成原始字节,UNKNOWN 和控制 token 不写入文本,最后再把 还原为空格。

文本
  ↓ Normalize / 必要的 PreTokenize
UTF-8 Symbol
  ↓ 按模型分数合并
pieces
  ↓ 查询词表
token ids
  ↓ 拼接与 byte 恢复
文本

至此,SentencePiece 的训练与推理形成闭环:Normalize 统一字符表示,PreTokenize 确定合并边界,Counter 学习词表,Tokenizer 使用同一模型完成编码和解码。

配套实现:Ismantic/PieceTokenizer

Tokenizer:BytePiece

引言

BytePieceCounter 是 PieceTokenizer 的训练组件,用于从原始语料中生成候选 piece、统计权重并裁剪词表。

本文的基本算法参考了苏剑林的实现,本项目将其改写为 C++,并增加 UTF-8 切分边界约束。

核心任务:

输入:大量原始文本语料 输出:训练好的 BytePiece 模型,包含:

  • 词汇表:不同粒度的 subword pieces
  • 计数权重:每个 piece 经重分词后得到的频次

基本思路:

  • 统计阶段:收集文本中所有字节级 N-gram 的统计信息
  • 标注阶段:使用动态规划找到最优分词方案
  • 剪枝阶段:移除低频或冗余的词汇,优化词汇表大小
  • 重分词阶段:用保留词表重新切分候选,回收被裁剪 piece 的计数

BytePieceCounter 在训练过程中的剪枝阶段需要使用 BytePieceTokenizer 对被裁剪片段重新切分。

统计

N-gram 统计

N-gram 模型基于有限阶马尔可夫假设:当前符号的概率只依赖于前面的 $N-1$ 个符号。这里的符号是字节,不是词。

$$ \begin{aligned} P(w_1w_2\ldots w_m) &= \prod_i P(w_i\mid w_{i-N+1}\ldots w_{i-1}) \end{aligned} $$

BytePieceCounter 先统计字节级 N-gram,再利用这些局部条件概率给候选 piece 评分。最终词表中的 piece 必须在 UTF-8 字符边界切分,但 piece 内部的滑动 N-gram 可以从任意字节位置开始。

核心数据结构

BytePieceCounter 使用一组哈希表存储不同长度的 N-gram 统计:

std::vector<std::unordered_map<std::string, float_t>> N_;

结构解释

  • N_[i]:存储长度为 i 字节的所有子串及其统计值
  • N_[0]:空字符串(用于归一化)
  • N_[1]:所有 1-gram(单字节)
  • N_[2]:所有 2-gram(字节对)
  • N_[3]:所有 3-gram(三字节组合)
  • N_[6]:最长统计的 N-gram

示例初始化

N_.clear();
N_.resize(max_piece_count_ + 1);  // max_piece_count_ = 6
N_[0][""] = 0;  // 空字符串初始化

为什么选择字节级 N-gram?

  1. 语言无关性:任何 UTF-8 文本都能统一处理
  2. 完备性:保证 100% 覆盖,不存在未知字符
  3. 细粒度模式:能够发现字符内部和跨字符的统计规律

具体实现

void CountRawSegments(const std::vector<std::string>& segments) {
    // 初始化N_数组
    N_.clear();
    N_.resize(max + 1);
    N_[0][""] = 0;  // 空字符串计数

    // 对每个文本的每个位置,统计所有可能长度的子串
    for (const auto& text : segments) {
        for (size_t i = 0; i < text.length(); ++i) {
            for (size_t j = 0; j <= max; ++j) {
                if (i + j <= text.length()) {
                    std::string k = text.substr(i, j);
                    N_[j][k] += 1;  // 长度为j的子串k的计数+1
                }
            }
        }
    }
}

统计示例

文本:"南京市长江大桥"
UTF-8 字节序列:[E5,8D,97,E4,BA,AC,E5,B8,82,E9,95,BF,E6,B1,9F,E5,A4,A7,E6,A1,A5]
总字节数:21

填充N_数组:
N_[0]: {"": 21}  # 在21个字节的文本中,空字符串在每个位置都出现一次

N_[1] (1-Gram/单字节):
{E5:3, 8D:1, 97:1, E4:1, BA:1, AC:1, B8:1, 82:1, E9:1, 95:1, BF:1,
 E6:2, B1:1, 9F:1, A4:1, A7:1, A1:1, A5:1}

N_[2] (2-Gram/字节对):
{E58D:1, 8D97:1, 97E4:1, E4BA:1, BAAC:1, ACE5:1, E5B8:1, B882:1,
 82E9:1, E995:1, 95BF:1, BFE6:1, E6B1:1, B19F:1, 9FE5:1, E5A4:1,
 A4A7:1, A7E6:1, E6A1:1, A1A5:1}

N_[3](3-Gram,以下只列从字符边界开始的部分):
{E58D97:1, E4BAAC:1, E5B882:1, E995BF:1, E6B19F:1, E5A4A7:1, E6A1A5:1}
# 这些条目对应"南","京","市","长","江","大","桥"
# 实际统计还包含8D97E4、97E4BA等从字符内部开始的滑动窗口

N_[4](4-Gram,节选):
{E58D97E4:1, E4BAACE5:1, E5B882E9:1, E995BFE6:1, E6B19FE5:1, E5A4A7E6:1}

N_[5](5-Gram,节选):
{E58D97E4BA:1, E4BAACE5B8:1, E5B882E995:1, E995BFE6B1:1, E6B19FE5A4:1}

N_[6](6-Gram,节选):
{E58D97E4BAAC:1, E4BAACE5B882:1, E5B882E995BF:1, E995BFE6B19F:1, E6B19FE5A4A7:1}
# 这里列出的条目对应字符对"南京","京市","市长","长江","江大","大桥"
# 实际统计同样包含从字符内部开始的6字节滑动窗口

实际的 StreamCountRaw 会分批读取语料,先经过 PreTokenizer,再由多个线程执行上面的统计核心。得到的是出现次数,不是概率;随后用相邻阶的计数比估计条件概率:

目标:计算 $P(C\mid AB)=P(ABC)/P(AB)$

对数形式:$\log P(C\mid AB)=\log P(ABC)-\log P(AB)$

void PruneRaw() {
    // 为语料中未出现的字节加入回退伪计数
    for (int i = 0; i < 256; ++i) {
        std::string byte_str(1, static_cast<char>(i));
        if (N_[1].find(byte_str) == N_[1].end()) {
            N_[1][byte_str] = 1;
            N_[0][""] += 1;
        }
    }

    // 从最长 N-gram 开始向下处理
    for (int i = N_.size() - 1; i >= 0; --i) {
        std::unordered_map<std::string, float_t> pruned;

        // 1. 频率过滤 + 对数概率转换
        for (const auto& [k, v] : N_[i]) {
            if (k.length() == i &&
                v >= (i > 1 ? counter_spec_.min_count() : 0)) {
                pruned[k] = std::log(v);  // 先保存对数计数
            }
        }

        // 2. 计算条件概率
        if (i < N_.size() - 1) {
            std::unordered_map<std::string, float_t> next_pruned;
            for (const auto& [k, v] : N_[i + 1]) {
                std::string prefix = k.substr(0, i);  // 前i个字符
                auto it = pruned.find(prefix);
                if (it != pruned.end()) {
                    // log P(k|prefix) = log P(k) - log P(prefix)
                    next_pruned[k] = v - it->second;
                }
            }
            N_[i + 1] = std::move(next_pruned);
        }

        N_[i] = std::move(pruned);
    }
}

结果示例

修剪后的N_数组(对数概率形式):

N_[1]: 包含log P(byte)
{E5: log(3/259), E4: log(1/259), E6: log(2/259), E9: log(1/259), ...}
# 本例有18种已出现字节,另外238种字节各加入1次伪计数,因此分母为21+238=259

N_[2]: 包含log P(byte₂|byte₁)
{E58D: log P(8D|E5), 8D97: log P(97|8D), ...}

N_[3]: 包含log P(byte₃|byte₁byte₂)
{E58D97: log P(97|E58D), E4BAAC: log P(AC|E4BA), ...}
# 这一层特别重要,对应完整 UTF-8 字符的条件概率

N_[4]: 包含log P(byte₄|byte₁byte₂byte₃)
...

标注

状态空间

这里的状态空间服务于模型训练,与最终 BytePieceTokenizer 使用的词图动态规划不同:

BytePieceCounter 的状态

  • 共有 max 个状态:0, 1, 2, …, max-1
  • 状态 j < max-1 表示当前 token 已连续包含 j+1 个字节
  • 状态 max-1 是饱和状态,表示当前 token 至少包含 max 个字节
  • 从任意状态转移到状态 0,表示在当前字节之前切分,并以当前字节开始新 token

状态转移

void InitT() {
    int num_ = max;
    T_.resize(num_, std::vector<float_t>(num_, -INF));

    for (int i = 0; i < num_; ++i) {
        // 转移到状态0:在下一个字节开始新token
        T_[i][0] = 0;

        // 转移到状态i+1:当前token继续增长
        if (i + 1 < num_) {
            T_[i][i + 1] = 0;
        }

        // 最高状态可以自环(保持最大长度)★ 关键设计
        if (i == num_ - 1) {
            T_[i][i] = 0;
        }
    }
}

转移规则解释

  • T[i][0] = 0:在下一个字节开始新 token
  • T[i][i+1] = 0:当前 token 继续增长
  • T[max-1][max-1] = 0:在饱和状态使用滑动 N-gram 继续给长 token 评分

自环机制:支持任意长度 piece

数学表示: 设max = 6,则状态转移允许:

状态5 → 状态5 (自环)

这意味着 token 长度达到 6 字节后可以保持在状态 5,并使用长度为 6 的滑动窗口继续评分。

概率计算公式: 对于字节序列 \(x_1,\ldots,x_L\),当 \(L\ge 6\) 时:

P(piece) ≈ P(x₁)P(x₂|x₁)...P(x₆|x₁...x₅)
           × ∏ₜ₌₇ᴸ P(xₜ|xₜ₋₅...xₜ₋₁)

长度超过 6 后,每一步都用最近 6 个字节的计数除以其 5 字节前缀计数,估计下一字节的条件概率。

实际示例

假设piece = "ABCDEFGH"(长度8字节)

P(ABCDEFGH) = P(A)P(B|A)P(C|AB)P(D|ABC)P(E|ABCD)P(F|ABCDE)P(G|BCDEF)P(H|CDEFG)

分解过程:
1. 'A' → 初始状态0,使用N_[1]["A"]
2. 'B' → 状态0→状态1,使用N_[2]["AB"] - N_[1]["A"]
3. 'C' → 状态1→状态2,使用N_[3]["ABC"] - N_[2]["AB"]
4. 'D' → 状态2→状态3,使用N_[4]["ABCD"] - N_[3]["ABC"]
5. 'E' → 状态3→状态4,使用N_[5]["ABCDE"] - N_[4]["ABCD"]
6. 'F' → 状态4→状态5,使用N_[6]["ABCDEF"] - N_[5]["ABCDE"]
7. 'G' → 状态5→状态5,使用N_[6]["BCDEFG"] - N_[5]["BCDEF"] ★ 自环
8. 'H' → 状态5→状态5,使用N_[6]["CDEFGH"] - N_[5]["CDEFG"] ★ 自环

虽然 N-gram 统计只到 6-gram,但状态 5 的自环可以使用滑动窗口继续评估更长的 piece。piece 最终仍会受到 max_piece_size 的裁剪约束,默认不超过 18 字节。

状态转移表示例(max=6):

T矩阵(简化表示,0表示允许转移,-∞表示不允许):

    →  0  1  2  3  4  5
从 ↓
 0     0  0  -∞ -∞ -∞ -∞
 1     0  -∞ 0  -∞ -∞ -∞
 2     0  -∞ -∞ 0  -∞ -∞
 3     0  -∞ -∞ -∞ 0  -∞
 4     0  -∞ -∞ -∞ -∞ 0
 5     0  -∞ -∞ -∞ -∞ 0  ★ 自环允许无限增长

具体实现

BytePieceCounter 的核心是一个动态规划算法:在字节级统计的基础上实现字符级切分,因此需要引入 UTF-8 边界约束。

UTF-8 位置预处理

虽然 N-gram 统计是字节级的,但最终切分必须保持字符完整。算法首先检测 UTF-8 字符边界:

// UTF-8 位置预处理:标记每个字节在 UTF-8 字符中的位置
std::vector<int> utf8_position(num, 0);
int i = 0;
while (i < num) {
    int char_length = ustr::OneUTF8Size(text.data() + i);

    // 标记 UTF-8 字符的每个字节位置
    for (int j = 0; j < char_length && i + j < num; ++j) {
        utf8_position[i + j] = j;  // 0=首字节, 1=第二字节, 2=第三字节
    }
    i += char_length;
}

UTF-8 位置标记示例

文本:"南京"
字节:[E5, 8D, 97, E4, BA, AC]
位置: 0   1   2   3   4   5
utf8_position: [0, 1, 2, 0, 1, 2]
                ↑     ↑
              字符边界  字符边界

解释:
- 位置0,1,2:属于字符"南",分别是第1,2,3字节
- 位置3,4,5:属于字符"京",分别是第1,2,3字节
- 只有utf8_position[i]==0的位置是字符边界,可以作为切分点

UTF-8 约束

utf8_position[i] 记录字节 i 在当前 UTF-8 字符中的偏移。状态转移需要满足三条规则:

  1. 状态 0 只能出现在字符首字节,避免从字符内部开始 token。
  2. 当前状态和前一状态不能小于对应的字符内偏移,保证路径覆盖完整字符。
  3. 普通 N-gram 的起点必须位于字符边界;状态 5 自环使用滑动 6-gram 时不受此限制。

因此,piece 的起止位置始终落在 UTF-8 字符边界,而长 piece 内部的滑动窗口仍可跨越字符边界。

转移示例

仍以“南京”为例。状态从 0 开始,表示当前 token 已包含一个字节:

位置:       0   1   2   3   4   5
字节:      E5  8D  97  E4  BA  AC
字符内偏移: 0   1   2   0   1   2

处理“南”的三个字节时,合法路径依次为:

位置 0:状态 0,token 从 E5 开始
位置 1:状态 0 → 1,继续读入 8D
位置 2:状态 1 → 2,继续读入 97

位置 1 不能进入状态 0,否则 token 会从续字节 8D 开始;位置 2 也不能处于状态 0 或 1,因为这样的 token 无法覆盖“南”的完整三个字节。

来到位置 3,即“京”的首字节 E4,有两种合法选择:

状态 2 → 0:在“南”和“京”之间切分
状态 2 → 3:不切分,让当前 token 继续增长

第一条路径产生“南 / 京”,第二条路径则可能产生“南京”。两条路径都会进入动态规划,由累计 N-gram 得分决定最终结果。反过来,若一个普通窗口的起点落在 8D97 这样的续字节上,即使状态编号能够连接,也会被 ngram_start 检查排除。

当 token 达到 6 字节后,状态进入 5。若继续读取第三个汉字,状态 5 可以自环,此时 6-gram 窗口会向前滑动,起点允许落在字符内部;这只是评分窗口移动,token 本身仍从原来的字符边界开始。

动态规划框架

基于上述约束机制,动态规划算法的整体结构如下:

std::vector<std::string> Tokenize(const std::string& text) const {
    const int num = text.length();
    if (num == 0) return {};

    // 1. UTF-8 位置预处理(已完成)
    std::vector<int> utf8_position = PreprocessUTF8(text);

    // 2. 节点评分矩阵:scores[i][j] = 在字节位置i处于状态j的得分
    std::vector<std::vector<float_t>> scores(num,
        std::vector<float_t>(max, -INF));

    // 3. 路径记录矩阵
    std::vector<std::vector<int>> routes(num - 1,
        std::vector<int>(max, 0));

核心思想:寻找一条穿越状态空间的最优路径,使得总概率最大化。

节点评分填充

// 3. 填充节点评分(基于 N-gram 统计)
    for (int j = 0; j < max; ++j) {
        for (int i = j; i < num; ++i) {
            // 状态 0 只能从 UTF-8 字符边界开始
            if (j == 0 && utf8_position[i] > 0) continue;

            std::string piece = text.substr(i - j, j + 1);
            if (j + 1 < N_.size()) {
                auto it = N_[j + 1].find(piece);
                if (it != N_[j + 1].end()) {
                    scores[i][j] = it->second;  // 使用N-gram 概率
                }
            }
        }
    }

这一步提取可用的 N-gram 得分,并排除从 UTF-8 字符内部开始的新 token。更完整的状态合法性在随后的转移阶段检查。

动态规划状态转移

关键是过滤掉不合理的转移 (某些状态不需要转移以及某些状态之间不能转移)。

// 4. 动态规划核心:寻找最优路径
    for (int i = 1; i < num; ++i) {
        for (int curr_j = 0; curr_j < max; ++curr_j) {
            // 当前状态的 UTF-8 约束检查
            if (curr_j < utf8_position[i]) continue;

            int best_prev_j = -1;
            float_t best_score = -INF;

            for (int prev_j = 0; prev_j < max; ++prev_j) {
                // 前一位置的 UTF-8 约束
                if (prev_j < utf8_position[i-1]) continue;

                // 状态转移约束(基于T矩阵)
                if (T_[prev_j][curr_j] == -INF) continue;

                // 普通窗口必须从字符边界开始;饱和状态自环除外
                bool sliding = prev_j == max - 1 && curr_j == max - 1;
                int ngram_start = i - curr_j;
                if (!sliding && ngram_start > 0 &&
                    utf8_position[ngram_start] > 0) continue;

                // 计算转移得分
                float_t score = scores[i-1][prev_j] + T_[prev_j][curr_j] + scores[i][curr_j];

                if (score > best_score) {
                    best_score = score;
                    best_prev_j = prev_j;
                }
            }

            if (best_prev_j != -1) {
                routes[i-1][curr_j] = best_prev_j;
                scores[i][curr_j] = best_score;
            } else {
                scores[i][curr_j] = -INF;  // 无有效转移路径
            }
        }
    }

curr_jprev_j 的检查保证状态覆盖当前字符已经经过的字节数;ngram_start 的检查则让非滑动窗口从字符边界开始。这两类约束分别负责路径合法性和候选 piece 的起点合法性。

最优路径回溯

// 5. 找到最后位置的最佳状态
    int best_last_state = 0;
    float_t best_score = -INF;
    for (int j = 0; j < max; ++j) {
        if (j >= utf8_position[num - 1] && scores[num - 1][j] > best_score) {
            best_score = scores[num - 1][j];
            best_last_state = j;
        }
    }

    // 6. 回溯构建最优路径
    std::vector<int> opt_route(num);
    int curr_pos = num - 1;
    int curr_state = best_last_state;

    while (curr_pos >= 0) {
        opt_route[curr_pos] = curr_state;
        if (curr_pos > 0) {
            curr_state = routes[curr_pos-1][curr_state];
            curr_pos--;
        } else {
            break;
        }
    }

    // 7. 根据路径提取tokens
    std::vector<int> split_points;
    split_points.push_back(0);

    for (int i = 1; i < opt_route.size(); ++i) {
        // 只在 UTF-8 首字节处切分
        if (opt_route[i] == 0 && utf8_position[i] == 0) {
            split_points.push_back(i);
        }
    }
    split_points.push_back(num);

    // 8. 构建最终token序列
    std::vector<std::string> tokens;
    for (size_t i = 0; i < split_points.size() - 1; ++i) {
        tokens.push_back(text.substr(split_points[i],
                                   split_points[i + 1] - split_points[i]));
    }

    return tokens;
}

裁剪

第一次标注会产生大量候选 piece。PrunePieces 先按 max_piece_sizemin_count 分成保留集合与裁剪集合,再把被裁剪 piece 的计数重新分配给保留词表:

for (const auto& [piece, count] : pieces) {
    if (piece.length() <= counter_spec_.max_piece_size() &&
        count >= counter_spec_.min_count()) {
        keep[piece] = count;
    } else {
        drop[piece] = count;
    }
}

for (const auto& [piece, count] : SplitPieces(keep, drop)) {
    keep[piece] += count;
}

SplitPieces 用保留集合构建临时 BytePieceTokenizer。构造函数会把计数归一化为对数概率,然后用最大概率路径重新切分待裁剪内容:

Str2Int SplitPieces(const Str2Int& keep, const Str2Int& drop) {
    std::unordered_map<std::string, float_t> dict;
    for (const auto& [piece, count] : keep) {
        dict.emplace(piece, static_cast<float_t>(count));
    }
    BytePieceTokenizer tokenizer(dict);

    Str2Int counter;
    for (const auto& [piece, count] : drop) {
        for (const auto& token : tokenizer.Tokenize(piece)) {
            counter[token] += count;
        }
    }
    return counter;
}

随后,算法反复用当前词表切分其全部 piece,直到词表大小不再变化。最后若候选数仍超过目标 vocab_size,则排序截断:单字节候选优先,其余候选主要按计数降序排列。这里的单字节候选是普通候选;真正保证任意输入可编码的是模型初始化时单独加入的 256 个 BYTE 类型元词条。

计数再分配示例

假设keep = {"南京", "市", "长江", "大桥"}
     drop = {"南京市", "市长江", "长江大桥"}

重分词过程:
"南京市" → tokenizer.Tokenize("南京市") → ["南京", "市"]
"市长江" → tokenizer.Tokenize("市长江") → ["市", "长江"]
"长江大桥" → tokenizer.Tokenize("长江大桥") → ["长江", "大桥"]

结果统计:
counter = {"南京":1, "市":2, "长江":2, "大桥":1}

最终更新:
keep["南京"] += 1  # 原频率 + 重分词贡献
keep["市"] += 2
keep["长江"] += 2
keep["大桥"] += 1

示例

下面用“南京”串联统计、标注和裁剪过程。为便于展示,令 max = 3,状态 0、1、2 分别表示当前 token 已包含 1、2、至少 3 个字节。状态 2 是饱和状态,可以通过自环继续扩展 token。

统计结果

“南京”的 UTF-8 字节序列为:

位置:  0   1   2   3   4   5
字节: E5  8D  97  E4  BA  AC
边界:  ✓           ✓

假设语料统计得到以下对数条件概率:

log P(E5)          = -1.0
log P(8D | E5)     = -0.2
log P(97 | E5 8D)  = -0.2
log P(E4)          = -1.1
log P(BA | E4)     = -0.2
log P(AC | E4 BA)  = -0.2

log P(E4 | 8D 97)  = -0.1
log P(BA | 97 E4)  = -0.1

前三项描述“南”,中间三项可以独立描述“京”,最后两项用于状态 2 自环后跨越两个字符的滑动窗口。窗口可以从字符内部开始,但 token 只能在位置 0、3、6 切分。

比较候选路径

路径一将两个汉字分别作为 piece:

切分:  南 | 京
状态:  0 1 2 | 0 1 2
得分:(-1.0 - 0.2 - 0.2) + (-1.1 - 0.2 - 0.2)
     = -2.9

路径二让状态 2 保持自环,将“南京”作为一个 piece:

切分:  南京
状态:  0 1 2 2 2 2
窗口: E5 8D 97
       8D 97 E4
       97 E4 BA
       E4 BA AC
得分:-1.0 - 0.2 - 0.2 - 0.1 - 0.1 - 0.2
     = -1.8

因为 -1.8 > -2.9,动态规划选择“南京”。回溯得到的状态序列是:

位置: 0 1 2 3 4 5
状态: 0 1 2 2 2 2

序列中没有再次出现状态 0,因此位置 0 到文本末尾构成一个完整 token。

裁剪与计数再分配

假设对小型语料完成标注后得到:

南京    8
京      4
市      5
南京市  1
京市    1

min_count = 2,则“南京市”和“京市”进入待裁剪集合。临时 BytePieceTokenizer 使用保留词表重新切分它们:

南京市 → 南京 / 市
京市   → 京 / 市

相应计数被转移到仍然保留的 piece。重分词完成后,模型保存 piece 及其计数权重;BytePieceTokenizer 加载词表时再进行归一化:

P(piece) = count(piece) / Z
Z = Σ count(piece)

此外,模型会单独保存 256 个 BYTE 类型元词条。普通词表未覆盖某段输入时,编码过程会回退到这些字节词条,从而保证任意字节序列都可表示。至此,字节 N-gram 统计、最优路径标注、词表裁剪和字节回退形成了完整闭环。

配套实现:Ismantic/PieceTokenizer

番外篇:词向量 W2V

模型定义

W2V 是一种高效的词向量学习模型,它能够把词汇映射到低维稠密向量空间中,使语义相似的词在向量空间中距离较近。W2V 包含 Skip-Gram 和 CBOW 两种模型架构,以及负采样和分层 Softmax 两种常用训练方法。本文专注于 CBOW 与基于霍夫曼树的分层 Softmax。

CBOW

CBOW 模型的核心思想是:通过上下文词汇预测中心词。给定一个词序列,CBOW 将目标词周围的上下文词作为输入,预测中间的目标词。

具体来说,对于句子中的词 \(w_t\) ,我们使用其前后各 \(c\) 个词 \({w_{t-c}, …, w_{t-1}, w_{t+1}, …, w_{t+c}}\) 作为上下文,来预测 \(w_t\) 。

Softmax

标准的神经网络语言模型中,输出层通常使用 Softmax 函数来计算词汇表中每个词的概率:

$$ P(w_o|context) = \frac{\exp(v_{w_o}^T \cdot v_c)} {\sum_{w=1}^{V} \exp(v_w^T \cdot v_c)} $$

其中:

  • \(v_{w_o}\) 是输出词 \(w_o\) 的向量表示
  • \(v_c\) 是由上下文得到的隐藏层输出
  • \(V\) 是词汇表大小

这种方法的计算时间复杂度为 \(O(V)\),当词汇表包含数十万甚至数百万个词时,计算成本变得极其昂贵。

Context

CBOW 模型中,给定上下文词集合 \(C = {w_{c_1}, w_{c_2}, …, w_{c_t}}\),隐藏层输出通过平均上下文词向量得到:

$$ v_c = \frac{1}{|C|} \sum_{c \in C} v_{w_c} $$

其中 \(v_{w_c}\) 是上下文词 \(w_c\) 的输入向量表示。

霍夫曼树

W2V引入了分层Softmax (Hierarchical Softmax) 技术,通过霍夫曼树把计算复杂度由 \(O(V)\) 降低到 \(O(\log V)\)。

霍夫曼树的定义

霍夫曼树是一种二叉树,具有以下性质

  • 叶子节点代表词汇表中的词,全部叶子节点对应词汇表中的全部词
  • 根节点到叶子节点的固定路径能一一对应到具体的词
  • 词的概率计算能转化为沿着这个路径的决策序列
  • 内部节点对应着进行一次二分类决策

霍夫曼树的Softmax计算

词 \(w\) 对应一个叶子节点。设从根节点到该叶子节点的路径依次为 \(n_1,n_2,\ldots,n_{L(w)}\),其中 \(L(w)\) 表示路径包含的节点数。因此,这条路径包含 \(L(w)-1\) 条边,也就是 \(L(w)-1\) 次二分类决策。除最后的叶子节点 \(n_{L(w)}\) 外,每个内部节点 \(n_j\) 都有一个参数向量 \(\theta_{n_j}\)。

概率计算公式 $$ P(w|context) = \prod_{j=1}^{L(w)-1} \sigma(I(n_j, n_{j+1}) \cdot \theta_{n_j}^T \cdot v_c) $$

其中:

  • \(\sigma(x) = \frac{1}{1+\exp(-x)}\) 是 Sigmoid 函数
  • \(I(n_j,n_{j+1})\) 是路径指示函数(左分支为1,右分支为-1)
  • \(v_c\) 是隐藏层输出

注:最后参与计算的是内部节点 \(n_{L(w)-1}\),叶子节点 \(n_{L(w)}\) 不包含参数,也不会参与点积计算。

内部节点的物理意义

霍夫曼树中的每个内部节点实际上在进行一次二分类决策

  • \(\sigma(\theta_{n_j}^T v_c)\) 表示选择左分支的概率
  • \(1 - \sigma(\theta_{n_j}^T v_c) = \sigma(-\theta_{n_j}^T v_c)\) 表示选择右分支的概率

因此,到达目标词 \(w\) 的概率就是沿路径进行所有正确决策的概率乘积。

时间复杂度分析

传统Softmax的复杂度

  • 需要计算所有 \(V\) 个词的 \(\exp(v_w^T \cdot v_c)\)
  • 需要对所有 \(V\) 个值求和作为分母
  • 总计算量:\(O(V \cdot m)\),其中 \(m\) 是向量维度

霍夫曼树Softmax的复杂度

  • 只需要沿着一条路径计算 \(L(w)-1\) 个 Sigmoid 函数
  • 霍夫曼树不保证每个词的路径长度都是 \(O(\log V)\);极端情况下,个别低频词的路径可能很长
  • 霍夫曼编码最小化的是按照词频加权的平均路径长度,其平均编码长度接近词频分布的信息熵
  • 对常见的词频分布,平均路径长度通常可近似看作 \(O(\log V)\),因此平均计算量通常写作 \(O(\log V \cdot m)\)

霍夫曼树的构建

  1. 初始化:将词汇表中的每个词作为叶子节点,节点权重设为词频
  2. 迭代合并:重复选择两个权重最小的节点进行合并,新节点权重为两个子节点权重之和
  3. 路径编码:为每条从根节点到叶子节点的路径分配二进制编码
    • 左分支编码为1,右分支编码为-1
    • 高频词自然获得较短的路径,低频词获得较长的路径

这种构建方式会让高频词更晚加入到树中,距离根更近,确保了计算效率的最优化:高频词由于路径短,计算快速;低频词虽然路径长,但由于出现频率低,总体上仍然高效。

以下给出一个示例:

基本设定

  • 词汇表{A, B, C, D}
  • 词频统计A:8, B:6, C:3, D:1
  • 编码规则
    • 左分支 = 1(正类)
    • 右分支 = -1(负类)

霍夫曼树构建全流程

初始化节点

节点频次类型
n1A8叶子节点
n2B6叶子节点
n3C3叶子节点
n4D1叶子节点

迭代合并过程

  1. 第一轮合并
    合并最低频的 C(3)D(1) → 创建内部节点 Node2(4)
Node2 (频次=4)
/ \
C D
  1. 第二轮合并
    合并 Node2(4)B(6) → 创建内部节点 Node1(10)
 Node1 (频次=10)
 /   \
B     Node2
    /   \
   C     D
  1. 最终合并
    合并 Node1(10)A(8) → 创建根节点 Root(18)
    Root (频次=18)
   /    \
  A     Node1
       /    \
      B     Node2
           /    \
          C      D

完整路径编码表

路径编码序列决策次数计算示例
ARoot → A左(1)1σ(θ_Root·v_c)
BRoot → Node1 → B右(-1)→左(1)2σ(-θ_Root·v_c) × σ(θ_Node1·v_c)
CRoot → Node1 → Node2 → C右→右→左3σ(-θ_Root·v_c) × σ(-θ_Node1·v_c) × σ(θ_Node2·v_c)
DRoot → Node1 → Node2 → D右→右→右3σ(-θ_Root·v_c) × σ(-θ_Node1·v_c) × σ(-θ_Node2·v_c)

霍夫曼树的证明

接下来还要证明以下公式,是霍夫曼树Softmax能取代标准Softmax的基本要求:

$$ \begin{aligned} \sum_{w=1}^V P(w|context) &= \sum_{w=1}^V \prod_{j=1}^{L(w)-1} \sigma(I(n_j, n_{j+1}) \cdot \theta_{n_j}^T v_c) \\ &= 1 \end{aligned} $$

基本定义

霍夫曼树结构

  • 内部节点:每个内部节点 \(n\) 包含参数向量 \(\theta_n\),用于二元决策
  • 叶子节点:代表词汇表中的单词 \(w\),总数为 \(V\)

概率计算 对于隐藏层输出 \(h\) 和内部节点 \(n\):

  • 向左子节点移动概率:

$$ p(\text{left}|n) = \sigma(\theta_n \cdot v_c) = \frac{1}{1+e^{-\theta_n \cdot v_c}} $$

  • 向右子节点移动概率:

$$ p(\text{right}|n) = \sigma(-\theta_n \cdot v_c) = 1 - \sigma(\theta_n \cdot v_c) $$

其中 \(\sigma(x)\) 为 sigmoid 函数,满足: $$ \sigma(x) + \sigma(-x) = 1 $$

归纳法证明

基础证明(树高度=1) 单层霍夫曼树含1个内部节点和2个叶子节点:

$$ \begin{aligned} p(w_1) &= \sigma(\theta \cdot v_c) \\ p(w_2) &= \sigma(-\theta \cdot v_c) \\ \sum_{i=1}^2 p(w_i) &= \sigma(\theta \cdot v_c) + \sigma(-\theta \cdot v_c) = 1 \end{aligned} $$

归纳假设 假设对于高度 \(=k\) 的霍夫曼树,所有叶子节点概率和为1:

$$ \sum_{w \in \text{Leaves}_k} p(w) = 1 $$

归纳步骤(高度=k+1) 考虑根节点及其左右子树:

  1. 根节点决策概率:

$$ \begin{cases} p_L = \sigma(\theta_{root} \cdot v_c) \\ p_R = \sigma(-\theta_{root} \cdot v_c) \end{cases} $$

  1. 根据归纳假设:
    • 左子树叶子概率和 \(= 1\) → 贡献 \(p_L \times 1\)
    • 右子树叶子概率和 \(= 1\) → 贡献 \(p_R \times 1\)
  2. 整体概率和:

$$ \begin{aligned} p_L + p_R &= \sigma(\theta_{root} \cdot v_c) + \sigma(-\theta_{root} \cdot v_c) \\ &= 1 \end{aligned} $$

递归性质证明

这里需要区分两种概率:从整棵树的根节点出发计算的全局概率,以及已经到达某个子树根节点 \(n\) 后的条件概率。对于以 \(n\) 为根的任意子树,都有:

$$ \sum_{w \in \operatorname{Leaves}(n)} P(w\mid n,context)=1 $$

如果 \(n\) 是内部节点,设其左右孩子分别为 \(n_L\) 和 \(n_R\),那么:

$$ \begin{aligned} \sum_{w \in \operatorname{Leaves}(n)}P(w\mid n,context) &=P(n_L\mid n,context) \sum_{w\in\operatorname{Leaves}(n_L)}P(w\mid n_L,context) \\ &\quad+P(n_R\mid n,context) \sum_{w\in\operatorname{Leaves}(n_R)}P(w\mid n_R,context) \\ &=P(n_L\mid n,context)+P(n_R\mid n,context) \\ &=1 \end{aligned} $$

从根节点应用这一递归关系,就能得到所有叶子节点的全局概率之和为 1。相应地,如果使用从整棵树根节点出发的全局概率,那么某个子树中所有叶子的概率之和等于到达该子树根节点的概率,而不一定等于 1。

示例验证

3个叶子节点的霍夫曼树:

   Root
  /    \
 A      w3
/  \
w1 w2

概率计算:

$$ \begin{aligned} p(w_1) &= p(A) \times p(\text{left}|A) \\ p(w_2) &= p(A) \times p(\text{right}|A) \\ p(w_3) &= p(\text{right}|Root) \\ \sum_{i=1}^3 p(w_i) &= p(A)[p(\text{left}|A)+p(\text{right}|A)] + p(w_3) \\ &= p(A)\times 1 + p(w_3) \\ &= \sigma(\theta_{Root}\cdot v_c) + \sigma(-\theta_{Root}\cdot v_c) \\ &= 1 \end{aligned} $$

关键结论

通过以下性质保证归一化:

  1. 局部归一化: \(\forall n,\ p(\text{left}|n)+p(\text{right}|n)=1\)
  2. 递归累乘: \(p(w) = \prod_{\text{path to }w} p(\text{branch})\)
  3. 树结构完备性:每个样本必被分配到唯一叶子节点

因此分层 Softmax 满足:

$$ \sum_{w=1}^V p(w|v_c) = 1 $$

目标函数

对于训练样本 \((context, w)\),我们希望最大化条件概率 \(P(w|context)\)。采用最大似然估计,目标函数为:

$$ \mathcal{L} = \log P(w|context) = \log \prod_{j=1}^{L(w)-1} \sigma(I(n_j, n_{j+1}) \cdot \theta_{n_j}^T v_c) $$

$$ \mathcal{L} = \sum_{j=1}^{L(w)-1} \log \sigma(I(n_j, n_{j+1}) \cdot \theta_{n_j}^T v_c) $$

对于整个训练语料库,总的目标函数为:

$$ \mathcal{L}{\mathrm{total}} = \sum{(context, w) \in D} \sum_{j=1}^{L(w)-1} \log \sigma(I(n_j, n_{j+1}) \cdot \theta_{n_j}^T v_{\mathrm{context}}) $$

其中 \(D\) 是训练数据集, \(v_{context}\) 是对应上下文的隐藏层输出。

梯度推导

CBOW + 霍夫曼树模型中,需要推导的参数包括:

  • 输入词向量 \(v_{w_c}\) :每个词作为上下文时的向量表示
  • 霍夫曼树节点向量 \(\theta_{n_j}\) : 霍夫曼树内部节点的参数向量

对霍夫曼树节点参数的梯度

对于路径上的节点 \(n_j\),我们需要计算目标函数对 \(\theta_{n_j}\) 的梯度。

单个样本的梯度计算

$$ \frac{\partial \mathcal{L}}{\partial \theta_{n_j}} = \frac{\partial}{\partial \theta_{n_j}} \log \sigma(I(n_j, n_{j+1}) \cdot \theta_{n_j}^T \cdot v_c) $$

由Sigmoid函数的导数性质 \(\frac{d}{dx}\log \sigma(x) = 1 - \sigma(x)\):

$$ \frac{\partial}{\partial \theta_{n_j}} \log \sigma(I(n_j, n_{j+1}) \cdot \theta_{n_j}^T \cdot v_c) = [1 - \sigma(I(n_j, n_{j+1}) \cdot \theta_{n_j}^T v_c)] \cdot I(n_j, n_{j+1}) \cdot v_c $$

梯度的物理意义

  • 当 \(\sigma(I(n_j, n_{j+1}) \cdot \theta_{n_j}^T \cdot v_c) \approx 1\) 时(预测正确),梯度较小
  • 当 \(\sigma(I(n_j, n_{j+1}) \cdot \theta_{n_j}^T \cdot v_c) \approx 0\) 时(预测错误),梯度较大

对隐藏层输出的梯度

隐藏层输出 \(v_c\) 连接到路径上的所有节点,因此其梯度是所有节点梯度的累加:

$$ \frac{\partial \mathcal{L}}{\partial v_c} = \sum_{j=1}^{L(w)-1} \frac{\partial}{\partial v_c} \log \sigma(I(n_j, n_{j+1}) \cdot \theta_{n_j}^T \cdot v_c) $$

$$ \frac{\partial \mathcal{L}}{\partial v_c} = \sum_{j=1}^{L(w)-1} [1 - \sigma(I(n_j, n_{j+1}) \cdot \theta_{n_j}^T \cdot v_c)] \cdot I(n_j, n_{j+1}) \cdot \theta_{n_j} $$

对输入词向量的梯度

由于隐藏层输出是上下文词向量的平均:

$$ v_c = \frac{1}{|C|} \sum_{c \in C} v_{w_c} $$

因此:

$$ \frac{\partial v_c}{\partial v_{w_c}} = \frac{1}{|C|} $$

由链式法则,对每个上下文词向量的梯度为:

$$ \frac{\partial \mathcal{L}}{\partial v_{w_c}} = \frac{\partial \mathcal{L}}{\partial v_c} \cdot \frac{\partial v_c}{\partial v_{w_c}} = \frac{1}{|C|} \cdot \frac{\partial \mathcal{L}}{\partial v_c} $$

完整的梯度表达式

综合以上推导,完整的梯度表达式为:

霍夫曼树节点参数梯度: $$ \frac{\partial \mathcal{L}}{\partial \theta_{n_j}} = [1 - \sigma(I(n_j, n_{j+1}) \cdot \theta_{n_j}^T v_c)] \cdot I(n_j, n_{j+1}) \cdot v_c $$

上下文词向量梯度: $$ \frac{\partial \mathcal{L}}{\partial v_{w_c}} = \frac{1}{|C|} \sum_{j=1}^{L(w)-1} [1 - \sigma(I(n_j, n_{j+1}) \cdot \theta_{n_j}^T \cdot v_c)] \cdot I(n_j, n_{j+1}) \cdot \theta_{n_j} $$

参数更新

使用梯度上升法(因为我们要最大化似然函数)更新参数:

霍夫曼树节点参数更新: $$ \theta_{n_j} \leftarrow \theta_{n_j} {}+ \alpha \cdot [1 - \sigma(I(n_j, n_{j+1}) \cdot \theta_{n_j}^T v_c)] \cdot I(n_j, n_{j+1}) \cdot v_c $$

输入词向量更新: $$ v_{w_c} \leftarrow v_{w_c} {}+ \frac{\alpha}{|C|} \sum_{j=1}^{L(w)-1} [1 - \sigma(I(n_j, n_{j+1}) \cdot \theta_{n_j}^T v_c)] \cdot I(n_j, n_{j+1}) \cdot \theta_{n_j} $$

其中 \(\alpha\) 是学习率。

实际实现时,必须先使用更新前的内部节点向量 \(\theta_{n_j}\) 累加 \(\partial \mathcal{L}/\partial v_c\),再更新上下文词向量。内部节点参数也应当根据同一次前向计算得到的梯度更新。如果一边修改 \(\theta_{n_j}\),一边使用修改后的值继续计算输入词向量梯度,代码执行的就不再是上面推导出的同一次梯度更新。常见实现会先把隐藏层梯度完整累加到一个临时向量中,然后统一更新所有上下文词向量。

训练实现

Wavec 的 FastText::CBOW 对应上面的推导。syn0 保存输入词向量,syn1 保存霍夫曼树内部节点向量;neu1 是上下文平均向量,neu1e 累加隐藏层梯度:

for (int word : context) {
    for (int c = 0; c < vec_size; ++c) {
        neu1[c] += syn0[word * vec_size + c];
    }
}
for (int c = 0; c < vec_size; ++c) {
    neu1[c] /= context.size();
}

对路径上的每个内部节点,代码先用尚未更新的 syn1 累加 neu1e,再更新该节点参数。霍夫曼编码 code=0 对应 \(\sigma(f)\),code=1 对应 \(1-\sigma(f)\),所以实现中的

float g = (1 - code - p) * alpha;

与前文使用 \(I\in{1,-1}\) 的梯度表达式等价。全部路径节点处理完成后,neu1e / context.size() 被加到每个上下文词向量上,正好对应平均操作产生的 \(1/|C|\)。

训练循环还包含几项公式之外的工程策略:

  • 动态窗口:每个中心词从 1window 随机选择实际窗口大小,减少固定边界造成的偏差。
  • 高频词下采样:降低高频中心词进入训练的概率,让信息量更高的词获得更多有效更新。
  • 学习率衰减:学习率随每个线程的处理进度从 start_alpha 线性下降,并设置 0.0001 的下限。
  • 异步训练:各线程处理不同的文档区间,并共享词向量参数。这种 Hogwild 式更新省去了频繁的梯度合并与全局锁,适合参数规模大而单次更新稀疏的场景;代价是不同线程的更新顺序不固定,训练结果会存在轻微差异。

这些策略不改变 CBOW 与分层 Softmax 的目标函数,而是围绕样本选择、训练进度和参数更新方式改善实际训练效率。

配套实现:Ismantic/Wavec

番外篇:LDA 与 SparseLDA

主题定义

LDA(Latent Dirichlet Allocation,潜在狄利克雷分配)是一种经典的概率主题模型,用来发现文档集合中的潜在主题结构。它假设每篇文档由多个主题混合而成,而每个主题都是词汇表上的一个概率分布。

LDA 的生成过程:

  • 主题—词语分布:对每个主题 \(t\),从狄利克雷先验中采样词语分布 \(\phi_t \sim \operatorname{Dirichlet}(\beta)\)
  • 文档分布:对每篇文档 \(m\),从狄利克雷先验中采样主题分布 \(\theta_m \sim \operatorname{Dirichlet}(\alpha)\)
  • 生成词语:对文档中的每个位置 \(n\)
    • 根据 \(\theta_m\) 采样主题 \(z_{m,n}\)
    • 根据该主题的词语分布 \(\phi_{z_{m,n}}\) 采样词语 \(w_{m,n}\)

这里的 \(\theta_m\) 和 \(\phi_t\) 都是隐变量,不是人工预设的固定分布。训练 LDA 的过程,就是根据观察到的词语反推这些隐变量。

到底主题是什么

实际上,主题就是词汇表上的概率分布。每个主题代表一个语义相关的词语集合。

示例主题:

主题 0(科技):

P(w|z=0) = {
  "手机": 0.12,  "苹果": 0.08,  "电脑": 0.10,
  "软件": 0.07,  "网络": 0.09,  "数据": 0.06,
  "算法": 0.05,  "程序": 0.04,  "系统": 0.03,
  "应用": 0.04,  "技术": 0.05,  ...
}

主题 1(体育):

P(w|z=1) = {
  "足球": 0.15,  "比赛": 0.12,  "球员": 0.10,
  "进球": 0.08,  "教练": 0.07,  "联赛": 0.09,
  "冠军": 0.06,  "球队": 0.11,  "场地": 0.04,
  "训练": 0.05,  "战术": 0.03,  ...
}

主题 2(美食):

P(w|z=2) = {
  "苹果": 0.06,  "牛肉": 0.09,  "餐厅": 0.08,
  "味道": 0.10,  "烹饪": 0.07,  "食材": 0.11,
  "美味": 0.09,  "菜谱": 0.05,  "香甜": 0.04,
  "新鲜": 0.06,  "营养": 0.05,  ...
}

举个生成的例子

假设有这么一个文档,其主题分布为:

P(z|文档) = {主题0: 0.6, 主题1: 0.1, 主题2: 0.3}

过程演示:

假设有这么一个文档:

最新的手机配备了先进的苹果芯片,其内置软件能够高效处理海量数据。
通过智能算法和网络连接,这个程序运行流畅。
今天的比赛很精彩,现场还有美味的苹果和各种食物,味道很棒。

该文档包含的关键词是:["手机", "苹果", "软件", "数据", "算法", "网络", "程序", "比赛", "美味", "味道"]

那么这些关键词的生成过程大概是这样:

  1. 第0个词:采样主题 → 主题0(科技),采样词语 → “手机”
  2. 第1个词:采样主题 → 主题0(科技),采样词语 → “苹果”(指苹果公司/芯片)
  3. 第2个词:采样主题 → 主题0(科技),采样词语 → “软件”
  4. 第3个词:采样主题 → 主题0(科技),采样词语 → “数据”
  5. 第4个词:采样主题 → 主题0(科技),采样词语 → “算法”
  6. 第5个词:采样主题 → 主题0(科技),采样词语 → “网络”
  7. 第6个词:采样主题 → 主题0(科技),采样词语 → “程序”
  8. 第7个词:采样主题 → 主题1(体育),采样词语 → “比赛”
  9. 第8个词:采样主题 → 主题2(美食),采样词语 → “美味”
  10. 第9个词:采样主题 → 主题2(美食),采样词语 → “味道”

这个例子展示的是词语及其主题的生成关系,而不是一段有顺序的文本。LDA 是词袋模型,只建模词语在文档中的共现,不建模局部词序,因此不能像自回归语言模型一样生成连贯句子。

注意到,“苹果”可以同时出现在多个主题中。LDA 会利用整篇文档中的词语共现,为“苹果”的每一次出现推断一个主题,因此能够在一定程度上反映多义性。不过,LDA 是词袋模型,不使用局部词序,也不能保证完成严格的词义消歧。

关键公式

LDA 模型可以通过 Gibbs 采样来推断每个词语的主题分配。关键的采样公式为:

$$ P(z_{m,n}=t \mid \mathbf{z}{-m,n}, \mathbf{w}) \propto \frac{(\alpha + n{t|m}^{-m,n})(\beta + n_{w|t}^{-m,n})} {\sum_v(\beta + n_{v|t}^{-m,n})} $$

该公式表示:对于文档 m 中第 n 个词语,其主题分配为 t 的概率正比于右侧的表达式。

完整的文档—主题因子还包含分母 \(n_{\cdot|m}^{-m,n}+K\alpha\)。对同一个词语枚举候选主题 \(t\) 时,这个分母保持不变,因此可以在未归一化的 Gibbs 采样权重中省略。词语—主题因子的分母则随 \(t\) 变化,不能省略。

分项解释:

  • 第一部分 \((\alpha + n_{t|m}^{-m,n})\):文档-主题倾向

    • 表示去除当前词语后,文档 m 中分配给主题 t 的词语数量
    • \(\alpha\) 是平滑参数,避免零概率
  • 第二部分 \((\beta + n_{w|t}^{-m,n})\):词语-主题倾向

    • 表示去除当前词语后,全部语料中词语 w 被分配给主题 t 的次数
    • \(\beta\) 是平滑参数
  • 分母部分 \(\sum_v(\beta + n_{v|t}^{-m,n})\):归一化项

    • 表示主题 t 的总词语数量,用于归一化

直观理解: 假设当前词语是“苹果“,候选主题是“科技“:

  • 如果当前文档中已有很多词语被分配给“科技“主题(第一部分)
  • 且“苹果“在“科技“主题中出现概率高(第二部分)
  • 那么“苹果“被分配给“科技“主题的概率就很高

SparseLDA

标准 LDA 每次采样需要遍历全部主题,时间复杂度为 \(O(K)\),其中 \(K\) 是主题数。当主题数很大时,这会成为性能瓶颈。SparseLDA 通过将采样概率分解为三个“桶”来减少实际需要遍历的主题。

数学推导

由标准的 LDA 采样公式开始:

$$ P(z_{m,n}=t) \propto \frac{(\alpha+n_{t|m}^{-m,n})(\beta+n_{w|t}^{-m,n})} {\beta V+n_{\cdot|t}^{-m,n}} $$

对记号简化,定义:

  • \(n_{t|m}\) :文档 m 中分配给主题 t 的词数
  • \(n_{w|t}\) :词语 w 被分配给主题 t 的次数
  • \(n_{·|t}\) :主题 t 的总词数

第一步:展开分子

把分子展开: $$ (\alpha+n_{t|m})(\beta+n_{w|t}) = \alpha\beta + \alpha n_{w|t} + n_{t|m}\beta + n_{t|m}n_{w|t} $$

第二步:组织成三个部分

把展开式重新组织: $$ \frac{\alpha\beta + \alpha n_{w|t} + n_{t|m}\beta + n_{t|m}n_{w|t}} {\beta V + n_{\cdot|t}} $$

$$ \frac{\alpha\beta}{\beta V + n_{\cdot|t}} {}+ \frac{n_{t|m}\beta}{\beta V + n_{\cdot|t}} {}+ \left( \frac{\alpha n_{w|t}}{\beta V + n_{\cdot|t}} {}+ \frac{n_{t|m}n_{w|t}}{\beta V + n_{\cdot|t}} \right) $$

第三步:得到三个桶

桶 s (平滑项): $$ s_t = \frac{\alpha\beta}{\beta V + n_{\cdot|t}} $$

桶 r (文档项): $$ r_t = \frac{n_{t|m}\beta}{\beta V + n_{\cdot|t}} $$

桶 q (词语项): $$ q_t = \frac{\alpha n_{w|t}}{\beta V + n_{\cdot|t}} {}+ \frac{n_{t|m}n_{w|t}}{\beta V + n_{\cdot|t}} = \frac{(\alpha + n_{t|m})n_{w|t}}{\beta V + n_{\cdot|t}} $$

平滑桶 \(s\) 仍然包含全部 \(K\) 个主题;文档桶 \(r\) 只包含当前文档中出现过的主题;词语桶 \(q\) 只包含当前词语被分配过的主题。后两个集合通常远小于 \(K\)。平滑桶的总质量可以维护起来,并且采样通常只需遍历实际落入的那个桶,因此平均开销会明显低于每次完整遍历全部主题。

具体算法

采样步骤:

  1. 计算桶的总质量: $$ S = \sum_t s_t {}+ \sum_{t:n_{t|m}>0} r_t {}+ \sum_{t:n_{w|t}>0} q_t $$

  2. 按比例选择桶:

    • 随机数生成 \(u \sim \text{Uniform}(0,S)\)
    • 若 \(u < S_s\), 选择桶 s
    • 若 \(S_s \leq u < S_s + S_r\),选择桶 r
    • 否则选择桶 q
  3. 桶内采样:

    • 根据选中的桶,对应的候选主题中概率采样

平滑桶的总质量 \(S_s=\sum_t s_t=\sum_t\frac{\alpha\beta}{\beta V+n_{\cdot|t}}\) 需要缓存,否则每个词语都重新求和仍要遍历全部主题。理论上,每次主题移动只改变旧、新两个主题的总计数,可以用 \(O(1)\) 增量精确更新;Semat 当前实现选择在每轮迭代开始时通过 UpdateCache 重新计算一次,轮内使用该缓存值,因此这里采用的是近似更新。

多线程

仅仅降低单次采样的时间复杂度还不够,还可以进一步利用多线程。观察 SparseLDA 的采样公式可以发现:合理划分文档和词汇块,能够消除大部分文档—主题计数与词语—主题计数的写冲突,但全局主题计数仍然需要单独处理。

回顾下采样公式: \(P(z_{m,n}=t) \propto \frac{(\alpha+n_{t|m})(\beta+n_{w|t})}{\beta V + n_{·|t}}\)

计数依赖分析:

  • 更新文档 m 中词语 w的主题时,会影响:
    • \(n_{t|m}\) : 文档m的主题计数
    • \(n_{w|t}\) : 词语w的主题计数
    • \(n_{·|t}\) : 主题的总计数

竞争条件问题: 若多个线程同时操作,会造成这些数据出现不一致:

  • 同一文档的不同词语:会竞争修改 \(n_{*|m}\)
  • 不同文档的相同词语:会竞争修改 \(n_{w|*}\)
  • 任意词语:都会竞争修改 \(n_{·|*}\) (这个竞争无法避免)

一种直接做法是为共享计数加锁,但锁竞争可能抵消并行带来的收益。

N-Queen 方案 类似于 N 皇后问题中皇后之间不能相互攻击的约束,让不同线程操作的文档块和词汇块互不相交,可以避免这两类局部计数之间的竞争。

核心思想:

  • 把文档和词汇表都分块
  • 不同线程负责不同的文档-词汇块组合
  • 通过巧妙的调度避免竞争

局部计数无冲突条件: 线程 i 处理文档块 \(D_i\) 和词汇块 \(V_i\),线程 j 处理文档块 \(D_j\) 和词汇块 \(V_j\)。满足以下条件时,文档—主题计数和词语—主题计数不会发生写冲突:

  • \(D_i \cap D_j = \varnothing\)(文档块不相交)
  • \(V_i \cap V_j = \varnothing\)(词汇块不相交)

这并不意味着整个算法已经完全无锁。Semat 的分块保证同一轮中的线程不会同时修改同一文档的 nm[m],也不会同时修改同一词语的 nv[w];但所有线程仍会修改共享的 nvsum[t]。当前代码没有为这些读写使用原子操作或锁,在 C++ 内存模型中属于数据竞争,而不仅是统计意义上的“陈旧计数”。若要得到定义明确的并行实现,应将 nvsum 改为原子计数,或者累计线程局部增量并在每轮调度后合并。

分块策略示例(4线程,8文档,8词汇):

第1轮:

        词0-1  词2-3  词4-5  词6-7
文档0-1  [T0]   ×     ×     ×
文档2-3   ×    [T1]   ×     ×  
文档4-5   ×     ×    [T2]   ×
文档6-7   ×     ×     ×    [T3]

第2轮(列循环右移):

        词0-1  词2-3  词4-5  词6-7
文档0-1   ×     ×     ×    [T0]
文档2-3  [T1]   ×     ×     ×
文档4-5   ×    [T2]   ×     ×
文档6-7   ×     ×    [T3]   ×

第3轮(继续右移):

        词0-1  词2-3  词4-5  词6-7
文档0-1   ×     ×    [T0]   ×
文档2-3   ×     ×     ×    [T1]
文档4-5  [T2]   ×     ×     ×
文档6-7   ×    [T3]   ×     ×

第4轮(完成覆盖):

        词0-1  词2-3  词4-5  词6-7
文档0-1   ×    [T0]   ×     ×
文档2-3   ×     ×    [T1]   ×
文档4-5   ×     ×     ×    [T2]
文档6-7  [T3]   ×     ×     ×

Perplexity

困惑度(Perplexity)用于衡量模型对未见文档的预测能力。标准评估应在留出的验证集或测试集上进行,而不是直接复用训练语料及其当前主题分配。

基本思想: LDA假设每个词语的生成分为两步:

  1. 根据文档的主题分布选择一个主题 t
  2. 根据主题 t 的词语分布选择一个词语 w

主题是隐变量,因此计算词语的边缘概率时,不能只使用当前分配到的一个主题,而要对全部主题求和。对于文档 \(m\) 中位置 \(i\) 的词语:

$$ P(w_{m,i}\mid m) = \sum_{t=1}^{K}P(t\mid m)P(w_{m,i}\mid t) $$

其中:

  • \(P(t|m) = \frac{n_{t|m}+\alpha}{n_m + K\alpha}\) : 文档 m 中主题 t 的概率
  • \(P(w|t) = \frac{n_{w|t}+\beta}{n_t + V\beta}\) : 主题 t 中词语 w 的概率

对数似然函数:

$$ \log P(\text{全部词语}) = \sum_{m=1}^{M}\sum_{i=1}^{N_m} \log\left( \sum_{t=1}^{K}P(t\mid m)P(w_{m,i}\mid t) \right) $$

困惑度(Perplexity) 定义:

$$ \text{Perplexity} = \exp\left( {}-\frac{\log P(\text{全部词语})}{\text{总词数}} \right) $$

其中总词数为测试集中所有文档的词语数量。困惑度可以直观理解为模型在每个位置面对的平均有效候选数;在相同数据集和评估方法下,数值越低,表示模型的预测能力越好。实际评估时,还需要为测试文档推断文档—主题分布,并避免使用待评估词语本身的信息。

训练过程中的近似指标

标准困惑度需要在测试文档上推断主题分布,并对全部主题求和,计算成本较高。训练过程中为了快速观察模型变化,也可以利用每个词语当前的主题分配,计算一个近似得分:

$$ \log P(z_{m,i}\mid m)+\log P(w_{m,i}\mid z_{m,i}) $$

将所有词语的得分相加,再按词数取负平均并求指数,就能得到一个随训练过程变化的指标。它利用的是训练语料及其当前主题分配,适合观察同一次训练是否趋于稳定;但它没有对隐含主题进行边缘化,也没有评估未见文档,因此不能代替前面定义的标准困惑度。

配套实现:Ismantic/Semat