TensorFlow模型变量值分配异常:转Jax模型时的调试求助
问题描述
从仓库加载TensorFlow模型:
model = tf.saved_model.load(folder)
目标是将该模型复刻到Jax中,需要确认加载的变量(权重和偏置)值是否正确。通过以下方式可获取第i个变量的值:
vars = model.variables print(vars[i].numpy())
但将这些值分配到Jax网络后无法得到正确结果,因此尝试调试特定层的输出,确保权重和偏置一致。尝试了以下代码:
numpy_vars = [v.numpy() for v in vars] # Eager模式下执行 with tf.compat.v1.Session(graph = graph) as sess: tvars = tf.compat.v1.trainable_variables() tf.compat.v1.variables_initializer(vars).run() # 初始化vars或tvars for v, tv in zip(numpy_vars, tvars): tv.assign(v) print(tvars[0].eval()) # 输出图模式下的变量值 print('------------------------------') print(numpy_vars[0])
发现输出值与预期不符(仅形状一致),怀疑model.graph中存在初始化操作。若将tv.assign(v)替换为sess.run(tv.assign(v)),则会报错:
TypeError: Argument `fetch` = <tf.Variable 'UnreadVariable' shape=(11, 256) dtype=float32> has invalid type "_UnreadVariable" must be a string or Tensor. (Can not convert a _UnreadVariable into a Tensor or Operation.)
请问如何正确分配这些变量值,使其在图执行过程中保持固定?
解决方案
核心问题分析
你遇到的问题根源在于Eager模式变量与图模式变量的混淆,以及图模式下变量初始化、赋值的执行逻辑错误:
- 你初始化的是Eager模式下加载的
model.variables,而非当前Session图中的tvars,导致后续赋值被初始化操作覆盖; tv.assign(v)在图模式下只是创建赋值操作,未实际执行;直接sess.run(tv.assign(v))报错是因为tv.assign(v)返回的是变量本身,而Session需要的是可执行的Tensor或操作。
正确实现步骤
以下是图模式下正确加载并固定变量值的代码:
import tensorflow as tf # 1. Eager模式下加载模型并提取变量值 model = tf.saved_model.load(folder) numpy_vars = [v.numpy() for v in model.variables] # 2. 构建图或加载目标图(确保tvars与model.variables顺序/形状一致) graph = tf.Graph() with graph.as_default(): # 这里需要定义和原模型一致的层结构,保证trainable_variables的顺序、数量匹配 # ...(你的模型结构定义代码)... tvars = tf.compat.v1.trainable_variables() # 3. 在Session中正确初始化并赋值 with tf.compat.v1.Session(graph=graph) as sess: # 初始化当前图中的变量(必须是tvars,而非原模型的Eager变量) sess.run(tf.compat.v1.variables_initializer(tvars)) # 执行赋值操作:遍历每个变量,运行assign操作的结果 for np_v, tv in zip(numpy_vars, tvars): # assign返回赋值后的变量,需通过sess.run执行该操作才能生效 sess.run(tv.assign(np_v)) # 验证赋值结果 loaded_values = sess.run(tvars) print("图模式变量值:", loaded_values[0]) print("------------------------------") print("原模型变量值:", numpy_vars[0])
关键注意事项
- 变量匹配:必须确保
tvars(图模式下的可训练变量)与model.variables的顺序、形状、数量完全一致,否则赋值会错位; - 初始化对象:Session中初始化的是当前图的
tvars,而非原模型的Eager变量; - 赋值执行:
tv.assign(np_v)仅创建赋值操作,必须通过sess.run()执行该操作才能生效,不能直接run变量本身; - 顺序不能颠倒:先初始化图变量,再执行赋值,避免初始化覆盖赋值结果。
替代简化方案(推荐)
如果只是为了验证变量值是否一致,无需使用图模式,直接在Eager模式下复刻TensorFlow模型并对比输出:
# Eager模式下构建复刻的TensorFlow模型 def build_replica_model(): # 定义和原模型完全一致的网络结构 # ... return model_replica model_replica = build_replica_model() # 将原模型变量值赋值给复刻模型 for orig_var, replica_var in zip(model.variables, model_replica.variables): replica_var.assign(orig_var.numpy()) # 输入相同测试数据,对比原模型和复刻模型的输出 test_input = tf.random.normal([1, input_shape]) orig_output = model(test_input) replica_output = model_replica(test_input) print(tf.reduce_all(tf.abs(orig_output - replica_output) < 1e-6)) # 验证输出是否一致,应返回True
这种方式更简单,能直接验证变量是否正确赋值,之后再将变量值迁移到Jax模型即可。
内容的提问来源于stack exchange,提问作者Pablo
相关产品推荐
相关产品推荐

