TensorFlow实现句子语义相似度报错RuntimeError求助
解决Universal Sentence Encoder加载时的RuntimeError问题
这个错误的核心原因是你的代码用的是TensorFlow 1.x的图模式API,但当前环境默认开启了TensorFlow 2.x的即时执行(Eager Execution)。hub.Module()是TF1时代的模块加载方式,依赖计算图,而Eager模式下没有全局计算图,所以会抛出"Exporting/importing meta graphs is not supported..."的错误。
下面给你两种可行的解决方案,推荐第二种更贴合TF2的主流用法:
方案一:禁用Eager Execution,沿用TF1风格代码
只需要在导入TensorFlow之后,添加一行代码关闭Eager模式,就能让原来的代码正常运行:
import tensorflow as tf # 添加这行,禁用Eager Execution tf.compat.v1.disable_eager_execution() import tensorflow_hub as hub import numpy as np from sklearn.metrics.pairwise import cosine_similarity # get cosine similairty matrix def cos_sim(input_vectors): similarity = cosine_similarity(input_vectors) return similarity # get topN similar sentences def get_top_similar(sentence, sentence_list, similarity_matrix, topN): # find the index of sentence in list index = sentence_list.index(sentence) # get the corresponding row in similarity matrix similarity_row = np.array(similarity_matrix[index, :]) # get the indices of top similar indices = similarity_row.argsort()[-topN:][::-1] return [sentence_list[i] for i in indices] module_url = "https://tfhub.dev/google/universal-sentence-encoder/2" # Import the Universal Sentence Encoder's TF Hub module embed = hub.Module(module_url) # Reduce logging output. tf.logging.set_verbosity(tf.logging.ERROR) sentences_list = [ # phone related 'My phone is slow', 'My phone is not good', 'I need to change my phone. It does not work well', 'How is your phone?', # age related 'What is your age?', 'How old are you?', 'I am 10 years old', # weather related 'It is raining today', 'Would it be sunny tomorrow?', 'The summers are here.' ] with tf.Session() as session: session.run([tf.global_variables_initializer(), tf.tables_initializer()]) sentences_embeddings = session.run(embed(sentences_list)) similarity_matrix = cos_sim(np.array(sentences_embeddings)) sentence = "It is raining today" top_similar = get_top_similar(sentence, sentences_list, similarity_matrix, 3) # printing the list using loop for x in range(len(top_similar)): print(top_similar[x])
方案二:改用TF2.x兼容的API(推荐)
TF2.x推荐使用更简洁的hub.load()或者Keras层的方式加载模型,不需要手动管理tf.Session(),代码更简洁易维护:
import tensorflow as tf import tensorflow_hub as hub import numpy as np from sklearn.metrics.pairwise import cosine_similarity # get cosine similairty matrix def cos_sim(input_vectors): similarity = cosine_similarity(input_vectors) return similarity # get topN similar sentences def get_top_similar(sentence, sentence_list, similarity_matrix, topN): index = sentence_list.index(sentence) similarity_row = np.array(similarity_matrix[index, :]) indices = similarity_row.argsort()[-topN:][::-1] return [sentence_list[i] for i in indices] # 使用TF2兼容的方式加载模型 module_url = "https://tfhub.dev/google/universal-sentence-encoder/4" # 推荐用v4版本,更适配TF2 embed_model = hub.load(module_url) sentences_list = [ # phone related 'My phone is slow', 'My phone is not good', 'I need to change my phone. It does not work well', 'How is your phone?', # age related 'What is your age?', 'How old are you?', 'I am 10 years old', # weather related 'It is raining today', 'Would it be sunny tomorrow?', 'The summers are here.' ] # 直接生成句向量,不需要Session sentences_embeddings = embed_model(sentences_list).numpy() similarity_matrix = cos_sim(sentences_embeddings) sentence = "It is raining today" top_similar = get_top_similar(sentence, sentences_list, similarity_matrix, 3) for item in top_similar: print(item)
改动说明:
- 使用
hub.load()替代hub.Module(),加载TF2兼容的模型版本(这里用了v4,你也可以选其他TF2兼容的版本) - 不需要手动初始化变量和Session,直接调用模型生成向量,用
.numpy()把Tensor转成Numpy数组 - 代码更简洁,符合TF2的编程风格
内容的提问来源于stack exchange,提问作者user12907213
相关产品推荐
相关产品推荐

