为daft.PGM类添加偏方法的实现是否符合Python风格?求优化建议
改进Daft PGM自定义节点方法的建议
你的代码已经实现了简化节点添加的核心需求,以下是针对代码健壮性、可维护性和Python规范的改进建议:
1. 修正partialmethod的方法引用
原代码中使用dag.add_node作为partialmethod的目标方法,在类定义阶段dag尚未完全初始化,可能引发潜在问题。建议直接引用父类daft.PGM.add_node:
obsNode = partialmethod(daft.PGM.add_node, aspect=2.2, fontsize=10, plot_params={'facecolor': 'cadetblue'})
2. 遵循PEP8命名规范
将驼峰式命名(如obsNode)改为蛇形命名(如add_obs_node),更符合Python代码风格,提升可读性:
add_obs_node = partialmethod(daft.PGM.add_node, aspect=2.2, fontsize=10, plot_params={'facecolor': 'cadetblue'}) add_dec_node = partialmethod(daft.PGM.add_node, aspect=2.2, fontsize=10, shape="rectangle", plot_params={'facecolor': 'thistle'})
3. 提取公共参数到类常量
将重复的公共参数(aspect=2.2、fontsize=10)定义为类常量,便于统一修改:
class DAG(daft.PGM): DEFAULT_ASPECT = 2.2 DEFAULT_FONTSIZE = 10 add_obs_node = partialmethod(daft.PGM.add_node, aspect=DEFAULT_ASPECT, fontsize=DEFAULT_FONTSIZE, plot_params={'facecolor': 'cadetblue'}) # 其他节点方法同理
4. 简化__init__方法
子类未添加额外初始化逻辑时,可省略__init__方法,Python会自动调用父类的__init__:
class DAG(daft.PGM): # 直接定义节点方法,无需重写__init__
5. 动态生成节点方法(可选)
如果需要新增多种节点类型,可通过字典配置动态生成方法,减少重复代码:
class DAG(daft.PGM): DEFAULT_ASPECT = 2.2 DEFAULT_FONTSIZE = 10 NODE_TYPES = { "obs": {"plot_params": {"facecolor": "cadetblue"}}, "dec": {"shape": "rectangle", "plot_params": {"facecolor": "thistle"}}, "det": {"alternate": True, "plot_params": {"facecolor": "aliceblue"}}, "lat": {"plot_params": {"facecolor": "aliceblue"}} } def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) # 动态生成节点方法 for node_type, params in self.NODE_TYPES.items(): method_name = f"add_{node_type}_node" setattr(self, method_name, partialmethod( daft.PGM.add_node, aspect=self.DEFAULT_ASPECT, fontsize=self.DEFAULT_FONTSIZE, **params ))
6. 支持参数覆盖
partialmethod允许在调用时传入新参数覆盖默认值,例如调整某个观测节点的字体大小:
pgm.add_obs_node("sb", "Start\nBalance", 1, 4, fontsize=12)
改进后的完整代码示例
import matplotlib.pyplot as plt import daft from functools import partialmethod class DAG(daft.PGM): DEFAULT_ASPECT = 2.2 DEFAULT_FONTSIZE = 10 add_obs_node = partialmethod(daft.PGM.add_node, aspect=DEFAULT_ASPECT, fontsize=DEFAULT_FONTSIZE, plot_params={'facecolor': 'cadetblue'}) add_dec_node = partialmethod(daft.PGM.add_node, aspect=DEFAULT_ASPECT, fontsize=DEFAULT_FONTSIZE, shape="rectangle", plot_params={'facecolor': 'thistle'}) add_det_node = partialmethod(daft.PGM.add_node, aspect=DEFAULT_ASPECT, fontsize=DEFAULT_FONTSIZE, alternate=True, plot_params={'facecolor': 'aliceblue'}) add_lat_node = partialmethod(daft.PGM.add_node, aspect=DEFAULT_ASPECT, fontsize=DEFAULT_FONTSIZE, plot_params={'facecolor': 'aliceblue'}) pgm = DAG(node_fc="aliceblue", dpi=150, alternate_style="outer") pgm.add_obs_node("sb", "Start\nBalance", 1, 4) pgm.add_dec_node("ba", "Bet\nAmount", 1, 3) pgm.add_det_node("w", "Winnings", 2.7, 3) pgm.add_lat_node("cf", "Coin\nFlip", 2.7, 2) pgm.add_det_node("nb", "New\nBalance", 2.7, 4) pgm.add_edge("sb", "ba") pgm.add_edge("ba", "w") pgm.add_edge("cf", "w") pgm.add_edge("w", "nb") pgm.add_edge("sb", "nb") pgm.render() plt.show()
输出效果

内容的提问来源于stack exchange,提问作者Adam Fleischhacker
相关产品推荐
相关产品推荐

