Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
33 changes: 33 additions & 0 deletions configs/packaged_inference.yaml
Original file line number Diff line number Diff line change
@@ -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
1 change: 1 addition & 0 deletions configs/paths/local.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -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/
Expand Down
1 change: 1 addition & 0 deletions configs/paths/shared.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -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/
Expand Down
Empty file added inference/s2bms/v1/.gitkeep
Empty file.
111 changes: 111 additions & 0 deletions notebooks/20-GT-prep_inference_data.ipynb
Original file line number Diff line number Diff line change
@@ -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
}
Loading
Loading