3.1 Keras 与排程室的分工:高级 API 的执行真相


文档摘要

3.1 Keras 与排程室的分工:高级 API 的执行真相 本节摘要:Keras 不是与 TensorFlow 并列的另一个框架,而是 TensorFlow 官方的高级建模接口:层负责声明前向结构,compile 负责装配训练规程,fit 负责逐批调度执行。本节用加州房价回归跑通第一条装配线,讲清高级 API 各组件与第 1 章机制(张量、变量、图、求导)的一一对应关系——知道它包了什么、没包什么,才能判断什么时候该往下钻。 学习目标 阅读完本节,你应当能够: 说出 Keras 层、模型、损失、优化器四个组件各自封装了第 1 章的哪个机制; 用最少代码在加州房价上完成一次完整训练并解读输出; 判断"这个需求 fit 能不能直接做",给出往下钻的判断标准。

3.1 Keras 与排程室的分工:高级 API 的执行真相

本节摘要:Keras 不是与 TensorFlow 并列的另一个框架,而是 TensorFlow 官方的高级建模接口:层负责声明前向结构,compile 负责装配训练规程,fit 负责逐批调度执行。本节用加州房价回归跑通第一条装配线,讲清高级 API 各组件与第 1 章机制(张量、变量、图、求导)的一一对应关系——知道它包了什么、没包什么,才能判断什么时候该往下钻。

学习目标

阅读完本节,你应当能够:

  1. 说出 Keras 层、模型、损失、优化器四个组件各自封装了第 1 章的哪个机制;
  2. 用最少代码在加州房价上完成一次完整训练并解读输出;
  3. 判断"这个需求 fit 能不能直接做",给出往下钻的判断标准。

两行命令包掉了什么

先看最短路径:加载第 2 章已标准化的加州房价数据,建模、编译、训练、评估,总共四行核心调用。

import tensorflow as tf import numpy as np from sklearn.datasets import fetch_california_housing housing = fetch_california_housing() X = housing.data.astype("float32") y = housing.target.astype("float32") model = tf.keras.Sequential([ tf.keras.layers.Input(shape=(8,)), tf.keras.layers.Dense(64, activation="relu"), tf.keras.layers.Dense(1), ]) model.compile(optimizer=tf.keras.optimizers.Adam(0.001), loss="mse") history = model.fit(X, y, epochs=5, batch_size=64, validation_split=0.2, verbose=0) print(f"final loss: {history.history['loss'][-1]:.4f}") # 输出示例:final loss: 0.4917 # 数值随初始化浮动,五轮内从约 1.2 降到 0.5 上下 loss, = model.evaluate(X[-4128:], y[-4128:], verbose=0) print(f"eval loss: {loss:.4f}") # 输出示例:eval loss: 0.5053

这四行在排程室里的完整剧本:Sequential 把三个层登记成前向图,每个 Dense 在首次调用时创建权重 Variable(1.6 节的机制);compile 把损失与优化器装配进一张"训练步规程";fit 用 1.5 节的 tf.function 把训练步跟踪成图,然后逐批执行——前向调 Dense 链算预测,GradientTape 沿图回放求梯度,优化器 assign 写变量;evaluate 在不更新变量的模式下跑同一套前向。

图 12 四行代码与四层机制的对应

图 12 四行代码与四层机制的对应

边界:fit 能做什么、不能做什么

高级 API 的边界感决定你的技术路线。fit 直接支持的:单输入单输出、多输入多输出(传字典或元组)、数据集输入、验证集、类权重、回调。它不便支持的:单个批次内的复杂逻辑(如对抗训练的交替更新)、非标准损失组合的精细控制、多优化器并行。判断口诀:流程是"逐批前向求导更新"这个形状的,fit 能装;形状不同的训练流程,去 5.1 节手写。

# 边界实验:验证 fit 确实在图模式下跑 import time big_x = np.random.rand(50000, 8).astype("float32") big_y = np.random.rand(50000).astype("float32") model.compile(optimizer="adam", loss="mse", run_eagerly=False) # 默认图模式 t0 = time.time() model.fit(big_x, big_y, epochs=1, batch_size=64, verbose=0) graph_mode_s = time.time() - t0 model.compile(optimizer="adam", loss="mse", run_eagerly=True) # 强制逐行执行 t0 = time.time() model.fit(big_x, big_y, epochs=1, batch_size=64, verbose=0) eager_mode_s = time.time() - t0 print(f"graph {graph_mode_s:.1f}s vs eager {eager_mode_s:.1f}s") # 参考输出:graph 1.4s vs eager 8.2s # run_eagerly 是理解 1.5 节的实验开关,生产环境保持默认 False

何时往下钻

三个信号提示该往下钻了。信号一,需要看中间量:调试时想知道某层的输出分布、梯度的范数,别在 fit 外围猜,用 3.7 节自定义回调或 5.1 节手写循环。信号二,训练形状不同:GAN 的两个网络交替更新、对比学习的成对采样,fit 装不下。信号三,性能细节:混合精度、自定义梯度,都要绕过高级 API 直接操作图。反过来说,只要 fit 还装得下,就别手写——高级 API 少写的是代码,更是出错面。

⚠️ 常见坑:在 fit 之前忘了 compile,报错信息直指模型未编译,但新手常去查数据形状。看到 "compile your model" 字样,先补 compile 再查别的。

💡 关键直觉:把 Keras 组件当成排程室的岗位表——层是登记员、compile 是规程起草人、fit 是调度员、回调是巡检员。出问题时先判断"该找哪个岗位",排查路径立刻清晰。

本节要点回顾

  • Keras 是官方高级接口:层、compile、fit 分别封装图登记、规程装配、逐批调度。
  • 四行最短路径:建模、编译、训练、评估,背后是第 1 章全套机制。
  • run_eagerly 开关:性能对比证明 fit 默认在图模式执行。
  • 边界口诀:标准形状的训练 fit 能装,形状不同去手写循环。
  • 三个下钻信号:看中间量、非标流程、性能细节。

下节把"建模"这一步展开成三种方式。


作者与出处
原作者: 灏天文库
来源:灏天文库
整理: 灏天文库整理
由灏天文库平台收录,内容或由平台用户上传,仅供学习交流
发布者: 作者: 灏天文库 转发
评论区 (0)
U