From e13d266ea5e83feb01b4e7506140c13bb24a16ab Mon Sep 17 00:00:00 2001 From: Ben Kantor <888122+bkntr@users.noreply.github.com> Date: Thu, 22 May 2025 13:32:57 +0300 Subject: [PATCH] Fix model loading on machines without gpu --- src/brainways/model/model_utils.py | 12 +++++++++--- src/brainways/utils/_tests/test_setup.py | 8 ++++++++ 2 files changed, 17 insertions(+), 3 deletions(-) create mode 100644 src/brainways/utils/_tests/test_setup.py diff --git a/src/brainways/model/model_utils.py b/src/brainways/model/model_utils.py index 1d921aa..f148d00 100644 --- a/src/brainways/model/model_utils.py +++ b/src/brainways/model/model_utils.py @@ -40,11 +40,17 @@ def load_model(model_dir: Path) -> SiameseModel: # Load state dict state_dict_path = model_dir / "state_dict.pt" - state_dict = torch.load(state_dict_path, weights_only=True) + state_dict = torch.load( + state_dict_path, + map_location="cpu", + weights_only=True, + ) model.load_state_dict(state_dict) - # Set model to evaluation mode and move to GPU model.eval() - model.to("cuda") + + # Move model to GPU if available + if torch.cuda.is_available(): + model.to("cuda") return model diff --git a/src/brainways/utils/_tests/test_setup.py b/src/brainways/utils/_tests/test_setup.py new file mode 100644 index 0000000..8494b0b --- /dev/null +++ b/src/brainways/utils/_tests/test_setup.py @@ -0,0 +1,8 @@ +from brainways.utils.setup import BrainwaysSetup + + +def test_setup(): + BrainwaysSetup( + atlas_names=["whs_sd_rat_39um", "allen_mouse_25um"], + progress_callback=lambda x: None, + ).run()