Conditional MediumUNet Diffusion Model (Flax NNX, JAX)

A class-conditional, from-scratch DDPM-style diffusion model trained on Tiny ImageNet ($64 imes 64 imes 3$, 200 classes), built using JAX and Flax NNX. The model architecture (ConditionalUNet) predicts the clean image from a noisy input, conditioned simultaneously on a sinusoidally embedded continuous noise level and a class label embedding.

Model Details

  • Architecture: Medium U-Net with Class Conditioning (featuring sinusoidal time embeddings, learnable class embeddings, strided convolutions for downsampling, transposed convolutions for upsampling, and functional residual skip connections)
  • Conditioning:
    1. Scalar noise level amount ∈ [0, 1] via sinusoidal positional embeddings.
    2. Class index label ∈ [0, 199] via dense embedding vectors injected into the residual blocks.
  • Framework: JAX / Flax NNX
  • Training data: Tiny ImageNet (200 classes) — full 100,000-image Tiny ImageNet training set
  • Optimizer: AdamW with warmup-cosine-decay LR schedule, gradient clipping (global norm 1.0)
  • Epochs: 40
  • Batch size: 256

Exported Formats

This repository bundles multiple export formats of the same trained weights:

File / Folder Format Notes
model_checkpoint/ Orbax checkpoint Native JAX/Flax NNX state, load with orbax.checkpoint
conditional_medium_unet.mlir MLIR StableHLO dialect, via jax.export
conditional_medium_unet.stablehlo StableHLO bytecode Portable compiled representation (with 100% baked-in weights)

Usage (JAX / Flax NNX Native)

import orbax.checkpoint as ocp
from flax import nnx

# Recreate the model skeleton, then restore the trained parameters
model = ConditionalUNet(in_features=3, out_features=3, num_classes=200, rngs=nnx.Rngs(0))
graphdef, abstract_state = nnx.split(model)

checkpointer = ocp.Checkpointer(ocp.StandardCheckpointHandler())
state = checkpointer.restore("model_checkpoint", abstract_state)
model = nnx.merge(graphdef, state)

⚡ How to use the StableHLO Model (Production Inference)

The ultimate advantage of this pipeline is that you don't need the Python architecture class anymore, nor do you need to manage raw weight parameters! Because the model was exported via a functional static JIT trace, all trained weights are fully baked directly into the StableHLO bytecode binary. The compiled graph acts as a completely autonomous black-box function.

Here is how to properly load and execute the pre-compiled StableHLO graph to perform a full iterative diffusion reverse loop using pure JAX with the HQ quadratic schedule (80 steps):

import jax
import jax.numpy as jnp
import numpy as np

# 1. Load the pre-compiled StableHLO bytecode from file
with open("conditional_medium_unet.stablehlo", "rb") as f:
    stablehlo_bytecode = f.read()

# 2. Deserialize it back into a callable JAX artifact and JIT compile it
loaded_artifact = jax.export.deserialize(stablehlo_bytecode)
compiled_fn = jax.jit(loaded_artifact.call)

# 3. Initialize pure normal noise (batch_size=1, 64x64, 3 channels)
key = jax.random.PRNGKey(42)
x = jax.random.normal(key, shape=(1, 64, 64, 3))

# Choose class label (e.g., 42 for 'mantis')
CHOSEN_CLASS = 42
sample_label = jnp.array([CHOSEN_CLASS], dtype=jnp.int32)

# 4. Set up the HQ Quadratic Schedule (80 steps)
n_steps = 80
steps_line = jnp.linspace(1.0, 0.0, n_steps + 1)
ts = steps_line ** 2

print("🚀 Running conditional inference loop via independent StableHLO runtime...")
for i in range(n_steps):
    t_curr = ts[i]
    t_next = ts[i + 1]
    current_amount = jnp.array([t_curr], dtype=jnp.float32)

    # 🔥 LIGHTNING-FAST HARDWARE INFERENCE:
    # Pass image, time step, and class label. No model state or weights needed!
    pred_clean = compiled_fn(x, current_amount, sample_label)
    pred_clean = jnp.clip(pred_clean, -1.0, 1.0)

    if i == n_steps - 1:
        x = pred_clean
        break

    # DDIM Mathematical Correction
    eps_estimated = (x - (1.0 - t_curr) * pred_clean) / (t_curr + 1e-8)
    x = (1.0 - t_next) * pred_clean + t_next * eps_estimated

# Rescale from [-1, 1] back to [0, 1] for image display
final_image = np.array((x + 1.0) / 2.0)[0]
print("✅ Image generated successfully! Output shape:", final_image.shape)

Technical Parameters & Design Choices

  • Color & Resolution: Native support for $64 imes 64$ RGB images, a massive step up in spatial complexity compared to standard MNIST benchmarks.
  • Class-Conditional Latent Guidance: Embedding-driven injection allows targeted generation of specific objects among 200 distinct classes, turning the network from an unconditional sampler into a controllable asset.
  • Production Readiness: Ideal for specialized target applications on low-power edge nodes or cross-language backends (C++, Rust) via the native StableHLO runtime interface.
Downloads last month

-

Downloads are not tracked for this model. How to track
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support