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

如何在Keras中手动利用模型权重与偏置完成回归预测?

手动实现Keras回归模型预测及结果对比

1. 批量提取所有层的权重与偏置

先一次性获取所有4层的权重矩阵和偏置向量,方便后续逐层计算:

import numpy as np
import time

# 遍历模型所有层,提取权重(w)和偏置(b)
layers_params = []
for layer in model.layers:
    weights, biases = layer.get_weights()
    layers_params.append((weights, biases))

2. 手动实现逐层前向传播

按照模型的层级顺序,依次计算每一层的输出,注意对应层的激活函数:

def manual_predict(input_data):
    current_output = input_data
    for idx, (w, b) in enumerate(layers_params):
        # 计算线性变换:Z = WX + B
        linear_output = current_output @ w + b
        # 前3层应用ReLU激活,最后一层无激活
        if idx < 3:
            current_output = np.maximum(linear_output, 0)  # ReLU激活实现
        else:
            current_output = linear_output
    return current_output

# 计算手动预测结果
manual_results = manual_predict(x_test)
# 获取Keras模型的预测结果(关闭日志输出)
keras_results = model.predict(x_test, verbose=0)

3. 验证手动计算与Keras预测的一致性

由于浮点运算存在微小精度误差,使用np.allclose而非严格相等来验证:

# 检查结果是否匹配
print("手动计算与Keras预测结果是否一致:", np.allclose(manual_results, keras_results, atol=1e-6))
# 查看最大误差值
print("预测结果最大误差:", np.max(np.abs(manual_results - keras_results)))

4. 统计单个样本的预测耗时

循环遍历测试集的每个样本,记录手动计算的耗时并统计:

sample_time_list = []
for single_sample in x_test:
    # 确保样本形状适配模型输入((1, 173))
    sample_input = single_sample.reshape(1, -1)
    start = time.time()
    # 执行手动预测
    manual_predict(sample_input)
    end = time.time()
    sample_time_list.append(end - start)

# 输出耗时统计
print("单个样本平均预测耗时:", np.mean(sample_time_list), "秒")
print("单个样本最大预测耗时:", np.max(sample_time_list), "秒")
print("单个样本最小预测耗时:", np.min(sample_time_list), "秒")

关键注意点

  • 确保x_test的形状为(样本数量, 173),若维度不符需提前用reshape调整
  • Keras的Dense层计算逻辑为output = activation(dot(input, kernel) + bias),其中kernel就是我们提取的权重矩阵,因此矩阵乘法顺序是input @ weights而非weights @ input
  • 手动实现的ReLU要与Keras默认行为一致(即max(0, x)),避免激活函数差异导致结果偏差

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.05 06:07:44