def save_mae_reconstruction_preview(
samples: list[dict[str, Any]],
output: MarsMAEOutput,
out_path: pathlib.Path | str,
*,
patch_size: int,
max_items: int = 4,
) -> pathlib.Path:
"""Save original/masked/reconstructed MAE panels for a small batch."""
items = min(max_items, len(samples), output.reconstruction.shape[0])
if items <= 0:
raise ValueError("No samples provided.")
images = torch.stack([sample["image"] for sample in samples[:items]], dim=0)
masked_inputs = build_masked_input_image(
images,
output.patch_valid_mask[:items],
output.visible_mask[:items],
patch_size=patch_size,
)
_, composites = build_reconstruction_composite(
images,
output.reconstruction[:items],
output.patch_valid_mask[:items],
output.visible_mask[:items],
patch_size=patch_size,
)
fig, axes = plt.subplots(items, 4, figsize=(16, 3.8 * items))
if items == 1:
axes = np.array([axes])
for idx in range(items):
sample = samples[idx]
md = sample["metadata"]
row_axes = axes[idx]
row_axes[0].imshow(_to_display_rgb(images[idx]), interpolation="nearest")
row_axes[0].set_title(f"input\n{md['obs_id']}", fontsize=9)
row_axes[1].imshow(_to_display_rgb(masked_inputs[idx]), interpolation="nearest")
row_axes[1].set_title("masked encoder input", fontsize=9)
row_axes[2].imshow(_to_display_rgb(composites[idx]), interpolation="nearest")
row_axes[2].set_title(
f"reconstruction composite\nloss={float(output.loss):.4f}",
fontsize=9,
)
overlay = _build_patch_state_overlay(
images[idx],
output.patch_valid_mask[idx],
output.visible_mask[idx],
patch_size=patch_size,
)
row_axes[3].imshow(overlay, interpolation="nearest")
row_axes[3].set_title("patch states", fontsize=9)
for ax in row_axes:
ax.axis("off")
fig.suptitle("MarsCLIP Stage A MAE preview", fontsize=12)
fig.tight_layout()
out = pathlib.Path(out_path)
out.parent.mkdir(parents=True, exist_ok=True)
fig.savefig(out, dpi=160)
plt.close(fig)
return out