Skip to content

Commit ebecbac

Browse files
committed
fix: GPU selection on both CUDA and MPS
Signed-off-by: Luca Foppiano <luca@foppiano.org>
1 parent 0de0df2 commit ebecbac

7 files changed

Lines changed: 74 additions & 10 deletions

File tree

delft/sequenceLabelling/data_loader.py

Lines changed: 12 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -22,6 +22,17 @@
2222
from delft.utilities.Utilities import len_until_first_pad, truncate_batch_values
2323

2424

25+
def _worker_init_fn(worker_id):
26+
# Each fork-mode worker inherits the parent's LMDB env handle, which is unsafe.
27+
# Re-open per worker so each gets an independent reader-locktable slot.
28+
info = torch.utils.data.get_worker_info()
29+
if info is None:
30+
return
31+
embeddings = getattr(info.dataset, "embeddings", None)
32+
if embeddings is not None and hasattr(embeddings, "reopen_lmdb"):
33+
embeddings.reopen_lmdb()
34+
35+
2536
def collate_fn(batch):
2637
"""
2738
Custom collate function to handle variable-length sequences.
@@ -431,6 +442,7 @@ def create_dataloader(
431442
num_workers=num_workers,
432443
pin_memory=pin_memory,
433444
collate_fn=collate_fn,
445+
worker_init_fn=_worker_init_fn if num_workers > 0 else None,
434446
)
435447

436448

delft/sequenceLabelling/tagger.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,7 @@
55

66
from delft.sequenceLabelling.data_loader import create_dataloader
77
from delft.utilities.Tokenizer import tokenizeAndFilter
8+
from delft.utilities.Utilities import pick_device
89

910

1011
class Tagger(object):
@@ -13,7 +14,7 @@ def __init__(self, model, model_config, embeddings=None, preprocessor=None, devi
1314
self.preprocessor = preprocessor
1415
self.model_config = model_config
1516
self.embeddings = embeddings
16-
self.device = device if device else torch.device("cpu")
17+
self.device = pick_device(device)
1718

1819
def tag(self, texts, output_format, features=None):
1920
if output_format == "json":

delft/sequenceLabelling/trainer.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -19,6 +19,7 @@
1919
from delft.sequenceLabelling.config import ModelConfig, TrainingConfig
2020
from delft.sequenceLabelling.evaluation import classification_report
2121
from delft.sequenceLabelling.preprocess import Preprocessor
22+
from delft.utilities.Utilities import pick_device
2223

2324
# Default file names
2425
DEFAULT_WEIGHT_FILE_NAME = "model_weights.pt"
@@ -150,7 +151,7 @@ def __init__(
150151

151152
# Set device
152153
if device is None:
153-
self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
154+
self.device = pick_device()
154155
else:
155156
self.device = torch.device(device)
156157

delft/sequenceLabelling/wrapper.py

Lines changed: 2 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -23,6 +23,7 @@
2323
from delft.sequenceLabelling.data_loader import create_dataloader
2424
from delft.sequenceLabelling.evaluation import classification_report
2525
from delft.sequenceLabelling.models import get_model
26+
from delft.utilities.Utilities import pick_device
2627
from delft.sequenceLabelling.preprocess import Preprocessor, prepare_preprocessor
2728
from delft.sequenceLabelling.trainer import (
2829
CONFIG_FILE_NAME,
@@ -97,10 +98,7 @@ def __init__(
9798
self.nb_workers = nb_workers
9899

99100
# Set device
100-
if device is None:
101-
self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
102-
else:
103-
self.device = torch.device(device)
101+
self.device = pick_device(device)
104102

105103
word_emb_size = 0
106104
self.embeddings = None

delft/textClassification/data_loader.py

Lines changed: 12 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,17 @@
1010
from delft.textClassification.preprocess import to_indices_single, to_vector_single
1111

1212

13+
def _worker_init_fn(worker_id):
14+
# Each fork-mode worker inherits the parent's LMDB env handle, which is unsafe.
15+
# Re-open per worker so each gets an independent reader-locktable slot.
16+
info = torch.utils.data.get_worker_info()
17+
if info is None:
18+
return
19+
embeddings = getattr(info.dataset, "embeddings", None)
20+
if embeddings is not None and hasattr(embeddings, "reopen_lmdb"):
21+
embeddings.reopen_lmdb()
22+
23+
1324
class TextClassificationDataset(Dataset):
1425
"""
1526
Dataset for text classification.
@@ -132,6 +143,7 @@ def create_dataloader(
132143
batch_size=batch_size,
133144
shuffle=shuffle,
134145
num_workers=num_workers,
146+
worker_init_fn=_worker_init_fn if num_workers > 0 else None,
135147
)
136148

137149
return loader

delft/textClassification/wrapper.py

Lines changed: 2 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -13,6 +13,7 @@
1313
from delft.textClassification.trainer import Trainer
1414
from delft.utilities.Embeddings import Embeddings, load_resource_registry
1515
from delft.utilities.misc import print_parameters, to_wandb_table
16+
from delft.utilities.Utilities import pick_device
1617

1718
# File names for saving/loading
1819
PREPROCESSOR_FILE = "preprocessor.json"
@@ -87,10 +88,7 @@ def __init__(
8788
self.report_to_wandb = report_to_wandb
8889
self.wandb = None
8990

90-
if device is None:
91-
self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
92-
else:
93-
self.device = torch.device(device)
91+
self.device = pick_device(device)
9492

9593
self.registry = load_resource_registry(os.path.join(DELFT_PROJECT_DIR, "resources-registry.json"))
9694

delft/utilities/Utilities.py

Lines changed: 42 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,12 +1,14 @@
11
# some convenient methods for all models
22
import numpy as np
33
import regex as re
4+
import torch
45

56
# seed is fixed for reproducibility
67
from numpy.random import seed
78

89
seed(7)
910
import argparse
11+
import os
1012
import os.path
1113
import shutil
1214
from urllib.parse import urlparse
@@ -16,6 +18,46 @@
1618
from tqdm import tqdm
1719

1820

21+
def best_device() -> torch.device:
22+
"""
23+
Pick the best available compute device: CUDA → MPS (Apple Silicon) → CPU.
24+
25+
When MPS is selected, transparently enables CPU fallback for ops that don't
26+
have MPS kernels yet, so unsupported ops degrade gracefully instead of
27+
raising NotImplementedError mid-training.
28+
"""
29+
if torch.cuda.is_available():
30+
return torch.device("cuda")
31+
if torch.backends.mps.is_available():
32+
os.environ.setdefault("PYTORCH_ENABLE_MPS_FALLBACK", "1")
33+
return torch.device("mps")
34+
return torch.device("cpu")
35+
36+
37+
def pick_device(device=None) -> torch.device:
38+
"""
39+
Resolve a device (auto-pick when device=None, else honor the caller's choice)
40+
and print a one-line summary so the user can see which compute is in use.
41+
"""
42+
if device is None:
43+
d = best_device()
44+
elif isinstance(device, torch.device):
45+
d = device
46+
else:
47+
d = torch.device(device)
48+
49+
if d.type == "cuda":
50+
n = torch.cuda.device_count()
51+
name = torch.cuda.get_device_name(d)
52+
plural = "" if n == 1 else "s"
53+
print(f"Running on {d} ({name}); {n} CUDA device{plural} available")
54+
elif d.type == "mps":
55+
print(f"Running on {d} (Apple Silicon GPU via Metal Performance Shaders)")
56+
else:
57+
print(f"Running on {d}")
58+
return d
59+
60+
1961
def truncate_batch_values(batch_values: list, max_sequence_length: int) -> list:
2062
return [row[:max_sequence_length] for row in batch_values]
2163

0 commit comments

Comments
 (0)