如何高效调用数组中对应对象的get_log_probability方法并传入匹配特征值
解决方案
基础场景实现(单组x_temp和theta配对)
对于长度相等的x_temp和theta列表,直接用zip配对对应位置的元素,通过列表推导式一行即可完成调用和结果收集:
x_temp = [3.4, 2.1, 3.3, 6.6] theta = [obj1, obj2, obj3, obj4] log_probs = [theta_j.get_log_probability(x_j)[0] for theta_j, x_j in zip(theta, x_temp)]
该写法自动对齐两个列表的0号索引到最后一个索引,完全匹配你需要的一一对应调用逻辑,无需手动维护索引。
原predict函数优化
你现有的三层显式循环可以通过推导式压缩为仅1层遍历输入样本的循环,代码更简洁易读:
import numpy as np def predict(self, X): y_hat = [] for new_x in X: # 批量计算每个类别的总对数概率 prob_classes = [ np.log(self._pi[i]) + sum( theta_j.get_log_probability(new_x_j)[0] for theta_j, new_x_j in zip(theta_row, new_x) ) for i, theta_row in enumerate(self._theta) ] # 选取后验概率最高的类别 y_hat.append(self._classes[np.argmax(prob_classes)]) return y_hat
如果你的get_log_probability方法支持批量传入特征值,还可以进一步优化:提前按特征维度归集所有样本的特征值,批量调用对应theta对象的方法减少函数调用开销,适合样本量较大的场景。
内容的提问来源于stack exchange,提问作者Loco_Vegano
相关产品推荐
相关产品推荐

