咨询PyTorch中_lstm_impl未指定的两个模板参数来源
关于PyTorch _lstm_impl模板参数推导的问题解答
你观察得很仔细!这个问题的核心是C++的函数模板参数自动推导特性,我来给你拆解清楚:
首先先把你提到的代码片段贴出来方便对照:
函数声明
template< template<typename,typename> class LayerT, template<typename,typename> class BidirLayerT, typename cell_params, typename io_type > std::tuple<io_type, Tensor, Tensor> _lstm_impl( const io_type& input, const std::vector<cell_params>& params, const Tensor& hx, const Tensor& cx, int64_t num_layers, double dropout_p, bool train, bool bidirectional ) { /* ... */ }
调用代码
auto results = _lstm_impl<FullLayer, FullBidirectionalLayer>( input, params, hx[0], hx[1], num_layers, dropout_p, train, bidirectional );
为什么只传2个模板参数就可以?
这两个未显式指定的模板参数cell_params和io_type,是编译器通过函数实参自动推导出来的:
- io_type的推导:函数第一个形参是
const io_type& input,调用时你传入了变量input,编译器会直接取input的实际类型,作为io_type的具体类型(比如在PyTorch里这里通常是Tensor或者对应的适配类型)。 - cell_params的推导:函数第二个形参是
const std::vector<cell_params>& params,调用时传入的params是一个vector容器,编译器会从这个vector的元素类型,自动推导出cell_params的具体类型。
那为什么前两个模板参数LayerT和BidirLayerT必须显式指定?因为它们是模板模板参数(也就是template<typename,typename> class这种类型的参数)——这种参数代表的是一个"模板类",而不是具体的类型,编译器没办法通过函数实参推导出你要用哪个模板类,所以必须手动指定FullLayer和FullBidirectionalLayer这两个具体的模板类。
简单来说,C++编译器会帮你"补全"那些能从函数参数里推断出来的模板参数,不用你一个个写全~
内容的提问来源于stack exchange,提问作者Zack
相关产品推荐
相关产品推荐

