You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在LibTorch中实现PyTorch的张量切片赋值与矩阵索引存储

解决LibTorch中三维张量下三角位置赋值问题

核心问题分析

你之前的LibTorch代码出错,主要是两个原因:

  1. 索引数量不匹配:三维张量cholesky需要三个维度的索引,你只传了trl[0]和trl[1]两个,缺了第一个batch维度的索引;
  2. 维度对应错误:没有明确指定第一个维度的全部元素,导致索引和赋值张量的维度无法对齐。

正确实现方案

有两种简洁的方式实现和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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.11 12:45:42