如何保存TensorFlow Probability中训练好的Transformed Distribution
保存TFP TransformedDistribution的两种可行方案
方案1:使用TensorFlow SavedModel格式(推荐)
这是最简便的方案,无需手动管理参数,适配绝大多数使用场景。
保存代码
import tensorflow as tf import tensorflow_probability as tfp tfd = tfp.distributions # 将分布封装为tf.Module适配SavedModel存储规范 class TrainedDistWrapper(tf.Module): def __init__(self, transformed_dist): self.dist = transformed_dist # 定义采样方法的输入签名,确保可被序列化 @tf.function(input_signature=[tf.TensorSpec(shape=[], dtype=tf.int32)]) def sample(self, sample_size): return self.dist.sample(sample_size) # 可按需添加log_prob等其他需要用到的方法 @tf.function(input_signature=[tf.TensorSpec(shape=[None, *YOUR_DATA_SHAPE], dtype=tf.float32)]) def log_prob(self, x): return self.dist.log_prob(x) # 传入你训练完成的transformed_distribution wrapper = TrainedDistWrapper(transformed_distribution) # 保存到本地路径 save_path = "./trained_tfp_distribution" tf.saved_model.save(wrapper, save_path)
加载代码
# 加载已保存的分布 loaded_wrapper = tf.saved_model.load(save_path) # 直接调用采样方法,和训练时的使用逻辑完全一致 generated_samples = loaded_wrapper.sample(SAMPLE_SIZE).numpy()
注意:如果你用到了自定义的Bijector或者Keras层,需要在加载环境中保留对应的类定义,建议使用tf.keras.utils.register_keras_serializable装饰自定义类,避免加载时报错。
方案2:手动保存参数、重建分布
适合需要灵活调整分布结构的场景,需要保证训练和加载时的基础结构完全对齐。
保存参数
import numpy as np # 提取分布所有可训练参数 dist_params = {var.name: var.numpy() for var in transformed_distribution.trainable_variables} # 保存为numpy压缩文件 np.savez("./tfp_dist_params.npz", **dist_params)
加载重建
# 1. 按照训练时的代码,完全一致地初始化base分布和bijector base_distribution = # 与训练时的初始化逻辑完全相同 stacked_ffjord = # 与训练时的bijector初始化逻辑完全相同 # 2. 重建TransformedDistribution reconstructed_dist = tfd.TransformedDistribution( distribution=base_distribution, bijector=stacked_ffjord ) # 3. 赋值保存的参数 loaded_params = np.load("./tfp_dist_params.npz") for var in reconstructed_dist.trainable_variables: var.assign(loaded_params[var.name]) # 正常采样使用 generated_samples = reconstructed_dist.sample(SAMPLE_SIZE).numpy()
内容的提问来源于stack exchange,提问作者kamka
相关产品推荐
相关产品推荐

