TensorFlow Recommenders模型加载Checkpoint后保存遇DecodeError问题求助
解决TensorFlow Recommenders BruteForce模型保存的DecodeError问题
问题原因
从epoch checkpoint加载BruteForce模型后保存时出现protobuf解析错误,根源是BruteForce通过index_from_dataset构建索引后,内部张量的维度名称可能包含非UTF-8编码的二进制数据,checkpoint加载后这些元数据被保留,导致模型序列化时protobuf无法解析。
解决方案
不要直接保存/加载整个BruteForce模型,而是拆分保存查询模型、候选模型以及候选嵌入数据,加载时重新构建BruteForce索引:
1. 训练阶段拆分保存
# 单独保存查询模型和候选模型 query_model.save("query_model") candidate_model.save("candidate_model") # 保存候选标签与对应的嵌入向量(用numpy存储) import numpy as np # 生成候选嵌入数据 candidate_labels = [] candidate_embeddings = [] for label_batch in unique_labels.batch(50): embedding_batch = candidate_model(label_batch) candidate_labels.extend(label_batch.numpy()) candidate_embeddings.extend(embedding_batch.numpy()) # 保存到本地 np.save("candidate_labels.npy", candidate_labels) np.save("candidate_embeddings.npy", candidate_embeddings)
2. 加载时重新构建BruteForce模型
import tensorflow as tf import tensorflow_recommenders as tfrs import numpy as np # 加载查询模型和候选模型 query_model = tf.keras.models.load_model("query_model") candidate_model = tf.keras.models.load_model("candidate_model") # 初始化BruteForce层 brute_model = tfrs.layers.factorized_top_k.BruteForce(query_model) # 从保存的numpy数据重建索引数据集 labels = np.load("candidate_labels.npy") embeddings = np.load("candidate_embeddings.npy") index_dataset = tf.data.Dataset.from_tensor_slices((labels, embeddings)).batch(50) # 构建索引 brute_model.index_from_dataset(index_dataset) # 现在可以正常保存模型 brute_model.save(destination_path)
替代方案:直接加载查询模型权重
如果必须基于checkpoint恢复,只加载查询模型的权重,再重新构建索引:
# 加载查询模型权重(假设训练时保存了query_model的checkpoint) query_model.load_weights("query_model_checkpoint") # 重新构建BruteForce并生成索引 brute_model = tfrs.layers.factorized_top_k.BruteForce(query_model) brute_model.index_from_dataset(tf.data.Dataset.zip((unique_labels.batch(50), unique_labels.batch(50).map(candidate_model)))) # 保存模型 brute_model.save(destination_path)
内容的提问来源于stack exchange,提问作者Shubham
相关产品推荐
相关产品推荐

