TF1迁移TF2触发str与VariableScope拼接TypeError求解
TensorFlow V1 迁移 V2 运行报错
TypeError: can only concatenate str (not "VariableScope") to str 排查 问题背景
- 项目原本基于TensorFlow V1开发,迁移到TensorFlow V2环境时已修复绝大多数代码兼容问题,运行时抛出上述类型错误,未检索到可行解决方案。
相关代码片段
wdmSystem.py 入口调用代码
sess = tf.compat.v1.Session() sess, outMetrics = cft.train(sess, tf.compat.v1.train.AdamOptimizer, loss, metricsDict, trainingParam, feedDictFun, debug=True)
training.py 核心 train 函数
def train(sess, optimizer, loss, metricsDict, trainingParam, feedDictFun, debug=False): optimizer = optimizer(learning_rate=trainingParam.learningRate) meanMetricOpsDict, updateOps, resetOps = create_mean_metrics(metricsDict) trainStep, accumulateOps, zeroOps = accumulatedOptimizer(optimizer, loss, tf.compat.v1.trainable_variables(), trainingParam.nMiniBatches) init = tf.compat.v1.global_variables_initializer() if trainingParam.summaries: s = [tf.compat.v1.summary.scalar(name, metric) for name,metric in meanMetricOpsDict.items()] metricSummaries = tf.compat.v1.summary.merge(s) summaries_dir = os.path.join(trainingParam.path, 'tboard', trainingParam.summaryString) os.makedirs(summaries_dir, exist_ok=True) trainWriter = tf.compat.v1.summary.FileWriter(summaries_dir + '/train', sess.graph) else: trainWriter = None sess.run(init) saver = tf.compat.v1.train.Saver() checkpoint_path = os.path.join(trainingParam.path,'checkpoint',trainingParam.filename,'best') if not os.path.exists(checkpoint_path): os.makedirs(checkpoint_path) else: print("Restoring...", flush=True) saver.restore(sess=sess,save_path=checkpoint_path) bestLoss = 100000 bestAcc = 0 lastImprovement = 0 sess.run([resetOps, zeroOps]) for epoch in range(1, trainingParam.nEpochs+1): sess.run(resetOps) for batch in range(1,trainingParam.nBatches+1): sess.run(zeroOps) if debug: print('batch: ', batch, end=' | ', flush=True) print('miniBatch: ', end='', flush=True) for miniBatch in range(1,trainingParam.nMiniBatches+1): if debug: print(miniBatch, end=', ') feedDict = feedDictFun(trainingParam) sess.run([accumulateOps, updateOps], feed_dict=feedDict) if debug: print('', flush=True) sess.run(trainStep) outMetrics = sess.run(list(meanMetricOpsDict.values()), feed_dict=feedDict) outMetrics = { key:val for key,val in zip(list(meanMetricOpsDict.keys()), outMetrics) } if trainingParam.summaries: outMetricSummaries = sess.run(metricSummaries, feed_dict=feedDict) trainWriter.add_summary(outMetricSummaries, epoch) earlyStoppingMetric = outMetrics[trainingParam.earlyStoppingMetric] if earlyStoppingMetric < bestLoss: bestLoss = earlyStoppingMetric lastImprovement = epoch saver.save(sess=sess, save_path=checkpoint_path) if epoch%trainingParam.displayStep == 0: outString = 'epoch: {:04d}'.format(epoch) for key, value in outMetrics.items(): outString += ' - {}: {:.4f}'.format(key, value) print(outString, flush=True) saver.restore(sess=sess,save_path=checkpoint_path) sess.run(resetOps) for _ in range(trainingParam.evalBatches): feedDict = feedDictFun(trainingParam) sess.run(updateOps, feed_dict=feedDict) outMetrics = sess.run(list(meanMetricOpsDict.values()), feed_dict=feedDict) outMetrics = { key:val for key,val in zip(list(meanMetricOpsDict.keys()), outMetrics) } return sess, outMetrics
training.py create_mean_metrics 函数
def create_mean_metrics(metricsDict): meanMetricOpsDict = {} updateOps = [] resetOps = [] for name, tensor in metricsDict.items(): lossOp, updateOp, resetOp = create_reset_metric(tf.compat.v1.metrics.mean, name, tensor) meanMetricOpsDict[name] = lossOp updateOps.append(updateOp) resetOps.append(resetOp) return meanMetricOpsDict, updateOps, resetOps
training.py create_reset_metric 函数(报错位置)
def create_reset_metric(metric, scope='reset_metrics', *metric_args, **metric_kwargs): """ 指标重置逻辑参考官方issue实现 """ with tf.compat.v1.variable_scope(scope) as scope: metric_op, update_op = metric(*metric_args, **metric_kwargs) vars = tf.compat.v1.get_variable(scope, collections=tf.compat.v1.GraphKeys.LOCAL_VARIABLES) reset_op = tf.compat.v1.variables_initializer(vars) return metric_op, update_op, reset_op
完整报错追踪栈
Traceback (most recent call last): File "/home/osama/PycharmProjects/claude-master/tf_wdmSystem-learning.py", line 222, in <module> sess, outMetrics = cft.train(sess, tf.compat.v1.train.AdamOptimizer, loss, metricsDict, trainingParam, feedDictFun, debug=True) File "/home/osama/PycharmProjects/claude-master/claude/claudeflow/training.py", line 38, in train meanMetricOpsDict, updateOps, resetOps = create_mean_metrics(metricsDict) File "/home/osama/PycharmProjects/claude-master/claude/claudeflow/training.py", line 20, in create_mean_metrics lossOp, updateOp, resetOp = create_reset_metric(tf.compat.v1.metrics.mean, name, tensor) File "/home/osama/PycharmProjects/claude-master/claude/claudeflow/training.py", line 11, in create_reset_metric vars = tf.compat.v1.get_variable(scope, collections=tf.compat.v1.GraphKeys.LOCAL_VARIABLES) File "/home/osama/PycharmProjects/claude-master/venv/lib/python3.10/site-packages/tensorflow/python/ops/variable_scope.py", line 1616, in get_variable return get_variable_scope().get_variable( File "/home/osama/PycharmProjects/claude-master/venv/lib/python3.10/site-packages/tensorflow/python/ops/variable_scope.py", line 1308, in get_variable full_name = self.name + "/" + name if self.name else name TypeError: can only concatenate str (not "VariableScope") to str Process finished with exit code 1
问题根因与修复方案
- 根因:
with tf.compat.v1.variable_scope(scope) as scope语句会将原本字符串类型的入参scope覆盖为VariableScope上下文对象,后续调用tf.compat.v1.get_variable时第一个参数要求传入字符串类型的变量名,实际传入了VariableScope实例,触发字符串拼接类型错误。同时原代码获取指标对应变量的逻辑也存在错误,tf.compat.v1.get_variable是用于创建/获取单个指定名称变量的接口,不能直接传入作用域对象获取作用域下的所有变量。 - 修复方法:
- 修改变量作用域上下文的赋值变量名,避免覆盖原字符串类型的scope入参
- 改用
tf.compat.v1.get_collection接口获取当前指标作用域下的所有本地变量,用于生成重置操作
修复后的
create_reset_metric代码如下:
def create_reset_metric(metric, scope='reset_metrics', *metric_args, **metric_kwargs): # 上下文返回值使用独立变量名,不覆盖原scope字符串参数 with tf.compat.v1.variable_scope(scope) as var_scope: metric_op, update_op = metric(*metric_args, **metric_kwargs) # 拉取当前作用域下所有属于本地变量集合的指标变量 metric_vars = tf.compat.v1.get_collection( tf.compat.v1.GraphKeys.LOCAL_VARIABLES, scope=var_scope.name ) reset_op = tf.compat.v1.variables_initializer(metric_vars) return metric_op, update_op, reset_op
内容的提问来源于stack exchange,提问作者Osama Ismail
相关产品推荐
相关产品推荐

