You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何获取对应最小总误差的Weight1、Weight2、bias1、bias2

寻找最小总误差对应的权重与偏置

先理清楚你的场景:

  • Weight1、Weight2、bias1、bias2都是随机生成的参数列表,存储结构如下:
    • list1= [[(list of Weight1)], [(list of Weight2)], [(list of bias1)], [list of bias2]]
    • list2= [[(list of Weight1)], [(list of Weight2)], [(list of bias1)], [list of bias2]]
    • list3= [[(list of Weight1)], [(list of Weight2)], [(list of bias1)], [list of bias2]]
  • popSize=3

你需要找到这三组参数中,对应tot_error最小的那一组Weight1、Weight2、bias1、bias2对吧?看你提供的代码,目前的问题是只对误差值做了排序,但没把误差和对应的参数关联起来,所以没法定位到具体是哪一组参数。我来帮你调整代码,实现需求:

import numpy as np

def findGStar(Weight1, Weight2, bias1, bias2):
    z1 = X_trainNorm.dot(Weight1) + bias1
    a1 = np.tanh(z1)
    z2 = a1.dot(Weight2) + bias2
    target = np.reshape(y_trainNorm, (-1, 1))
    # 用NumPy原生方法计算总绝对误差,比循环求和高效得多
    tot_error = np.sum(np.abs(z2 - target))
    return tot_error

# 先把你的三个参数列表组合成一个数组(假设你还没做这一步)
vector = [list1, list2, list3]
popSize = 3

# 存储每个参数组对应的(总误差,参数组)元组
error_param_pairs = []
for i in range(popSize):
    current_params = vector[i]
    error = findGStar(current_params[0], current_params[1], current_params[2], current_params[3])
    error_param_pairs.append( (error, current_params) )

# 按总误差从小到大排序
error_param_pairs.sort(key=lambda item: item[0])

# 取出最小误差对应的参数组
min_total_error = error_param_pairs[0][0]
best_params = error_param_pairs[0][1]

# 拆解出对应的权重和偏置
best_Weight1 = best_params[0]
best_Weight2 = best_params[1]
best_bias1 = best_params[2]
best_bias2 = best_params[3]

# 输出结果
print(f"最小总误差值: {min_total_error}")
print(f"对应的Weight1: {best_Weight1}")
print(f"对应的Weight2: {best_Weight2}")
print(f"对应的bias1: {best_bias1}")
print(f"对应的bias2: {best_bias2}")

几个关键的调整说明:

  1. 关联误差与参数:原来的代码只收集了误差值,排序后完全不知道这个最小误差对应哪组参数。现在我们把误差和参数打包成元组,排序后依然能一一对应。
  2. 优化误差计算:用np.sum(np.abs(z2 - target))替代手动循环求和,既简洁又高效,这也是NumPy的标准用法。
  3. 清晰的参数提取:排序后直接从第一个元素中拿到最小误差和对应的参数组,再拆解出你需要的四个参数,逻辑清晰易懂。

这样修改后,就能精准定位到总误差最小的那一组权重和偏置了。

内容的提问来源于stack exchange,提问作者Hana

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.27 06:52:59