Hi,
This patch:
diff --git a/ddw/fit_model.py b/ddw/fit_model.py
index 65fe88c..9e73176 100644
--- a/ddw/fit_model.py
+++ b/ddw/fit_model.py
@@ -233,7 +233,7 @@ def fit_model(
strategy = pl.strategies.DDPStrategy(
process_group_backend=distributed_backend,
find_unused_parameters=False, # setting this to true gave a warning that it might slow things down
- ) if len(devices) > 1 else None
+ ) if len(devices) > 1 else "auto"
trainer = pl.Trainer(
max_epochs=num_epochs,
accelerator="gpu",
@@ -246,7 +246,7 @@ def fit_model(
logger=logger,
callbacks=callbacks,
detect_anomaly=True,
- resume_from_checkpoint=resume_from_checkpoint, # for pytorch-lightning < 2.0
+ # resume_from_checkpoint=resume_from_checkpoint, # for pytorch-lightning < 2.0
)
# setup dataloaders
@@ -271,7 +271,7 @@ def fit_model(
if val_data_exists and resume_from_checkpoint is None:
trainer.validate(lit_unet, val_dataloader)
trainer.fit(
- #ckpt_path=resume_from_checkpoint, # for pytorch-lightning >= 2.0
+ ckpt_path=resume_from_checkpoint, # for pytorch-lightning >= 2.0
model=lit_unet,
train_dataloaders=fitting_dataloader,
val_dataloaders=val_dataloader,
diff --git a/ddw/utils/unet.py b/ddw/utils/unet.py
index 6ad1eb2..fea34f1 100644
--- a/ddw/utils/unet.py
+++ b/ddw/utils/unet.py
@@ -93,19 +93,36 @@ class LitUnet3D(pl.LightningModule):
# if scheduler is not None:
# scheduler.step()
+ @staticmethod
+ def _unwrap_dataloader(dataloader):
+ """Return a concrete dataloader from Lightning dataloader containers.
+
+ Older PyTorch Lightning versions may expose dataloaders through a
+ CombinedLoader-like object with a `.loaders` attribute. Newer versions
+ expose the dataloader directly through `trainer.train_dataloader` and
+ `trainer.val_dataloaders`. This helper keeps both cases working.
+ """
+ if hasattr(dataloader, "loaders"):
+ dataloader = dataloader.loaders
+ if isinstance(dataloader, dict):
+ dataloader = next(iter(dataloader.values()))
+ if isinstance(dataloader, (list, tuple)):
+ dataloader = dataloader[0]
+ return dataloader
+
def update_subtomo_missing_wedges(self):
"""
Update the missing wedges of model input subtomos.
"""
# we don't want to rotate the subtomos when updating them, so we create new dataloader objects with rotate_subtomos=False
datasets = []
- train_loader = self.trainer.train_dataloader.loaders
+ train_loader = self._unwrap_dataloader(self.trainer.train_dataloader)
train_set = train_loader.dataset
train_set.rotate_subtomos = False
datasets.append(train_set)
# val_dataloaders may be None
if self.trainer.val_dataloaders is not None:
- val_loader = self.trainer.val_dataloaders[0]
+ val_loader = self._unwrap_dataloader(self.trainer.val_dataloaders)
val_set = val_loader.dataset
val_set.rotate_subtomos = False
datasets.append(val_set)
@@ -153,8 +170,9 @@ class LitUnet3D(pl.LightningModule):
"""
Updates the average model input mean and standard deviation used to normalize the sub-tomograms.
"""
+ train_loader = self._unwrap_dataloader(self.trainer.train_dataloader)
loc, scale = get_avg_model_input_mean_and_std_from_dataloader(
- dataloader=self.trainer.train_dataloader, verbose=True
+ dataloader=train_loader, verbose=True
)
# update normalization in unet
together with these changes to requirements.txt
@@ -1,16 +1,19 @@
certifi>=2017.4.17
matplotlib==3.8.4
mrcfile==1.5.0
pandas==2.2.1
pexpect==4.9.0
-pytorch-lightning==1.8.0.post1
-PyYAML==6.0.1
+pytorch-lightning>=2.6.0, <3
+pyyaml*
scikit-image==0.22.0
scipy==1.13.0
tqdm==4.65.0
typer==0.16.0
typer-cli==0.12.0
typer-config==1.4.0
typer-slim==0.12.0
-typing_extensions==4.9.0
-urllib3<3,>=1.21.1
+typing-extensions*
+urllib3>=1.21.1, <3
+torch>=2.8.0, <3
+torchvision>=0.23.0, <0.24
+torchaudio>=2.8.0, <3
seem to be enough for Blackwell suppport.
Notice that minimum and maximum cuda capability supported by used version of PyTorch is (7.0) - (12.0)
Hi,
This patch:
together with these changes to requirements.txt
seem to be enough for Blackwell suppport.
Notice that minimum and maximum cuda capability supported by used version of PyTorch is (7.0) - (12.0)