Views
No views yet
1import torch
2from pathlib import Path
3
4# Load model checkpoint
5model_path = "astropt/090M/ckpt.pt"
6device = "cuda" if torch.cuda.is_available() else "cpu"
7
8# Load state dict
9checkpoint = torch.load(model_path, map_location=device)
10# Initialize your model architecture here
11# model.load_state_dict(checkpoint)1from datasets import load_dataset
2import torch
3
4# Load dataset
5dataset = load_dataset(
6 "msiudek/astroPT_euclid_dataset",
7 split="train_batch_1",
8 streaming=True
9)
10
11# Run inference
12model.eval()
13with torch.no_grad():
14 for sample in dataset:
15 vis_image = sample['VIS_image'] # 224×224
16 vis_image = torch.tensor(vis_image, dtype=torch.float32)
17 vis_image = vis_image.unsqueeze(0).unsqueeze(0) # [1, 1, 224, 224]
18
19 # Get embeddings
20 embeddings = model(vis_image)1@article{Siudek2025,
2 title={AstroPT: Astronomical Physics Transformers for Multi-modal Learning},
3 author={Siudek, M and others},
4 journal={Euclid Collaboration},
5 eprint={2503.15312},
6 archivePrefix={arXiv},
7 year={2025},
8 url={https://ui.adsabs.harvard.edu/abs/2025arXiv250315312E/abstract}
9}