alphauniverse / scripts /warmup.py
Hanbang Wu
hella changes
bb5022d
Raw
History Blame Contribute Delete
11.6 kB
"""
uv run python -m scripts.warmup
"""
from collections import defaultdict
from collections.abc import Sequence
from functools import cache
import numpy as np
import pandas as pd
import torch
import torch.nn.functional as F
from aion import AION
from aion.codecs import CodecManager
from aion.modalities import (
HSCAG,
HSCAI,
HSCAR,
HSCAY,
HSCAZ,
DESISpectrum,
HSCImage,
HSCMagG,
HSCMagI,
HSCMagR,
HSCMagY,
HSCMagZ,
HSCShape11,
HSCShape12,
HSCShape22,
LegacySurveyEBV,
LegacySurveyFluxG,
LegacySurveyFluxI,
LegacySurveyFluxR,
LegacySurveyFluxW1,
LegacySurveyFluxW2,
LegacySurveyFluxW3,
LegacySurveyFluxW4,
LegacySurveyFluxZ,
LegacySurveyImage,
LegacySurveyShapeE1,
LegacySurveyShapeE2,
LegacySurveyShapeR,
SDSSSpectrum,
)
from datasets import Dataset
from sklearn.decomposition import PCA
from sklearn.preprocessing import StandardScaler
from tqdm import tqdm
from app.helper import (
CACHE_DIR,
CROP_PX,
DES_BANDS,
DESI_MATCH_COLUMN,
DESI_SPECTRUM_COLUMN,
EMB_DIM,
FLUX_COLUMN,
GALAXY_PATH,
HSC,
HSC_BANDS,
HSC_FLUX_COLUMN,
HSC_MATCH_COLUMN,
LS,
MORPHOLOGIES,
N_GALAXIES,
N_PATCHES,
PATCH_PATHS,
PATCH_POINTS_PATHS,
POINTS_PATHS,
SDSS_MATCH_COLUMN,
SDSS_SPECTRUM_COLUMN,
TOKENS_PATH,
ProjectionMethod,
SearchMethod,
get_dataset,
get_labels,
patch_matrix,
)
DEVICE = torch.accelerator.current_accelerator(check_available=True) or torch.device(
"cpu"
)
_LS_SCALARS = [
(LegacySurveyEBV, f"EBV{LS}"),
(LegacySurveyFluxG, f"FLUX_G{LS}"),
(LegacySurveyFluxR, f"FLUX_R{LS}"),
(LegacySurveyFluxI, f"FLUX_I{LS}"),
(LegacySurveyFluxZ, f"FLUX_Z{LS}"),
(LegacySurveyFluxW1, f"FLUX_W1{LS}"),
(LegacySurveyFluxW2, f"FLUX_W2{LS}"),
(LegacySurveyFluxW3, f"FLUX_W3{LS}"),
(LegacySurveyFluxW4, f"FLUX_W4{LS}"),
(LegacySurveyShapeR, f"SHAPE_R{LS}"),
(LegacySurveyShapeE1, f"SHAPE_E1{LS}"),
(LegacySurveyShapeE2, f"SHAPE_E2{LS}"),
]
_HSC_SCALARS = [
(HSCAG, f"a_g{HSC}"),
(HSCAR, f"a_r{HSC}"),
(HSCAI, f"a_i{HSC}"),
(HSCAZ, f"a_z{HSC}"),
(HSCAY, f"a_y{HSC}"),
(HSCMagG, f"g_cmodel_mag{HSC}"),
(HSCMagR, f"r_cmodel_mag{HSC}"),
(HSCMagI, f"i_cmodel_mag{HSC}"),
(HSCMagZ, f"z_cmodel_mag{HSC}"),
(HSCMagY, f"y_cmodel_mag{HSC}"),
(HSCShape11, f"i_sdssshape_shape11{HSC}"),
(HSCShape22, f"i_sdssshape_shape22{HSC}"),
(HSCShape12, f"i_sdssshape_shape12{HSC}"),
]
_SPECTRUM_FIELDS = {
"flux": ("flux", np.float32),
"ivar": ("ivar", np.float32),
"wavelength": ("lambda", np.float32),
"mask": ("mask", bool),
}
IMAGE_TOKEN_KEY = LegacySurveyImage.token_key
N_ENCODER_TOKENS = (
LegacySurveyImage.num_tokens
+ HSCImage.num_tokens
+ DESISpectrum.num_tokens
+ SDSSSpectrum.num_tokens
+ len(_LS_SCALARS)
+ len(_HSC_SCALARS)
)
@cache
def _get_codec() -> CodecManager:
return CodecManager(device=DEVICE)
@cache
def get_model() -> AION:
model = AION.from_pretrained("polymathic-ai/aion-base").to(DEVICE).eval()
model.requires_grad_(False)
return model
@cache
def _encoder_view() -> Dataset:
columns = [
FLUX_COLUMN,
HSC_FLUX_COLUMN,
DESI_SPECTRUM_COLUMN,
SDSS_SPECTRUM_COLUMN,
HSC_MATCH_COLUMN,
DESI_MATCH_COLUMN,
SDSS_MATCH_COLUMN,
*(column for _, column in _LS_SCALARS),
*(column for _, column in _HSC_SCALARS),
]
return get_dataset().select_columns(columns).with_format("numpy")
def _tensor(values, dtype=None) -> torch.Tensor:
return torch.from_numpy(np.asarray(values, dtype)).to(DEVICE)
def _present(distances) -> np.ndarray:
return np.flatnonzero(np.isfinite(np.asarray(distances, np.float64)))
def _encode(modality) -> torch.Tensor:
encoded = _get_codec().encode(modality)[modality.token_key]
return encoded.reshape(len(encoded), -1)
def _image(modality, rows, members, bands) -> list:
if not len(members):
return []
by_band = [
{band["band"].upper(): band["flux"] for band in row} for row in rows[members]
]
flux = np.stack([[row[band] for band in bands] for row in by_band])
top = (flux.shape[-2] - CROP_PX) // 2
left = (flux.shape[-1] - CROP_PX) // 2
crop = flux[..., top : top + CROP_PX, left : left + CROP_PX]
return [(members, _encode(modality(flux=_tensor(crop, np.float32), bands=bands)))]
def _spectrum(modality, rows, members) -> list:
groups = defaultdict(list)
for i in members:
groups[len(rows[i])].append(i)
parts = []
for group in groups.values():
samples = {
argument: _tensor(
[[sample[field] for sample in rows[i]] for i in group], dtype
)
for argument, (field, dtype) in _SPECTRUM_FIELDS.items()
}
parts.append((group, _encode(modality(**samples))))
return parts
def _scalar(modality, values, members) -> list:
if not len(members):
return []
value = np.nan_to_num(np.asarray(values, np.float32)[members])
return [(members, _encode(modality(value=_tensor(value))))]
def _slot(modality, n, parts) -> tuple[torch.Tensor, torch.Tensor]:
tokens = torch.zeros(n, modality.num_tokens, dtype=torch.long, device=DEVICE)
mask = torch.ones(n, modality.num_tokens, dtype=torch.bool, device=DEVICE)
for members, encoded in parts:
index = _tensor(members, np.int64)
tokens[index] = encoded.to(DEVICE, torch.long)
mask[index] = False
return tokens, mask
def tokenize_rows(indices: Sequence[int]) -> tuple[dict, dict]:
rows = _encoder_view()[list(indices)]
n = len(rows[FLUX_COLUMN])
everything = np.arange(n)
hsc = _present(rows[HSC_MATCH_COLUMN])
encoded = {
LegacySurveyImage: _image(
LegacySurveyImage, rows[FLUX_COLUMN], everything, DES_BANDS
),
HSCImage: _image(HSCImage, rows[HSC_FLUX_COLUMN], hsc, HSC_BANDS),
DESISpectrum: _spectrum(
DESISpectrum,
rows[DESI_SPECTRUM_COLUMN],
_present(rows[DESI_MATCH_COLUMN]),
),
SDSSSpectrum: _spectrum(
SDSSSpectrum,
rows[SDSS_SPECTRUM_COLUMN],
_present(rows[SDSS_MATCH_COLUMN]),
),
**{m: _scalar(m, rows[column], everything) for m, column in _LS_SCALARS},
**{m: _scalar(m, rows[column], hsc) for m, column in _HSC_SCALARS},
}
tokens, mask = {}, {}
for modality, parts in encoded.items():
tokens[modality.token_key], mask[modality.token_key] = _slot(modality, n, parts)
return tokens, mask
def embed(tokens: dict, mask: dict) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
model = get_model()
with torch.no_grad(), torch.autocast(device_type=DEVICE.type, dtype=torch.bfloat16):
enc_tokens, enc_emb, enc_mask, mod_mask = model.embed_inputs(
tokens, mask=mask, num_encoder_tokens=N_ENCODER_TOKENS
)
ctx = model._encode(enc_tokens, enc_emb, enc_mask)
valid = ~enc_mask.squeeze(1)
galaxy = (ctx.float() * valid[..., None]).sum(dim=1) / valid.sum(
dim=1, keepdim=True
)
is_image = mod_mask == model.modality_info[IMAGE_TOKEN_KEY]["id"]
patches = (len(ctx), N_PATCHES, -1)
return (
ctx[is_image].reshape(patches),
enc_tokens[is_image].reshape(patches),
galaxy,
)
def build_embeddings() -> None:
outputs = {
method: np.lib.format.open_memmap(
PATCH_PATHS[method],
mode="w+",
dtype=np.float32,
shape=(
N_GALAXIES * N_PATCHES,
2 * EMB_DIM if method == SearchMethod.CODEBOOK_AND_ENCODED else EMB_DIM,
),
)
for method in SearchMethod
}
galaxies = np.lib.format.open_memmap(
GALAXY_PATH, mode="w+", dtype=np.float32, shape=(N_GALAXIES, EMB_DIM)
)
token_ids = np.lib.format.open_memmap(
TOKENS_PATH, mode="w+", dtype=np.int32, shape=(N_GALAXIES, N_PATCHES)
)
BATCH_SIZE = 32
for start in tqdm(range(0, N_GALAXIES, BATCH_SIZE), desc="encode"):
tokens, mask = tokenize_rows(range(start, min(start + BATCH_SIZE, N_GALAXIES)))
context, codebook, galaxy = embed(tokens, mask)
codebook = F.normalize(codebook.reshape(-1, EMB_DIM), dim=-1)
context = context.reshape(-1, EMB_DIM)
block = {
"codebook": codebook,
"encoded": context,
"codebook_and_encoded": torch.cat(
[codebook, F.normalize(context, dim=-1)], dim=-1
),
}
lo = start * N_PATCHES
for method, values in block.items():
values = values.float().cpu().numpy()
outputs[method][lo : lo + len(values)] = values
galaxies[start : start + len(galaxy)] = galaxy.cpu().numpy()
ids = tokens[IMAGE_TOKEN_KEY].cpu().numpy()
token_ids[start : start + len(ids)] = ids
for array in (*outputs.values(), galaxies, token_ids):
array.flush()
def _project(X: np.ndarray, method: ProjectionMethod) -> np.ndarray:
match method:
case "autoencoder":
from umap.parametric_umap import ParametricUMAP
scaled = StandardScaler().fit_transform(X)
coords = ParametricUMAP().fit_transform(scaled)
return StandardScaler().fit_transform(coords)
case "pca":
latents = StandardScaler().fit_transform(X)
return PCA(n_components=2).fit_transform(latents)
case "umap":
from embedding_atlas.projection import compute_projection
df = compute_projection(
pd.DataFrame({"embedding": list(X)}),
inputs="embedding",
modality="vector",
x="x",
y="y",
neighbors="neighbors",
)
return df[["x", "y"]].to_numpy()
def build_points(method: ProjectionMethod) -> None:
coords = _project(np.load(GALAXY_PATH, mmap_mode="r"), method)
category = get_labels()
pd.DataFrame(
{
"id": np.arange(len(coords)),
"x": coords[:, 0],
"y": coords[:, 1],
"category": category,
"morphology": np.take(MORPHOLOGIES, category),
}
).to_parquet(POINTS_PATHS[method], index=False)
def build_patch_points(method: ProjectionMethod) -> None:
per_galaxy = N_PATCHES + 1
kind = np.zeros((N_GALAXIES, per_galaxy), np.int8)
kind[:, N_PATCHES] = 1
patches = patch_matrix(SearchMethod.ENCODED).reshape(N_GALAXIES, N_PATCHES, EMB_DIM)
galaxies = np.load(GALAXY_PATH, mmap_mode="r").reshape(N_GALAXIES, 1, EMB_DIM)
stacked = np.concatenate([patches, galaxies], axis=1).reshape(-1, EMB_DIM)
coords = _project(stacked, method)
category = np.repeat(get_labels(), per_galaxy)
pd.DataFrame(
{
"id": np.arange(len(coords)),
"x": coords[:, 0],
"y": coords[:, 1],
"category": category,
"morphology": np.take(MORPHOLOGIES, category),
"galaxy": np.repeat(np.arange(N_GALAXIES), per_galaxy),
"kind": kind.ravel(),
}
).to_parquet(PATCH_POINTS_PATHS[method], index=False)
if __name__ == "__main__":
CACHE_DIR.mkdir(parents=True, exist_ok=True)
build_embeddings()
for projection in ProjectionMethod:
build_points(projection)
build_patch_points(projection)