NEPA
Collection
6 items • Updated • 12
How to use SixAILab/nepa-dit-xl-400k with Diffusers:
pip install -U diffusers transformers accelerate
import torch
from diffusers import DiffusionPipeline
# switch to "mps" for apple devices
pipe = DiffusionPipeline.from_pretrained("SixAILab/nepa-dit-xl-400k", dtype=torch.bfloat16, device_map="cuda")
prompt = "Astronaut in a jungle, cold color palette, muted colors, detailed, 8k"
image = pipe(prompt).images[0]Configuration Parsing Warning:In UNKNOWN_FILENAME: "diffusers._class_name" must be a string
Class-conditional latent diffusion model on ImageNet-1K at 256x256.
LlamaNepaModel, 28 layers, d=1152). It reads [class token, <boi>, noisy latent]
and emits one embedding per latent patch (patch size 4 on the 32x32x4 SD latent).Flux2Transformer2DModel, 10 double + 18 single blocks), conditioned on the NEPA
embeddings, trained with flow matching + REPA. Weights are the EMA at step 400k.stabilityai/sd-vae-ft-ema decoder (f8, 4 channels); latents share the SD VAE encoder space.FlowMatchEulerDiscreteScheduler), classifier-free guidance with the null class.Everything needed to run the model (encoder + pipeline code) ships in this repo; load with trust_remote_code=True.
pip install "diffusers>=0.36" "transformers>=5.3" torch safetensors
Tested with transformers 5.3.0 / diffusers 0.36.0 and transformers 5.16.1 / diffusers 0.40.0 (identical outputs).
import torch
from diffusers import DiffusionPipeline
pipe = DiffusionPipeline.from_pretrained(
"SixAILab/nepa-dit-xl-400k", trust_remote_code=True, torch_dtype=torch.bfloat16
).to("cuda")
class_ids = pipe.get_label_ids(["golden retriever", "macaw"]) # or pass ints in [0, 999] directly
images = pipe(
class_labels=class_ids,
guidance_scale=3.6,
num_inference_steps=96,
generator=torch.Generator("cuda").manual_seed(0),
).images
images[0].save("golden_retriever.png")
__call__ arguments:
| arg | default | meaning |
|---|---|---|
class_labels |
required | list of ImageNet-1K ids (one image per id) |
guidance_scale |
3.6 | CFG weight; 1.0 disables guidance |
num_inference_steps |
96 | Euler ODE steps |
generator |
None | torch.Generator for reproducibility |
output_type |
"pil" |
"pil", "np", "pt" or "latent" |
Output resolution is fixed at 256x256. bf16 is recommended (all reported numbers use it); fp16 (e.g. on a T4) and fp32 also work.
guidance_scale around 1.4-1.5; larger values
(3-4) trade diversity for fidelity, as with DiT.guidance_scale > 1.