如何在LibTorch中沿0维度获取3D Tensor的最大值?
在LibTorch中对3D张量沿第0维度取最大值的正确实现
你的代码存在两个关键问题,导致编译错误和功能不符合预期:
问题1:错误调用torch::max
你使用torch::max({target_q_T}, 0)是对张量列表求最大值,而非对单个张量沿指定维度计算最大值。正确调用应直接传入目标张量,无需用花括号包裹。
问题2:无法直接输出返回的tuple
LibTorch中,当指定维度调用torch::max时,返回的是std::tuple<at::Tensor, at::Tensor>(对应Python返回的values和indices),C++标准输出std::cout没有默认支持输出这种tuple类型,必须分别访问其中的元素。
修正后的完整代码
#include <torch/torch.h> #include <iostream> int main() { auto target_q_T = torch::rand({5, 10, 1}); // 正确调用:对单个张量沿第0维度取最大值 auto max_result = torch::max(target_q_T, 0); // 提取最大值张量和对应的索引张量 auto max_values = std::get<0>(max_result); auto max_indices = std::get<1>(max_result); // 分别输出结果 std::cout << "最大值张量:\n" << max_values << "\n\n"; std::cout << "索引张量:\n" << max_indices << "\n"; return 0; }
代码说明
torch::max(target_q_T, 0):与Python中torch.max(target_q_np, 0)行为完全一致,返回包含最大值和对应索引的tuple。std::get<0>(max_result):获取tuple中的第一个元素,即最大值张量,对应Python返回结果的values。std::get<1>(max_result):获取tuple中的第二个元素,即最大值所在的索引张量,对应Python返回结果的indices。
如果仅需最大值张量,也可以直接通过torch::max(target_q_T, 0).values()获取,效果与std::get<0>相同。
内容的提问来源于stack exchange,提问作者Ant
相关产品推荐
相关产品推荐

