LibTorch(PyTorch) C++环境下将at::Tensor转换为double的方法问询
实现方法
在LibTorch中,单元素标量张量转C++基础类型可以直接调用at::Tensor的模板方法item<T>(),你要转double类型直接用loss.item<double>()即可。
你给出的原始代码还存在括号未闭合的语法错误,修正后的完整可运行代码如下:
#include <torch/torch.h> #include <cstdlib> int main() { auto const input1(torch::randn({28*28})); auto const input2(torch::randn({28*28})); double const lossThreshold{0.05}; auto const loss{torch::nn::functional::mse_loss(input1, input2)}; // 转换为double后比较 return loss.item<double>() > lossThreshold ? EXIT_FAILURE : EXIT_SUCCESS; }
注意事项
item<T>()仅适用于只有单个元素的张量,你当前场景下mse_loss默认返回单元素标量张量,可直接使用;如果张量包含多个元素,调用该方法会抛出异常。- 若需要严格保证精度,可先将张量显式转换为double类型再取值:
loss.to(torch::kDouble).item<double>(),避免原张量为float类型时出现精度误差。
内容的提问来源于stack exchange,提问作者Raashid
相关产品推荐
相关产品推荐

