You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何保存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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.10.03 19:39:00