在TensorFlow中求解指定四次多项式方程根的高效方法问询
在TensorFlow中求解四次多项式方程的高效方案
问题明确
需要对给定张量r的每个元素,求解以下四次方程的根:
k4 * x**4 + k3 * x**3 + k2 * x**2 + k1 * x - r = 0
其中固定系数为:
k1 = 339.749 k2 = -31.988 k3 = 48.275 k4 = -7.201
最终输出一个张量,每个元素对应r中元素所定义方程里实部最小的解的实部。
实现思路
四次方程存在解析求根公式(费拉里方法),结合TensorFlow的向量化运算特性,可避免循环、提升计算效率,核心步骤如下:
1. 标准化方程
先将原方程转换为首一多项式(最高次项系数为1),形式为:
$$x^4 + a x^3 + b x^2 + c x + d = 0$$
其中系数通过固定值计算得到:
- $a = k3 / k4$
- $b = k2 / k4$
- $c = k1 / k4$
- $d = -r / k4$
2. 费拉里方法求解四次方程
通过变量替换消去三次项,构造并求解三次预解方程,再将四次方程分解为两个二次方程求解,最后转换回原变量的根。
3. 提取目标结果
对每个r元素对应的所有根,提取其实部后筛选出最小值,作为最终结果。
TensorFlow完整代码实现
import tensorflow as tf # 固定系数 k1 = 339.749 k2 = -31.988 k3 = 48.275 k4 = -7.201 # 预计算首一多项式的固定系数 a = tf.constant(k3 / k4, dtype=tf.float32) b = tf.constant(k2 / k4, dtype=tf.float32) c = tf.constant(k1 / k4, dtype=tf.float32) a_sq = tf.square(a) a_cu = tf.pow(a, 3) a_qu = tf.pow(a, 4) p = tf.constant(b - 3 * a_sq / 8, dtype=tf.float32) q = tf.constant(c - a * b / 2 + a_cu / 8, dtype=tf.float32) def solve_quartic(r): # 计算首一多项式的d项 d = -r / k4 s = d - a * c / 4 + a_sq * b / 16 - 3 * a_qu / 256 # 构造三次预解方程的系数 cubic_coeffs = tf.stack([ tf.constant(1.0, dtype=tf.float32), p / 2, (tf.square(p) / 16) - (s / 4), -tf.square(q) / 64 ], axis=-1) # 求解三次方程,筛选实根(处理数值误差,虚部接近0视为实根) cubic_roots = tf.math.polynomial.roots(cubic_coeffs) real_mask = tf.abs(tf.math.imag(cubic_roots)) < 1e-6 z0 = tf.math.real(tf.boolean_mask(cubic_roots, real_mask))[:, 0] z0 = tf.expand_dims(z0, axis=-1) # 计算二次方程的系数,避免除以0 sqrt_term = tf.sqrt(2 * z0 + p) sqrt_term = tf.where(tf.abs(sqrt_term) < 1e-6, 1e-6, sqrt_term) q_over_2sqrt = q / (2 * sqrt_term) # 求解第一个二次方程 quad1_coeffs = tf.stack([1.0, sqrt_term, z0 + q_over_2sqrt], axis=-1) quad1_roots = tf.math.polynomial.roots(quad1_coeffs) # 求解第二个二次方程 quad2_coeffs = tf.stack([1.0, -sqrt_term, z0 - q_over_2sqrt], axis=-1) quad2_roots = tf.math.polynomial.roots(quad2_coeffs) # 合并所有根并转换回原变量x all_y_roots = tf.concat([quad1_roots, quad2_roots], axis=-1) all_x_roots = all_y_roots - a / 4 # 提取实部并取最小值 roots_real = tf.math.real(all_x_roots) min_real_part = tf.reduce_min(roots_real, axis=-1) return min_real_part # 测试示例 r_test = tf.constant([0.0, 100.0, -50.0], dtype=tf.float32) result = solve_quartic(r_test) print(result.numpy())
关键注意点
- 数值稳定性:计算平方根和除法时加入小阈值(1e-6),避免除以0或数值溢出,可根据精度需求调整。
- 向量化运算:所有操作基于TensorFlow张量运算,自动支持批量处理
r的多个元素,无需显式循环。 - 结果处理:无需区分实根与复根,直接提取所有根的实部后取最小值即可满足需求。
内容的提问来源于stack exchange,提问作者Merlo98765
相关产品推荐
相关产品推荐

