如何在tf.estimator.Estimator中正确集成Beholder插件?
问题描述
这是Beholder插件,它支持可视化所有可训练变量(针对超深度网络有合理限制)。
我的问题:我正在使用tf.estimator.Estimator类进行训练,但发现Beholder插件与Estimator API兼容不佳。
我的代码如下:
# tf.data输入流水线设置 def dataset_input_fn(train=True): filenames = ... # 训练文件 if not train: filenames = ... # 测试文件 dataset = tf.data.TFRecordDataset(filenames), "GZIP") # ... 后续处理直至 ... iterator = batched_dataset.make_one_shot_iterator() return iterator.get_next() def train_input_fn(): return dataset_input_fn(train=True) def test_input_fn(): return dataset_input_fn(train=False) # 模型函数 def cnn(features, labels, mode, params): # 构建模型 # 为`ModeKeys.PREDICT`提供Estimator规格 if mode == tf.estimator.ModeKeys.PREDICT: return tf.estimator.EstimatorSpec( mode=mode, predictions={"sentiment": y_pred_cls}) eval_metric_ops = { "accuracy": accuracy_op, "precision": precision_op, "recall": recall_op } normal_summary_hook = tf.train.SummarySaverHook( 100, summary_op=summary_op) return tf.estimator.EstimatorSpec( mode=mode, loss=cost_op, train_op=train_op, eval_metric_ops=eval_metric_ops, training_hooks=[normal_summary_hook] ) classifier = tf.estimator.Estimator(model_fn=cnn, params=..., model_dir=...) classifier.train(input_fn=train_input_fn, steps=1000) ev = classifier.evaluate(input_fn=test_input_fn, steps=1000) tf.logging.info("Loss: {}".format(ev["loss"])) tf.logging.info("Precision: {}".format(ev["precision"])) tf.logging.info("Recall: {}".format(ev["recall"])) tf.logging.info("Accuracy: {}".format(ev["accuracy"]))
我不清楚该在这个架构中何处添加beholder钩子。若在cnn函数中作为训练钩子添加:
return tf.estimator.EstimatorSpec( mode=mode, loss=dnn.cost, train_op=dnn.train_op, eval_metric_ops=eval_metric_ops, training_hooks=[normal_summary_hook, beholder_hook] )
会出现错误:InvalidArgumentError: You must feed a value for placeholder tensor 'Placeholder' with dtype uint8 and shape [?,?,?]。
若尝试用tf.train.MonitoredTrainingSession来配置classifier,训练能正常进行,但Beholder插件无任何日志记录。查看标准输出发现会先后创建两个会话,似乎tf.estimator.Estimator分类器会先终止现有会话,再启动自己的会话。
请问有没有解决办法?
解决方案
太懂这种框架兼容的头疼了!Estimator的会话管理确实比较封闭,直接用Beholder的默认钩子容易踩坑,咱们一步步来搞定:
1. 核心问题拆解
- 直接加原始Beholder钩子报错:是因为它默认会尝试读取输入占位符的数据,但Estimator的
input_fn生成的是基于tf.data的迭代器,占位符的生命周期和普通会话不一样,钩子拿不到有效值。 - 用
MonitoredTrainingSession无效:Estimator内部会自己创建并管理MonitoredSession,外部的会话配置会被完全覆盖,Beholder根本没机会接入。
2. 自定义适配Estimator的Beholder钩子
我们需要写一个轻量的适配层,让Beholder能跟着Estimator的会话生命周期走,只专注于监控可训练变量,避开输入占位符的坑:
from tensorboard.plugins.beholder.beholder import Beholder import tensorflow as tf class BeholderEstimatorHook(tf.train.SessionRunHook): def __init__(self, trainable_vars, logdir): # 初始化Beholder,只传入要监控的变量和日志目录 self.beholder = Beholder(variables=trainable_vars, logdir=logdir) self.logdir = logdir def after_create_session(self, session, coord): # Estimator创建会话后,把会话传给Beholder self.beholder.set_session(session) def before_run(self, run_context): # 每次训练步前,让Beholder记录变量状态 self.beholder.update(session=run_context.session) # 不需要获取额外的张量,返回空的SessionRunArgs return tf.train.SessionRunArgs([])
3. 在模型函数中集成自定义钩子
修改你的cnn模型函数,在训练模式下创建并添加这个自定义钩子:
def cnn(features, labels, mode, params): # --- 原有的模型构建代码不变 --- # 构建模型层、计算损失、指标等... # 获取所有可训练变量 trainable_vars = tf.trainable_variables() # 初始化自定义Beholder钩子 beholder_hook = BeholderEstimatorHook( trainable_vars=trainable_vars, logdir=params['model_dir'] # 必须和Estimator的model_dir一致 ) # 原有的Summary钩子 normal_summary_hook = tf.train.SummarySaverHook( 100, summary_op=summary_op) if mode == tf.estimator.ModeKeys.TRAIN: return tf.estimator.EstimatorSpec( mode=mode, loss=cost_op, train_op=train_op, eval_metric_ops=eval_metric_ops, training_hooks=[normal_summary_hook, beholder_hook] ) # --- 原有的PREDICT、EVAL模式处理不变 --- if mode == tf.estimator.ModeKeys.PREDICT: return tf.estimator.EstimatorSpec( mode=mode, predictions={"sentiment": y_pred_cls}) # EVAL模式的处理...
4. 确保Estimator配置正确
创建Estimator时,一定要把model_dir传入params,这样Beholder的日志会写入Estimator的日志目录,TensorBoard才能识别:
model_dir = "./your_model_logs" classifier = tf.estimator.Estimator( model_fn=cnn, params={ 'model_dir': model_dir, # 你的其他模型参数... }, model_dir=model_dir )
5. 验证效果
启动训练后,用TensorBoard打开同一个日志目录:
tensorboard --logdir=./your_model_logs
打开Beholder面板就能看到所有可训练变量的可视化数据了!
额外提示
如果需要监控输入数据,可以在before_run方法里添加获取输入张量的逻辑,但要确保从features里获取,而不是直接用占位符:
def before_run(self, run_context): # 如果你需要监控输入特征,可以这样获取 features_tensor = run_context.session.graph.get_tensor_by_name("your_features_tensor_name:0") features_value = run_context.session.run(features_tensor) self.beholder.update(session=run_context.session, feed_dict={...}) return tf.train.SessionRunArgs([])
不过一般来说,Beholder主要用来监控变量,这一步不是必须的。
内容的提问来源于stack exchange,提问作者Insectatorious

