Skip to content
← All tracks

JAX Stochasticity

JAX

Distributions, reparameterization, sampling techniques, MCMC, gradient estimators. Randomness as a controllable resource.

This track is written for JAX, which isn't the mode you're browsing in.

0 / 25 solved
  1. 1. Not solved yet. Uniform Sampling
  2. 2. Not solved yet. Normal Sampling with mean and std
  3. 3. Not solved yet. Bernoulli Mask Sampling
  4. 4. Not solved yet. Categorical Sampling
  5. 5. Not solved yet. Reparameterization Trick: Gaussian
  6. 6. Not solved yet. Gumbel-Softmax
  7. 7. Not solved yet. Dirichlet Sampling
  8. 8. Not solved yet. Decoding Temperature
  9. 9. Not solved yet. Top-k Logit Masking
  10. 10. Not solved yet. Nucleus (Top-p) Masking
  11. 11. Not solved yet. Gumbel Argmax (Categorical via Trick)
  12. 12. Not solved yet. Metropolis-Hastings Step
  13. 13. Not solved yet. Log Acceptance Ratio
  14. 14. Not solved yet. HMC Leapfrog Step
  15. 15. Not solved yet. REINFORCE Gradient Estimator
  16. 16. Not solved yet. Reparameterization Gradient
  17. 17. Not solved yet. REINFORCE with Baseline
  18. 18. Not solved yet. Batched Sampling via vmap
  19. 19. Not solved yet. Random Walk via lax.scan
  20. 20. Not solved yet. Importance Sampling
  21. 21. Not solved yet. Beta Distribution Sampling
  22. 22. Not solved yet. Multivariate Normal Sampling
  23. 23. Not solved yet. Random Permutation
  24. 24. Not solved yet. Poisson Sampling
  25. 25. Not solved yet. ELBO for Gaussian VI

Check yourself

4 questions · one attempt each

These do not count toward finishing the track. They are here to catch the things that are easy to read past.

0 / 4

A dropout layer reuses the same key on every call. What does the model see?

key = jax.random.key(0)
# every layer, every step, calls:
jax.random.bernoulli(key, 0.5, shape)
Question 1 of 4