# 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.](https://ambiqai.github.io/helia-edge/assets/mae-reconstruction.svg)

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

With a constructed `MaskedAutoencoder` named `model` and input samples `x`:

```python
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.

[Constructor API](https://ambiqai.github.io/helia-edge/reference/api/helia_edge/trainers/mask_autoencoder/)
[Complete CPU fixture](https://github.com/AmbiqAI/helia-edge/blob/main/tests/trainers/mae_parity_fixture.py)

## Inspect or customize the training step

Access reconstruction pairs for a native training loop

```python
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:

```sh
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.
