SpectralGPT / scripts /result.py
zhangrenchao's picture
Update SpectralGPT model package
7a2d30b verified
Raw
History Blame Contribute Delete
3.21 kB
import argparse
import json
from pathlib import Path
import numpy as np
import yaml
import matplotlib.pyplot as plt
def rgb(image):
array = np.clip(image[[3, 2, 1]], 0, 1).transpose(1, 2, 0)
return (array * 255).astype(np.uint8)
def main():
parser = argparse.ArgumentParser(description="Evaluate and visualize reconstruction")
parser.add_argument("--config", default="conf/config.yaml")
args = parser.parse_args()
with open(args.config, encoding="utf-8") as handle:
config = yaml.safe_load(handle)
output_dir = Path(config["runtime"]["output_dir"])
reconstruction_path = output_dir / "reconstruction.npz"
if not reconstruction_path.exists():
raise FileNotFoundError(
f"Missing inference output: {reconstruction_path}. "
"Run `python scripts/inference.py` first."
)
with np.load(reconstruction_path) as data:
inputs = data["inputs"]
reconstructions = data["reconstructions"]
predictions = data["predictions"]
mask_images = data["mask_images"]
data_source = str(data["data_source"]) if "data_source" in data.files else "unknown"
protocol = str(data["protocol"]) if "protocol" in data.files else "unknown"
normalization = str(data["normalization"]) if "normalization" in data.files else "unknown"
denominator = max(float(mask_images.sum()), 1.0)
mse = float((((inputs - predictions) ** 2) * mask_images).sum() / denominator)
mae = float((np.abs(inputs - predictions) * mask_images).sum() / denominator)
psnr = float(-10 * np.log10(max(mse, 1e-12)))
per_band_denominator = np.maximum(mask_images.sum(axis=(0, 2, 3)), 1)
spectral_rmse = np.sqrt((((inputs - predictions) ** 2) * mask_images).sum(axis=(0, 2, 3)) / per_band_denominator)
dot = (inputs * reconstructions).sum(axis=1)
norms = np.linalg.norm(inputs, axis=1) * np.linalg.norm(reconstructions, axis=1)
pixel_mask = (mask_images > 0).any(axis=1) & (norms > 1e-8)
sam = np.arccos(np.clip(dot / np.maximum(norms, 1e-8), -1, 1))
metrics = {"masked_mse": mse, "masked_mae": mae, "masked_psnr_db": psnr,
"masked_spectral_angle_deg": float(np.degrees(sam[pixel_mask]).mean()),
"data_source": data_source, "protocol": protocol,
"normalization": normalization,
"per_band_rmse": spectral_rmse.tolist()}
with open(output_dir / "metrics.json", "w", encoding="utf-8") as handle:
json.dump(metrics, handle, indent=2)
masked = inputs[0] * (1.0 - mask_images[0])
figure, axes = plt.subplots(1, 4, figsize=(13, 3.5))
for axis, image, title in zip(axes, [inputs[0], masked, predictions[0], reconstructions[0]],
["Input", "Visible tokens", "MAE prediction", "Composite"]):
axis.imshow(rgb(image))
axis.set_title(title)
axis.axis("off")
figure.tight_layout()
figure.savefig(output_dir / "reconstruction.png", dpi=140)
plt.close(figure)
print(json.dumps(metrics, indent=2))
print(f"saved: {output_dir / 'metrics.json'}")
print(f"saved: {output_dir / 'reconstruction.png'}")
if __name__ == "__main__":
main()