Note
Go to the end to download the full example code.
Production workflows#
This example demonstrates callbacks, checkpointing, and serialization features for production use with long-running analyses.
Imports#
import logging
import sys
import tempfile
from pathlib import Path
import numpy as np
from latents.callbacks import CheckpointCallback, LoggingCallback, ProgressCallback
from latents.gfa import GFAFitConfig, GFAModel
from latents.gfa.config import GFASimConfig
from latents.gfa.simulation import load_simulation, save_simulation, simulate
from latents.observation import ObsParamsHyperPrior
# Configure logging to stdout (sphinx-gallery captures stdout, not stderr)
# force=True ensures fresh config even if another example ran first
logging.basicConfig(
level=logging.INFO, format="%(message)s", stream=sys.stdout, force=True
)
Setup: simulate data#
We use a simple simulation for demonstration. See Simulating from the GFA model for details on the generative model.
sim_config = GFASimConfig(
n_samples=50,
y_dims=np.array([8, 8]),
x_dim=3,
snr=1.0,
random_seed=0,
)
hyperprior = ObsParamsHyperPrior(
a_alpha=1.0, b_alpha=1.0, a_phi=1.0, b_phi=1.0, beta_d=1.0
)
sim_result = simulate(sim_config, hyperprior)
Y = sim_result.observations
Callback overview#
Callbacks allow custom actions at specific points during fitting:
on_fit_start(ctx)— called after initializationon_fit_end(ctx, reason)— called when fitting completeson_iteration_end(ctx, iteration, lb, lb_prev)— called each iterationon_flag_changed(ctx, flag, value, iteration)— called on status changeson_x_dim_pruned(ctx, n_removed, x_dim_remaining, iteration)— called when ARD prunes
Implement any subset of these methods (duck typing). Three callbacks are
provided: ProgressCallback,
LoggingCallback, and
CheckpointCallback.
ProgressCallback#
ProgressCallback displays a tqdm progress bar
during fitting. Best for interactive terminal sessions.
progress_cb = ProgressCallback(desc="Fitting GFA")
print(f"ProgressCallback: {progress_cb}")
ProgressCallback: ProgressCallback(desc='Fitting GFA')
Note
ProgressCallback works best in terminals. In notebooks and documentation builds, tqdm output can be verbose. We don’t run it in this example.
LoggingCallback#
LoggingCallback emits structured log events
to Python’s logging system. Configure logging to see output (done in imports).
config = GFAFitConfig(
x_dim_init=5,
fit_tol=1e-6,
max_iter=500,
random_seed=0,
prune_x=True,
)
model = GFAModel(config=config)
model.fit(Y, callbacks=[LoggingCallback()])
fit.started [n_samples=50; y_dims=[8 8]; x_dim=5]
fit.x_dim_pruned [n_removed=1; x_dim_remaining=4; iteration=15]
fit.x_dim_pruned [n_removed=1; x_dim_remaining=3; iteration=18]
fit.max_iter [iteration=500]
GFAModel(fitted=True, config=GFAFitConfig(x_dim_init=5, fit_tol=1e-06, max_iter=500, prune_x=True, prune_tol=1e-07, save_x=False, save_c_cov=False, save_fit_progress=True, random_seed=0, min_var_frac=0.001))
Log output shows fit start, dimension pruning events, and convergence.
For debugging, set level=logging.DEBUG to see more detail.
CheckpointCallback#
CheckpointCallback saves model state periodically
during fitting. Useful for long runs where you want recovery from interruption.
We use a temporary directory for this example.
with tempfile.TemporaryDirectory() as tmpdir:
checkpoint_dir = Path(tmpdir)
checkpoint_cb = CheckpointCallback(
save_dir=checkpoint_dir,
every_n_iter=100, # Save every 100 iterations
save_initial=True, # Save before iteration 0
save_final=True, # Save after convergence
max_checkpoints=2, # Keep only 2 periodic checkpoints
)
model2 = GFAModel(config=config)
model2.fit(Y, callbacks=[LoggingCallback(), checkpoint_cb])
# List saved checkpoints
checkpoints = sorted(checkpoint_dir.glob("*.safetensors"))
print(f"\nCheckpoints saved to {checkpoint_dir}:")
for cp in checkpoints:
print(f" {cp.name}")
fit.started [n_samples=50; y_dims=[8 8]; x_dim=5]
fit.checkpoint_saved [path=/tmp/tmpfqzvr7ws/checkpoint_init.safetensors; iteration=1]
fit.x_dim_pruned [n_removed=1; x_dim_remaining=4; iteration=15]
fit.x_dim_pruned [n_removed=1; x_dim_remaining=3; iteration=18]
fit.checkpoint_saved [path=/tmp/tmpfqzvr7ws/checkpoint_iter_000100.safetensors; iteration=100]
fit.checkpoint_saved [path=/tmp/tmpfqzvr7ws/checkpoint_iter_000200.safetensors; iteration=200]
fit.checkpoint_saved [path=/tmp/tmpfqzvr7ws/checkpoint_iter_000300.safetensors; iteration=300]
fit.checkpoint_saved [path=/tmp/tmpfqzvr7ws/checkpoint_iter_000400.safetensors; iteration=400]
fit.checkpoint_saved [path=/tmp/tmpfqzvr7ws/checkpoint_iter_000500.safetensors; iteration=500]
fit.max_iter [iteration=500]
fit.checkpoint_saved [path=/tmp/tmpfqzvr7ws/checkpoint_final.safetensors; iteration=500]
Checkpoints saved to /tmp/tmpfqzvr7ws:
checkpoint_final.safetensors
checkpoint_init.safetensors
checkpoint_iter_000400.safetensors
checkpoint_iter_000500.safetensors
The max_checkpoints parameter limits disk usage by deleting older
periodic checkpoints. Initial, final, and interrupt checkpoints are
always kept.
For interrupt handling, CheckpointCallback installs a SIGINT handler
that saves state before exiting when you press Ctrl+C.
Combining callbacks#
Multiple callbacks can be used together. They execute in list order.
fit.started [n_samples=50; y_dims=[8 8]; x_dim=5]
fit.checkpoint_saved [path=/tmp/tmpja9e51jf/checkpoint_init.safetensors; iteration=1]
fit.x_dim_pruned [n_removed=1; x_dim_remaining=4; iteration=15]
fit.x_dim_pruned [n_removed=1; x_dim_remaining=3; iteration=18]
fit.checkpoint_saved [path=/tmp/tmpja9e51jf/checkpoint_iter_000200.safetensors; iteration=200]
fit.checkpoint_saved [path=/tmp/tmpja9e51jf/checkpoint_iter_000400.safetensors; iteration=400]
fit.max_iter [iteration=500]
fit.checkpoint_saved [path=/tmp/tmpja9e51jf/checkpoint_final.safetensors; iteration=500]
Model serialization#
Save fitted models with save() and reload with
load(). Uses safetensors format for security
(no arbitrary code execution on load).
with tempfile.TemporaryDirectory() as tmpdir:
model_path = Path(tmpdir) / "model.safetensors"
# Save the fitted model
model.save(model_path)
print(f"Model saved to {model_path.name}")
# Load into a new instance
loaded_model = GFAModel.load(model_path)
# Verify posteriors match
C_original = model.obs_posterior.C.mean
C_loaded = loaded_model.obs_posterior.C.mean
print(f"Posteriors match: {np.allclose(C_original, C_loaded)}")
# Loaded model is ready for inference
X_inferred = loaded_model.infer_latents(Y)
print(f"Inferred latents shape: {X_inferred.mean.shape}")
Model saved to model.safetensors
Posteriors match: True
Inferred latents shape: (3, 50)
Resume interrupted fits#
If a fit is interrupted (e.g., timeout, Ctrl+C with checkpoint), you can
resume from a saved checkpoint with resume_fit().
with tempfile.TemporaryDirectory() as tmpdir:
model_path = Path(tmpdir) / "partial.safetensors"
# Fit with early stopping (simulate interruption)
# save_x=True is required for resume_fit to work without manual recomputation
short_config = GFAFitConfig(
x_dim_init=5,
fit_tol=1e-10, # Very tight tolerance
max_iter=50, # Stop early
random_seed=0,
prune_x=True,
save_x=True,
)
model_partial = GFAModel(config=short_config)
model_partial.fit(Y, callbacks=[LoggingCallback()])
model_partial.save(model_path)
iterations_before = len(model_partial.tracker.lb)
print(f"\nIterations before resume: {iterations_before}")
print(f"Converged: {model_partial.flags.converged}")
# Resume from checkpoint
resumed = GFAModel.load(model_path)
resumed.resume_fit(Y, max_iter=5000, callbacks=[LoggingCallback()])
iterations_after = len(resumed.tracker.lb)
print(f"Iterations after resume: {iterations_after}")
print(f"Converged: {resumed.flags.converged}")
fit.started [n_samples=50; y_dims=[8 8]; x_dim=5]
fit.x_dim_pruned [n_removed=1; x_dim_remaining=4; iteration=15]
fit.x_dim_pruned [n_removed=1; x_dim_remaining=3; iteration=18]
fit.max_iter [iteration=50]
Iterations before resume: 50
Converged: False
fit.started [n_samples=50; y_dims=[8 8]; x_dim=3]
fit.converged [iteration=2196]
Iterations after resume: 2196
Converged: True
The tracker appends iterations from the resumed fit, preserving the full optimization history.
Simulation serialization#
For reproducibility, save simulation results with
save_simulation() and reload with
load_simulation().
with tempfile.TemporaryDirectory() as tmpdir:
sim_path = Path(tmpdir) / "simulation.safetensors"
# Save complete simulation (observations, latents, parameters)
save_simulation(sim_path, sim_result)
print(f"Simulation saved to {sim_path.name}")
# Reload
loaded_sim = load_simulation(sim_path)
# Verify
obs_match = np.allclose(sim_result.observations.data, loaded_sim.observations.data)
print(f"Observations match: {obs_match}")
latents_match = np.allclose(sim_result.latents.data, loaded_sim.latents.data)
print(f"Latents match: {latents_match}")
Simulation saved to simulation.safetensors
Observations match: True
Latents match: True
This saves the complete snapshot: config, hyperprior, sampled parameters,
latents, and observations. For smaller files when you only need
reproducibility, use save_simulation_recipe()
which saves just the config and hyperprior (requires random_seed).
Total running time of the script: (0 minutes 1.069 seconds)