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

高级 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 是调度员、回调是巡检员。出问题时先判断"该找哪个岗位",排查路径立刻清晰。
下节把"建模"这一步展开成三种方式。