[Tensorflow] 使用 model.save_weights() 保存 Keras Subclassed Model
import numpy as np import matplotlib.pyplot as plt import os import time import tensorflow as tf tf.enable_eager_execution() # create data X = np.linspace(-1, 1, 5000) np.random.shuffle(X) y = 0.5 * X + 2 + np.random.normal(0, 0.05, (5000,)) # plot data plt.scatter(X, y) plt.show() # split data X_train, y_train = X[:4000], y[:4000] X_test, y_test = X[4000:], y[4000:] # tf.data BATCH_SIZE = 32 BUFFER_SIZE = 512 dataset = tf.data.Dataset.from_tensor_slices((X_train, y_train)).batch(BATCH_SIZE).shuffle(BUFFER_SIZE) # subclassed model UNITS = 1 class Model(tf.keras.Model): def __init__(self): super(Model, self).__init__() self.fc = tf.keras.layers.Dense(units=UNITS) def call(self, inputs): return self.fc(inputs) model = Model() optimizer = tf.train.AdamOptimizer() # loss function def loss_function(real, pred): return tf.losses.mean_squared_error(labels=real, predictions=pred) EPOCHS = 30 checkpoint_dir = './save_subclassed_keras_model_training_checkpoints' if not os.path.exists(checkpoint_dir): os.makedirs(checkpoint_dir) # training loop for epoch in range(EPOCHS): start = time.time() epoch_loss = 0 for (batch, (x, y)) in enumerate(dataset): x = tf.cast(x, tf.float32) y = tf.cast(y, tf.float32) x = tf.expand_dims(x, axis=1) y = tf.expand_dims(y, axis=1) # print(x) # tf.Tensor([...], shape=(BATCH_SIZE, 1), dtype=float32) # print(y) # tf.Tensor([...], shape=(BATCH_SIZE, 1), dtype=float32) with tf.GradientTape() as tape: predictions = model(x) # print(predictions) # tf.Tensor([...], shape=(BATCH_SIZE, 1), dtype=float32) batch_loss = loss_function(real=y, pred=predictions) grads = tape.gradient(batch_loss, model.variables) optimizer.apply_gradients(zip(grads, model.variables), global_step=tf.train.get_or_create_global_step()) epoch_loss += batch_loss if (batch + 1) % 10 == 0: print('Epoch {} Batch {} Loss {:.4f}'.format(epoch + 1, batch + 1, batch_loss/int(x.shape[0]))) print('Epoch {} Loss {:.4f}'.format(epoch + 1, epoch_loss/len(X_train))) print('Time taken for 1 epoch {} sec\n'.format(time.time() - start)) # save checkpoint checkpoint_prefix = os.path.join(checkpoint_dir, 'ckpt') if (epoch + 1) % 10 == 0: model.save_weights(checkpoint_prefix.format(epoch=epoch), overwrite=True) _model = Model() _model.load_weights(tf.train.latest_checkpoint(checkpoint_dir)) _model.build(input_shape=tf.TensorShape([BATCH_SIZE, 1])) _model.summary() test_dataset = tf.data.Dataset.from_tensor_slices(X_test).batch(1) for (batch, x) in enumerate(test_dataset): x = tf.cast(x, tf.float32) x = tf.expand_dims(x, axis=1) print(x) predictions = _model(x) print(predictions) exit()
[Tensorflow] 使用 model.save_weights() 保存 Keras Subclassed Model的更多相关文章
- [Tensorflow] 使用 model.save_weights() 保存 / 加载 Keras Subclassed Model
在 parameters.py 中,定义了各类参数. # training data directory TRAINING_DATA_DIR = './data/' # checkpoint dire ...
- [Tensorflow] 使用 tf.train.Checkpoint() 保存 / 加载 keras subclassed model
在 subclassed_model.py 中,通过对 tf.keras.Model 进行子类化,设计了两个自定义模型. import tensorflow as tf tf.enable_eager ...
- keras系列︱Sequential与Model模型、keras基本结构功能(一)
引自:http://blog.csdn.net/sinat_26917383/article/details/72857454 中文文档:http://keras-cn.readthedocs.io/ ...
- Keras(一)Sequential与Model模型、Keras基本结构功能
keras介绍与基本的模型保存 思维导图 1.keras网络结构 2.keras网络配置 3.keras预处理功能 模型的节点信息提取 config = model.get_config() 把mod ...
- [Model] LeNet-5 by Keras
典型的卷积神经网络. 数据的预处理 Keras傻瓜式读取数据:自动下载,自动解压,自动加载. # X_train: array([[[[ 0., 0., 0., ..., 0., 0., 0.], [ ...
- xen 保存快照的实现之 —— device model 状态保存
xen 保存快照的实现之 —— device model 状态保存 实现要点: 设备状态保存在 /var/lib/xen/qemu-save.x 文件这个文件由 qemu-dm 产生,也由 qemu- ...
- AI - TensorFlow - 示例05:保存和恢复模型
保存和恢复模型(Save and restore models) 官网示例:https://www.tensorflow.org/tutorials/keras/save_and_restore_mo ...
- 如何保存Keras模型
我们不推荐使用pickle或cPickle来保存Keras模型 你可以使用model.save(filepath)将Keras模型和权重保存在一个HDF5文件中,该文件将包含: 模型的结构,以便重构该 ...
- Python之TensorFlow的模型训练保存与加载-3
一.TensorFlow的模型保存和加载,使我们在训练和使用时的一种常用方式.我们把训练好的模型通过二次加载训练,或者独立加载模型训练.这基本上都是比较常用的方式. 二.模型的保存与加载类型有2种 1 ...
随机推荐
- hadoop(九) - hbase shell命令及Java接口
一. shell命令 1. 进入hbase命令行 ./hbase shell 2. 显示hbase中的表 list 3. 创建user表,包括info.data两个列族 create 'user' ...
- .net Core使用Orcle官方驱动连接数据库 C#参考教程 http://www.csref.cn
.net Core使用Orcle官方驱动连接数据库 最近在研究.net Core,因为公司的项目用到的都是Oracle数据库,所以简单试一下.net Core怎样连接Oracle. Oracle官 ...
- PHP mysql 连接ipV6地址
需要在PHP页面中通过ipv6连接数据库,但是发现无论是用mysql_connect还是mysqli_connect,如果host是ipv6格式,就不能正常连接,会提示“php_network_get ...
- 20170623_oracle_优化与体系结构
一般优化技巧 建议不用"*"代替所有列名 删除所有数据用TRUNCATE代替DELETE 用NOT EXISTS 代替NOT IN 用EXISTS代替IN 用EXISTS代替DIS ...
- 在Ubuntu 12.04 LTS下成功访问Windows域共享(mount //192.168.1.102/share -o user=DOMIAN\\user,pass=passwd /mnt)
Ubuntu 12.04 LTS下成功访问Windows域共享: 1,在命令行模式下 mount //192.168.1.102/share -o user=DOMIAN\\user,pass=pas ...
- 对于系统盘升级windows10怕空间不够,还是打算继续卸载一些软件
本来是打算从其他盘压缩,然后扩展,可是怕把磁盘给弄坏了,然后就保存原来的,就是看升级的推送什么时候来了.
- ubuntu 16.04 Eclipse 图标显示为 ?(已解决)
这个问题挺好解决: sudo gedit /usr/share/applications/eclipse.desktop在这个文件中将Icon=/home/soyo/eclipse/icon.xpm, ...
- 3-1 vue生存指南 - todolist实现-数据渲染
由于Vue.js作者是中国人,会说汉语,所以国内生态会更好一点.Vue.js作者是尤雨溪,
- c++ class does not name a type (转载)
转载:http://blog.csdn.net/typename/article/details/7173550 declare class does not name a type 出现这个编译错误 ...
- Thinkphp模板标签if和eq的区别和比较
在TP模板语言中.if和eq都可以用于变量的比较.总结以下几点: 1.两个变量的比较: <if condition=”$item.group_id eq $one.group_id”> & ...