如何优化NB类中get_bs函数的嵌套for循环?求紧凑写法示例
优化NB类中
get_bs函数的嵌套循环 问题描述
我希望重写NB类中的get_bs函数,优化其中的嵌套for循环使其更紧凑。考虑将嵌套拆分为两个单行嵌套循环(例如for i,j in...的形式),这是否可行?恳请提供示例及相关建议,谢谢。
原get_bs函数代码:
def get_bs(self): sigSumYes = [0] * self.num_elements sigSumNo = [0] * self.num_elements for i in range(self.num_elements - 1): for j in range(self.num_diabetesyes): diff_yes = self.diabetes_yes[i][j] - self.a_diabetesyes[i] sigSumYes[i] += m.pow(diff_yes, 2) self.b_diabetesyes[i] = m.sqrt(sigSumYes[i] / (len(self.diabetes_yes[i]) - 1)) for j in range(self.num_diabetesno): diff_no = self.diabetes_no[i][j] - self.a_diabetesno[i] sigSumNo[i] += m.pow(diff_no, 2) self.b_diabetesno[i] = m.sqrt(sigSumNo[i] / (len(self.diabetes_no[i]) - 1))
解答:可行且能大幅简化代码
你的思路完全可行,而且我们可以通过生成器表达式/列表推导式结合内置sum函数来彻底替换内层循环,让代码更紧凑、可读性更强,同时利用Python内置函数的优化提升性能。
方式1:替换内层循环为生成器求和(最简洁)
我们可以把每个特征下的平方差求和逻辑用一行生成器表达式实现,避免手动遍历j索引。这样整个get_bs函数可以简化为:
def get_bs(self): for i in range(self.num_elements - 1): # 处理diabetes_yes的标准差 mean_yes = self.a_diabetesyes[i] sum_sq_diff_yes = sum(m.pow(x - mean_yes, 2) for x in self.diabetes_yes[i]) self.b_diabetesyes[i] = m.sqrt(sum_sq_diff_yes / (len(self.diabetes_yes[i]) - 1)) # 处理diabetes_no的标准差 mean_no = self.a_diabetesno[i] sum_sq_diff_no = sum(m.pow(x - mean_no, 2) for x in self.diabetes_no[i]) self.b_diabetesno[i] = m.sqrt(sum_sq_diff_no / (len(self.diabetes_no[i]) - 1))
方式2:用enumerate同时遍历索引和数据(对应你提到的for i,j in...形式)
如果你想显式遍历每个特征的索引和对应的数据列表,可以用enumerate搭配zip,把索引和两组数据同时取出,代码结构更直观:
def get_bs(self): # 遍历每个特征的索引i,以及对应的yes/no数据列表 for i, (yes_data, no_data) in enumerate(zip(self.diabetes_yes[:-1], self.diabetes_no[:-1])): mean_yes = self.a_diabetesyes[i] sum_sq_diff_yes = sum(m.pow(x - mean_yes, 2) for x in yes_data) self.b_diabetesyes[i] = m.sqrt(sum_sq_diff_yes / (len(yes_data) - 1)) mean_no = self.a_diabetesno[i] sum_sq_diff_no = sum(m.pow(x - mean_no, 2) for x in no_data) self.b_diabetesno[i] = m.sqrt(sum_sq_diff_no / (len(no_data) - 1))
这里用[:-1]截取前n-1组数据,和原代码中range(self.num_elements -1)的逻辑保持一致。
额外优化建议
- 避免重复计算长度:原代码中多次调用
len(self.diabetes_yes[i]),可以提前赋值给变量,减少重复计算开销。 - 简化数学运算:如果
m是math模块,(x - mean)**2比m.pow(x - mean, 2)更简洁直观,性能也相当。 - 移除冗余变量:原代码中的
sigSumYes和sigSumNo数组完全可以去掉,因为我们直接通过生成器求和得到结果,不需要预先初始化累加数组。 - 类型一致性检查:确保
self.a_diabetesyes[i]和diabetes_yes[i]中的元素类型一致(比如都是Decimal或float),避免不必要的类型转换。
结合以上建议的最终优化版本:
def get_bs(self): for i in range(self.num_elements - 1): yes_data = self.diabetes_yes[i] count_yes = len(yes_data) mean_yes = self.a_diabetesyes[i] sum_sq_diff_yes = sum((x - mean_yes)**2 for x in yes_data) self.b_diabetesyes[i] = m.sqrt(sum_sq_diff_yes / (count_yes - 1)) no_data = self.diabetes_no[i] count_no = len(no_data) mean_no = self.a_diabetesno[i] sum_sq_diff_no = sum((x - mean_no)**2 for x in no_data) self.b_diabetesno[i] = m.sqrt(sum_sq_diff_no / (count_no - 1))
这个版本既保留了逻辑清晰性,又彻底消除了嵌套的for j循环,代码更紧凑且性能更优。
内容的提问来源于stack exchange,提问作者Rafiul Nakib
相关产品推荐
相关产品推荐

