如何基于Meshgrid通用计算多变量二次多项式的值?
问题描述
我有几个变量x、y、z:
import numpy as np # Variables x = np.linspace(-1.0, 1.0, num=5, endpoint=True) y = np.linspace(-1.0, 1.0, num=5, endpoint=True) z = np.linspace(-1.0, 1.0, num=5, endpoint=True) # And the corresponding meshgrid x_v, y_v, z_v = np.meshgrid(x, y, z)
想要计算系数为C的二次多项式在Meshgrid每个点上的值。为此编写了一个通用函数,可返回任意数量参数对应的多项式项:
def pol_terms(*args): # -- Calculate terms of the polynomial -- # Linear part (b0, b11*x1, b12*x2, ..., b1k*xk, where k=len(args)) entries = [1] + [pt for pt in args] # Combinations part (b12*x12, b13*x13, ..., bk-1k*xk-1k) n_args = len(args) for i in range(n_args): for j in range(i+1, n_args): entries.append(args[i]*args[j]) # Quadratic part entries += [pt**2 for pt in args] return np.array([entries])
对于单组变量值,可以通过以下方式计算多项式值:
X = pol_terms(1,2,3) X.dot(C)
但不知道如何对Meshgrid生成的x_v、y_v、z_v执行同样操作。对pol_terms使用np.vectorize无效,因为该函数返回数组;也不想手动遍历所有维度,希望方案具备通用性(支持2到5个变量)。
解决方案
核心思路是修改多项式项生成函数,让它直接支持数组输入,利用NumPy的广播机制批量处理所有网格点,避免循环和低效的np.vectorize。
修改后的多项式项生成函数
import numpy as np def pol_terms_array(*args): # 确保所有输入数组形状一致 shape = args[0].shape # 常数项:生成与输入同形状的全1数组 entries = [np.ones(shape)] # 一次项:直接加入输入的变量数组 entries.extend(args) # 交叉项:遍历所有i<j的组合,计算对应变量的乘积 n_args = len(args) for i in range(n_args): for j in range(i+1, n_args): entries.append(args[i] * args[j]) # 二次项:每个变量的平方 entries.extend([pt**2 for pt in args]) # 将所有项堆叠成 (n_terms, ...) 形状的数组,方便后续点积计算 return np.stack(entries, axis=0)
计算所有网格点的多项式值
假设系数C是长度与多项式项数匹配的一维数组(比如3个变量时,项数为1+3+3+3=10,则C长度为10),执行以下代码即可:
# 生成多项式项数组,形状为 (n_terms, 5,5,5) terms = pol_terms_array(x_v, y_v, z_v) # 计算每个点的多项式值,结果形状为 (5,5,5) poly_values = np.tensordot(C, terms, axes=1)
正确性验证
可以选取单个网格点与批量计算结果对比:
# 取某单个网格点的变量值 single_x, single_y, single_z = x_v[0,0,0], y_v[0,0,0], z_v[0,0,0] # 用原函数计算单个点值 single_point_terms = pol_terms(single_x, single_y, single_z) single_value = single_point_terms.dot(C) # 对比批量计算的对应位置值 assert np.isclose(single_value, poly_values[0,0,0])
通用性说明
该函数支持2-5个变量的场景,只要输入的网格数组形状一致即可自动适配。比如处理2个变量的情况:
x2 = np.linspace(-1,1,5) y2 = np.linspace(-1,1,5) x2_v, y2_v = np.meshgrid(x2, y2) # 2个变量时,项数为1+2+1+2=6,C_2d需为长度6的一维数组 terms_2d = pol_terms_array(x2_v, y2_v) poly_2d = np.tensordot(C_2d, terms_2d, axes=1)
效率说明
修改后的函数完全基于NumPy向量化操作,比手动循环或np.vectorize效率高得多,尤其适合大尺寸网格数据的计算。
内容的提问来源于stack exchange,提问作者Roman J.
相关产品推荐
相关产品推荐

