Generation API
High-level, checkpoint-driven generation for discrete and continuous latent models. These classes wrap a trained model together with its tokenizer so you can go from a checkpoint to decoded images in a few lines. For low-level sampling primitives (schedulers, guidance, early stopping), see Sampling.
Discrete Latent Generation
- class medlatents.generation.DiscreteLatentGenerator(model_type, model, tokenizer, device)[source][source]
Bases:
objectUnified generator for discrete latent generative models.
Supports: - autoreg: Autoregressive generation - maskgit: Iterative parallel generation (MaskGIT) - flow: Discrete flow matching generation - diffusion: D3PM discrete diffusion generation
- Parameters:
model_type (
Literal['autoreg','maskgit','flow','diffusion'])model (
Module)tokenizer (
DiscreteTokenizer)device (
device)
- __init__(model_type, model, tokenizer, device)[source][source]
- Parameters:
model_type (
Literal['autoreg','maskgit','flow','diffusion'])model (
Module)tokenizer (
DiscreteTokenizer)device (
device)
- classmethod from_checkpoints(model_type, model_path, tokenizer_path, device=None, weights_only=True)[source][source]
Load generator from model and tokenizer checkpoints.
- generate(num_samples, seq_length, temperature=1.0, top_k=None, num_steps=None, seed=None, **kwargs)[source][source]
Generate samples based on model type.
- Parameters:
num_samples (
int) – Number of samples to generateseq_length (
int) – Length of sequencestemperature (
float, default:1.0) – Sampling temperaturetop_k (
int|None, default:None) – Top-k sampling parameternum_steps (
int|None, default:None) – Number of iterative steps (maskgit/flow/diffusion)seed (
int|None, default:None) – Random seed for reproducibility (None for random)**kwargs – Additional model-specific parameters
- Return type:
- Returns:
Generated token sequences [num_samples, seq_length]
- generate_volumes(num_samples, seq_length, output_dir, temperature=1.0, top_k=None, num_steps=None, seed=None, **kwargs)[source][source]
Generate samples and save as NIfTI volumes.
Continuous Latent Generation
- class medlatents.generation.ContinuousLatentGenerator(model_type, model, tokenizer, device, diffusion=None, flow=None)[source][source]
Bases:
objectGenerator for continuous latent diffusion and flow models.
- Parameters:
model_type (
Literal['diffusion','flow'])model (
Module)tokenizer (
ContinuousTokenizer|None)device (
device)diffusion (
ContinuousGaussianDiffusion|None, default:None)flow (
RectifiedFlow|None, default:None)
- __init__(model_type, model, tokenizer, device, diffusion=None, flow=None)[source][source]
- Parameters:
model_type (
Literal['diffusion','flow'])model (
Module)tokenizer (
ContinuousTokenizer|None)device (
device)diffusion (
ContinuousGaussianDiffusion|None, default:None)flow (
RectifiedFlow|None, default:None)
- generate(num_samples, latent_shape=None, guidance_scale=1.0, condition=None, null_condition=None, num_steps=None, temperature=1.0, initial_noise=None, model_kwargs=None, seed=None)[source][source]
- Parameters:
- Return type:
- decode_latents(latents, **decode_kwargs)[source][source]
Decode latents to reconstruction using tokenizer.
- generate_volumes(num_samples, output_dir, latent_shape=None, guidance_scale=1.0, condition=None, null_condition=None, num_steps=None, temperature=1.0, initial_noise=None, model_kwargs=None, affine=None, seed=None)[source][source]