如何在Python中创建数学函数类实现优化算法可视化?
错误根因
你遇到的两个报错分别对应两个基础问题:
- 类的普通实例方法默认会把实例自身作为第一个参数传入,你定义的
objective1n1没有加self参数,调用fcts.fct_1(5)时,实际传入的参数是(类实例, 5),就会触发参数数量不匹配的报错。 - 初始化类时你直接执行了函数
self.objective1n1(x),此时变量x未定义,自然会报变量不存在的错误。你需要的是把函数本身赋值给实例属性,而不是提前执行函数获取计算结果。
修复方案
给不需要访问实例属性的目标函数加上@staticmethod装饰器,就可以避免自动传入self参数;初始化时直接赋值函数对象,不要加括号执行即可。修复后完整的可运行代码如下,已经封装了你需要的不同维度绘图、切换函数的能力:
import numpy as np import matplotlib.pyplot as plt from mpl_toolkits.mplot3d import Axes3D class FunctionVisualisation: # 静态方法不需要访问实例属性,不会自动传入self参数 @staticmethod def objective1n1(x): return x**5.0 - 2*x**4 - 0.5*x**3 + 4*x**2 - x @staticmethod def objective2n1(x, y): return x**5.0 - 2*y**4 - 15*x**3 + 4*y**2 - x def __init__(self): # 直接赋值函数对象,不要加括号执行 self.fct_1 = self.objective1n1 self.fct_2 = self.objective2n1 # 存储当前选中的函数、维度配置,方便后续切换 self.current_func = None self.dim = None # 封装1维函数绘图方法 def plot_1d(self, func=None, x_range=(-5, 5), num_points=100): self.current_func = func if func is not None else self.fct_1 self.dim = 1 x = np.linspace(x_range[0], x_range[1], num_points) y = self.current_func(x) plt.figure(figsize=(8,4)) plt.plot(x, y) plt.xlabel('x') plt.ylabel('f(x)') plt.grid(True) plt.show() # 封装2维函数绘图方法 def plot_2d(self, func=None, x_range=(-5,5), y_range=(-5,5), num_points=50): self.current_func = func if func is not None else self.fct_2 self.dim = 2 x = np.linspace(x_range[0], x_range[1], num_points) y = np.linspace(y_range[0], y_range[1], num_points) X, Y = np.meshgrid(x, y) Z = self.current_func(X, Y) fig = plt.figure(figsize=(10,6)) ax = fig.add_subplot(111, projection='3d') ax.plot_surface(X, Y, Z, cmap='viridis') ax.set_xlabel('x') ax.set_ylabel('y') ax.set_zlabel('f(x,y)') plt.show()
使用示例
# 实例化类 fcts = FunctionVisualisation() # 测试单个点的函数调用 print(fcts.fct_1(5)) # 输出计算结果 395.0 print(fcts.fct_2(1,2)) # 输出计算结果 -59.0 # 绘制默认1维函数图像 fcts.plot_1d() # 绘制默认2维函数图像 fcts.plot_2d() # 切换自定义函数也很方便,比如新增一个二次函数 def custom_1d(x): return x**2 + 2*x + 1 fcts.plot_1d(func=custom_1d)
关于numpy向量化计算的说明
你看到的numpy示例是提前生成了x的采样点数组,再向量化计算结果,上面的实现已经完全兼容这种用法:
x = np.linspace(-4*np.pi,4*np.pi,50) y = fcts.fct_1(x) # 直接批量计算所有采样点的函数值,和你参考的示例效果完全一致
内容的提问来源于stack exchange,提问作者lpnorm
相关产品推荐
相关产品推荐

