tensorflow / tensorflow/probability

DenseVariational on Apple Silicon with GPU

Open
#1,612 0 comments 0 reactions 0 assignees View on GitHub

Nobody has claimed this yet.

Dominant language
Jupyter Notebook
Stars
4.4k
Forks
1.1k
PR merge metrics
No merged PRs in 30d

Description

Issue:

The dense variational layer does not appear to return random values when used with Apple M1 GPU but works fine with CPU only.

Expected Behaviour:

When calling the model on the same input data, the output should be random (as the weights are sampled from probability distributions).

To reproduce:

The following code is taken from this stackoverflow post.

This code generates synthetic data

import numpy as np
import tensorflow as tf
import tensorflow_probability as tfp
import matplotlib.pyplot as plt
tfd = tfp.distributions
tfpl = tfp.layers

physical_devices = tf.config.list_physical_devices('CPU')
print(physical_devices)  

# the following line sets the deivce
# tf.config.set_visible_devices([], 'GPU') # disable GPU

x_train = np.linspace(-1, 2, 5000)[:, np.newaxis]
y_train = np.power(x_train, 3) + 0.1*(2+x_train)*np.random.randn(5000)[:, np.newaxis]

plt.scatter(x_train, y_train, alpha=0.1)
plt.show()  

The following section defines the model.

def prior(kernel_size, bias_size, dtype = None):
    
    n = kernel_size + bias_size
    
    prior_model = tf.keras.Sequential([
        
        tfpl.DistributionLambda(
        
            lambda t: tfd.MultivariateNormalDiag(loc = tf.zeros(n)  ,  scale_diag = tf.ones(n)
                                                
                                                ))
        
    ])
    
    return(prior_model)


def posterior(kernel_size, bias_size, dtype = None):
    
    n = kernel_size + bias_size
    
    posterior_model = tf.keras.Sequential([
        
        tfpl.VariableLayer(tfpl.MultivariateNormalTriL.params_size(n)  , dtype = dtype),   
        
        tfpl.MultivariateNormalTriL(n)  
        
        
    ])
    
    return(posterior_model)

x_in = tf.keras.layers.Input(shape = (1,))

x =   tfpl.DenseVariational(units=tfpl.IndependentNormal.params_size(1),
                          make_prior_fn=prior,
                          make_posterior_fn=posterior,
                          kl_weight=1/x_train.shape[0])(x_in)

y_out =    tfpl.DenseVariational(units=1,
                  make_prior_fn=prior,
                  make_posterior_fn=posterior,
                  kl_weight=1/x_train.shape[0])(x)

model = tf.keras.Model(inputs = x_in, outputs = y_out)

def nll(y_true, y_pred):
    dist = tfp.distributions.Normal(loc=y_pred, scale=1.0)
    return tf.reduce_sum(-dist.log_prob(y_true))

model.compile(loss=nll, optimizer= 'Adam')
model.summary()
history = model.fit(x_train, y_train, epochs=1, verbose=False)

The subsequent code block tries to plot the trajectories of different predictions but when it has unexpected behaviour on a GPU.

predicted = [model(x_train) for _ in range(100)]
for i, res in enumerate(predicted):
                plt.plot(x_train, res , alpha=0.1)
plt.scatter(x_train, y_train, alpha=0.1)
plt.show()

Contributor guide

Open the contributing guide

First steps

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. Open a pull request that references the issue number.

Research direction

The issue's model definition and repeated prediction calls are the only entry points provided; start by running the reproduction with the CPU and Apple M1 GPU settings shown. Done means the same input produces the expected varying predictions on the GPU, with the behavior compared against the CPU result.

Written by the indexing model from the issue text.

Assessment

Tech stack
python
Domain
machine-learning
Issue type
Bug
Difficulty
4/5
Estimated time
3-5 days
Activity status
Stale
Clarity
Needs clarification
Newbie friendliness
25/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.