diff --git a/configs/packaged_inference.yaml b/configs/packaged_inference.yaml new file mode 100644 index 00000000..42166d4a --- /dev/null +++ b/configs/packaged_inference.yaml @@ -0,0 +1,33 @@ +defaults: + - paths: ${oc.env:STORAGE_MODE,local} + - extras: default + - hydra: default + - data: s2bms_prediction + - _self_ + +# task name, determines output directory path +task_name: "inference" +tags: ["dev"] +seed: 12345 + +geo_data: + aef_avr: + path: ${paths.root_dir}/inference/s2bms/avr_aef_256.csv # can be a subset or just one entry + enable_nans: true +# dimension: +# dtype: + +inference_ckpt_path: ${paths.root_dir}/inference/s2bms/v1/example.ckpt + +device: mps + +# Merging +#predictive_ckpt_path: ${paths.checkpoint_root}/s2bms_prediction_ckpt/qiqe3b3d_epoch_033.ckpt +#alignment_ckpt_path: ${paths.checkpoint_root}/s2bms_alignment_ckpt/5dwqrxbk_epoch_018.ckpt +#training_order: ["prediction_model", "alignment_model"] +# +## If set, inference.py will save a merged inference checkpoint you can reload with +## `inference_ckpt_path`. +#save_ckpt: true +##save_inference_ckpt_path: ${paths.checkpoint_root}/s2bms_inference_model/${now:%Y-%m-%d}_${now:%H-%M-%S}.csv +#save_inference_ckpt_path: ${paths.root_dir}/inference/v1/example.ckpt diff --git a/configs/paths/local.yaml b/configs/paths/local.yaml index c7e598e3..e915ca4f 100644 --- a/configs/paths/local.yaml +++ b/configs/paths/local.yaml @@ -9,6 +9,7 @@ data_dir: ${oc.env:DATA_DIR,${paths.root_dir}/data/} cache_dir: ${oc.env:CACHE_DIR,${paths.data_dir}/cache} checkpoint_root: ${paths.data_dir}/checkpoints/ checkpoint_dir: ${paths.checkpoint_root}/${oc.select:logger.wandb.project,other}_ckpt +outputs_dir: ${paths.root_dir}/outputs # path to logging directory log_dir: ${paths.root_dir}/logs/ diff --git a/configs/paths/shared.yaml b/configs/paths/shared.yaml index 4a8e1326..9ba1f2db 100644 --- a/configs/paths/shared.yaml +++ b/configs/paths/shared.yaml @@ -9,6 +9,7 @@ data_dir: ${oc.env:DATA_DIR,${oc.env:SHARED_ROOT,${paths.root_dir}}/data} cache_dir: ${oc.env:SHARED_CACHE,${paths.data_dir}/cache} checkpoint_root: ${paths.data_dir}/checkpoints/ checkpoint_dir: ${paths.checkpoint_root}/${oc.select:logger.wandb.project,other}_ckpt +outputs_dir: ${paths.root_dir}/outputs # path to logging directory log_dir: ${oc.env:SHARED_ROOT,${paths.root_dir}}/logs/ diff --git a/inference/s2bms/v1/.gitkeep b/inference/s2bms/v1/.gitkeep new file mode 100644 index 00000000..e69de29b diff --git a/notebooks/20-GT-prep_inference_data.ipynb b/notebooks/20-GT-prep_inference_data.ipynb new file mode 100644 index 00000000..2f624289 --- /dev/null +++ b/notebooks/20-GT-prep_inference_data.ipynb @@ -0,0 +1,111 @@ +{ + "cells": [ + { + "cell_type": "code", + "execution_count": null, + "id": "0", + "metadata": {}, + "outputs": [], + "source": [ + "import os\n", + "\n", + "import pandas as pd\n", + "import torch\n", + "\n", + "os.chdir(\"..\")\n", + "os.getcwd()" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "1", + "metadata": {}, + "outputs": [], + "source": [ + "pth = \"data/s2bms/splits/s2bms_union_val_test.pth\"\n", + "split_indices = torch.load(pth, weights_only=False)" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "2", + "metadata": {}, + "outputs": [], + "source": [ + "df_model_ready = pd.read_csv(\"data/s2bms/model_ready_s2bms.csv\")\n", + "df_model_ready[\"split\"] = None\n", + "\n", + "df_model_ready.loc[\n", + " df_model_ready[\"name_loc\"].isin(list(split_indices[\"val_indices\"])), \"split\"\n", + "] = \"val\"\n", + "df_model_ready.loc[\n", + " df_model_ready[\"name_loc\"].isin(list(split_indices[\"test_indices\"])), \"split\"\n", + "] = \"test\"\n", + "df_model_ready.loc[\n", + " df_model_ready[\"name_loc\"].isin(list(split_indices[\"train_indices\"])), \"split\"\n", + "] = \"train\"\n", + "df_model_ready = df_model_ready.dropna(subset=[\"split\"])\n", + "df_model_ready[\"split\"].value_counts(dropna=False)\n", + "\n", + "keep_columns = [\n", + " i for i in df_model_ready.columns if i in [\"split\", \"name_loc\", \"lon\", \"lat\"] or \"target_\" in i\n", + "]\n", + "df_model_ready[keep_columns]\n", + "\n", + "os.makedirs(\"inference/s2bms\", exist_ok=True)\n", + "df_model_ready.to_csv(\"inference/s2bms/targets.csv\", index=False)" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "3", + "metadata": {}, + "outputs": [], + "source": [ + "df = pd.read_csv(\"data/s2bms/eo/avr_aef_256.csv\")\n", + "df[\"split\"] = None\n", + "\n", + "df.loc[df[\"name_loc\"].isin(list(split_indices[\"val_indices\"])), \"split\"] = \"val\"\n", + "df.loc[df[\"name_loc\"].isin(list(split_indices[\"test_indices\"])), \"split\"] = \"test\"\n", + "df.loc[df[\"name_loc\"].isin(list(split_indices[\"train_indices\"])), \"split\"] = \"train\"\n", + "df = df.dropna(subset=[\"split\"])\n", + "df[\"split\"].value_counts(dropna=False)" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "4", + "metadata": {}, + "outputs": [], + "source": [ + "df = df.merge(df_model_ready[[\"name_loc\", \"lon\", \"lat\"]], on=\"name_loc\", how=\"left\")\n", + "df.to_csv(\"inference/s2bms/avr_aef_256.csv\")" + ] + } + ], + "metadata": { + "kernelspec": { + "display_name": "Python 3", + "language": "python", + "name": "python3" + }, + "language_info": { + "codemirror_mode": { + "name": "ipython", + "version": 2 + }, + "file_extension": ".py", + "mimetype": "text/x-python", + "name": "python", + "nbconvert_exporter": "python", + "pygments_lexer": "ipython2", + "version": "2.7.6" + } + }, + "nbformat": 4, + "nbformat_minor": 5 +} diff --git a/src/models/inference_model.py b/src/models/inference_model.py index d904072c..dc61201a 100644 --- a/src/models/inference_model.py +++ b/src/models/inference_model.py @@ -10,12 +10,14 @@ from src.models.components.geo_encoders.base_geo_encoder import BaseGeoEncoder from src.models.components.metrics.metrics_wrapper import MetricsWrapper from src.models.components.pred_heads.linear_pred_head import BasePredictionHead +from src.models.components.projectors.base_projector import BaseProjector from src.models.components.text_encoders.base_text_encoder import ( BaseTextEncoder, ) from src.utils import RankedLogger from src.utils.errors import FileNotSpecified from src.utils.logging_utils import log_model_loading +from utils.errors import IllegalArgumentCombination log = RankedLogger(__name__, rank_zero_only=True) @@ -23,13 +25,14 @@ class InferenceModel(BaseModel): def __init__( self, - geo_encoder: BaseGeoEncoder | None, - geo_encoder_prediction: BaseGeoEncoder | None, - text_encoder: BaseTextEncoder | None, - prediction_head: BasePredictionHead | None, - num_classes: int | None, + geo_encoder: BaseGeoEncoder, + text_encoder: BaseTextEncoder, + prediction_head: BasePredictionHead, + num_classes: int, + geo_adapter: BaseProjector | None = None, + text_adapter: BaseProjector | None = None, metrics: MetricsWrapper | None = None, - ks: list[int] | None = [5, 10, 15], + ks: list[int] | None = None, match_to_geo: bool = True, **kwargs, ) -> None: @@ -58,12 +61,12 @@ def __init__( tabular_dim=None, ) - if geo_encoder_prediction: - self.geo_encoder_prediction = geo_encoder_prediction + self.geo_adapter = geo_adapter + self.text_adapter = text_adapter # Params from alignment model self.match_to_geo = match_to_geo - self.ks = ks + self.ks = ks or [5, 10, 15] @override def _setup(self, stage: str) -> None: @@ -71,32 +74,40 @@ def _setup(self, stage: str) -> None: if stage != "inference": raise ValueError(f"Trying to {stage} inference model") - print("-------Model------------") + log.info("-------Model------------") # Configure encoders - if hasattr(self, "geo_encoder"): + if self.geo_encoder: self.geo_encoder.setup() - if hasattr(self, "text_encoder"): + if self.geo_adapter: + self.geo_adapter.set_input_dim(self.geo_encoder.output_dim) + self.geo_adapter.setup() + + if self.text_encoder: self.text_encoder.setup() - if hasattr(self, "geo_encoder_prediction"): - self.geo_encoder_prediction.setup() - - if hasattr(self, "text_encoder") and hasattr(self, "geo_encoder"): - # Configure optional extra projection so text embeddings match geo embeddings. - if self.text_encoder.output_dim != self.geo_encoder.output_dim: - if self.match_to_geo: - self.text_encoder.add_projector(projected_dim=self.geo_encoder.output_dim) - else: - self.geo_encoder.add_projector(projected_dim=self.text_encoder.output_dim) + if self.text_adapter: + self.text_adapter.set_input_dim(self.text_encoder.output_dim) + self.text_adapter.setup() + + # Sanity check for dimension matching + geo_branch_dim = ( + self.geo_adapter.output_dim if self.geo_adapter else self.geo_encoder.output_dim + ) + text_branch_dim = ( + self.text_adapter.output_dim if self.text_adapter else self.text_encoder.output_dim + ) + + if geo_branch_dim != text_branch_dim: + raise IllegalArgumentCombination( + "Provided prediction and alignment model checkpoints are not mergeable" + ) + # Configure prediction head - if hasattr(self, "prediction_head") and self.prediction_head.net is None: + if self.prediction_head and self.prediction_head.net is None: if self.num_classes is None: raise ValueError( "InferenceModel requires `num_classes` to build the prediction head." ) - if hasattr(self, "geo_encoder_prediction"): - input_dim = self.geo_encoder_prediction.output_dim - else: - self.geo_encoder.output_dim + input_dim = self.geo_encoder.output_dim self.prediction_head.set_dim(input_dim=input_dim, output_dim=self.num_classes) self.prediction_head.setup() print("------------------------") @@ -112,17 +123,19 @@ def _step( def forward_text(self, text: list[str]) -> torch.Tensor: batch = {"text": text} - return self.text_encoder(batch, "train") + text_feats = self.text_encoder(batch, "train") + if self.text_adapter: + text_feats = self.text_adapter(text_feats) + return text_feats def forward_geo(self, batched_eo) -> Tuple[torch.Tensor, torch.Tensor | None]: geo_feats = self.geo_encoder(batched_eo) - pred = None - if hasattr(self, "prediction_head"): - if hasattr(self, "geo_encoder_prediction"): - geo_feats_pr = self.geo_encoder_prediction(batched_eo) - pred = self.prediction_head(geo_feats_pr) - else: - pred = self.prediction_head(geo_feats) + + if self.prediction_head: + pred = self.prediction_head(geo_feats) + + if self.geo_adapter: + geo_feats = self.geo_adapter(geo_feats) return geo_feats, pred @override @@ -130,21 +143,15 @@ def forward( self, batch: Dict[str, torch.Tensor], ) -> Tuple[torch.Tensor | None, torch.Tensor | None, torch.Tensor | None]: - if hasattr(self, "geo_encoder"): - geo_feats, pred = self.forward_geo(batch) - if hasattr(self, "text_encoder"): - text_feats = self.text_encoder(batch, "test") + geo_feats, pred = self.forward_geo(batch) + text_feats = self.text_encoder(batch, "test") # Change dtype of geo data if it does not match text dtype - if ( - hasattr(self, "text_encoder") - and hasattr(self, "geo_encoder") - and geo_feats.dtype != text_feats.dtype - ): + if geo_feats.dtype != text_feats.dtype: geo_feats = geo_feats.to(text_feats.dtype) - return pred, geo_feats, text_feats + return pred, geo_feats, text_feats def concept_similarities(self, geo_embeds, concept=None) -> torch.Tensor: # Get concept embeddings @@ -157,9 +164,6 @@ def concept_similarities(self, geo_embeds, concept=None) -> torch.Tensor: elif self.concept_embeds is not None: concept_embeds = self.concept_embeds - else: - with torch.no_grad(): - concept_embeds = self.text_encoder({"text": self.concepts}, mode="train") # Similarity geo_embeds = F.normalize(geo_embeds, dim=1) @@ -225,19 +229,13 @@ def merge_inference_model(cfg, save_ckpt=False) -> InferenceModel | None: align_ckpt = torch.load(align_ckpt_path, map_location="cpu", weights_only=False) inference_hparams = {} - geo_encoder_prediction_flag = False # Sanity check: ensure geo encoder configs match. if pred_ckpt["hyper_parameters"].get("geo_encoder") != align_ckpt["hyper_parameters"].get( "geo_encoder" ): - log.warning("Geo encoder configs differ between checkpoints; results may be invalid.") - if input("Do you want to proceed? y/n").lower() != "y": - return None - if input("Use separate geo_encoders for prediction and alignment? y/n").lower() == "y": - inference_hparams["geo_encoder_prediction"] = pred_ckpt["hyper_parameters"].get( - "geo_encoder" - ) - geo_encoder_prediction_flag = True + raise IllegalArgumentCombination( + "Geo encoder configs differ between checkpoints; results may be invalid." + ) pred_trainable_modules = pred_ckpt["hyper_parameters"].get("trainable_modules", []) align_trainable_modules = align_ckpt["hyper_parameters"].get("trainable_modules", []) @@ -265,43 +263,50 @@ def merge_inference_model(cfg, save_ckpt=False) -> InferenceModel | None: model: InferenceModel = hydra.utils.instantiate(inference_hparams) model.setup("inference") - # Load alignment weights first (text encoder). - text_state = { - k: v for k, v in align_ckpt["state_dict"].items() if k.startswith("text_encoder.") - } - if text_state: - res = model.load_state_dict(text_state, strict=False) - log_model_loading("text_encoder", res) - - if geo_encoder_prediction_flag: - geo_state = { - k: v for k, v in align_ckpt["state_dict"].items() if k.startswith("geo_encoder.") - } - geo_pred_state = { - k: v for k, v in pred_ckpt["state_dict"].items() if k.startswith("geo_encoder.") - } - if geo_pred_state: - res = model.load_state_dict(geo_pred_state, strict=False) - log_model_loading("geo_encoder_pred", res) - elif cfg.training_order[0] == "prediction_model" and not geo_align_encoder_trained: - geo_state = { - k: v for k, v in pred_ckpt["state_dict"].items() if k.startswith("geo_encoder.") + collected_states = {} + + # Get text encoder + collected_states.update( + { + k: v + for k, v in align_ckpt["state_dict"].items() + if k.startswith("text_encoder.") + and k != "text_encoder.model.text_model.embeddings.position_ids" } + ) + + # Get text adapter + collected_states.update( + {k: v for k, v in align_ckpt["state_dict"].items() if k.startswith("text_adapter.")} + ) + + # Get geo_encoder + if cfg.training_order[0] == "prediction_model" and not geo_align_encoder_trained: + collected_states.update( + {k: v for k, v in pred_ckpt["state_dict"].items() if k.startswith("geo_encoder.")} + ) else: - geo_state = { - k: v for k, v in align_ckpt["state_dict"].items() if k.startswith("geo_encoder.") - } - if geo_state: - res = model.load_state_dict(geo_state, strict=False) - log_model_loading("geo_encoder", res) - - # Load prediction head weights from predictive ckpt. - head_state = { - k: v for k, v in pred_ckpt["state_dict"].items() if k.startswith("prediction_head.") - } - if head_state: - res = model.load_state_dict(head_state, strict=False) - log_model_loading("Predictive_head", res) + collected_states.update( + { + k: v + for k, v in align_ckpt["state_dict"].items() + if k.startswith("geo_encoder.") or k.startswith("geo_adapter.") + } + ) + + # Get geo_adapter + collected_states.update( + {k: v for k, v in align_ckpt["state_dict"].items() if k.startswith("geo_adapter.")} + ) + + # Get prediction head weights from predictive ckpt. + collected_states.update( + {k: v for k, v in pred_ckpt["state_dict"].items() if k.startswith("prediction_head.")} + ) + + # Load collected states + res = model.load_state_dict(collected_states, strict=False) + log_model_loading("Inference Model", res) # Save model if save_ckpt: diff --git a/src/packaged_inference.py b/src/packaged_inference.py new file mode 100644 index 00000000..591d810c --- /dev/null +++ b/src/packaged_inference.py @@ -0,0 +1,124 @@ +import os +from typing import Optional + +import hydra +import pandas as pd +import rootutils +import torch +import torch.nn.functional as F +from dotenv import load_dotenv +from omegaconf import DictConfig + +from src.models.inference_model import load_inference_model, merge_inference_model +from src.utils import extras + +rootutils.setup_root(__file__, indicator=".project-root", pythonpath=True) +load_dotenv() + +# Disable tokenizers parallelism to avoid warnings when using multiprocessing +if os.environ.get("TOKENIZERS_PARALLELISM") is None: + os.environ["TOKENIZERS_PARALLELISM"] = "false" + + +def get_model(cfg): + # If a merged inference ckpt is provided, just load it. + inference_ckpt_path = cfg.get("inference_ckpt_path") + if inference_ckpt_path: + model = load_inference_model(inference_ckpt_path, cfg.paths.cache_dir) + else: + model = merge_inference_model(cfg, save_ckpt=True) + + return model + + +def get_geo_data(cfg, device="cpu"): + import pandas as pd + + params = cfg["geo_data"] + modality = list(params.keys())[0] + params = params[modality] + dtype = params.get("dtype", torch.float32) + dim = params.get("dimension", 64 if modality == "aef_avr" else 128) + + path = params.get("path", KeyError(f"Please specify {modality} path to csv file")) + assert os.path.exists(path), FileNotFoundError(f"{path} does not exist.") + df = pd.read_csv(path) + + # Filter out locations without data for embeddings + emb_cols = [f"emb_{i}" for i in range(dim)] + + geo_data = torch.tensor(df[emb_cols].to_numpy(), dtype=dtype, device=device) + return { + "eo": {modality: geo_data}, + "name_loc": df.name_loc.to_list(), + "lat": df.lat.to_list(), + "lon": df.lon.to_list(), + "split": df.split.to_list(), + } + + +@torch.no_grad() +@hydra.main(version_base="1.3", config_path="../configs", config_name="packaged_inference.yaml") +def main(cfg: DictConfig, save_results: bool = False) -> Optional[float]: + """Main entry point for training. + + :param cfg: DictConfig configuration composed by Hydra. + :param save_results: Whether to save inference results. + :return: Optional[float] with optimized metric value. + """ + # apply extra utilities + # (e.g. ask for tags if none are provided in cfg, print cfg tree, etc.) + extras(cfg) + + model = get_model(cfg) + model.to(cfg.device) + + # TODO: do what you need with the inference model + # Supply text in a list. Examples + text = ["Forested area", "Area with water bodies near-by"] + # or + # text = [input('Enter a concept')] + # or + # with open('../outputs/example_concepts.txt', 'r') as f: + # text = f.readlines() + # text = [t.strip() for t in text] + + text_embeds = model.forward_text(text) + + # supply geo data as csv file + if cfg.get("geo_data"): + b = get_geo_data( + cfg, device=cfg.get("device", "cpu") + ) # batch has name_loc, lat, lon and split arguments + geo_embeds, pred = model.forward_geo(b) + + geo_embeds = F.normalize(geo_embeds, dim=1) + text_embeds = F.normalize(text_embeds, dim=1) + similarity_matrix = geo_embeds @ text_embeds.T + + results = torch.cat([similarity_matrix, pred], dim=1) + + if cfg.get("save_output") or save_results: + df = pd.DataFrame( + results.cpu(), + columns=text + [f"target_{i}" for i in range(pred.shape[-1])], + index=b.get("name_loc"), + ) + df.reset_index(inplace=True, names=["name_loc"]) + path = cfg.get("save_output") + os.makedirs(os.path.dirname(path), exist_ok=True) + df.to_csv(path) + print(f"Cosine similarities are saved to {path}") + else: + names = b["name_loc"] + for i in range(len(names)): + print(f"Location {names[i]} similarity with:") + for j, t in enumerate(text): + print(f" - {t}: {similarity_matrix[i, j]:.2f} similarity") + print(f"And prediction vector: {pred[i].detach().cpu().numpy()}") + + return + + +if __name__ == "__main__": + main()