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.
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.
Train with Keras
Section titled “Train with Keras”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.
Inspect or customize the training step
Section titled “Inspect or customize the training step”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:
python tests/trainers/compare_mae_backends.py \ --tensorflow-python /path/to/tf-env/bin/python \ --torch-python /path/to/torch-env/bin/pythonCI 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.