用 C++ 模板元编程实现 Advent of Code
原文由 Nelson Elhage 于 发布,订阅该博客
今年十二月,我一时兴起,决定看看能用纯编译期的 C++ 元编程完成 Advent of Code 的多少天题目。
写这篇文章时,我已经完成了两天,不确定还能不能继续做下去。不过,这已经比我昨天的计划多做了一天,而昨天的计划又比我第一次尝试后以为自己能做到的要多。所以,走着瞧吧。
不过,第一天的题目已经足够有趣,值得写一篇短文分享。你可以在 GitHub 上找到代码,下面我也会对第一天的解法做一次带注释的详解!
基础类型
我们先定义一个基础的编译期 list 类型,还会用到 literal<v> 来把值提升到类型系统中。我决定不用 constexpr——那会让事情变得太简单——但允许自己使用基本的算术运算以及编译期的整数和布尔值。
template <typename... Elts>
struct list {};
template <auto v>
struct literal {
constexpr static auto value = v;
};前置准备:读取输入
理想情况下,为了追求极致的纯粹性,我会用新的 C++23 #embed 指令来读取输入;但就本文写作时而言,似乎很难找到支持它的编译器。因此,我退而求其次,用 xxd -i 做了一个 polyfill。我们可以执行 xxd -i < input.txt > input.i,然后通过 #include 和可变参数模板来引入输入。我们还会用一点预处理器技巧,通过 -D 编译选项来传入输入文件。
template <auto... Elt>
struct read_input {
using type = list<literal<char(Elt)>...>;
};
#define _QUOTED(x) #x
#define QUOTED(x) _QUOTED(x)
#ifndef INPUT
// Default value to keep my IDE environment happy
#define INPUT /dev/null
#endif
using problem = read_input<
#include QUOTED(INPUT)
>::type;折叠
C++ 模板元编程很大程度上就是一门函数式语言,因此我们在列表上定义的主要操作自然是 fold。这是一种相对标准的构造:
template <template<typename, typename> typename Fn,
typename Init, typename List>
struct fold {};
template <template<typename, typename> typename Fn,
typename Init>
struct fold<Fn, Init, list<>> {
using type = Init;
};
template <template<typename, typename> typename Fn,
typename Init, typename El1, typename... Elts>
struct fold<Fn, Init, list<El1, Elts...>> {
using type = fold<
Fn,
typename Fn<Init, El1>::type,
list<Elts...>>::type;
};更多辅助工具
我们再引入几个辅助工具,就差不多可以开始做第一部分了。我们会添加一个 nil 哨兵、一个基本的 if 条件分支,以及一个 or_else 函数——它会取参数中第一个非 nil 的值:
struct nil;
template <typename T>
struct is_nil { using type = literal<false>; };
template <>
struct is_nil<nil> { using type = literal<true>; };
template <typename Cond, typename Then, typename Else>
struct if_else {};
template <typename Then, typename Else>
struct if_else<literal<true>, Then, Else> { using type = Then; };
template <typename Then, typename Else>
struct if_else<literal<false>, Then, Else> { using type = Else; };
template <typename L, typename R>
struct or_else { using type = L; };
template <typename R>
struct or_else<nil, R> { using type = R; };第一部分
如果你没有做过 Advent of Code,这里快速回顾一下题目:对于输入文件中的每一行,我们需要找到该行中的第一个和最后一个数字,将它们拼成一个两位数(即“校准值”),然后输出文件中所有校准值之和。
为了效率,我们会用一次对输入的 fold 来实现它。在遍历过程中,我们会维护一个由 (目前为止的总和,当前行的第一个数字,当前行目前为止看到的最后一个数字) 组成的状态元组:
template <typename Accum, typename FirstDigit, typename LastDigit>
struct State {
using accum = Accum;
using first = FirstDigit;
using last = LastDigit;
};
using empty_state = State<literal<0>, nil, nil>;fold 函数通过一些简单的规则来推进状态:
- 如果字符是换行符,就计算这一行的校准值,加到累加器中,并重置按行维护的状态。
- 如果字符是数字,则无条件将“最后一个数字”更新为该数字
- 并且,如果这一行还没有出现过数字,也同时更新第一个数字。
这些都是对已有辅助工具相当直接的运用:
template <typename T>
using as_digit = if_else<
literal<T::value >= '0' && T::value <= '9'>,
literal<T::value - '0'>,
nil>;
template <typename In, typename Char>
struct Fn {
using digit = as_digit<Char>::type;
struct IfNewline {
using type = State<
literal<In::accum::value + 10 * In::first::value + In::last::value>,
nil,
nil>;
};
using type = if_else<
literal<Char::value == '\n'>,
IfNewline,
if_else<
typename is_nil<digit>::type,
In,
typename if_else<
typename is_nil<typename In::first>::type,
State<typename In::accum, digit, digit>,
State<typename In::accum, typename In::first, digit>
>::type
>>::type::type;
};其中一个小技巧是,为了让外层的 if_else 具有惰性求值,我们是对它的结果求 ::type,而不是对它的参数求值。这样编译器就只会对实际被选中的那个分支进行求值。
试一试
为了方便,我们允许一点点运行时行为来打印结果:
int main() {
using answer = solve<problem>::type;
printf("%d\n", answer::value);
}有了它,我们现在就可以验证题目中给出的初始测试用例:
$ xxd -i < test.txt > test.i && \
c++ -DINPUT=test.i -std=c++20 part1.cc -o out/part1 && \
out/part1
142但如果我们尝试用真实输入,就会遇到问题:
$ xxd -i < input.txt > input.i && \
c++ -DINPUT=input.i -std=c++20 part1.cc -o out/part1 && \
out/part1
part1.cc:8:1: fatal error: recursive template instantiation exceeded maximum depth of 1024
using as_digit = if_else<
[cut C++ template spew]
part1.cc:8:1: note: use -ftemplate-depth=N to increase recursive template instantiation depth这里的根本问题在于,我们的 fold 需要对输入中的每个字符都进行一次递归,而真实输入的长度超过了 20,000 个字符!我们可以按照提示尝试增大 -ftemplate-depth;在我的机器上,把这个值设得足够大确实能让程序编译通过,但 clang 会给出关于栈容量的可怕警告,而且编译非常慢——需要 90 秒以上。
我们需要换一种方法。
优化 fold
这里的问题相当根本;递归的模板展开基本上就是模板元编程中唯一可用的手段。不过,事实证明我们还有一招可用。
list<> 是一个可变参数模板,这意味着我们可以将其全部内容作为 C++ 的参数包来操作。只要能让参数保持为参数包的形式,我们就能利用编译器内部相对快速的路径,效率会高得多。
仔细阅读上面的 cppreference 页面,我们会注意到 C++17 加入了“折叠表达式”——一种将参数包展开为对参数包进行折叠运算的表达式的方法,所用的运算符可由用户自选!这和我们想做的事惊人地相似,只不过它工作在值层面,使用的是值层面的运算符。
不过,事实证明,只要巧妙地运用 auto、decltype 和一个包装结构体,这就不是问题——我们可以在对折叠表达式进行类型检查的过程中构造出执行类型层面计算的类型。
编译器仍然需要处理一个嵌套极深的表达式,但在我的测试中,这比 20,000 层深的模板要高效得多:
template<template<typename, typename> typename Fn>
struct fold_helper {
template <typename T>
struct F {
using type = T;
template <typename R>
auto operator<<(F<R>) {
return F<typename Fn<T, R>::type>{};
};
};
};
template <template<typename, typename> typename Fn,
typename Init, typename... Elts>
struct fold<Fn, Init, list<Elts...>> {
template <typename T>
using F = fold_helper<Fn>::template F<T>;
using type = decltype((F<Init>{} << ... << F<Elts>{}))::type;
};结构体 fold_helper<Fn>::F 将我们的类型层面函数转换成可以通过表达式层面运算符求值的形式。表达式 fold_helper<Fn>::F<A>{} << fold_helper<Fn>::F<B>{} 会求值为类型为 Fn<A, B>::type 的值。有了它,我们就可以把列表元素展开成使用该转换的折叠表达式,并用 decltype 来获取输出的类型。在整个过程中,不会有任何代码被实际生成——更不会——天哪——被实际执行。
这招奏效了!使用这个 fold 实现,上面第一部分的解法在我的 M1 Air 上大约一秒就能编译完成,并给出了正确答案!
$ xxd -i < input.txt > input.i && \
time c++ -fbracket-depth=25000 -DINPUT=input.i -std=c++20 \
part1.cc -o out/part1 && \
out/part1
c++ -fbracket-depth=25000 -DINPUT=input.i -std=c++20 part1.cc -o out/part1 1.10s user 0.06s system 101% cpu 1.141 total
54390第二部分
在第二部分中,数字也可能是用单词拼写出来的——比如 seven82683 的校准值就是“73”。
我们先引入一个小小的抽象。在 fold 中处理换行是可行的,但有点麻烦。我们可以写一个辅助工具,把按换行符分割的逻辑抽象出来,让我们只需关心单行的处理。
该接口将以对行进行 fold 的形式呈现,而不是比如说具体化出一个列表的列表;保持这种流式处理的方式对性能至关重要。这个辅助工具本身也是 fold 的一个相当直接的应用:
template <typename Head, typename Tail>
struct pair {
using head = Head;
using tail = Tail;
};
template <template<typename, typename> typename Fn, typename Delim>
struct fold_lines_f {
template <typename In, typename Elt>
struct F{
using type = pair<
typename append<typename In::head, Elt>::type,
typename In::tail
>;
};
template <typename A>
struct F<A, Delim>{
using type = pair<
list<>,
typename Fn<typename A::tail, typename A::head>::type
>;
};
};
template<template<typename, typename> typename Fn,
typename Init, typename L,
typename Delim = literal<'\n'>>
struct fold_lines {
using type = fold<
fold_lines_f<Fn, Delim>::template F,
pair<list<>, Init>,
L>::type::tail;
};有了它,我们现在需要计算单行内的校准值……
手动编译一个状态机
我们将再次通过对每一行做一次 fold 来实现校准值的计算。
为了匹配各种可能的数字,我们将维护一个有限状态机匹配器,并为每个 <状态, 字符> 对编写转移规则——这本质上就是某些正则表达式实现在底层会产生的 DFA。我们会用目前已看到的“相关”子串来命名状态;例如,Sseve 表示我们刚刚看到了“seve”;在这种情况下,如果遇到 n,就说明我们找到了数字 7。
注意,许多数字单词中包含的字母也可能是另一个数字单词的开头,我们需要在转移中处理这种情况;例如,如果我们处于状态 Sseve 并看到了 i,就需要转移到 Sei,因为我们可能正面对字符串“seveight”,需要处理其中的 eight。幸运的是,任意两个数字之间没有更复杂的重叠,所以状态机仍然相当简单。
我们先从状态定义开始,用 S0 表示“没有相关前缀”:
struct S0 {};
// one
struct So {};
struct Son {};
// two
struct St {};
struct Stw {};
// three
struct Sth {};
struct Sthr {};
struct Sthre {};
// four
struct Sf {};
struct Sfo {};
struct Sfou {};
// five
struct Sfi {};
struct Sfiv {};
// six
struct Ss {};
struct Ssi {};
// seven
struct Sse {};
struct Ssev {};
struct Sseve {};
// eight
struct Se {};
struct Sei {};
struct Seig {};
struct Seigh {};
// nine
struct Sn {};
struct Sni {};
struct Snin {};接下来我们可以定义转移规则。我们会大量使用通配符,并依靠偏特化的相对特异性来选择正确的规则。
如果没有其他规则匹配,就回到 S0:
template <typename St, typename El>
struct next_state { using type = S0; };而如果我们看到了任意数字单词的首字母——且没有更具体的规则匹配——就可以直接进入对应的状态。
// Initial letters
template<typename S> struct next_state<S, literal<'o'>> { using type = So; };
template<typename S> struct next_state<S, literal<'t'>> { using type = St; };
template<typename S> struct next_state<S, literal<'f'>> { using type = Sf; };
template<typename S> struct next_state<S, literal<'s'>> { using type = Ss; };
template<typename S> struct next_state<S, literal<'e'>> { using type = Se; };
template<typename S> struct next_state<S, literal<'n'>> { using type = Sn; };然后我们定义规则内部和规则之间的转移。其中大多数都很直接——例如 (So, 'n') → Son,但对于那些可能会“中途”跳到另一个数字匹配过程中的情况,我们就得小心了。
// one
template<> struct next_state<So, literal<'n'>> { using type = Son; };
template<> struct next_state<Son, literal<'e'>> { using type = Se; };
template<> struct next_state<Son, literal<'i'>> { using type = Sni; };
// two
template<> struct next_state<St, literal<'w'>> { using type = Stw; };
// three
template<> struct next_state<St, literal<'h'>> { using type = Sth; };
template<> struct next_state<Sth, literal<'r'>> { using type = Sthr; };
template<> struct next_state<Sthr, literal<'e'>> { using type = Sthre; };
template<> struct next_state<Sthre, literal<'i'>> { using type = Sei; };
// four
template<> struct next_state<Sf, literal<'o'>> { using type = Sfo; };
template<> struct next_state<Sfo, literal<'u'>> { using type = Sfou; };
template<> struct next_state<Sfo, literal<'n'>> { using type = Son; };
// five
template<> struct next_state<Sf, literal<'i'>> { using type = Sfi; };
template<> struct next_state<Sfi, literal<'v'>> { using type = Sfiv; };
// six
template<> struct next_state<Ss, literal<'i'>> { using type = Ssi; };
// seven
template<> struct next_state<Ss, literal<'e'>> { using type = Sse; };
template<> struct next_state<Sse, literal<'v'>> { using type = Ssev; };
template<> struct next_state<Sse, literal<'i'>> { using type = Sei; };
template<> struct next_state<Ssev, literal<'e'>> { using type = Sseve; };
template<> struct next_state<Sseve, literal<'i'>> { using type = Sei; };
// eight
template<> struct next_state<Se, literal<'i'>> { using type = Sei; };
template<> struct next_state<Sei, literal<'g'>> { using type = Seig; };
template<> struct next_state<Seig, literal<'h'>> { using type = Seigh; };
// nine
template<> struct next_state<Sn, literal<'i'>> { using type = Sni; };
template<> struct next_state<Sni, literal<'n'>> { using type = Snin; };我们还没有定义何时算“匹配”了一条规则,以及匹配成功后该做什么;事实证明,对于这个问题,把“匹配成功”定义为一个独立于状态转移的函数,会让两个函数都简单得多。用更专业的说法,我们本质上是在为这个匹配器实现电气工程师们所说的米利型状态机。
match_digit 函数的签名与 next_state 类似——state, input → digit | nil——它告诉我们在某个特定位置匹配到了哪个数字(如果有的话):
template <typename State, typename El>
struct match_digit { using type = nil; };
template <> struct match_digit<Son, literal<'e'>> { using type = literal<1>; };
template <> struct match_digit<Stw, literal<'o'>> { using type = literal<2>; };
template <> struct match_digit<Sthre, literal<'e'>> { using type = literal<3>; };
template <> struct match_digit<Sfou, literal<'r'>> { using type = literal<4>; };
template <> struct match_digit<Sfiv, literal<'e'>> { using type = literal<5>; };
template <> struct match_digit<Ssi, literal<'x'>> { using type = literal<6>; };
template <> struct match_digit<Sseve, literal<'n'>> { using type = literal<7>; };
template <> struct match_digit<Seigh, literal<'t'>> { using type = literal<8>; };
template <> struct match_digit<Snin, literal<'e'>> { using type = literal<9>; };
template<typename S> struct match_digit<S, literal<'0'>> { using type = literal<0>; };
template<typename S> struct match_digit<S, literal<'1'>> { using type = literal<1>; };
template<typename S> struct match_digit<S, literal<'2'>> { using type = literal<2>; };
template<typename S> struct match_digit<S, literal<'3'>> { using type = literal<3>; };
template<typename S> struct match_digit<S, literal<'4'>> { using type = literal<4>; };
template<typename S> struct match_digit<S, literal<'5'>> { using type = literal<5>; };
template<typename S> struct match_digit<S, literal<'6'>> { using type = literal<6>; };
template<typename S> struct match_digit<S, literal<'7'>> { using type = literal<7>; };
template<typename S> struct match_digit<S, literal<'8'>> { using type = literal<8>; };
template<typename S> struct match_digit<S, literal<'9'>> { using type = literal<9>; };(我们本可以用主模板定义中的一个条件分支来替代最后 10 个特化,但我觉得这个版本虽然有点冗长,却更直观一些。)
有了这些类型,我们现在就可以像第一部分那样,用一个非常相似的 fold 来计算单行的校准值。与之前相比,我们不再需要跟踪总的累加器——我们会用 fold_lines 把它放到外层循环中处理——但我们确实需要保留匹配器的状态,因此我们的状态类看起来非常相似:
template <typename MatchState, typename First, typename Last>
struct CalibrationState {
using state = MatchState;
using first = First;
using last = Last;
};要推进一个元素,我们需要计算输出状态和匹配到的数字(如果有的话)。我们还会用一些辅助工具来简化“首次/最后一次看到的数字”的计算:
template <typename State, typename El>
struct LineF {
using new_state = next_state<typename State::state, El>::type;
using digit = match_digit<typename State::state, El>::type;
using type = CalibrationState<
new_state,
typename or_else<typename State::first, digit>::type,
typename or_else<digit, typename State::last>::type
>;
};有了这个 fold 函数,计算校准值就很容易了:
template<typename Line>
struct calibration {
using acc = fold<LineF, CalibrationState<S0, nil, nil>, Line>::type;
using type = literal<acc::first::value*10 + acc::last::value>;
};我们甚至可以写一些快速的测试用例——C++ 模板元编程对内联测试有很好的支持。
static_assert(is_same<
typename calibration<typename read_input<
's', 'e', 'v', 'e', 'n', '8', '2', '6', '8', '3'
>::type>::type,
literal<73>>::value);整合起来
对行的折叠只是对校准值做一个简单的求和——到这一步,几乎是 trivial 的。
template<typename State, typename Line>
struct Fn {
using lineval = calibration<Line>::type;
using type = literal<State::value + lineval::value>;
};
using answer = fold_lines<Fn, literal<0>, problem>::type;成功了!而且,在我的笔记本上,编译“仅”需约 1.5 秒!
结论
过去我也零星做过一些 C++ 模板元编程——例如,我曾为 C++ 写过一个简单的 x86 汇编器——但这是我第一次真正接触 C++17 的模板元编程,并真正一头扎进深水区。
而且,这种感触我以前也听说过,但在做完这道题之后……老实说,我不得不承认现代 C++ 模板已经是一门相当像样的函数式编程语言了,甚至还有一些不错的特性!而发现折叠表达式这个技巧也让我惊喜不已,它甚至能让列表处理变得相对高效。
我……仍然不确定还会再做多少天,或者还会写多少 C++,但最终我还是相当享受这次练习!
随机一篇博客
评论
登录后参与讨论