如何在LibTorch中实现PyTorch的张量切片赋值与矩阵索引存储
解决LibTorch中三维张量下三角位置赋值问题
核心问题分析
你之前的LibTorch代码出错,主要是两个原因:
- 索引数量不匹配:三维张量
cholesky需要三个维度的索引,你只传了trl[0]和trl[1]两个,缺了第一个batch维度的索引; - 维度对应错误:没有明确指定第一个维度的全部元素,导致索引和赋值张量的维度无法对齐。
正确实现方案
有两种简洁的方式实现和PyTorch完全一致的逻辑:
方案1:使用LibTorch索引语法(推荐)
需要包含torch/index.h头文件,利用Slice()对应PyTorch中的:,直接复刻PyTorch的写法:
#include <torch/torch.h> #include <torch/index.h> using namespace torch::indexing; int main() { int64_t batch_size = 2; int64_t mat_size = 3; auto cholesky = torch::zeros({batch_size, mat_size, mat_size}); auto trl = torch::tril_indices(mat_size, mat_size, -1); // 获取下三角位置索引(不含对角线) auto cholesky_vector = torch::randn({batch_size, trl.size(1)}); // 形状[batch_size, K],K是下三角元素数量 // 对应PyTorch的 cholesky[:, trl[0], trl[1]] = cholesky_vector cholesky.index_put_({Slice(), trl[0], trl[1]}, cholesky_vector); return 0; }
方案2:手动生成广播后的索引
如果不想用索引语法,也可以手动生成每个维度的索引,确保长度和赋值张量的展平长度一致:
#include <torch/torch.h> int main() { int64_t batch_size = 2; int64_t mat_size = 3; auto cholesky = torch::zeros({batch_size, mat_size, mat_size}); auto trl = torch::tril_indices(mat_size, mat_size, -1); auto cholesky_vector = torch::randn({batch_size, trl.size(1)}); // 生成batch维度的索引:每个batch元素重复K次,展平后长度为batch_size*K auto batch_idx = torch::arange(batch_size).unsqueeze(1).expand({-1, trl.size(1)}).flatten(); // 行、列索引重复batch_size次 auto row_idx = trl[0].repeat(batch_size); auto col_idx = trl[1].repeat(batch_size); // 展平赋值张量 auto flat_vector = cholesky_vector.flatten(); // 执行赋值 cholesky.index_put_({batch_idx, row_idx, col_idx}, flat_vector); return 0; }
验证说明
两种方案都能实现:对每个batch的cholesky张量,将cholesky_vector的对应元素赋值到trl[0]和trl[1]指定的下三角位置,和PyTorch的运行结果完全一致。
内容的提问来源于stack exchange,提问作者Ant
相关产品推荐
相关产品推荐

