Skip to content

clip.visualize_marsclip_mae

visualize_marsclip_mae

Preview utilities for Stage A MAE masked/reconstructed MarsCLIP patches.

save_mae_reconstruction_preview

save_mae_reconstruction_preview(samples: list[dict[str, Any]], output: MarsMAEOutput, out_path: Path | str, *, patch_size: int, max_items: int = 4) -> pathlib.Path

Save original/masked/reconstructed MAE panels for a small batch.

Source code in src/clip/visualize_marsclip_mae.py
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