argparse需在if name == "__main__":下尽早解析参数,用于配置设备(os.environ/CUDA_VISIBLE_DEVICES须在import tensorflow前)、数据集、模型及Checkpoint恢复;所有超参必须显式源自args,避免硬编码,并验证实际生效。

argparse 怎么和 TensorFlow 训练脚本一起用
直接在训练脚本开头加 argparse.ArgumentParser,把超参、路径、设备等变量全变成命令行可传的参数。TensorFlow 本身不干涉参数解析,关键是你得在 tf.function 或训练循环前把参数读进来并传进去。
常见错误是把 argparse 放在模型定义之后、或塞进 @tf.function 里——这会报 ValueError: Cannot convert a symbolic Tensor,因为 argparse 是 Python 运行时逻辑,不能进图构建阶段。
- 所有
argparse相关代码必须在if __name__ == "__main__":下,且在tf.data.Dataset构建、模型实例化之前执行 - 避免用
args直接构造tf.constant或参与图结构定义(比如tf.range(args.batch_size));改用tf.TensorShape或运行时传入input_signature - 推荐把
args封装成配置字典,再传给训练函数,方便后续扩展(比如加 wandb 或 tensorboard 日志路径)
训练时怎么指定 GPU / CPU 和内存限制
TensorFlow 默认自动选 GPU,但命令行参数可以强制控制设备和显存。注意:这不是 argparse 能直接搞定的,得靠 os.environ 或 tf.config 配合。
典型错误是先 import tensorflow 再设 os.environ["CUDA_VISIBLE_DEVICES"] ——必须在 import tensorflow 之前设置,否则无效,GPU 还是会被自动占用。
立即学习“Python免费学习笔记(深入)”;
快速生成专业的 Python 脚本和应用代码。一键创建完整项目结构,支持CLI、API、爬虫、Bot、Django等多种项目类型,包含完整的项目结构、配置文件、依赖管理、测试、README和文档。
- 在
argparse解析完后、import tensorflow前插入:os.environ["CUDA_VISIBLE_DEVICES"] = str(args.gpu_id)(值为-1表示禁用 GPU) - 若需限制显存增长,用
tf.config.experimental.set_memory_growth,但只能对已识别的 GPU 生效,所以得在os.environ设置后、import 后立即调用 - CPU 线程数可通过
tf.config.threading.set_intra_op_parallelism_threads控制,适合在多核机器上压测数据加载瓶颈
怎么让训练脚本支持 --resume 和 --checkpoint_dir
恢复训练不是简单 load_weights 就完事,得同步恢复 optimizer 状态、epoch 计数、学习率调度器。TensorFlow 的 tf.train.Checkpoint 是标准方案,但命令行参数要能驱动它。
容易踩的坑是 checkpoint 路径拼错,或者没把 optimizer 和 model 一起塞进同一个 Checkpoint 实例,导致恢复后 optimizer 状态丢失,loss 突然飙升。
-
--checkpoint_dir必须是完整路径(如./logs/exp01),不要带文件名;Checkpoint会自动找最新的ckpt-xxx.index - 恢复时先创建
Checkpoint实例,再调用restore(checkpoint_manager.latest_checkpoint);如果latest_checkpoint是None,说明没找到,就跳过恢复 - 务必检查
restore()返回的status对象是否assert_consumed()成功,否则可能部分变量没加载上(尤其模型结构有改动时)
训练脚本跑起来后怎么验证参数真被用了
别靠猜。最直接的办法是在 main 函数开头加一句 print(f"Using batch_size={args.batch_size}, lr={args.lr}"),再配合日志输出实际初始化的模型参数量、dataset.cardinality() 值、甚至 tf.config.list_physical_devices("GPU") 结果。
更隐蔽的问题是某些参数看似传进去了,但被硬编码覆盖了。比如脚本里写了 batch_size = 32,又从 args 读了 args.batch_size 却没赋值过去。
- 所有超参变量必须显式来源于
args,禁止出现裸数字或字符串字面量替代(除非是调试用的固定 fallback) - 建议加个
args_dict = vars(args)打印,确认 argparse 没把类型搞错(比如把"1e-4"当字符串而不是 float) - 如果用了
tf.keras.callbacks.ModelCheckpoint,它的filepath字符串里也得插args.run_id之类,否则多个实验会互相覆盖

















