Skip to content
heliaEDGE
User guide
HELIA

Masked-autoencoder training

A masked autoencoder learns from the input itself: hide some patches, reconstruct them, and compare the predictions with the original values at those hidden positions.

Masked reconstruction: keep original targets, encode visible patches, decode with mask tokens, and compute loss only at masked positions.

The encoder sees visible patches. The decoder receives their encoded representation plus mask tokens. MaskedAutoencoder returns the target and predicted patches at matching masked positions.

With a constructed MaskedAutoencoder named model and input samples x:

model.compile(loss="mse", optimizer="adam", jit_compile=False)
model.fit(x, epochs=10)

Targets come from x, so pass no separate y or sample_weight. Both are rejected. The forward path uses public Keras operations; gradient application has separate TensorFlow and PyTorch implementations.

To construct model, connect a patch layer, patch encoder, encoder and decoder. The API describes their contracts; the CPU parity fixture provides a complete small construction with synthetic inputs.

Access reconstruction pairs for a native training loop
targets, predictions = model.reconstruction_targets(x, training=True)

Use these tensors with your native TensorFlow or PyTorch objective and optimizer.

call() returns the same Reconstruction(targets, predictions) named tuple. calculate_loss(x, test=False) returns ReconstructionLoss(loss, targets, predictions) using the compiled objective. Both retain tuple unpacking. Keras layers receive the training flag; plain single-argument callables retain their original calling convention and must own their state behavior. Use serializable Keras layers to save complete models.

Mask randomness and saving a model

MaskedPatchEncoder2D(seed=...) tracks a Keras seed generator. Inference disables Dropout/BatchNorm training behavior but still samples masks. Same-seed random streams are tested within each backend; cross-backend numerical checks use fixed masks, inputs and weights. Whole RNG/sampler resume and identical stochastic training trajectories are not promised. Save/reload tests cover full-model weights, controlled-mask predictions and optimizer slots/iteration on each backend.

Check TensorFlow and PyTorch parity

Build the model before constructing a native Torch optimizer so every parameter is present. The regression fixtures in tests/trainers/ exercise native TF/Torch optimization. Run their forward/gradient/update comparison with:

Terminal window
python tests/trainers/compare_mae_backends.py \
--tensorflow-python /path/to/tf-env/bin/python \
--torch-python /path/to/torch-env/bin/python

CI compares evidence produced in the separate backend jobs at rtol=1e-5, atol=1e-6. This is a small float32 CPU fixture, not certification of mixed precision, Torch compilation, distribution or MAE export.