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

循环内无法通过赋值覆盖二维numpy.ndarray的问题求助

三对角矩阵QR分解中NumPy数组赋值失效问题解决

问题现象

实现基于Givens旋转的三对角矩阵QR分解时,循环中执行X[i] = 形状匹配的np.ndarray赋值语句后,数组切片X[i]的值未发生改变。赋值常量(如X[i]=10)有效,但赋值数组无效果。

原代码

import numpy as np
def qr_tridiagonal(T: np.ndarray):
    m, n = T.shape
    X = T.copy()
    Qt = np.identity(m)
    for i in range(n-1):
        ai = X[i, i]
        ak = X[i+1, i]
        c = ai/(ai**2 + ak**2)**.5
        s = ak/(ai**2 + ak**2)**.5
        # Givens rotation
        tmp1 = c*X[i] + s*X[i+1]
        tmp2 = c*X[i+1] - s*X[i]
        print("tmp1 before:", tmp1)
        print("X[i] before:", X[i])
        X[i] = tmp1
        X[i+1] = tmp2
        print("tmp1 after:", tmp1)
        print("X[i] after:", X[i])
        print()

        print(X)

    return Qt.T, X


A = np.array([[1, 1, 0, 0], [1, 1, 1, 0], [0, 1, 1, 1], [0, 0, 1, 1]])
Q, R = qr_tridiagonal(A)

异常输出(前4行)

tmp1 before: [1.41421356 1.41421356 0.70710678 0.        ]
X[i] before: [1 1 0 0]
tmp1 after: [1.41421356 1.41421356 0.70710678 0.        ]
X[i] after: [1 1 0 0]

原因分析

问题根源在于数组类型不匹配:

  • 输入数组A是整数类型(默认int32或int64),复制得到的X也继承了整数类型。
  • Givens旋转计算得到的tmp1、tmp2是浮点数类型,赋值给整数数组切片时,NumPy会自动将浮点数截断为整数。
  • 示例中tmp1的元素(如1.414)截断后为1,与原X[i]的整数值完全一致,因此看起来赋值未生效。

解决方案

只需确保X为浮点类型即可,有两种实现方式:

方式1:复制时转换为浮点类型

修改X = T.copy()为:

X = T.copy().astype(np.float64)

方式2:初始化输入数组为浮点类型

修改A的初始化代码:

A = np.array([[1, 1, 0, 0], [1, 1, 1, 0], [0, 1, 1, 1], [0, 0, 1, 1]], dtype=np.float64)

修改后的代码示例(方式1)

import numpy as np
def qr_tridiagonal(T: np.ndarray):
    m, n = T.shape
    X = T.copy().astype(np.float64)  # 转换为浮点类型
    Qt = np.identity(m)
    for i in range(n-1):
        ai = X[i, i]
        ak = X[i+1, i]
        c = ai/(ai**2 + ak**2)**.5
        s = ak/(ai**2 + ak**2)**.5
        # Givens rotation
        tmp1 = c*X[i] + s*X[i+1]
        tmp2 = c*X[i+1] - s*X[i]
        print("tmp1 before:", tmp1)
        print("X[i] before:", X[i])
        X[i] = tmp1
        X[i+1] = tmp2
        print("tmp1 after:", tmp1)
        print("X[i] after:", X[i])
        print()

        print(X)

    return Qt.T, X


A = np.array([[1, 1, 0, 0], [1, 1, 1, 0], [0, 1, 1, 1], [0, 0, 1, 1]])
Q, R = qr_tridiagonal(A)

修正后的输出(前4行)

tmp1 before: [1.41421356 1.41421356 0.70710678 0.        ]
X[i] before: [1. 1. 0. 0.]
tmp1 after: [1.41421356 1.41421356 0.70710678 0.        ]
X[i] after: [1.41421356 1.41421356 0.70710678 0.        ]

内容的提问来源于stack exchange,提问作者Linnore

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.09 17:35:25