I'm trying to build a tf.keras.Sequential model using tfp.layers.DistributionLambda. I'm following the DistributionLambda example, but would want to replace the tfd.Normal with a variable-containing tfd.TransformedDistribution with RealNVP bijector.
import tensorflow as tf
import tensorflow_probability as tfp
model = tf.keras.Sequential((
tf.keras.layers.Lambda(lambda x: tf.shape(x)[-1]),
tfp.layers.DistributionLambda(lambda t: (
tfp.distributions.TransformedDistribution(
distribution=(
tfp.distributions.MultivariateNormalDiag(loc=tf.zeros(t))),
bijector=tfp.bijectors.RealNVP(
num_masked=2,
shift_and_log_scale_fn=tfp.bijectors.real_nvp_default_template(
hidden_layers=[32, 32]))))),
))
x = tf.random.uniform((5, 3))
distribution = model(x)
However, this fails with the following:
The layer cannot safely ensure proper Variable reuse across multiple calls, and consquently this behavior is disallowed for safety. Lambda layers are not well suited to stateful computation; instead, writing a subclassed Layer is the recommend way to define layers with Variables.
Notice that the RealNVP bijector's variables have to be initialized within the bijector, unlike e.g. the Normal distribution's variables in DistributionLambda example, in which they are created at the top-level of Sequential.
I wonder if there's a way to use DistributionLambda with such a setup where the variables have to be created within the distribution? If so, what's the correct way to handle the variables inside the DistributionLambda layer? If not, what would be the recommended way of building a model like this?