- 人工智能
- 深度学习
- 机器学习
- 强化学习
【免费下载链接】TensorLayer
Deep Learning and Reinforcement Learning Library for Scientists and Engineers
本指南以 TensorLayer 的tensorlayer.files模块(对应文档 docs/modules/files.rst)为主线,系统讲解三类核心能力:加载 14 种经典公开数据集、用 hdf5 / npz / npz_dict 等多种格式保存与恢复网络权重、以及文件/文件夹级的工程工具函数。读完本文,你将掌握从"下载数据集 → 训练 → 保存模型 → 恢复模型"的完整数据闭环,并理解每种保存格式的底层实现差异,能够在自己的项目中按需选择最合适的方案。
模块总览:tensorlayer.files的三大能力
tensorlayer.files是 TensorLayer 面向"数据与模型文件"的通用工具箱,其功能定位正如模块文档首页所概括:"A collections of helper functions to work with dataset. Load benchmark dataset, save and restore model, save and load variables."(一组处理数据集的辅助函数:加载基准数据集、保存和恢复模型、保存和加载变量)。
从源码入口 tensorlayer/files/init.py 可以确认,该模块由两部分聚合而成:
- 数据集加载器:位于 tensorlayer/files/dataset_loaders/,每个数据集一个独立文件(如
mnist_dataset.py、cifar10_dataset.py、voc_dataset.py等); - 通用工具函数:集中在 tensorlayer/files/utils.py(约 2900 行),涵盖模型权重保存/恢复、变量保存/恢复、文件与文件夹操作等。
全部函数通过tl.files.xxx对外暴露,功能可划分为三大类:
| 类别 | 代表性函数 | 用途 |
|---|---|---|
| 数据集加载 | load_mnist_dataset、load_cifar10_dataset、load_voc_dataset等 14 个加载器 +download_file_from_google_drive | 自动下载并解析公开基准数据集 |
| 模型保存与恢复 | save_npz/load_npz、save_npz_dict/load_and_assign_npz_dict、save_weights_to_hdf5/load_hdf5_to_weights_in_order/load_hdf5_to_weights | 以 npz、hdf5 等格式保存/恢复网络权重 |
| 变量与文件工具 | save_any_to_npy/load_npy_to_any、file_exists、del_file、load_file_list、maybe_download_and_extract、natural_keys等 | 工程级的文件系统辅助能力 |
数据集加载:14 个基准数据集一键获取
所有加载器都遵循同一设计哲学:首次调用时自动下载并解析,之后直接复用本地缓存。它们共享底层下载与解压机制maybe_download_and_extract,因此统一返回 NumPy 数组或列表形式的数据,可直接喂给 TensorLayer 网络训练。下表汇总了各加载器的签名、默认参数与返回值(依据 tensorlayer/files/utils.py 及各 dataset_loaders 文件):
| 函数 | 默认参数 | 返回内容 |
|---|---|---|
load_mnist_dataset | shape=(-1, 784), path='data' | X_train, y_train, X_val, y_val, X_test, y_test(50000/10000/10000 划分) |
load_fashion_mnist_dataset | shape=(-1, 784), path='data' | 同上结构,Fashion-MNIST 数据 |
load_cifar10_dataset | shape=(-1, 32, 32, 3), path='data', plotable=False | X_train, y_train, X_test, y_test(60000 张 32×32 彩色图,10 类) |
load_cropped_svhn | path='data', include_extra=True | X_train, y_train, X_test, y_test(include_extra=True时把 531131 张 extra 图并入训练集) |
load_ptb_dataset | path='data' | train_data, valid_data, test_data, vocab_size(整型词序列) |
load_matt_mahoney_text8_dataset | path='data' | list of str(text8 原始词列表) |
load_imdb_dataset | path='data', nb_words=None, skip_top=0, maxlen=None, test_split=0.2, seed=113, start_char=1, oov_char=2, index_from=3 | X_train, y_train, X_test, y_test(整型词序列) |
load_nietzsche_dataset | path='data' | str(尼采全集文本) |
load_wmt_en_fr_dataset | path='data' | train_path, dev_path(英法翻译语料目录) |
load_flickr25k_dataset | tag='sky', path='data', n_threads=50, printable=False | list of array(按 tag 过滤的图像) |
load_flickr1M_dataset | tag='sky', size=10, path='data', n_threads=50, printable=False | list of array(size取 1~10,10 表示全量 100 万张) |
load_cyclegan_dataset | filename='summer2winter_yosemite', path='data' | im_train_A, im_train_B, im_test_A, im_test_B(CycleGAN 风格迁移数据) |
load_celebA_dataset | path='data' | list of str(CelebA 图片路径,经 Google Drive 下载) |
load_voc_dataset | path='data', dataset='2012', contain_classes_in_person=False | 10 个返回值(图像/语义分割/实例分割/标注文件列表、类别、Darknet 格式标注等) |
load_mpii_pose_dataset | path='data', is_16_pos_only=False | img_train_list, ann_train_list, img_test_list, ann_test_list(人体姿态估计) |
使用示例:MNIST 与 CIFAR-10
MNIST 是最常用的入门数据。加载器内部把官方 60000 张训练图切分为 50000 训练 + 10000 验证,与 10000 张测试图一起返回:
import tensorlayer as tl # 展平向量形式(默认),适合全连接网络 X_train, y_train, X_val, y_val, X_test, y_test = tl.files.load_mnist_dataset(shape=(-1, 784), path='data') # 单通道图像形式,适合卷积网络 X_train, y_train, X_val, y_val, X_test, y_test = tl.files.load_mnist_dataset(shape=(-1, 28, 28, 1))仓库中的 examples/basic_tutorials/tutorial_mnist_simple.py 等 6 个 MNIST 教程均采用这一调用方式。shape参数由内部解析器直接用于np.frombuffer(...).reshape(shape),像素值统一缩放到[0, 1](源码见 utils.py 的_load_mnist_dataset)。
CIFAR-10 的加载支持两种通道顺序,并可选可视化抽查:
# NHWC 通道顺序(TensorFlow 惯例) X_train, y_train, X_test, y_test = tl.files.load_cifar10_dataset(shape=(-1, 32, 32, 3), plotable=False) # NCHW 通道顺序 X_train, y_train, X_test, y_test = tl.files.load_cifar10_dataset(shape=(-1, 3, 32, 32))面向 NLP 与特殊任务的加载器
- PTB 语言建模数据:
load_ptb_dataset()返回整型词 ID 序列与词表大小,内部借助tl.nlp.build_vocab/tl.nlp.words_to_word_ids完成分词建表,约 929k 训练词、10k 词表(见 tensorlayer/files/dataset_loaders/ptb_dataset.py); - text8 词向量数据:
load_matt_mahoney_text8_dataset()返回原始词列表,可直接用于 Word2Vec 类任务,仓库测试 tests/test_nlp.py 中即用它作为数据源; - IMDB 情感分类:
load_imdb_dataset提供nb_words(词表上限)、maxlen(最大序列长度截断)、test_split(测试集比例)、skip_top(忽略最高频词)等 Keras 风格参数,默认以start_char=1标记序列起始、oov_char=2标记词表外词; - VOC 目标检测:
load_voc_dataset(dataset="2012")会解析 XML 标注,额外产出 Darknet 格式的标注字符串(class_id x_centre y_centre width height比例格式)与 TensorFlow Object Detection 风格的标注字典,共 10 个返回值,可直接对接 examples/data_process/tutorial_tf_dataset_voc.py 的 TFRecord 流水线; - CelebA / CycleGAN 等图像数据:
load_celebA_dataset依赖download_file_from_google_drive从 Google Drive 拉取(需自行安装tqdm与requests,源码在 utils.py 的download_file_from_google_drive);load_cyclegan_dataset则按trainA/trainB/testA/testB四目录返回未配对图像,并把灰度图自动扩成三通道。
模型保存与恢复:npz / npz_dict / hdf5 三套方案
模块文档明确给出了选型建议:"TensorFlow provides.ckptfile format to save and restore the models, while we suggest to use standard python file formathdf5to save models for the sake of cross-platform. Other file formats such as.npzare also available."—— 即 TensorFlow 原生.ckpt可用,但推荐使用跨平台的hdf5,此外也支持npz。以下完整继承文档中的核心示例并逐行注释:
## 1) 以 .h5(hdf5)格式保存模型 tl.files.save_weights_to_hdf5('model.h5', network.all_weights) # 按顺序恢复模型权重 tl.files.load_hdf5_to_weights_in_order('model.h5', network.all_weights) # 按名称恢复模型权重 tl.files.load_hdf5_to_weights('model.h5', network.all_weights) ## 2) 以 .npz 格式保存模型 tl.files.save_npz(network.all_weights, name='model.npz') # 恢复方式一:先加载再手动分配 load_params = tl.files.load_npz(name='model.npz') tl.files.assign_weights(sess, load_params, network) # 恢复方式二:一步完成"加载 + 分配" tl.files.load_and_assign_npz(sess=sess, name='model.npz', network=network) ## 3) 部分参数分配(迁移学习 / 预训练微调常用) # 只分配第 1 个参数 tl.files.assign_weights(sess, [load_params[0]], network) # 只分配前 3 个参数 tl.files.assign_weights(sess, load_params[:3], network)注:该示例中的
sess参数对应 TensorFlow 1.x 的 Session 用法。在当前 TensorFlow 2.x 源码实现中,assign_weights(weights, network)直接对network.all_weights[idx]调用.assign(param),返回赋值操作列表(见 utils.py 的assign_weights),不再需要 Session,调用方式为tl.files.assign_weights(load_params, network)。
npz 系列:列表式与字典式两种存法
列表式(保存顺序,恢复顺序):
save_npz(save_list=None, name='model.npz'):内部先经tf_variables_to_numpy把 TensorFlow 变量批量转为 NumPy 数组,再以np.savez(name, params=...)存入,数据统一挂在params键下;load_npz(path='', name='model.npz'):np.load(...)['params']返回按保存顺序排列的参数列表;assign_weights(weights, network):将参数列表按序赋给network.all_weights,返回赋值操作列表,支持切片实现"只恢复部分层";load_and_assign_npz(name=None, network=None):合并前两步,文件不存在时返回False并记录错误日志。
字典式(保存名称,按名恢复):
save_npz_dict(save_list=None, name='model.npz'):以每个张量的tensor.name为键、数值为值写入 npz;load_and_assign_npz_dict(name='model.npz', network=None, skip=False):按名称把权重分配回网络。skip参数控制名称不匹配时的行为——True时跳过并告警,False时抛出RuntimeError。源码还会先检查 npz 内是否存在重复键(重复则抛异常),保证恢复过程的确定性。
测试 tests/files/test_utils_saveload.py 对上述两种 npz 格式做了完整的"保存 → 篡改权重 → 恢复 → 校验误差 < 1e-7"回环验证,可作为用法参考。
hdf5 系列:按顺序 vs 按名称
hdf5 保存的核心实现在_save_weights_to_hdf5_group(utils.py):以层名建组,在根属性layer_names记录层名列表,每层内部再用weight_names属性 + 同名 dataset 存储权重矩阵。这一结构同时支撑了两种恢复策略:
load_hdf5_to_weights_in_order(filepath, network):按顺序恢复。要求网络层顺序与保存文件一致;若文件比网络多出冗余层,只要前面匹配则自动忽略多余部分(源码_load_weights_from_hdf5_group_in_order按索引逐层对应);load_hdf5_to_weights(filepath, network, skip=False):按名称恢复。通过layer_index = {layer.name: layer}建立名称索引实现按名查找,skip控制名称缺失时是跳过还是抛错;对 BatchNorm 层还有专门的squeeze()兼容处理(针对维度不匹配的历史文件)。
两者的关键差异在于对网络结构顺序的依赖程度:按名称恢复允许调整网络层顺序,更适合"加载预训练权重到结构略有变化的网络"的场景。此外 hdf5 保存/加载会校验layer_names属性是否存在,若文件不是 TL 保存的会抛出NameError提示。
进阶:如果希望连**网络结构(架构)**一起保存,模块还提供了
save_hdf5_graph/load_hdf5_graph(见 utils.py 顶部实现),它们把模型 config 写入 hdf5 属性并可附带权重,跨脚本重建整个模型;static_graph2net负责按层配置回放构建网络。加载时会比对保存时的 TensorFlow 与 TensorLayer 版本号,不一致则给出告警。这两个函数虽未出现在文档主索引中,但在源码中完整可用。
ckpt 兼容层
为兼容 TensorFlow 原生生态,模块同样保留了save_ckpt/load_ckpt(utils.py)以及load_and_assign_ckpt、ckpt_to_npz_dict(后者可把 ckpt 权重转为 npz 字典,rename_key=True时还能把xxx/w_w重命名为 TL 规范的xxx/filters:0)。不过源码注释指出 eager 模式下的 ckpt 保存尚未稳定实现,因此跨平台场景仍以 hdf5 为推荐方案。
任意变量保存:.npy格式
当需要保存的并非网络权重、而是训练曲线、统计量或任意 Python 对象时,使用save_any_to_npy/load_npy_to_any:
# 保存任意字典对象 tl.files.save_any_to_npy(save_dict={'data': ['a', 'b']}, name='test.npy') # 恢复 data = tl.files.load_npy_to_any(name='test.npy') print(data) # {'data': ['a','b']}实现上就是np.save/np.load(..., allow_pickle=True)的封装,加载时优先尝试.item()还原字典语义。适合保存超参数、日志等元信息,与模型权重文件分开管理。
文件与文件夹工具:工程级文件系统操作
这部分函数是数据流水线的"地基",在数据集加载器内部被广泛复用(例如几乎所有load_xxx_dataset都会先调用maybe_download_and_extract检查本地缓存)。逐个说明:
| 函数 | 签名要点 | 行为 |
|---|---|---|
file_exists(filepath) | 文件路径 | 等价os.path.isfile,返回布尔值 |
folder_exists(folderpath) | 文件夹路径 | 等价os.path.isdir,返回布尔值 |
del_file(filepath) | 文件路径 | 等价os.remove,删除单个文件 |
del_folder(folderpath) | 文件夹路径 | 等价shutil.rmtree,递归删除整个文件夹 |
read_file(filepath) | 文件路径 | 以文本模式读取并返回字符串 |
load_file_list(path=None, regx='\\.jpg', printable=True, keep_prefix=False) | 路径 + 正则 | 返回匹配正则的文件名列表;keep_prefix=True时返回带完整路径的列表;path=None时使用当前工作目录 |
load_folder_list(path="") | 文件夹路径 | 返回该目录下所有子文件夹的完整路径列表 |
exists_or_mkdir(path, verbose=True) | 文件夹路径 | 不存在则创建并返回False,已存在返回True |
maybe_download_and_extract(filename, working_directory, url_source, extract=False, expected_bytes=None) | 文件名 + 目录 + URL | 本地无文件时下载(带进度条),extract=True时自动解压 tar/zip;expected_bytes校验文件大小,不符则抛异常 |
load_file_list的正则过滤很有用,例如只取文件夹中的 npz 权重文件:
file_list = tl.files.load_file_list(path='checkpoints', regx='w1pre_[0-9]+\\.(npz)')maybe_download_and_extract是数据集加载的核心基础设施——MNIST 的 gz、CIFAR-10 的 tar.gz、PTB 的 tgz、CelebA 的 zip 全部经由它下载解压;它还通过expected_bytes做下载完整性校验(text8 即指定了 31344016 字节)。
排序与可视化辅助
人类可读的自然排序:natural_keys(text)解决"im2.jpg排在im11.jpg前面"这类字典序问题。配合list.sort(key=...)使用:
l = ['im1.jpg', 'im31.jpg', 'im11.jpg', 'im21.jpg', 'im03.jpg', 'im05.jpg'] l.sort(key=tl.files.natural_keys) # ['im1.jpg', 'im03.jpg', 'im05.jpg', 'im11.jpg', 'im21.jpg', 'im31.jpg']其实现基于re.split('(\d+)', text)把字符串切分为数字与非数字片段并做类型化比较(utils.py)。Flickr、CycleGAN、VOC 等图像数据加载器在拼接文件名列表时都依赖它保证顺序一致。
npz 权重可视化:npz_to_W_pdf(path=None, regx='w1pre_[0-9]+\\.(npz)')遍历匹配的 npz 文件,把第一个权重矩阵用tl.visualize.draw_weights绘制并导出为同名 PDF,适合快速检查卷积核/权重分布的训练变化。需注意visualize模块的绘图依赖(如 matplotlib)。
从示例与测试看最佳实践
仓库中的真实用法可作为落地方案参考:
- 训练结束保存权重:examples/basic_tutorials/tutorial_mnist_simple.py 训练完成后直接
network.save_weights('model.h5'); - 强化学习分网络保存:examples/reinforcement_learning/tutorial_A3C.py、tutorial_DDPG.py 等用
tl.files.save_npz(trainable_weights, name=...)分别保存 actor/critic 网络,用tl.files.save_weights_to_hdf5保存 Q 网络; - 数据集对接 TFRecord:examples/data_process/tutorial_tf_dataset_voc.py 演示了
load_voc_dataset→ TFRecord 的完整链路; - 回环正确性验证:tests/files/test_utils_saveload.py 对 hdf5 / npz / npz_dict 三种格式均执行"保存 → 篡改 → 恢复 → 断言数值误差 < 1e-7";tests/models/test_model_save.py 进一步覆盖了
skip加载、嵌套 VGG、LayerList 等复杂网络结构的保存恢复。
实践建议小结:跨平台部署与长期存档首选hdf5(save_weights_to_hdf5+load_hdf5_to_weights_in_order);快速保存/恢复单次实验结果用npz(save_npz+load_and_assign_npz);需要按权重名做选择性恢复(如迁移学习只加载骨干层)时用npz_dict或load_hdf5_to_weights并按需开启skip=True;模型结构也要持久化时,升级到save_hdf5_graph/load_hdf5_graph。结合 docs/modules/files.rst 的 API 清单与 tensorlayer/files/ 源码,即可按需组合出完整、可复现的数据与模型管理方案。
- 人工智能
- 深度学习
- 机器学习
- 强化学习
【免费下载链接】TensorLayer
Deep Learning and Reinforcement Learning Library for Scientists and Engineers
相关推荐
MMPose 三维人体网格恢复数据集准备指南:SMPL 模型、标注文件与六大数据集详解
MMPose 三维人体网格恢复数据集准备指南:SMPL 模型、标注文件与六大数据集详解 本文为 MMPose(OpenMMLab 人体姿态估计工具箱)三维人体网
计算机视觉人工智能深度学习一文解决!Intel RealSense .bag文件加载失败与数据恢复指南
一文解决!Intel RealSense .bag文件加载失败与数据恢复指南 在使用Intel® RealSense™ SDK(GitHub_Trending/
智能硬件音视频计算机视觉Garnet AOF文件修复:日志损坏恢复工具使用
Garnet AOF文件修复:日志损坏恢复工具使用 引言:AOF日志损坏的致命风险 在分布式缓存系统中,数据持久化是保障业务连续性的关键环节。Garnet作为微
缓存KV存储后端
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考