C++14下构造含N-1个中间激活、1个末尾激活的变长参数包
问题背景
核心需求为将2个输入参数[x, y]扩展为变长参数包[x, ...x, y],即重复N-1次第一个参数后拼接第二个参数。
当前在C++14环境下实现自定义多层感知机(MLP):除最后一层使用独立激活函数bar外,其余所有层共用激活函数foo,目标是让代码支持任意层数的模型。现有实现硬编码了激活函数对应的参数包,仅能在层数N = 5时正常编译。
需求说明
需要实现C++14兼容的工具模板,基于两个传入的激活函数生成指定形式的参数包,实现第一个参数向左重复N-1次、末尾拼接第二个参数的效果,修改computeMlp的调用逻辑后适配任意N值。工具模板预期行为如下:
template <size_t N, typename ActivationMid, typename ActivationLast> struct makeActivationSequence {}; // 待实现 // 预期输出规则 // makeActivationSequence<0>(foo, bar) -> [] // makeActivationSequence<1>(foo, bar) -> [bar] // makeActivationSequence<2>(foo, bar) -> [foo, bar] // makeActivationSequence<3>(foo, bar) -> [foo, foo, bar] // makeActivationSequence<4>(foo, bar) -> [foo, foo, foo, bar] // ...
- 约束:仅可使用C14标准特性,禁止使用
if constexpr等C17及以上版本语法。
原始复现代码
#include <cstddef> #include <utility> #include <cstdio> template <size_t LayerIndex, typename Activation> void computeIndexedLayer(const Activation& activation) { printf("Doing work for layer %zu, activated output %zu\n", LayerIndex, activation(LayerIndex)); } template <std::size_t... index, typename... Activation> void computeIndexedLayers(std::index_sequence<index...>, Activation&&... activation) { (void)std::initializer_list<int>{ (computeIndexedLayer<index + 1>(std::forward<Activation>(activation)), 0)... }; } template <size_t N, typename ActivationMid, typename ActivationLast> void computeMlp(ActivationMid&& mid, ActivationLast&& last) { // 此处硬编码了4个mid+1个last,仅支持N=5 computeIndexedLayers(std::make_index_sequence<N>(), std::forward<ActivationMid>(mid), std::forward<ActivationMid>(mid), std::forward<ActivationMid>(mid), std::forward<ActivationMid>(mid), std::forward<ActivationLast>(last) ); } int main() { computeMlp<5>([](const auto& x){ return x + 1;}, [](const auto& x){ return x * 1000;}); // 其他N值会因参数包长度不匹配编译失败 // computeMlp<4>([](const auto& x){ return x + 1;}, [](const auto& x){ return x * 1000;}); }
解决方案
复用std::index_sequence的递归构造思路,通过模板递归生成前N-1个中间激活函数的参数包,最后拼接末尾激活函数,全程兼容C++14标准。
完整实现代码
#include <tuple> #include <type_traits> // 递归辅助模板:逐层展开生成N个中间激活的参数占位 template <size_t Count, typename ActivationMid, typename... Tail> struct activation_seq_impl : activation_seq_impl<Count - 1, ActivationMid, ActivationMid, Tail...> {}; // 递归终止:Count归0时,将末尾激活拼接到参数包尾部 template <typename ActivationMid, typename... Tail> struct activation_seq_impl<0, ActivationMid, Tail...> { template <typename Mid, typename Last> static auto apply(Mid&&, Last&& last, Tail... tail) -> std::tuple<Tail..., Last> { return std::tuple<Tail..., Last>( std::forward<Tail>(tail)..., std::forward<Last>(last) ); } }; // 主模板:处理N>=1的场景 template <size_t N, typename Enable = void> struct makeActivationSequence { template <typename Mid, typename Last> auto operator()(Mid&& mid, Last&& last) const -> decltype(activation_seq_impl<N-1, Mid>::apply( std::forward<Mid>(mid), std::forward<Last>(last), std::forward<Mid>(mid) )) { return activation_seq_impl<N-1, Mid>::apply( std::forward<Mid>(mid), std::forward<Last>(last), std::forward<Mid>(mid) ); } }; // N=0特化:返回空参数包 template <size_t N> struct makeActivationSequence<N, typename std::enable_if<N == 0>::type> { template <typename Mid, typename Last> std::tuple<> operator()(Mid&&, Last&&) const { return std::tuple<>(); } }; // 辅助函数:将tuple展开为变长参数包传给computeIndexedLayers template <size_t... Is, typename Tuple> void call_compute_layers(std::index_sequence<Is...>, Tuple&& act_pack) { computeIndexedLayers( std::make_index_sequence<std::tuple_size<typename std::decay<Tuple>::type>::value>(), std::get<Is>(std::forward<Tuple>(act_pack))... ); } // 改造后的computeMlp,支持任意N值 template <size_t N, typename ActivationMid, typename ActivationLast> void computeMlp(ActivationMid&& mid, ActivationLast&& last) { auto act_pack = makeActivationSequence<N>()( std::forward<ActivationMid>(mid), std::forward<ActivationLast>(last) ); call_compute_layers(std::make_index_sequence<N>(), std::move(act_pack)); }
实现说明
- 递归构造参数包时通过模板继承自动堆叠N-1个中间激活的类型占位,终止条件处拼接最后一层激活,完美匹配需求的参数顺序
- 全程使用完美转发,传入的lambda等可调用对象不会产生多余拷贝,左值传入时自动保留引用语义
- 边界逻辑完全符合预期:N=0返回空包、N=1仅返回最后一层激活、N=k返回k-1个中间激活加1个末尾激活
- 无任何C17及以上语法,可直接在C14工具链下编译运行
内容的提问来源于stack exchange,提问作者Vasu
相关产品推荐
相关产品推荐

