From c8fc19f238524b961c470507a5b8f740e6c5d9ec Mon Sep 17 00:00:00 2001 From: Matthias Diener Date: Fri, 16 Feb 2024 19:37:16 -0600 Subject: [PATCH 1/6] add support for CupyArrayContext --- examples/wave.py | 5 ++++- mirgecom/array_context.py | 22 +++++++++++++++++++++- 2 files changed, 25 insertions(+), 2 deletions(-) diff --git a/examples/wave.py b/examples/wave.py index 8d75a68cc..7d0b51464 100644 --- a/examples/wave.py +++ b/examples/wave.py @@ -269,6 +269,8 @@ def rhs(t, w): help="use numpy-based eager actx.") parser.add_argument("--mpi", default=True, action=argparse.BooleanOptionalAction, help="use MPI") + parser.add_argument("--cupy", action="store_true", + help="use cupy-based eager actx.") args = parser.parse_args() casename = args.casename or "wave" @@ -276,7 +278,8 @@ def rhs(t, w): actx_class = get_reasonable_array_context_class(lazy=args.lazy, distributed=args.mpi, profiling=args.profiling, - numpy=args.numpy) + numpy=args.numpy, + cupy=args.cupy) if args.mpi: main_func = main diff --git a/mirgecom/array_context.py b/mirgecom/array_context.py index 84cdcf4b9..f3a7edf63 100644 --- a/mirgecom/array_context.py +++ b/mirgecom/array_context.py @@ -44,11 +44,31 @@ def get_reasonable_array_context_class(*, lazy: bool, distributed: bool, - profiling: bool, numpy: bool = False) -> Type[ArrayContext]: + profiling: bool, numpy: bool = False, + cupy: bool = False) -> Type[ArrayContext]: """Return a :class:`~arraycontext.ArrayContext` with the given constraints.""" if lazy and profiling: raise ValueError("Can't specify both lazy and profiling") + if numpy and cupy: + raise ValueError("Can't specify both numpy and cupy") + + if cupy: + if profiling: + raise ValueError("Can't specify both cupy and profiling") + if lazy: + raise ValueError("Can't specify both cupy and lazy") + + from warnings import warn + warn("The CupyArrayContext is still under development") + + if distributed: + from grudge.array_context import MPICupyArrayContext + return MPICupyArrayContext + else: + from grudge.array_context import CupyArrayContext + return CupyArrayContext + if numpy: if profiling: raise ValueError("Can't specify both numpy and profiling") From 963f293d91616cd11069e092a5b02e3f18b64d37 Mon Sep 17 00:00:00 2001 From: Matthias Diener Date: Fri, 16 Feb 2024 19:42:51 -0600 Subject: [PATCH 2/6] skip ocl init --- mirgecom/array_context.py | 14 +++++++++++++- 1 file changed, 13 insertions(+), 1 deletion(-) diff --git a/mirgecom/array_context.py b/mirgecom/array_context.py index f3a7edf63..96ad48e13 100644 --- a/mirgecom/array_context.py +++ b/mirgecom/array_context.py @@ -125,6 +125,18 @@ def actx_class_is_numpy(actx_class: Type[ArrayContext]) -> bool: return False +def actx_class_is_cupy(actx_class: Type[ArrayContext]) -> bool: + """Return True if *actx_class* is cupy-based.""" + try: + from grudge.array_context import CupyArrayContext + if issubclass(actx_class, CupyArrayContext): + return True + else: + return False + except ImportError: + return False + + def initialize_actx(actx_class: Type[ArrayContext], comm: Optional["Comm"]) \ -> ArrayContext: """Initialize a new :class:`~arraycontext.ArrayContext` based on *actx_class*.""" @@ -133,7 +145,7 @@ def initialize_actx(actx_class: Type[ArrayContext], comm: Optional["Comm"]) \ MPIPytatoArrayContext) # Special handling for NumpyArrayContext since it needs no CL context - if actx_class_is_numpy(actx_class): + if actx_class_is_numpy(actx_class) or actx_class_is_cupy(actx_class): if comm: return actx_class(mpi_communicator=comm) # type: ignore[call-arg] else: From 80ee778bd82e2d803378a1fba76098870749cf43 Mon Sep 17 00:00:00 2001 From: Matthias Diener Date: Fri, 16 Feb 2024 19:56:28 -0600 Subject: [PATCH 3/6] add to all examples --- examples/ablation-workshop.py | 5 ++++- examples/advection_diffusion_reaction.py | 5 ++++- examples/autoignition.py | 5 ++++- examples/blasius.py | 5 ++++- examples/combozzle.py | 5 ++++- examples/doublemach.py | 4 +++- examples/doublemach_physical_av.py | 4 +++- examples/heat-source.py | 5 ++++- examples/hotplate.py | 5 ++++- examples/lump.py | 5 ++++- examples/mixture.py | 5 ++++- examples/multiple-volumes.py | 5 ++++- examples/orthotropic-diffusion.py | 5 ++++- examples/poiseuille-multispecies.py | 5 ++++- examples/poiseuille.py | 5 ++++- examples/pulse-mixture.py | 5 ++++- examples/pulse-tpe.py | 5 ++++- examples/pulse.py | 5 ++++- examples/scalar-advdiff.py | 4 +++- examples/scalar-lump.py | 5 ++++- examples/sod.py | 5 ++++- examples/taylor-green.py | 5 ++++- examples/thermally-coupled.py | 5 ++++- examples/vortex.py | 5 ++++- 24 files changed, 93 insertions(+), 24 deletions(-) diff --git a/examples/ablation-workshop.py b/examples/ablation-workshop.py index 8fb990169..9703f4a98 100644 --- a/examples/ablation-workshop.py +++ b/examples/ablation-workshop.py @@ -1207,13 +1207,16 @@ def my_post_step(step, t, dt, state): dest="restart_file", nargs="?", action="store", help="simulation restart file") parser.add_argument("--casename", help="casename to use for i/o") + parser.add_argument("--cupy", action="store_true", + help="use cupy-based eager actx.") args = parser.parse_args() from warnings import warn warn("Automatically turning off DV logging. MIRGE-Com Issue(578)") from mirgecom.array_context import get_reasonable_array_context_class actx_class = get_reasonable_array_context_class( - lazy=args.lazy, distributed=True, profiling=args.profiling, numpy=args.numpy) + lazy=args.lazy, distributed=True, profiling=args.profiling, + numpy=args.numpy, cupy=args.cupy) logging.basicConfig(format="%(message)s", level=logging.INFO) if args.casename: diff --git a/examples/advection_diffusion_reaction.py b/examples/advection_diffusion_reaction.py index 2e82cd1bb..7fc459d31 100644 --- a/examples/advection_diffusion_reaction.py +++ b/examples/advection_diffusion_reaction.py @@ -331,13 +331,16 @@ def my_rhs(t, u): help="use numpy-based eager actx.") parser.add_argument("--restart_file", help="root name of restart file") parser.add_argument("--casename", help="casename to use for i/o") + parser.add_argument("--cupy", action="store_true", + help="use cupy-based eager actx.") args = parser.parse_args() lazy = args.lazy from mirgecom.array_context import get_reasonable_array_context_class actx_class = get_reasonable_array_context_class(lazy=args.lazy, distributed=True, profiling=args.profiling, - numpy=args.numpy) + numpy=args.numpy, + cupy=args.cupy) logging.basicConfig(format="%(message)s", level=logging.INFO) if args.casename: diff --git a/examples/autoignition.py b/examples/autoignition.py index 925e2a672..11a685f38 100644 --- a/examples/autoignition.py +++ b/examples/autoignition.py @@ -652,6 +652,8 @@ def my_rhs(t, state): help="use numpy-based eager actx.") parser.add_argument("--restart_file", help="root name of restart file") parser.add_argument("--casename", help="casename to use for i/o") + parser.add_argument("--cupy", action="store_true", + help="use cupy-based eager actx.") args = parser.parse_args() from warnings import warn warn("Automatically turning off DV logging. MIRGE-Com Issue(578)") @@ -667,7 +669,8 @@ def my_rhs(t, state): from mirgecom.array_context import get_reasonable_array_context_class actx_class = get_reasonable_array_context_class( - lazy=args.lazy, distributed=True, profiling=args.profiling, numpy=args.numpy) + lazy=args.lazy, distributed=True, profiling=args.profiling, + numpy=args.numpy, cupy=args.cupy) logging.basicConfig(format="%(message)s", level=logging.INFO) if args.casename: diff --git a/examples/blasius.py b/examples/blasius.py index 956f96b60..e03cff152 100644 --- a/examples/blasius.py +++ b/examples/blasius.py @@ -516,6 +516,8 @@ def my_post_step(step, t, dt, state): help="enable lazy evaluation [OFF]") parser.add_argument("--numpy", action="store_true", help="use numpy-based eager actx.") + parser.add_argument("--cupy", action="store_true", + help="use cupy-based eager actx.") args = parser.parse_args() @@ -531,7 +533,8 @@ def my_post_step(step, t, dt, state): from mirgecom.array_context import get_reasonable_array_context_class actx_class = get_reasonable_array_context_class( - lazy=args.lazy, distributed=True, profiling=args.profiling, numpy=args.numpy) + lazy=args.lazy, distributed=True, profiling=args.profiling, + numpy=args.numpy, cupy=args.cupy) logging.basicConfig(format="%(message)s", level=logging.INFO) if args.casename: diff --git a/examples/combozzle.py b/examples/combozzle.py index 0a2744ea2..57d62dee7 100644 --- a/examples/combozzle.py +++ b/examples/combozzle.py @@ -1254,6 +1254,8 @@ def dummy_rhs(t, state): parser.add_argument("--casename", help="casename to use for i/o") parser.add_argument("--tpe", action="store_true", help="Use tensor product elements (quads/hexes).") + parser.add_argument("--cupy", action="store_true", + help="use cupy-based eager actx.") args = parser.parse_args() from warnings import warn @@ -1270,7 +1272,8 @@ def dummy_rhs(t, state): from mirgecom.array_context import get_reasonable_array_context_class actx_class = get_reasonable_array_context_class( - lazy=args.lazy, distributed=True, profiling=args.profiling, numpy=args.numpy) + lazy=args.lazy, distributed=True, profiling=args.profiling, + numpy=args.numpy, cupy=args.cupy) logging.basicConfig(format="%(message)s", level=logging.INFO) if args.casename: diff --git a/examples/doublemach.py b/examples/doublemach.py index eb1484dbc..b6ed25713 100644 --- a/examples/doublemach.py +++ b/examples/doublemach.py @@ -443,6 +443,8 @@ def my_rhs(t, state): help="use numpy-based eager actx.") parser.add_argument("--restart_file", help="root name of restart file") parser.add_argument("--casename", help="casename to use for i/o") + parser.add_argument("--cupy", action="store_true", + help="use cupy-based eager actx.") args = parser.parse_args() from warnings import warn @@ -455,7 +457,7 @@ def my_rhs(t, state): from mirgecom.array_context import get_reasonable_array_context_class actx_class = get_reasonable_array_context_class(lazy=args.lazy, distributed=True, profiling=args.profiling, - numpy=args.numpy) + numpy=args.numpy, cupy=args.cupy) logging.basicConfig(format="%(message)s", level=logging.INFO) if args.casename: diff --git a/examples/doublemach_physical_av.py b/examples/doublemach_physical_av.py index f8b09173b..c5114c18e 100644 --- a/examples/doublemach_physical_av.py +++ b/examples/doublemach_physical_av.py @@ -716,6 +716,8 @@ def _my_rhs_phys_visc_div_av(t, state): help="use numpy-based eager actx.") parser.add_argument("--restart_file", help="root name of restart file") parser.add_argument("--casename", help="casename to use for i/o") + parser.add_argument("--cupy", action="store_true", + help="use cupy-based eager actx.") args = parser.parse_args() from warnings import warn @@ -730,7 +732,7 @@ def _my_rhs_phys_visc_div_av(t, state): actx_class = get_reasonable_array_context_class(lazy=args.lazy, distributed=True, profiling=args.profiling, - numpy=args.numpy) + numpy=args.numpy, cupy=args.cupy) logging.basicConfig(format="%(message)s", level=logging.INFO) if args.casename: diff --git a/examples/heat-source.py b/examples/heat-source.py index eb5a96c72..ff8ea0127 100644 --- a/examples/heat-source.py +++ b/examples/heat-source.py @@ -242,11 +242,14 @@ def my_post_step(step, t, dt, state): help="use numpy-based eager actx.") parser.add_argument("--restart_file", help="root name of restart file") parser.add_argument("--casename", help="casename to use for i/o") + parser.add_argument("--cupy", action="store_true", + help="use cupy-based eager actx.") args = parser.parse_args() from mirgecom.array_context import get_reasonable_array_context_class actx_class = get_reasonable_array_context_class( - lazy=args.lazy, distributed=True, profiling=args.profiling, numpy=args.numpy) + lazy=args.lazy, distributed=True, profiling=args.profiling, + numpy=args.numpy, cupy=args.cupy) logging.basicConfig(format="%(message)s", level=logging.INFO) if args.casename: diff --git a/examples/hotplate.py b/examples/hotplate.py index 6bbbaa0e1..5bca9ce4e 100644 --- a/examples/hotplate.py +++ b/examples/hotplate.py @@ -448,6 +448,8 @@ def my_rhs(t, state): help="use numpy-based eager actx.") parser.add_argument("--restart_file", help="root name of restart file") parser.add_argument("--casename", help="casename to use for i/o") + parser.add_argument("--cupy", action="store_true", + help="use cupy-based eager actx.") args = parser.parse_args() from warnings import warn @@ -460,7 +462,8 @@ def my_rhs(t, state): from mirgecom.array_context import get_reasonable_array_context_class actx_class = get_reasonable_array_context_class( - lazy=args.lazy, distributed=True, profiling=args.profiling, numpy=args.numpy) + lazy=args.lazy, distributed=True, profiling=args.profiling, + numpy=args.numpy, cupy=args.cupy) logging.basicConfig(format="%(message)s", level=logging.INFO) if args.casename: diff --git a/examples/lump.py b/examples/lump.py index 672fa5e09..d5233ef2d 100644 --- a/examples/lump.py +++ b/examples/lump.py @@ -383,6 +383,8 @@ def my_rhs(t, state): help="use numpy-based eager actx.") parser.add_argument("--restart_file", help="root name of restart file") parser.add_argument("--casename", help="casename to use for i/o") + parser.add_argument("--cupy", action="store_true", + help="use cupy-based eager actx.") args = parser.parse_args() from warnings import warn @@ -395,7 +397,8 @@ def my_rhs(t, state): from mirgecom.array_context import get_reasonable_array_context_class actx_class = get_reasonable_array_context_class( - lazy=args.lazy, distributed=True, profiling=args.profiling, numpy=args.numpy) + lazy=args.lazy, distributed=True, profiling=args.profiling, + numpy=args.numpy, cupy=args.cupy) logging.basicConfig(format="%(message)s", level=logging.INFO) if args.casename: diff --git a/examples/mixture.py b/examples/mixture.py index 4beba862c..8ead17514 100644 --- a/examples/mixture.py +++ b/examples/mixture.py @@ -446,6 +446,8 @@ def my_rhs(t, state): help="use numpy-based eager actx.") parser.add_argument("--restart_file", help="root name of restart file") parser.add_argument("--casename", help="casename to use for i/o") + parser.add_argument("--cupy", action="store_true", + help="use cupy-based eager actx.") args = parser.parse_args() from warnings import warn warn("Automatically turning off DV logging. MIRGE-Com Issue(578)") @@ -460,7 +462,8 @@ def my_rhs(t, state): from mirgecom.array_context import get_reasonable_array_context_class actx_class = get_reasonable_array_context_class( - lazy=args.lazy, distributed=True, profiling=args.profiling, numpy=args.numpy) + lazy=args.lazy, distributed=True, profiling=args.profiling, + numpy=args.numpy, cupy=args.cupy) logging.basicConfig(format="%(message)s", level=logging.INFO) if args.casename: diff --git a/examples/multiple-volumes.py b/examples/multiple-volumes.py index 8feed339a..289f7dade 100644 --- a/examples/multiple-volumes.py +++ b/examples/multiple-volumes.py @@ -392,6 +392,8 @@ def my_rhs(t, state): help="use numpy-based eager actx.") parser.add_argument("--restart_file", help="root name of restart file") parser.add_argument("--casename", help="casename to use for i/o") + parser.add_argument("--cupy", action="store_true", + help="use cupy-based eager actx.") args = parser.parse_args() from warnings import warn @@ -404,7 +406,8 @@ def my_rhs(t, state): from mirgecom.array_context import get_reasonable_array_context_class actx_class = get_reasonable_array_context_class( - lazy=args.lazy, distributed=True, profiling=args.profiling, numpy=args.numpy) + lazy=args.lazy, distributed=True, profiling=args.profiling, + numpy=args.numpy, cupy=args.cupy) logging.basicConfig(format="%(message)s", level=logging.INFO) if args.casename: diff --git a/examples/orthotropic-diffusion.py b/examples/orthotropic-diffusion.py index 84a8c6517..f1a075d8b 100644 --- a/examples/orthotropic-diffusion.py +++ b/examples/orthotropic-diffusion.py @@ -204,11 +204,14 @@ def my_post_step(step, t, dt, state): help="use numpy-based eager actx.") parser.add_argument("--restart_file", help="root name of restart file") parser.add_argument("--casename", help="casename to use for i/o") + parser.add_argument("--cupy", action="store_true", + help="use cupy-based eager actx.") args = parser.parse_args() from mirgecom.array_context import get_reasonable_array_context_class actx_class = get_reasonable_array_context_class( - lazy=args.lazy, distributed=True, profiling=args.profiling, numpy=args.numpy) + lazy=args.lazy, distributed=True, profiling=args.profiling, + numpy=args.numpy, cupy=args.cupy) logging.basicConfig(format="%(message)s", level=logging.INFO) if args.casename: diff --git a/examples/poiseuille-multispecies.py b/examples/poiseuille-multispecies.py index 570b310cb..7c6e6de2e 100644 --- a/examples/poiseuille-multispecies.py +++ b/examples/poiseuille-multispecies.py @@ -506,6 +506,8 @@ def my_rhs(t, state): help="use numpy-based eager actx.") parser.add_argument("--restart_file", help="root name of restart file") parser.add_argument("--casename", help="casename to use for i/o") + parser.add_argument("--cupy", action="store_true", + help="use cupy-based eager actx.") args = parser.parse_args() from warnings import warn @@ -518,7 +520,8 @@ def my_rhs(t, state): from mirgecom.array_context import get_reasonable_array_context_class actx_class = get_reasonable_array_context_class( - lazy=args.lazy, distributed=True, profiling=args.profiling, numpy=args.numpy) + lazy=args.lazy, distributed=True, profiling=args.profiling, + numpy=args.numpy, cupy=args.cupy) logging.basicConfig(format="%(message)s", level=logging.INFO) if args.casename: diff --git a/examples/poiseuille.py b/examples/poiseuille.py index a3838b5b7..b327c8b41 100644 --- a/examples/poiseuille.py +++ b/examples/poiseuille.py @@ -475,6 +475,8 @@ def my_rhs(t, state): help="use numpy-based eager actx.") parser.add_argument("--restart_file", help="root name of restart file") parser.add_argument("--casename", help="casename to use for i/o") + parser.add_argument("--cupy", action="store_true", + help="use cupy-based eager actx.") args = parser.parse_args() from warnings import warn @@ -487,7 +489,8 @@ def my_rhs(t, state): from mirgecom.array_context import get_reasonable_array_context_class actx_class = get_reasonable_array_context_class( - lazy=args.lazy, distributed=True, profiling=args.profiling, numpy=args.numpy) + lazy=args.lazy, distributed=True, profiling=args.profiling, + numpy=args.numpy, cupy=args.cupy) logging.basicConfig(format="%(message)s", level=logging.INFO) if args.casename: diff --git a/examples/pulse-mixture.py b/examples/pulse-mixture.py index e3d7b3803..8051e1062 100644 --- a/examples/pulse-mixture.py +++ b/examples/pulse-mixture.py @@ -420,6 +420,8 @@ def my_rhs(t, state): help="use numpy-based eager actx.") parser.add_argument("--restart_file", help="root name of restart file") parser.add_argument("--casename", help="casename to use for i/o") + parser.add_argument("--cupy", action="store_true", + help="use cupy-based eager actx.") args = parser.parse_args() from warnings import warn @@ -432,7 +434,8 @@ def my_rhs(t, state): from mirgecom.array_context import get_reasonable_array_context_class actx_class = get_reasonable_array_context_class( - lazy=args.lazy, distributed=True, profiling=args.profiling, numpy=args.numpy) + lazy=args.lazy, distributed=True, profiling=args.profiling, + numpy=args.numpy, cupy=args.cupy) logging.basicConfig(format="%(message)s", level=logging.INFO) if args.casename: diff --git a/examples/pulse-tpe.py b/examples/pulse-tpe.py index a05eab1c7..e7956a375 100644 --- a/examples/pulse-tpe.py +++ b/examples/pulse-tpe.py @@ -348,6 +348,8 @@ def my_rhs(t, state): parser.add_argument("--casename", help="casename to use for i/o") parser.add_argument("--tpe", action="store_true", help="Use tensor product (quad/hex) elements.") + parser.add_argument("--cupy", action="store_true", + help="use cupy-based eager actx.") args = parser.parse_args() from warnings import warn @@ -362,7 +364,8 @@ def my_rhs(t, state): from mirgecom.array_context import get_reasonable_array_context_class actx_class = get_reasonable_array_context_class( - lazy=args.lazy, distributed=True, profiling=args.profiling, numpy=args.numpy) + lazy=args.lazy, distributed=True, profiling=args.profiling, + numpy=args.numpy, cupy=args.cupy) logging.basicConfig(format="%(message)s", level=logging.INFO) if args.casename: diff --git a/examples/pulse.py b/examples/pulse.py index 74b5bb6aa..f6e84206c 100644 --- a/examples/pulse.py +++ b/examples/pulse.py @@ -342,6 +342,8 @@ def my_rhs(t, state): help="use numpy-based eager actx.") parser.add_argument("--restart_file", help="root name of restart file") parser.add_argument("--casename", help="casename to use for i/o") + parser.add_argument("--cupy", action="store_true", + help="use cupy-based eager actx.") args = parser.parse_args() from warnings import warn @@ -354,7 +356,8 @@ def my_rhs(t, state): from mirgecom.array_context import get_reasonable_array_context_class actx_class = get_reasonable_array_context_class( - lazy=args.lazy, distributed=True, profiling=args.profiling, numpy=args.numpy) + lazy=args.lazy, distributed=True, profiling=args.profiling, + numpy=args.numpy, cupy=args.cupy) logging.basicConfig(format="%(message)s", level=logging.INFO) if args.casename: diff --git a/examples/scalar-advdiff.py b/examples/scalar-advdiff.py index 75ddb1f77..7609d2c1c 100644 --- a/examples/scalar-advdiff.py +++ b/examples/scalar-advdiff.py @@ -465,6 +465,8 @@ def my_rhs(t, state): help="use numpy-based eager actx.") parser.add_argument("--restart_file", help="root name of restart file") parser.add_argument("--casename", help="casename to use for i/o") + parser.add_argument("--cupy", action="store_true", + help="use cupy-based eager actx.") args = parser.parse_args() from warnings import warn @@ -478,7 +480,7 @@ def my_rhs(t, state): from mirgecom.array_context import get_reasonable_array_context_class actx_class = get_reasonable_array_context_class(lazy=args.lazy, distributed=True, profiling=args.profiling, - numpy=args.numpy) + numpy=args.numpy, cupy=args.cupy) logging.basicConfig(format="%(message)s", level=logging.INFO) if args.casename: diff --git a/examples/scalar-lump.py b/examples/scalar-lump.py index 4a02af72f..c8625d212 100644 --- a/examples/scalar-lump.py +++ b/examples/scalar-lump.py @@ -393,6 +393,8 @@ def my_rhs(t, state): help="use numpy-based eager actx.") parser.add_argument("--restart_file", help="root name of restart file") parser.add_argument("--casename", help="casename to use for i/o") + parser.add_argument("--cupy", action="store_true", + help="use cupy-based eager actx.") args = parser.parse_args() from warnings import warn @@ -405,7 +407,8 @@ def my_rhs(t, state): from mirgecom.array_context import get_reasonable_array_context_class actx_class = get_reasonable_array_context_class( - lazy=args.lazy, distributed=True, profiling=args.profiling, numpy=args.numpy) + lazy=args.lazy, distributed=True, profiling=args.profiling, + numpy=args.numpy, cupy=args.cupy) logging.basicConfig(format="%(message)s", level=logging.INFO) if args.casename: diff --git a/examples/sod.py b/examples/sod.py index 9146e799d..2bc6ad643 100644 --- a/examples/sod.py +++ b/examples/sod.py @@ -388,6 +388,8 @@ def my_rhs(t, state): help="use numpy-based eager actx.") parser.add_argument("--restart_file", help="root name of restart file") parser.add_argument("--casename", help="casename to use for i/o") + parser.add_argument("--cupy", action="store_true", + help="use cupy-based eager actx.") args = parser.parse_args() from warnings import warn @@ -400,7 +402,8 @@ def my_rhs(t, state): from mirgecom.array_context import get_reasonable_array_context_class actx_class = get_reasonable_array_context_class( - lazy=args.lazy, distributed=True, profiling=args.profiling, numpy=args.numpy) + lazy=args.lazy, distributed=True, profiling=args.profiling, + numpy=args.numpy, cupy=args.cupy) logging.basicConfig(format="%(message)s", level=logging.INFO) if args.casename: diff --git a/examples/taylor-green.py b/examples/taylor-green.py index 7e087f0d4..02cef13ca 100644 --- a/examples/taylor-green.py +++ b/examples/taylor-green.py @@ -334,6 +334,8 @@ def my_rhs(t, state): help="turn on logging") parser.add_argument("--restart_file", help="root name of restart file") parser.add_argument("--casename", help="casename to use for i/o") + parser.add_argument("--cupy", action="store_true", + help="use cupy-based eager actx.") args = parser.parse_args() from warnings import warn @@ -346,7 +348,8 @@ def my_rhs(t, state): from mirgecom.array_context import get_reasonable_array_context_class actx_class = get_reasonable_array_context_class( - lazy=args.lazy, distributed=True, profiling=args.profiling, numpy=args.numpy) + lazy=args.lazy, distributed=True, profiling=args.profiling, + numpy=args.numpy, cupy=args.cupy) logging.basicConfig(format="%(message)s", level=logging.INFO) if args.casename: diff --git a/examples/thermally-coupled.py b/examples/thermally-coupled.py index 7ede40452..052612fb6 100644 --- a/examples/thermally-coupled.py +++ b/examples/thermally-coupled.py @@ -583,6 +583,8 @@ def my_rhs_and_gradients(t, state): help="use numpy-based eager actx.") parser.add_argument("--restart_file", help="root name of restart file") parser.add_argument("--casename", help="casename to use for i/o") + parser.add_argument("--cupy", action="store_true", + help="use cupy-based eager actx.") args = parser.parse_args() from warnings import warn @@ -595,7 +597,8 @@ def my_rhs_and_gradients(t, state): from mirgecom.array_context import get_reasonable_array_context_class actx_class = get_reasonable_array_context_class( - lazy=args.lazy, distributed=True, profiling=args.profiling, numpy=args.numpy) + lazy=args.lazy, distributed=True, profiling=args.profiling, + numpy=args.numpy, cupy=args.cupy) logging.basicConfig(format="%(message)s", level=logging.INFO) if args.casename: diff --git a/examples/vortex.py b/examples/vortex.py index ca8fce34d..f4351e825 100644 --- a/examples/vortex.py +++ b/examples/vortex.py @@ -401,6 +401,8 @@ def my_rhs(t, state): help="use numpy-based eager actx.") parser.add_argument("--restart_file", help="root name of restart file") parser.add_argument("--casename", help="casename to use for i/o") + parser.add_argument("--cupy", action="store_true", + help="use cupy-based eager actx.") args = parser.parse_args() from warnings import warn @@ -413,7 +415,8 @@ def my_rhs(t, state): from mirgecom.array_context import get_reasonable_array_context_class actx_class = get_reasonable_array_context_class( - lazy=args.lazy, distributed=True, profiling=args.profiling, numpy=args.numpy) + lazy=args.lazy, distributed=True, profiling=args.profiling, + numpy=args.numpy, cupy=args.cupy) logging.basicConfig(format="%(message)s", level=logging.INFO) if args.casename: From a720201de339f78ac1ce214fb2b84c03a693a43f Mon Sep 17 00:00:00 2001 From: Matthias Diener Date: Fri, 16 Feb 2024 20:26:33 -0600 Subject: [PATCH 4/6] pylint --- mirgecom/array_context.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/mirgecom/array_context.py b/mirgecom/array_context.py index 96ad48e13..2316c27d4 100644 --- a/mirgecom/array_context.py +++ b/mirgecom/array_context.py @@ -63,10 +63,12 @@ def get_reasonable_array_context_class(*, lazy: bool, distributed: bool, warn("The CupyArrayContext is still under development") if distributed: - from grudge.array_context import MPICupyArrayContext + from grudge.array_context import ( # pylint: disable=no-name-in-module + MPICupyArrayContext) return MPICupyArrayContext else: - from grudge.array_context import CupyArrayContext + from grudge.array_context import ( # pylint: disable=no-name-in-module + CupyArrayContext) return CupyArrayContext if numpy: From 65f2d1809fd6c9eed014d4845eaadfd196083d18 Mon Sep 17 00:00:00 2001 From: Matthias Diener Date: Mon, 3 Feb 2025 12:37:29 -0600 Subject: [PATCH 5/6] lint errors --- mirgecom/array_context.py | 9 +++++---- mirgecom/logging_quantities.py | 2 +- 2 files changed, 6 insertions(+), 5 deletions(-) diff --git a/mirgecom/array_context.py b/mirgecom/array_context.py index 1080fd186..d69e31c3c 100644 --- a/mirgecom/array_context.py +++ b/mirgecom/array_context.py @@ -64,11 +64,11 @@ def get_reasonable_array_context_class(*, lazy: bool, distributed: bool, warn("The CupyArrayContext is still under development") if distributed: - from grudge.array_context import ( # pylint: disable=no-name-in-module + from grudge.array_context import ( # pylint: disable=no-name-in-module # type: ignore[attr-defined] # noqa: E501 MPICupyArrayContext) return MPICupyArrayContext else: - from grudge.array_context import ( # pylint: disable=no-name-in-module + from grudge.array_context import ( # pylint: disable=no-name-in-module # type: ignore[attr-defined] # noqa: E501 CupyArrayContext) return CupyArrayContext @@ -335,7 +335,8 @@ def initialize_actx( if comm: actx_kwargs["mpi_communicator"] = comm - # Special handling for NumpyArrayContext/CupyArrayContext since they need no CL context + # Special handling for NumpyArrayContext/CupyArrayContext + # since they need no CL context if actx_class_is_numpy(actx_class) or actx_class_is_cupy(actx_class): if actx_class_is_numpy(actx_class): from grudge.array_context import MPINumpyArrayContext @@ -344,7 +345,7 @@ def initialize_actx( else: assert not issubclass(actx_class, MPINumpyArrayContext) else: - from grudge.array_context import MPICupyArrayContext + from grudge.array_context import MPICupyArrayContext # pylint: disable=no-name-in-module # type: ignore[attr-defined] # noqa: E501 if comm: assert issubclass(actx_class, MPICupyArrayContext) else: diff --git a/mirgecom/logging_quantities.py b/mirgecom/logging_quantities.py index 67cf65527..33631d76d 100644 --- a/mirgecom/logging_quantities.py +++ b/mirgecom/logging_quantities.py @@ -78,7 +78,7 @@ def initialize_logmgr(enable_logmgr: bool, logmgr.enable_save_on_sigterm() add_run_info(logmgr) - add_package_versions(logmgr) + # add_package_versions(logmgr) add_general_quantities(logmgr) add_simulation_quantities(logmgr) From 4eeca0648d2d24d2fdd0aa27c477a4362bb0e636 Mon Sep 17 00:00:00 2001 From: Matthias Diener Date: Mon, 3 Feb 2025 13:17:49 -0600 Subject: [PATCH 6/6] typing --- mirgecom/array_context.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/mirgecom/array_context.py b/mirgecom/array_context.py index d69e31c3c..29d4ed39e 100644 --- a/mirgecom/array_context.py +++ b/mirgecom/array_context.py @@ -64,11 +64,11 @@ def get_reasonable_array_context_class(*, lazy: bool, distributed: bool, warn("The CupyArrayContext is still under development") if distributed: - from grudge.array_context import ( # pylint: disable=no-name-in-module # type: ignore[attr-defined] # noqa: E501 + from grudge.array_context import ( # type: ignore[attr-defined] # pylint: disable=no-name-in-module # noqa: E501 MPICupyArrayContext) return MPICupyArrayContext else: - from grudge.array_context import ( # pylint: disable=no-name-in-module # type: ignore[attr-defined] # noqa: E501 + from grudge.array_context import ( # type: ignore[attr-defined] # pylint: disable=no-name-in-module # noqa: E501 CupyArrayContext) return CupyArrayContext @@ -125,7 +125,7 @@ def actx_class_is_numpy(actx_class: Type[ArrayContext]) -> bool: def actx_class_is_cupy(actx_class: Type[ArrayContext]) -> bool: """Return True if *actx_class* is cupy-based.""" try: - from grudge.array_context import CupyArrayContext + from grudge.array_context import CupyArrayContext # type: ignore[attr-defined] # noqa: E501 if issubclass(actx_class, CupyArrayContext): return True else: @@ -345,7 +345,7 @@ def initialize_actx( else: assert not issubclass(actx_class, MPINumpyArrayContext) else: - from grudge.array_context import MPICupyArrayContext # pylint: disable=no-name-in-module # type: ignore[attr-defined] # noqa: E501 + from grudge.array_context import MPICupyArrayContext # type: ignore[attr-defined] # pylint: disable=no-name-in-module # noqa: E501 if comm: assert issubclass(actx_class, MPICupyArrayContext) else: