Skip to content
Draft
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
7 changes: 3 additions & 4 deletions drugex/training/explorers/interfaces.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,12 +5,11 @@
from scipy.stats import gmean

import torch
from torch import nn
from tqdm import tqdm

from drugex import DEFAULT_GPUS, DEFAULT_DEVICE
from drugex.logs import logger
from drugex.training.interfaces import Model
from drugex.training.interfaces import Model, _maybe_data_parallel
from drugex.training.monitors import NullMonitor

class Explorer(Model, ABC):
Expand Down Expand Up @@ -241,7 +240,7 @@ def policy_gradient(self, loader):
The average loss of the agent
"""

net = nn.DataParallel(self.agent, device_ids=self.gpus)
net = _maybe_data_parallel(self.agent, self.device, self.gpus)
total_steps = len(loader)

for step_idx, src in enumerate(tqdm(loader, desc='Calculating policy gradient...', leave=False)):
Expand Down Expand Up @@ -310,7 +309,7 @@ def fit(self, train_loader, valid_loader=None, epochs=1000, patience=50, criteri
self.bestState = self.getModel()

n_iters = 1 if self.crover is None else 10
net = nn.DataParallel(self, device_ids=self.gpus)
net = _maybe_data_parallel(self, self.device, self.gpus)
logger.info(' ')

for it in range(n_iters):
Expand Down
9 changes: 5 additions & 4 deletions drugex/training/generators/graph_transformer.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@
from drugex.molecules.converters.dummy_molecules import dummyMolsFromFragments
from drugex.training.generators.utils import PositionwiseFeedForward, SublayerConnection, PositionalEncoding, tri_mask
from drugex.training.generators.interfaces import FragGenerator
from drugex.training.interfaces import _maybe_data_parallel
from drugex.utils import ScheduledOptim
from torch import optim

Expand Down Expand Up @@ -314,7 +315,7 @@ def trainNet(self, loader, epoch, epochs):
The training loss of the epoch
"""

net = nn.DataParallel(self, device_ids=self.gpus)
net = _maybe_data_parallel(self, self.device, self.gpus)
total_steps = len(loader)
current_step = 0
for src in tqdm(loader, desc='Iterating over training batches', leave=False):
Expand Down Expand Up @@ -360,7 +361,7 @@ def validateNet(self, loader, evaluator=None, no_multifrag_smiles=True, n_sample

valid_metrics = {}

net = nn.DataParallel(self, device_ids=self.gpus)
net = _maybe_data_parallel(self, self.device, self.gpus)
pbar = tqdm(loader, desc='Iterating over validation batches', leave=False)
smiles, frags = self.sample(pbar)
scores = self.evaluate(smiles, frags, evaluator=evaluator, no_multifrag_smiles=no_multifrag_smiles)
Expand Down Expand Up @@ -390,7 +391,7 @@ def sample(self, loader):
frags : list
List of fragments
"""
net = nn.DataParallel(self, device_ids=self.gpus)
net = _maybe_data_parallel(self, self.device, self.gpus)
frags, smiles = [], []
with torch.no_grad():
for src in loader:
Expand Down Expand Up @@ -437,4 +438,4 @@ def loaderFromFrags(self, frags, batch_size=32, n_proc=1):
def decodeLoaders(self, src, trg):
return self.voc_trg.decode(trg)
def iterLoader(self, loader):
return loader
return loader
6 changes: 3 additions & 3 deletions drugex/training/generators/interfaces.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,7 @@
from drugex.data.interfaces import DataSet
from drugex.logs import logger
from drugex.training.scorers.smiles import SmilesChecker
from drugex.training.interfaces import Model
from drugex.training.interfaces import Model, _maybe_data_parallel
from drugex.training.monitors import NullMonitor

class Generator(Model, ABC):
Expand Down Expand Up @@ -410,7 +410,7 @@ def generate(self, input_frags: List[str] = None, input_dataset: DataSet = None,

# Duplicate of self.sample to allow dropping molecules and progress bar on the fly
# without additional overhead caused by calling nn.DataParallel a few times
net = nn.DataParallel(self, device_ids=self.gpus)
net = _maybe_data_parallel(self, self.device, self.gpus)

if progress:
tqdm_kwargs.update({'total': num_samples, 'desc': 'Generating molecules'})
Expand Down Expand Up @@ -467,4 +467,4 @@ def generate(self, input_frags: List[str] = None, input_dataset: DataSet = None,
if not keep_frags:
df.drop('Frags', axis=1, inplace=True)

return df.round(3)
return df.round(3)
5 changes: 4 additions & 1 deletion drugex/training/generators/sequence_rnn.py
Original file line number Diff line number Diff line change
Expand Up @@ -43,7 +43,10 @@ def attachToGPUs(self, gpus):
-------
None
"""
self.device = torch.device(f'cuda:{gpus[0]}')
if gpus[0] == -1:
self.device = torch.device('cpu')
else:
self.device = torch.device(f'cuda:{gpus[0]}')
self.to(self.device)
self.gpus = (gpus[0],)

Expand Down
8 changes: 4 additions & 4 deletions drugex/training/generators/sequence_transformer.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@
from drugex.molecules.converters.dummy_molecules import dummyMolsFromFragments
from drugex.training.generators.utils import PositionalEmbedding, PositionwiseFeedForward, SublayerConnection, pad_mask, tri_mask
from drugex.training.generators.interfaces import FragGenerator
from drugex.training.interfaces import _maybe_data_parallel
from drugex.utils import ScheduledOptim


Expand Down Expand Up @@ -170,7 +171,7 @@ def trainNet(self, loader, epoch, epochs):
The loss value for the current epoch
"""

net = nn.DataParallel(self, device_ids=self.gpus)
net = _maybe_data_parallel(self, self.device, self.gpus)
total_steps = len(loader)
current_step = 0
for src, trg in tqdm(loader, desc='Iterating over training batches', leave=False):
Expand Down Expand Up @@ -215,7 +216,7 @@ def validateNet(self, loader, evaluator=None, no_multifrag_smiles=True, n_sample

valid_metrics = {}

net = nn.DataParallel(self, device_ids=self.gpus)
net = _maybe_data_parallel(self, self.device, self.gpus)
pbar = tqdm(loader, desc='Iterating over validation batches', leave=False)
smiles, frags = self.sample(pbar)
scores = self.evaluate(smiles, frags, evaluator=evaluator, no_multifrag_smiles=no_multifrag_smiles)
Expand Down Expand Up @@ -245,7 +246,7 @@ def sample(self, loader):
frags: `list`
A list of input fragments
"""
net = nn.DataParallel(self, device_ids=self.gpus)
net = _maybe_data_parallel(self, self.device, self.gpus)
frags, smiles = [], []
with torch.no_grad():
for src, _ in loader:
Expand Down Expand Up @@ -295,4 +296,3 @@ def decodeLoaders(self, src, trg):
def iterLoader(self, loader):
for _, src in loader:
yield src

66 changes: 66 additions & 0 deletions drugex/training/generators/tests.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,66 @@
"""Regression tests for CPU device handling in generators and explorers."""

import unittest
from unittest.mock import patch

import torch
from torch import nn

from drugex.training.generators.sequence_rnn import SequenceRNN
from drugex.training.interfaces import _maybe_data_parallel


class _Vocabulary:
"""Minimal vocabulary required to construct a sequence generator."""

size = 4
max_len = 8
tk2ix = {"GO": 0, "EOS": 1}


class DeviceHandlingTests(unittest.TestCase):
"""Verify CPU execution never constructs CUDA devices or DataParallel."""

def test_cpu_device_object_skips_data_parallel(self):
module = nn.Linear(2, 2)
with patch("drugex.training.interfaces.nn.DataParallel") as wrapper:
result = _maybe_data_parallel(module, torch.device("cpu"), (-1,))

self.assertIs(result, module)
wrapper.assert_not_called()

def test_cpu_device_string_skips_data_parallel(self):
module = nn.Linear(2, 2)
with patch("drugex.training.interfaces.nn.DataParallel") as wrapper:
result = _maybe_data_parallel(module, "cpu", (-1,))

self.assertIs(result, module)
wrapper.assert_not_called()

def test_cuda_device_uses_requested_ids(self):
module = nn.Linear(2, 2)
sentinel = object()
with patch(
"drugex.training.interfaces.nn.DataParallel", return_value=sentinel
) as wrapper:
result = _maybe_data_parallel(module, "cuda:1", (1, 2))

self.assertIs(result, sentinel)
wrapper.assert_called_once_with(module, device_ids=(1, 2))

def test_sequence_rnn_accepts_cpu_string_and_sentinel(self):
generator = SequenceRNN(
_Vocabulary(),
embed_size=2,
hidden_size=4,
device="cpu",
use_gpus=(-1,),
)

self.assertEqual(generator.device, torch.device("cpu"))
self.assertEqual(generator.gpus, (-1,))
self.assertEqual(next(generator.parameters()).device, torch.device("cpu"))


if __name__ == "__main__":
unittest.main()
15 changes: 14 additions & 1 deletion drugex/training/interfaces.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,18 @@
from torch import nn


def _maybe_data_parallel(
module: nn.Module,
device: torch.device | str,
device_ids: tuple[int, ...],
) -> nn.Module:
"""Wrap a module for multi-GPU execution when its device is CUDA."""
normalized_device = torch.device(device)
if normalized_device.type == "cpu":
return module
return nn.DataParallel(module, device_ids=device_ids)


class ModelEvaluator(ABC):
"""
A simple function to score a model based on the generated molecules and input fragments if applicable.
Expand Down Expand Up @@ -265,6 +277,7 @@ def updateDevices(self, device, gpus):
List of GPUs to use for the model.
"""

device = torch.device(device)
if device.type == 'cpu':
self.device = torch.device('cpu')
self.gpus = (-1,)
Expand Down Expand Up @@ -422,4 +435,4 @@ def getSaveModelOption(self) -> Literal['best', 'all', 'improvement']:
Literal['best', 'all', 'improvement']
The scheme implemented by the monitor to save model snapshots.
"""
pass
pass