如何在Tensor2Tensor中打印张量形状?GPU报错及无输出问题
解决Tensor2Transformer中打印张量形状的问题
针对你遇到的张量形状显示为“?”、手动新建Session触发GPU错误、tf.print无输出的问题,结合你使用tf.compat.v1.train.MonitoredSession和t2t-decoder启动的场景,给出以下可行方案:
方案一:借助MonitoredSession的run方法获取形状并解决设备问题
程序已通过MonitoredSession管理会话,不能单独新建tf.Session,需在MonitoredSession上下文内调用run获取张量形状,同时解决CPU/GPU不匹配问题:
- 调整启动命令强制使用CPU
启动t2t-decoder时添加--device=cpu参数,避免程序尝试分配GPU设备:
t2t-decoder --device=cpu # 保留其他原有参数
- 在MonitoredSession中获取并打印形状
找到代码中创建MonitoredSession的位置,在会话上下文内添加形状打印逻辑:
with tf.compat.v1.train.MonitoredSession( session_creator=tf.compat.v1.train.ChiefSessionCreator(...) ) as sess: # 原有解码/运行逻辑 # 获取张量x的动态形状 x_shape = sess.run(tf.shape(x)) print("张量x的形状:", x_shape) # 查看静态形状(可能包含?)可直接调用 print("张量x的静态形状:", x.get_shape())
方案二:让tf.print在MonitoredSession中生效
tf.print是计算图中的操作,必须被执行才会输出内容,需将其加入执行流或在会话中主动运行:
# 定义打印操作 print_shape_op = tf.print("张量x的动态形状:", tf.shape(x)) # 在MonitoredSession中执行该操作 with tf.compat.v1.train.MonitoredSession(...) as sess: # 主动运行打印操作 sess.run(print_shape_op) # 后续原有逻辑 # 或绑定打印操作与核心逻辑,确保每次执行都触发打印 with tf.control_dependencies([print_shape_op]): # 原有需要执行的操作,比如解码步骤 decoded_output = ...
补充说明
- 张量显示“?”表示这是动态维度,静态分析无法确定其大小,必须在运行时通过会话获取实际形状。
- 若需在代码中指定设备,可在张量定义时用
tf.device('/cpu:0')包裹,强制张量分配到CPU:
with tf.device('/cpu:0'): x = # 你的张量定义代码
内容的提问来源于stack exchange,提问作者Mahima Bathla
相关产品推荐
相关产品推荐

