TensorFlow中tf.contrib.summary新摘要评估方式及API使用疑问
嘿,我来帮你理清楚tf.contrib.summary这个新API的用法,尤其是你提到的和旧版的差异,还有record_summaries_every_n_global_steps的作用~
先搞懂新旧API的核心区别
旧版的tf.summary玩法很直接:先把所有要记录的摘要操作丢进去,用tf.summary.merge_all()合并成一个大操作,每次session.run这个合并操作,所有摘要就一股脑都写出来了——没有什么频率控制,要么全写,要么全不写。
而tf.contrib.summary(后来其实被整合到原生tf.summary里了,但咱们就按你说的这个版本聊)是靠上下文管理器来实现灵活的频率控制的,record_summaries_every_n_global_steps(n)就是用来给特定摘要加“频率锁”的工具。
针对你问的:循环里只run(summaries)的情况
假设你的代码大概是这个样子:
global_step = tf.train.get_or_create_global_step() # 准确率:没加频率控制,默认每次都记 tf.contrib.summary.scalar('accuracy', accuracy) # 损失:用上下文包起来,每10步记一次 with tf.contrib.summary.record_summaries_every_n_global_steps(10): tf.contrib.summary.scalar('loss', loss) # 获取所有摘要操作 summaries = tf.contrib.summary.all_summary_ops()
那当你在训练循环里每次跑session.run([summaries, global_step])的时候:
- 准确率的摘要每次迭代都会被写入,因为它没被频率控制的上下文包裹,属于“实时记录”的类型;
- 损失的摘要只有当global_step是10的倍数时才会触发写入——比如第10、20、30步,其他步骤里这个损失的写操作会直接被跳过,不会产生任何输出。
几个容易踩坑的细节要注意
- 一定要确保
global_step在每次迭代时都会自增!这个频率控制完全是靠global_step的数值来判断的,如果步数不动,那低频摘要永远都不会被写出来; - 如果你没显式创建global_step,它会自动用默认的全局步数,但我建议你显式创建并管理,避免莫名其妙的bug;
- 旧版的
merge_all()对应的新方法是tf.contrib.summary.all_summary_ops(),但注意这个方法返回的不是一个合并后的单一操作,而是一组操作,不过session.run的时候会自动处理它们的执行逻辑,不用你手动合并; - 要是你想临时关掉某些摘要,还可以用
tf.contrib.summary.never_record_summaries()这个上下文管理器,里面的摘要永远不会被写入。
给你个更清晰的实践写法
我平时写的时候会把不同频率的摘要分开管理,结构更清楚,不容易搞混:
global_step = tf.train.get_or_create_global_step() # 实时记录的摘要(每次迭代都写) with tf.contrib.summary.always_record_summaries(): tf.contrib.summary.scalar('accuracy', accuracy) tf.contrib.summary.image('input_samples', input_batch) # 低频记录的摘要(每10步写一次) with tf.contrib.summary.record_summaries_every_n_global_steps(10): tf.contrib.summary.scalar('loss', loss) tf.contrib.summary.histogram('layer_weights', model.layers[0].weights) summaries = tf.contrib.summary.all_summary_ops() # 训练循环 with tf.Session() as sess: # 别忘了初始化摘要相关的变量! tf.contrib.summary.initialize(graph=tf.get_default_graph()) sess.run(tf.global_variables_initializer()) for _ in range(100): # 运行摘要+训练步骤,同时更新global_step _, current_step = sess.run([summaries, global_step]) # 可选:打印个日志确认下 if current_step % 10 == 0: print(f"Step {current_step}: 损失摘要已记录")
这样写下来,哪些摘要什么时候写一目了然,也不容易出错。
内容的提问来源于stack exchange,提问作者Jakub Arnold
相关产品推荐
相关产品推荐

