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

使用xtensor实现Softmax层反向传播时触发维度错误问题求助

xtensor实现Softmax层反向传播时触发维度错误问题求助

我刚仔细看了你的代码和报错信息,这个问题的根源很明确:xt::row函数只能用于二维数组,但你传入的输入是一维的。

先看报错内容:

terminate called after throwing an instance of 'std::invalid_argument'
  what():  A row can only be accessed on an expression with exact two dimensions
Aborted (core dumped)

这个错误直接告诉你,你在调用xt::row(cached_Y, i)的时候,cached_Y并不是二维数组——因为你在main函数里传入的X是一维数组{1,2,3}(shape为(3,)),经过forward计算后的cached_Y自然也是一维的,用xt::row去访问一维数组的"行"肯定会触发维度校验错误。

下面给你两个可行的修复方案,你可以根据需求选择:


方案一:直接传入二维输入(最简单)

既然xt::row需要二维数组,那我们直接在main里把输入改成二维形式就行,同时要对应好Softmax的axis参数:

情况1:按行做Softmax(axis=1)

如果想对一行的多个类别计算Softmax,就把X定义为二维行向量:

int main(int argc, char* argv[]) {
    // 改成二维行向量,shape=(1,3)
    xt::xarray<double> X = {{1, 2, 3}};
    std::cout << "Input X: " << X << std::endl;

    // axis设为1,对列维度(每个类别)计算Softmax
    Softmax softmax_layer(1);
    xt::xarray<double> Y = softmax_layer.forward(X);
    std::cout << "Softmax Output Y: " << Y << std::endl;

    // DY也要对应二维行向量
    xt::xarray<double> DY = {{0.1, 0.2, 0.3}};
    std::cout << "Input DY: " << DY << std::endl;

    xt::xarray<double> DX = softmax_layer.backward(DY);
    std::cout << "Backward Output DX: " << DX << std::endl;

  return 0;
}

情况2:按列做Softmax(axis=0)

如果想对列维度计算Softmax,就把X定义为二维列向量:

int main(int argc, char* argv[]) {
    // 改成二维列向量,shape=(3,1)
    xt::xarray<double> X = {{1}, {2}, {3}};
    std::cout << "Input X: " << X << std::endl;

    // axis设为0,对行维度(每个类别)计算Softmax
    Softmax softmax_layer(0);
    xt::xarray<double> Y = softmax_layer.forward(X);
    std::cout << "Softmax Output Y: " << Y << std::endl;

    // DY对应二维列向量
    xt::xarray<double> DY = {{0.1}, {0.2}, {0.3}};
    std::cout << "Input DY: " << DY << std::endl;

    xt::xarray<double> DX = softmax_layer.backward(DY);
    std::cout << "Backward Output DX: " << DX << std::endl;

  return 0;
}

方案二:让Softmax类自动兼容一维输入(更通用)

如果想让你的Softmax层同时支持一维和二维输入,可以在forward函数里自动把一维输入扩展成二维,同时在backward里根据axis参数选择用row还是col访问元素:

修改后的Softmax类代码

class Softmax {
public:
    Softmax(int axis) : m_nAxis(axis) {}

    xt::xarray<double> forward(xt::xarray<double> X) {
        // 自动处理一维输入,转成对应维度的二维数组
        if (X.dimension() == 1) {
            if (m_nAxis == 0) {
                // 转成列向量 (n, 1)
                X = xt::reshape(X, {X.shape()[0], 1});
            } else {
                // 转成行向量 (1, n)
                X = xt::reshape(X, {1, X.shape()[0]});
            }
        }
        cached_Y = softmax(X, m_nAxis);
        return cached_Y;
    }

    xt::xarray<double> backward(const xt::xarray<double>& DY) {
        xt::xarray<double> DX = xt::zeros<double>(DY.shape());
        int axis = m_nAxis;
        int num_elements = cached_Y.shape()[axis];

        for (int i = 0; i < num_elements; i++) {
            xt::xarray<double> y;
            xt::xarray<double> dy_i;
            // 根据axis选择访问行还是列
            if (axis == 0) {
                y = xt::row(cached_Y, i);
                dy_i = xt::row(DY, i);
            } else {
                y = xt::col(cached_Y, i);
                dy_i = xt::col(DY, i);
            }
            // 展平成一维,方便计算雅克比矩阵
            y = xt::flatten(y);
            dy_i = xt::flatten(dy_i);
            // 计算Softmax的雅克比矩阵
            xt::xarray<double> J = xt::diag(y) - xt::linalg::outer(y, y);
            // 根据axis把结果赋值回DX对应位置
            if (axis == 0) {
                xt::row(DX, i) = xt::linalg::dot(J, dy_i);
            } else {
                xt::col(DX, i) = xt::linalg::dot(J, dy_i);
            }
        }

        return DX;
    }

private:
    int m_nAxis;
    xt::xarray<double> cached_Y;

    int positive_index(int idx, int size) {
        if (idx < 0) return idx + size;
        return idx;
    }

    xt::xarray<double> softmax(xt::xarray<double> X, int axis) {
        xt::svector<unsigned long> shape = X.shape();
        axis = positive_index(axis, shape.size());
        shape[axis] = 1;

        xt::xarray<double> Xmax = xt::amax(X, {axis});
        X = xt::exp(X - Xmax.reshape(shape));
        xt::xarray<double> SX = xt::sum(X, {axis});
        SX = SX.reshape(shape);
        X = X / SX;

        return X;
    }
};

这样修改后,你原来main里的一维输入代码不用改也能正常运行,同时也支持二维输入,兼容性更强。


备注:内容来源于stack exchange,提问作者user28212153

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.15 09:37:57