JiT-LDM 512x512 CXR Image Generator
This repository contains JiT-LDM diffusion flow-matching CXR image generative model, trained on ~1.9M belarusian screening CXR image dataset. This model designed for generating 512x512 Chest X-ray (CXR) images. The JiT was trained from scratch, applied directly to latent space of frozen pretrained VAE.
Model Details
- Developed by: Biomedical Image Analysis Laboratory (Lab225), United Institute of Informatics Problems, National Academy of Sciences of Belarus (UIIP NASB), Minsk.
- Architecture: Just image Transformer + VAE from Stable Diffusion 1.5
- Parameters: ~680 Million parameters (JiT) + ~49 Million parameters (VAE decoder).
- Primary Task: Class-conditioned generation of 512x512 CXR images
Data Pipeline & Preprocessing
- Image Preprocessing: The raw dataset consisted of high-dimensional 2-byte grayscale (single-channel) medical images. Each image was automatically cropped to the lung region and contrast-enhanced using a lung segmentation mask, then rescaled to a uniform resolution of 512x512 pixels and converted to 1-byte grayscale format.
- Latent Space Encoding: Prior to training, the entire preprocessed dataset was encoded into the latent space of a pre-trained Variational Autoencoder (VAE). The JiT model was subsequently trained directly on these latent representations.
- Class Labeling: Each image was annotated with one of 12 mutually exclusive class labels derived from three independent attributes: gender (male/female), age group (G1: ≤34 years; G2: 35-51 years; G3: ≥52 years), and abnormality status (healthy / with abnormalities). This yields a total of 2x3x2=12 classes.
Total number of images (latents) in each class:
| G1_F_Norm | G1_F_Abnorm | G1_M_Norm | G1_M_Abnorm | G2_F_Norm | G2_F_Abnorm | G2_M_Norm | G2_M_Abnorm | G3_F_Norm | G3_F_Abnorm | G3_M_Norm | G3_M_Abnorm |
|---|---|---|---|---|---|---|---|---|---|---|---|
| 351,834 | 33,034 | 383,224 | 20,136 | 246,065 | 21,134 | 202,270 | 20,479 | 237,792 | 129,884 | 164,249 | 94,673 |
- Training Schedule: The model was trained for 600 epochs on a single NVIDIA H800 GPU with a batch size of 144. A linear learning-rate warm-up was applied during the first epoch, increasing from 2e-6 to 2e-4, after which the learning rate remained constant at 2e-4. Mixed-precision training was employed using PyTorch autocast with float16 precision. The model supports classifier-free guidance.
Quantitative Metrics
Evaluation results on different training checkpoint highlighting the top performing generative scores. All these scores were computed on the class “Female, G1, Abnormal”, which contains ~33,000 images. This class was chosen because it exhibits substantial variability (due to the presence of abnormalities) and is large enough for stable metric estimation, yet not so large as to make computation prohibitively expensive. The brackets indicate which solver, how many steps, and what strength of the CFG were used to achieve the corresponding result.
| Model Checkpoint | FID↓ | KID↓ | CMMD↓ | Precision↑ | Recall↑ | F-score↑ |
|---|---|---|---|---|---|---|
| Epoch 100 | 4.56 (Heun3, 10 steps, scale 2.0) |
0.0055 (Heun3, 50 steps, scale 2.0) |
0.0188 (Heun3, 50 steps, scale 1.0) |
0.652 (DOPRI5, 25 steps, scale 2.0) |
0.639 (RK4, 10 steps, scale 1.0) |
0.624 (DOPRI5, 50 steps, scale 1.0) |
| Epoch 200 | 6.48 (Euler, 50 steps, scale 1.0) |
0.0098 (Euler, 50 steps, scale 1.0) |
0.0206 (Euler, 100 steps, scale 1.0) |
0.648 (Euler, 100 steps, scale 2.0) |
0.588 (RK4, 25 steps, scale 1.0) |
0.599 (Euler, 100 steps, scale 1.0) |
| Epoch 300 | 3.74 (Euler, 100 steps, scale 2.0) |
0.0043 (Euler, 100 steps, scale 2.0) |
0.0211 (Heun3, 50 steps, scale 1.0) |
0.658 (RK4, 50 steps, scale 2.0) |
0.633 (Heun3, 100 steps, scale 1.0) |
0.632 (Heun3, 100 steps, scale 1.0) |
| Epoch 400 | 6.55 (Euler, 100 steps, scale 1.0) |
0.0083 (Euler, 100 steps, scale 1.0) |
0.0205 (Heun3, 50 steps, scale 1.0) |
0.629 (Euler, 100 steps, scale 1.0) |
0.585 (RK4, 100 steps, scale 1.0) |
0.598 (RK4, 50 steps, scale 1.0) |
| Epoch 500 | 5.52 (Heun3, 25 steps, scale 2.0) |
0.0066 (Heun3, 25 steps, scale 2.0) |
0.0210 (Heun3, 50 steps, scale 1.0) |
0.646 (Heun3, 50 steps, scale 1.0) |
0.541 (RK4, 25 steps, scale 1.0) |
0.587 (Heun3, 50 steps, scale 1.0) |
| Epoch 600 | 8.93 (Euler, 50 steps, scale 2.0) |
0.0121 (Euler, 50 steps, scale 2.0) |
0.0322 (Euler, 50 steps, scale 1.0) |
0.584 (Euler, 100 steps, scale 1.0) |
0.551 (Euler, 100 steps, scale 1.0) |
0.567 (Euler, 100 steps, scale 1.0) |
Quick Start (How to Use)
# 'torchdiffeq' is a required non-standard library,
# If running on Google Colab or locally via Jupyter Notebook, run:
# !pip install torchdiffeq
# or install this library manually (see requirements.txt)
import torch
from transformers import AutoModel
device = "cuda" if torch.cuda.is_available() else "cpu"
# parameters for downloading the model and generating images
params = {
'checkpoint_number': 300, # available checkpoints are: 100, 200, 300, 400, 500, 600
'solver': 'euler', # choose any solver: 'euler', 'heun2', 'heun3', 'rk4', 'dopri5'
'steps': 100, # select the number of solver steps, usually from 10 to 100
'cfg_scale': 2.0, # select CFG scale, where 1.0 means no CFG is applied
'num_images': 4, # number of images to generate
'label': 1, # class label to be generated
}
# Load the custom Lab225 architecture directly from the hub
model = AutoModel.from_pretrained(
"lab225/jit-ldm-cxr-class-conditioned",
subfolder=f"checkpoint_{params['checkpoint_number']}",
trust_remote_code=True
)
model = model.to(device)
model.eval()
# image generation, returns PIL images
generated_images = model.sample(
torch.full((params['num_images'],), params['label']).to(device),
steps=params['steps'],
solver_name=params['solver'],
cfg_scale=params['cfg_scale']
)
- Downloads last month
- 14