如何在TensorFlow中打印完整(非截断)的张量?
解决TensorFlow中tf.Print输出张量被截断的问题
哈哈,我之前也踩过这个坑!你设置numpy的set_printoptions没用是因为tf.Print本身有自己的输出截断规则,和numpy的配置不搭边。要完整打印整个张量,关键是给tf.Print加上summarize参数,把它设为张量的总元素数或者一个足够大的值就行。
给你改好的代码:
import tensorflow as tf import numpy as np tensor = tf.constant(np.ones(999)) # 两种方式都行:要么用张量的实际大小,要么直接写个足够大的数 # 方式一:动态获取张量大小 with tf.Session() as sess: tensor_size = sess.run(tf.size(tensor)) tensor = tf.Print(tensor, [tensor], summarize=tensor_size) sess.run(tensor) # 方式二:直接指定一个比元素数大的数值(比如10000),更简单 # import tensorflow as tf # import numpy as np # tensor = tf.constant(np.ones(999)) # tensor = tf.Print(tensor, [tensor], summarize=10000) # with tf.Session() as sess: # sess.run(tensor)
另外提一句,要是你用的是TensorFlow 2.x版本,tf.Print已经被弃用了,官方推荐用tf.print(注意是小写的p),这个更灵活,设置summarize=-1就能直接打印所有元素:
import tensorflow as tf import numpy as np tensor = tf.constant(np.ones(999)) tf.print(tensor, summarize=-1) # summarize=-1表示输出全部元素
这样就能看到完整的张量内容啦!
内容的提问来源于stack exchange,提问作者Plezos
相关产品推荐
相关产品推荐

