Skip to content

Commit afbd708

Browse files
Copilotggorman
andcommitted
Simplify device validation implementation following Devito patterns
Co-authored-by: ggorman <5394691+ggorman@users.noreply.github.com>
1 parent 66277bd commit afbd708

3 files changed

Lines changed: 33 additions & 40 deletions

File tree

devito/passes/iet/langbase.py

Lines changed: 13 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -3,7 +3,6 @@
33
from abc import ABC
44

55
import cgen as c
6-
import numpy as np
76

87
from devito.data import FULL
98
from devito.ir import (DummyExpr, Call, Conditional, Expression, List, Prodder,
@@ -13,7 +12,7 @@
1312
from devito.passes import is_on_device
1413
from devito.passes.iet.engine import iet_pass
1514
from devito.symbolics import Byref, CondNe, SizeOf
16-
from devito.types.relational import Ge
15+
from sympy import Ge
1716
from devito.tools import as_list, is_integer, prod
1817
from devito.types import Symbol, QueueID, Wildcard
1918

@@ -63,7 +62,8 @@ def _get_num_devices(cls, platform):
6362
Get the number of accessible devices.
6463
Returns a tuple of (ngpus_symbol, call_to_get_num_devices).
6564
"""
66-
ngpus = Symbol(name='ngpus', dtype=np.int32)
65+
from devito.types import Symbol
66+
ngpus = Symbol(name='ngpus', dtype='int32')
6767
devicetype = as_list(cls[platform])
6868
call_ngpus = cls['num-devices'](devicetype, retobj=ngpus)
6969
return ngpus, call_ngpus
@@ -434,22 +434,19 @@ def _make_setdevice_seq(iet, nodes=()):
434434

435435
# Add device validation check
436436
ngpus, call_ngpus = self.langbb._get_num_devices(self.platform)
437-
438-
# Create validation: if deviceid >= num_devices, print error and exit
439-
validation_check = Conditional(
437+
438+
validation = Conditional(
440439
Ge(deviceid, ngpus),
441440
List(body=[
442-
Call('printf', ['"%s: Error - Requested device ID %d does not exist. '
443-
'Only %d device(s) available. Check CUDA_VISIBLE_DEVICES '
444-
'and container GPU configuration.\\n"',
445-
self.langbb['name'], deviceid, ngpus]),
441+
Call('printf', ['"%s: Error - device %d >= %d devices\\n"',
442+
self.langbb['name'], deviceid, ngpus]),
446443
Call('exit', [1])
447444
])
448445
)
449446

450447
device_setup = List(body=[
451448
call_ngpus,
452-
validation_check,
449+
validation,
453450
self.langbb['set-device']([deviceid] + devicetype)
454451
])
455452

@@ -468,21 +465,19 @@ def _make_setdevice_mpi(iet, objcomm, nodes=()):
468465

469466
ngpus, call_ngpus = self.langbb._get_num_devices(self.platform)
470467

471-
# Add device validation check for explicit device ID
472-
validation_check = Conditional(
468+
# Add device validation for explicit device ID
469+
validation = Conditional(
473470
Ge(deviceid, ngpus),
474471
List(body=[
475-
Call('printf', ['"%s: Error - Requested device ID %d does not exist. '
476-
'Only %d device(s) available. Check CUDA_VISIBLE_DEVICES '
477-
'and container GPU configuration.\\n"',
478-
self.langbb['name'], deviceid, ngpus]),
472+
Call('printf', ['"%s: Error - device %d >= %d devices\\n"',
473+
self.langbb['name'], deviceid, ngpus]),
479474
Call('exit', [1])
480475
])
481476
)
482477

483478
osdd_then = List(body=[
484479
call_ngpus,
485-
validation_check,
480+
validation,
486481
self.langbb['set-device']([deviceid] + devicetype)
487482
])
488483
osdd_else = self.langbb['set-device']([rank % ngpus] + devicetype)

tests/test_gpu_openacc.py

Lines changed: 9 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -208,20 +208,19 @@ def test_device_validation_error_message(self):
208208

209209
op = Operator(Eq(u.forward, u + 1), platform='nvidiaX', language='openacc')
210210

211-
# Check that the generated code contains device validation with informative error
211+
# Check that the generated code contains device validation
212212
code = str(op)
213-
213+
214214
# Should contain device count check
215215
assert 'acc_get_num_devices' in code, "Missing OpenACC device count check"
216-
216+
217217
# Should contain validation condition
218-
assert 'deviceid >= ngpus' in code, "Missing OpenACC device ID validation condition"
219-
220-
# Should contain helpful error message components
221-
assert 'does not exist' in code, "Missing 'does not exist' error message"
222-
assert 'CUDA_VISIBLE_DEVICES' in code, "Missing CUDA_VISIBLE_DEVICES guidance"
223-
assert 'container GPU configuration' in code, "Missing container guidance"
224-
218+
assert 'deviceid >= ngpus' in code, "Missing OpenACC device ID " + \
219+
"validation condition"
220+
221+
# Should contain error message
222+
assert 'Error - device' in code, "Missing error message"
223+
225224
# Should contain exit call to prevent undefined behavior
226225
assert 'exit(1)' in code, "Missing exit call on validation failure"
227226

tests/test_gpu_openmp.py

Lines changed: 11 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -25,7 +25,7 @@ def test_init_omp_env(self):
2525
assert 'if (deviceid != -1)' in init_code
2626
assert 'int ngpus = omp_get_num_devices()' in init_code
2727
assert 'if (deviceid >= ngpus)' in init_code
28-
assert 'does not exist' in init_code
28+
assert 'Error - device' in init_code
2929
assert 'omp_set_default_device(deviceid)' in init_code
3030

3131
@pytest.mark.parallel(mode=1)
@@ -36,13 +36,14 @@ def test_init_omp_env_w_mpi(self, mode):
3636

3737
op = Operator(Eq(u.forward, u.dx+1), language='openmp')
3838

39-
# With device validation, the MPI case also includes validation for explicit deviceid
39+
# With device validation, the MPI case also includes validation for explicit
40+
# deviceid
4041
init_code = str(op.body.init[0].body[0])
4142
assert 'if (deviceid != -1)' in init_code
4243
assert 'int ngpus = omp_get_num_devices()' in init_code
4344
# For MPI case with explicit deviceid, should have validation
4445
assert 'if (deviceid >= ngpus)' in init_code
45-
assert 'does not exist' in init_code
46+
assert 'Error - device' in init_code
4647
# Should still have MPI rank-based assignment in else clause
4748
assert 'int rank = 0' in init_code
4849
assert 'MPI_Comm_rank(comm,&rank)' in init_code
@@ -56,20 +57,18 @@ def test_device_validation_error_message(self):
5657

5758
op = Operator(Eq(u.forward, u.dx+1), language='openmp')
5859

59-
# Check that the generated code contains device validation with informative error
60+
# Check that the generated code contains device validation
6061
code = str(op)
61-
62+
6263
# Should contain device count check
6364
assert 'omp_get_num_devices()' in code, "Missing device count check"
64-
65+
6566
# Should contain validation condition
6667
assert 'deviceid >= ngpus' in code, "Missing device ID validation condition"
67-
68-
# Should contain helpful error message components
69-
assert 'does not exist' in code, "Missing 'does not exist' error message"
70-
assert 'CUDA_VISIBLE_DEVICES' in code, "Missing CUDA_VISIBLE_DEVICES guidance"
71-
assert 'container GPU configuration' in code, "Missing container guidance"
72-
68+
69+
# Should contain error message
70+
assert 'Error - device' in code, "Missing error message"
71+
7372
# Should contain exit call to prevent undefined behavior
7473
assert 'exit(1)' in code, "Missing exit call on validation failure"
7574

0 commit comments

Comments
 (0)