
MLX 线性代数模块完全指南mlx.core.linalg 的 19 个核心函数与底层实现解析【免费下载链接】mlxMLX: An array framework for Apple silicon项目地址: https://gitcode.com/GitHub_Trending/ml/mlx本文围绕 MLXApple silicon 上的数组框架官方文档 docs/src/python/linalg.rst 展开系统讲解mlx.core.linalg线性代数子模块中全部 19 个 APInorm、qr、svd、inv、cholesky、eig、lu、solve、det等的签名、用法、数值行为与底层 LAPACK 实现并结合 mlx/linalg.cpp、python/src/linalg.cpp 及 C/Python 测试源码帮助你快速在 MLX 中完成矩阵分解、特征值、范数与线性方程组求解等科学计算任务。mlx.core.linalg是 MLX 提供的线性代数例程集合覆盖向量/矩阵范数、QR / SVD / LU / Cholesky 分解、特征值分解、矩阵求逆含广义逆、行列式以及线性方程组求解。它遵循 NumPy 的np.linalg风格命名与行为约定同时针对 Apple silicon 的统一内存架构做了工程化实现绝大多数分解类操作在 CPU 后端通过 LAPACK 完成cholesky在 CUDA 后端可用norm则完全由 MLX 原语sum、abs、max、svd等组合而成可在默认设备上运行。读完本文你将能按需选择正确的函数与参数理解每个函数的设备限制与数据类型约束并学会用批处理维度同时处理多组矩阵运算。一、模块概览函数清单与快速选型原文档通过autosummary列出了mlx.core.linalg的全部公开接口。按功能分组如下功能类别函数关键特征范数norm同时支持向量范数与矩阵范数ord可传数值或字符串求逆inv、tri_inv、pinv、cholesky_inv普通逆、三角逆、Moore-Penrose 伪逆、Cholesky 加速求逆矩阵分解qr、svd、lu、lu_factor、choleskyQR、奇异值、LU含紧凑形式、Cholesky特征值eigvals、eig、eigvalsh、eigh一般矩阵 / 对称Hermitian矩阵的特征值与特征向量行列式det、slogdet行列式及其带符号对数形式方程组solve、solve_triangular一般线性系统、三角系统求解几何运算cross叉积沿指定轴尺寸 2 或 3从 python/src/linalg.cpp 末尾可见模块还提供matrix_norm作为norm的 Array API 标准别名parent_module.attr(matrix_norm) m.attr(norm)即mx.linalg.matrix_norm与mx.linalg.norm等价。快速选型建议求最小二乘/伪逆 →pinv要求稳定且可复现 →svd对称正定矩阵求逆 →cholesky_inv比inv更节省计算特征值只关心数值 →eigvals/eigvalsh需要特征向量 →eig/eigh多次求解同系数矩阵 → 先lu_factor复用分解再配合solve_triangular。二、公共约定设备、数据类型与批处理语义在深入各函数之前先明确贯穿整个模块的通用约束源码依据mlx/linalg.cpp 开头的校验函数与 python/src/linalg.cpp 的 Python 签名。1. 设备限制默认 CPU 流MLX 的线性代数分解例程大多基于 LAPACK目前在 GPUMetal后端尚未实现。源码中的两个检查函数决定了可用设备check_cpu_streamqr、svd、inv、tri_inv、pinv、cholesky_inv、eig、eigvals、eigvalsh、eigh、lu、lu_factor、solve、solve_triangular、det、slogdet均调用它若目标流为 GPU 会直接抛错并提示 Explicitly pass a CPU stream to run itcheck_cpu_or_cuda_stream仅cholesky使用若 GPU 上 CUDA 可用cu::is_available()则允许执行否则回退 CPU。因此调用这类函数时建议显式传入streammx.cpu例如import mlx.core as mx A mx.array([[2.0, 3.0], [1.0, 2.0]]) Q, R mx.linalg.qr(A, streammx.cpu)norm与cross不受此限制可在默认设备上运行其中矩阵 2-范数 / -2-范数 / 核范数ord2、-2、nuc内部依赖 SVD实际执行时会落到 CPU。2. 数据类型约束绝大多数函数qr、inv、pinv、cholesky、cholesky_inv、lu系列、det、slogdet仅接受float32/float64见check_floatsvd、eig、eigvals、eigh、eigvalsh额外支持complex64见check_float_or_complexdet/slogdet明确拒绝复数输入solve系列要求输入可提升promote到float32或float64整型输入会被拒绝或按at_least_float提升如det内部会astype到浮点后再计算。3. 批处理语义所有矩阵类函数都支持至少 2 维的输入当a.ndim 2时前a.ndim - 2个维度被视为批维度函数会逐矩阵对最后两维执行运算。例如svd文档明确说明“对前a.ndim - 2维的所有索引组合对最后两维逐一做 SVD”。norm则支持最多 2 个轴同时归约。三、norm向量/矩阵范数的完整语义norm是模块中最灵活的函数签名如下python/src/linalg.cpp 中的 nanobind 定义mx.linalg.norm(a, ordNone, axisNone, keepdimsFalse, *, streamNone)参数含义a输入数组。若axis为None而ord非None则a必须是 1-D 或 2-D若两者均为None返回a.flatten()的 2-范数ord范数阶数数值int/float或字符串fro、nuc、f均可默认None表示按给定轴计算 2-范数矩阵为 Frobenius 范数axisint或 2 元组指定计算范数的轴最多 2 个轴超过会抛ValueErrorkeepdims若为True被归约的轴在结果中保留为尺寸 1 的维度注意当不指定axis时keepdims会把结果 reshape 回与原输入相同的秩见 mlx/linalg.cpp 中norm对flatten与reshape的处理。ord 取值总表原文档实为 Python 绑定 docstring见 python/src/linalg.cpp给出的完整对应关系如下ord矩阵范数向量范数NoneFrobenius 范数2-范数fro/fFrobenius 范数--nuc核范数奇异值之和--infmax(sum(abs(x), axis1))行和最大值max(abs(x))-infmin(sum(abs(x), axis1))行和最小值min(abs(x))0--sum(x ! 0)非零元素个数1max(sum(abs(x), axis0))列和最大值sum(abs(x))-1min(sum(abs(x), axis0))列和最小值sum(abs(x)^(-1))^(-1)22-范数最大奇异值sum(abs(x)^2)^(1/2)-2最小奇异值sum(abs(x)^(-2))^(-1/2)其他数值--sum(abs(x)**ord)**(1/ord)其中 Frobenius 范数定义为||A||_F (sum_{i,j} abs(a_ij)^2)^(1/2)核范数为奇异值之和这两个字符串阶数只对矩阵2 个轴有意义对 1-D 输入会抛ValueError。另外文档特别提示ord 1的结果严格来说不是数学意义上的范数但可用于数值目的。源码实现要点norm在 mlx/linalg.cpp 中有三个重载对应ord的三种形态缺省 / 数值 / 字符串缺省ord直接走l2_norm即sqrt(sum(square(a), axis, keepdims))对复数输入改用abs(a)*abs(a)求和l2_norm中的issubdtype(a.dtype(), complexfloating)分支数值ord根据轴数分派到vector_norm1 个轴或matrix_norm2 个轴。vector_norm对0 / 1 / 2 / inf / -inf / 其他分别用not_equal、abs、平方和、max/min与power组合实现matrix_norm对±1、±inf用行列方向的sum(abs)max/min对2 / -2则先调用svd(a, false)取奇异值再取最大/最小这就是为什么 2-范数需在 CPU 上执行并处理了负轴、轴顺序与keepdims的展开逻辑字符串ordf/fro复用l2_normnuc用sum(svd(...)[0])累加奇异值其他字符串抛错。示例import mlx.core as mx from mlx.core import linalg as la a mx.arange(9) - 4 # array([-4, -3, -2, ..., 2, 3, 4], dtypeint32) b a.reshape((3, 3)) la.norm(a) # 2-范数: array(7.74597) la.norm(b) # 整个矩阵 2-范数与上相同先 flatten la.norm(b, fro) # Frobenius 范数: array(7.74597) la.norm(a, float(inf)) # max(abs(a)): array(4) la.norm(b, float(inf)) # 行和最大值: array(9) la.norm(b, 1) # 列和最大值: array(7) c mx.array([[1, 2, 3], [-1, 1, 4]]) la.norm(c, axis0) # 沿列0 轴的 2-范数: array([1.41421, 2.23607, 5]) la.norm(c, axis1) # 沿行1 轴的 2-范数: array([3.74166, 4.24264]) la.norm(c, ord1, axis1) # 沿行的 1-范数: array([6, 6]) m mx.arange(8).reshape(2, 2, 2) la.norm(m, axis(1, 2)) # 对每个批矩阵求 Frobenius 范数: array([3.74166, 11.225])数值正确性由 C 测试 tests/linalg_tests.cpp[mlx.core.linalg.norm]三个测试用例覆盖无 ord、数值 ord、字符串 ord 及负轴、keepdims形状与 Python 测试 python/tests/test_linalg.pytest_norm/test_complex_norm与np.linalg.norm逐项对比atol1e-5, rtol1e-6共同保障。四、矩阵分解QR、SVD、LU 与 Cholesky1. qr —— QR 分解Q, R mx.linalg.qr(a, *, streamNone)返回Q R a。要求输入至少 2 维且只支持float32/float64底层调用 LAPACK 的geqrforgqr见 mlx/backend/cpu/qrf.cpp输出形状由 mlx/linalg.cpp 中的qr构造Q取(..., M, min(M,N))R取(..., min(M,N), N)即经济形式分解。测试 tests/linalg_tests.cpp 验证了Q R A、Q正交Q^T Q I且R严格上三角。A mx.array([[2., 3.], [1., 2.]]) Q, R mx.linalg.qr(A, streammx.cpu) # Q ≈ [[-0.894427, -0.447214], [-0.447214, 0.894427]] # R ≈ [[-2.23607, -3.57771], [ 0, 0.447214]]2. svd —— 奇异值分解U, S, Vt mx.linalg.svd(a, compute_uvTrue, *, streamNone) S mx.linalg.svd(a, compute_uvFalse, *, streamNone)返回U、S、Vt满足A U diag(S) Vtcompute_uvFalse时只返回奇异值数组S。支持float32/float64/complex64批维度逐矩阵分解。形状约定U为(..., M, M)S为(..., min(M,N))Vt为(..., N, N)见 mlx/linalg.cpp 的svd。对复数输入奇异值以float32返回s_dtype a.dtype() complex64 ? float32 : a.dtype()。底层使用 LAPACKgesdd分治算法见 mlx/backend/cpu/svd.cpp 与 mlx/backend/cpu/lapack.h。3. lu / lu_factor —— LU 分解及其紧凑形式p, L, U mx.linalg.lu(a, *, streamNone) LU, pivots mx.linalg.lu_factor(a, *, streamNone)lu返回(p, L, U)满足A L[p, :] U2 维或mx.take_along_axis(L, p[..., None], axis-2) U高维批输入。注意p是置换索引而非置换矩阵这与scipy.linalg.lu默认返回置换矩阵不同如需构造完整置换矩阵文档给出了写法P mx.put_along_axis(mx.zeros_like(L), p[..., None], mx.array(1.0), axis-1)lu_factor返回紧凑的(LU, pivots)其中LU同时存放 L、U 因子pivots为uint32类型的行置换索引适合多次复用以减少重复分解。两者都要求至少 2 维、float32/float64且必须在 CPU 流上执行。实现上lu在 mlx/linalg.cpp 中由lu_factor的结果派生出L tril(LU, -1) eye、U triu(LU, 0)并对非方阵裁剪测试见test lu含 2x2、3x3 与批处理维度。4. cholesky —— Cholesky 分解L mx.linalg.cholesky(a, upperFalse, *, streamNone)对实对称正定文档措辞为 positive semi-definite矩阵返回三角因子upperFalse时L L.T aupperTrue时U.T U a。若输入非对称正定行为未定义源码 mlx/backend/cpu/cholesky.cpp 中 LAPACKpotrf返回非零info时仅对info 0抛错注释说明正定校验错误暂不抛出以免崩溃。这是模块中唯一在 CUDA 可用时允许 GPU 执行的分解check_cpu_or_cuda_stream。import mlx.core as mx A mx.array([[4.0, 2.0], [2.0, 3.0]]) L mx.linalg.cholesky(A) # L ≈ [[2, 0], [1, 1.41421]]五、求逆一族inv、tri_inv、pinv 与 cholesky_inv1. inv —— 矩阵求逆ainv mx.linalg.inv(a, *, streamNone)要求方阵至少 2 维返回a ainv ainv a I。实现上 mlx/backend/cpu/inverse.cpp 先做 LAPACKgetrfLU 分解再用getri完成求逆该文件注释说明借助(A⁻¹)ᵀ (Aᵀ)⁻¹恒等式规避列主序转置开销。只支持float32/float64且必须在 CPU 流运行。2. tri_inv —— 三角矩阵求逆ainv mx.linalg.tri_inv(a, upperFalse, *, streamNone)专门求三角矩阵的逆upperTrue表示上三角。底层调用 LAPACKtrtri并在求逆后显式把另一半三角清零见 mlx/backend/cpu/inverse.cpp 的tri_inv。它是solve_triangular与cholesky_inv的基石。3. pinv —— Moore-Penrose 伪逆aplus mx.linalg.pinv(a, *, streamNone)对任意形状矩阵含奇异矩阵计算广义逆满足a aplus a a。实现基于 SVD见 mlx/linalg.cpp 的pinv先做完整 SVD然后按 cutoff 阈值rcond 10 * max(m, n) * eps(dtype)过滤小奇异值对超过 cutoff 的奇异值取倒数rS where(S cutoff, 1/S, 0)最后V diag(rS) U^T。对空数组a.size() 0直接返回形状为输入转置形状的零矩阵与 NumPy 行为一致。C 测试覆盖方阵、m n与m n三种情形。4. cholesky_inv —— 利用 Cholesky 因子的求逆ainv mx.linalg.cholesky_inv(L, upperFalse, *, streamNone)输入是 Cholesky 因子L而非原矩阵A返回A⁻¹其中A L L.T。实现mlx/linalg.cpp先tri_inv求三角逆再按upper方向做一次matmul重组upperTrue时L_inv L_inv.T否则L_inv.T L_inv。相比直接inv(A)对对称正定问题通常计算量更低、数值更稳。文档同时提示若输入不是三角矩阵行为未定义。六、特征值分解eig / eigvals / eigh / eigvalsh这四个函数都要求方阵且至少 2 维必须在 CPU 流上执行validate_eig统一校验。1. eig / eigvals —— 一般矩阵允许复数w, v mx.linalg.eig(a, *, streamNone) # w: 特征值, v: 特征向量 w mx.linalg.eigvals(a, *, streamNone) # 仅特征值与 NumPy 的关键差异eig/eigvals的返回类型始终是complex64即使特征值全部为实数源码中输出 dtype 固定为complex64见 mlx/linalg.cppeig返回(w, v)列v[:, i]是对应第 i 个特征值的归一化右特征向量支持float32/float64/complex64输入底层调用 LAPACKgeevmlx/backend/cpu/eig.cpp对实矩阵会将成对共轭复特征值及其特征向量正确重组为复数输出。A mx.array([[1., -2.], [-2., 1.]]) w, v mx.linalg.eig(A, streammx.cpu) # w ≈ array([30j, -10j], dtypecomplex64) # v ≈ array([[0.7071070j, 0.7071070j], [-0.7071070j, 0.7071070j]], dtypecomplex64)2. eigh / eigvalsh —— 对称Hermitian矩阵w, v mx.linalg.eigh(a, UPLOL, *, streamNone) w mx.linalg.eigvalsh(a, UPLOL, *, streamNone)输入须为实对称或复 Hermitian 矩阵特征值按升序排列而eig/eigvals不保证顺序UPLO参数默认L指定使用矩阵的上三角U还是下三角L源码注释明确“假定输入对称仅使用所选三角不进行对称性检查”对复 Hermitian 输入特征值以float32返回底层调用 LAPACKsyevd实/heevd复分治算法见 mlx/backend/cpu/eigh.cpp 与 mlx/backend/cpu/lapack.h 中的INSTANTIATE_LAPACK_REAL(syevd)、INSTANTIATE_LAPACK_COMPLEX(heevd)。A mx.array([[1., -2.], [-2., 1.]]) w, v mx.linalg.eigh(A, streammx.cpu) # w ≈ array([-1., 3.], dtypefloat32) # v ≈ array([[ 0.707107, -0.707107], [ 0.707107, 0.707107]], dtypefloat32)七、行列式det 与 slogdetd mx.linalg.det(a, *, streamNone) sign, logabsdet mx.linalg.slogdet(a, *, streamNone)均要求方阵、至少 2 维、不支持复数输入且需 CPU 流slogdet返回(sign, logabsdet)sign ∈ {-1, 0, 1}logabsdet为行列式绝对值的自然对数奇异矩阵时sign 0、logabsdet -inf可用det sign * exp(logabsdet)重建行列式对数值极大/极小的行列式比直接det更稳定实现细节mlx/linalg.cppn 3时走det_raw_small解析式快速路径1x1、2x2 直接乘减3x3 按展开式避免 log/exp 往返n 3时走 LU 路径slogdet_impllu_factor取 U 的对角线用 pivot 错位计数pivot[i] ! i的个数与负对角元个数共同决定符号对角线含 0 则标记奇异det则对n 3由sign * exp(logabsdet)合成。A mx.array([[1., 2.], [3., 4.]]) mx.linalg.det(A, streammx.cpu) # array(-2, dtypefloat32) sign, logabsdet mx.linalg.slogdet(A, streammx.cpu) # sign ≈ array(-1), logabsdet ≈ array(0.693147) # ln(2)测试 tests/linalg_tests.cpp 的test det/test slogdet覆盖 1x1/2x2/3x3 快速路径、4x4 LU 路径单位阵、非方阵报错与奇异矩阵语义。八、线性方程组solve 与 solve_triangular1. solve —— 求解 AX Bx mx.linalg.solve(a, b, *, streamNone)返回唯一解x满足A x B。约束validate_solvea至少 2 维且必须方阵b至少 1 维a的最后维必须与b的倒数第二维b为 2 维以上时或最后维b为 1 维向量时匹配输入提升后的类型必须是float32或float64。实现上并非直接求逆而是 mlx/linalg.cpp 中三步走先lu(a)得到(pivots, L, U)再用argsort把置换索引还原为take_along_axis作用到b得到pb最后依次solve_triangular(L, pb, upperFalse)与solve_triangular(U, y, upperTrue)回代求解。批处理如(5, 3, 3)系数 (5, 3, 1)右端与多列右端(3, 2)均有测试覆盖test solve。2. solve_triangular —— 三角系统求解x mx.linalg.solve_triangular(a, b, upperFalse, *, streamNone)求解A X B其中A为三角矩阵upperTrue表示上三角。实现极简tri_inv(a, upper) b先求三角逆再矩阵乘因此对三角系统而言比通用solve更直接。C 测试test solve_triangluar对上下三角各给出一组解析解验证。九、cross沿轴叉积c mx.linalg.cross(a, b, axis-1, *, streamNone)计算两个数组沿指定轴默认-1的叉积该轴尺寸必须是 2 或 3尺寸为 2 时第三分量按 0 处理支持广播broadcast_shapes输出类型为两输入提升后的类型实现mlx/linalg.cpp将两输入沿轴split后按分量组合两个 3 维向量做标准的(a1*b2 - a2*b1, a2*b0 - a0*b2, a0*b1 - a1*b0)并针对输入为 2 维的情形补零分量。a mx.array([1., 2., 3.]) b mx.array([4., 5., 6.]) mx.linalg.cross(a, b) # array([-3, 6, -3])十、精度与测试保障如何验证这些实现模块的正确性由三层测试共同保障C 单元测试tests/linalg_tests.cpp逐函数验证数学性质——qr后Q正交、svd后U[:, :k] diag(S) Vt ≈ A且norm(S) norm(A, fro)、inv后A A⁻¹ I、cholesky后L L.T A与U.T U A、pinv满足两条 Moore-Penrose 条件、eigh特征向量正交且满足A v λ v、lu后L[P, :] U A、solve后A x b、det/slogdet的解析结果Python 测试python/tests/test_linalg.pytest_norm将mx.linalg.norm与np.linalg.norm在多种shape × ord × axis × keepdims组合下逐一对比atol1e-5, rtol1e-6test_complex_norm覆盖复数输入其余分解类测试同样以 NumPy 为参照双精度测试python/tests/test_double.py覆盖qr、svd、inv、tri_inv、cholesky、pinv、eigh、eigvalsh、lu、solve_triangular在float64下的行为MLX 对float64的支持可参考 docs/src/python/data_types.rst。十一、常见坑位与实战建议结合源码中的校验逻辑与测试归纳如下使用注意点务必传streammx.cpu除norm、cross、choleskyCUDA 可用时外其余函数在 GPU 流上会直接抛错批量场景建议先构造 CPU 流再批量调用dtype 先行提升qr/inv/cholesky等拒绝整型输入若从mx.arange等默认整型数组构造先用.astype(mx.float32)或mx.array(..., dtypemx.float32)提升维度与形状所有矩阵函数要求ndim 2inv、cholesky、特征值、行列式要求方阵norm的axis最多 2 个cross的轴尺寸必须为 2 或 3eig系列恒为复数输出即使实对称矩阵eig/eigvals也返回complex64需要实数特征值且矩阵对称请改用eigh/eigvalsh升序排列UPLO控制取上/下三角lu的 pivot 是索引重建时用take_along_axis(L, p[..., None], axis-2) U高维或直接L[p, :] U2 维不要把它当作置换矩阵直接相乘大规模 / 病态矩阵优先svd或pinv内置 cutoff 阈值10 * max(m,n) * eps需要对数尺度判定奇异程度时用slogdet。从实现层面看MLX 的 linalg 模块在 C 侧以 mlx/linalg.h 声明、mlx/linalg.cpp 实现Python 侧通过 python/src/linalg.cpp 的 nanobind 绑定暴露为mlx.core.linalg分解内核统一走 LAPACKmlx/backend/cpu/lapack.h 中封装了geqrf、orgqr、gesdd、potrf、getrf、getri、trtri、syevd、heevd、geev等例程同时 mlx/backend/cuda/primitives.cpp 中NO_GPU(Inverse)与NO_GPU_MULTI(LUF/QRF/SVD/Eig/Eigh)标记了这些算子当前的 CUDA 支持状态即 CUDA 后端同样暂未实现仅cholesky通过 cusolver 路径可用。了解这条“Python 绑定 → C 组合层 → LAPACK 内核 → 后端注册”的调用链有助于你在遇到设备或类型错误时快速定位问题来源。【免费下载链接】mlxMLX: An array framework for Apple silicon项目地址: https://gitcode.com/GitHub_Trending/ml/mlx创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考