diff --git a/docs/development/ADRs/next/0028-Plain-Builders-Instead-of-Factories.md b/docs/development/ADRs/next/0028-Plain-Builders-Instead-of-Factories.md new file mode 100644 index 0000000000..e5f4fbd0d0 --- /dev/null +++ b/docs/development/ADRs/next/0028-Plain-Builders-Instead-of-Factories.md @@ -0,0 +1,160 @@ +--- +tags: [backend, otf, toolchain, workflows, dependencies] +--- + +# Plain Builders Instead of factory-boy Factories + +- **Status**: valid +- **Authors**: Enrique González Paredes (@egparedes) +- **Created**: 2026-08-20 +- **Updated**: 2026-09-30 + +In the context of composing the GTFN and DaCe toolchains and their OTF compile +workflows, facing a production dependency on `factory-boy` — a test-data +library — whose `Trait` / `SubFactory` / `SelfAttribute` / `LazyAttribute` +machinery and stringly-typed `__`-path overrides are invisible to the type +checker and fail silently, we decided to replace the factory classes with +plain builder functions driven by one configuration object per toolchain +family and replaceable step builders, to achieve statically checked +composition, steps that agree on shared settings by construction, and one +fewer runtime dependency. + +## Context + +Every object the toolchain factories built was already a frozen dataclass. +`factory-boy` added a second, parallel construction language on top of them: + +- **Untyped.** The declarations are class attributes of a `Params` block, so + `mypy` cannot check them, and the factories needed `type: ignore` + suppressions to type-check at all. +- **Silently wrong.** Overrides are `__`-delimited strings resolved at + runtime. When a path does not resolve, nothing happens. This was not + hypothetical: an override meant to switch a backend to imperative code + generation (`otf_workflow__translation__use_imperative_backend=True`) never + reached the translation step, because a caching trait had wrapped it, so the + backend was the declarative one under another name for its whole life. +- **A runtime dependency for a test-time concern.** `factory-boy` was shipped + to every user to compose a handful of backends. + +Whatever replaces the factories must still meet two concerns ADR 0017 lists +for toolchain configuration: some settings must be configured in **several +components in sync** (the target device reaches the translation step, the +compiler and the allocator), and users need to **switch out or tweak nested +steps** (a translation step without the GTIR transforms, a different build +system). `factory-boy` met both, untyped: `SelfAttribute("..device_type")` for +the first, `SubFactory` plus `__` paths for the second. + +## Decision + +Factory classes are replaced by **plain builder functions**; `factory-boy` +moves to the test dependencies, where the IR test-data factories keep using it +for what it is designed for. The builders follow three rules. + +1. **Shared settings live in one configuration.** Each toolchain family has a + frozen config dataclass extending `backend.ToolchainConfig`. It holds the + settings that describe what the toolchain builds for — the device, the + build type, the cache lifetime, the data layout, translation caching — and + the family-specific settings that several steps must agree on. Its defaults + are read from `gt4py.next.config` when the config is created, which gives + ADR 0017's precedence: an explicit argument wins over the user + configuration, which wins over the builder default. Values derived from + shared settings are derived in exactly one place. + +2. **Every step is created by a step builder that receives the config.** A + step builder takes the config as its only positional argument and + **step-local settings only** as keyword arguments, so a shared setting + cannot be set for one step alone. A step is customized by passing a + different builder to the toolchain or compile-workflow builder: a + `functools.partial` of the default builder to change a step-local setting, + or any `Callable[[Config], Step]` to replace the step. Nested steps follow + the same pattern. + +3. **The toolchain builder owns the composition.** It calls the step builders + and then applies the wrappers, such as the translation cache, so a + customization always lands on the bare step. A fully custom step builder is + responsible for configuring its step consistently with the config it + receives; the builders do not validate what it returns. + +```python +make_gtfn_toolchain( + GTFNConfig(gpu=True), + name_postfix="_no_transforms", + translation=functools.partial(make_gtfn_translation, enable_itir_transforms=False), +) +``` + +The pre-existing flat-keyword `make_dace_backend` is kept as a deprecated +front end over the config-based builder, for existing callers. + +## Consequences + +- Composition is ordinary, statically checked Python: a misspelled step-local + setting, a value of the wrong type, or an attempt to set a shared setting + through a step builder is a type error, and a `TypeError` when the toolchain + is built. +- Default and partially customized steps agree on the shared settings by + construction, because they read them from the same config. Steps from fully + custom step builders are not checked. +- A step field that must agree with other steps should have no default, so a + builder that forgets to pass it fails instead of silently using the default. + Fields read by a single step, such as the build type of the build system, + may keep a default; a custom step builder that creates such a component + must pass the config value itself. +- A new shared setting is one config field, read where it is needed, instead + of a keyword argument threaded through every builder layer. +- Step builders run at build time and are not stored, so a `lambda` step + builder does not make the toolchain unpicklable. +- There is one flat config per toolchain family. It is not composable: a + toolchain assembled from sub-toolchains with different shared settings would + need a different structure. Nothing needs that today. +- Configuration is pulled by each step builder from the config rather than + pushed down by the parent, so a step builder's signature does not show which + shared settings it reads, and a setting a step gains later silently keeps its + step default until its builder reads it from the config. +- The construction logic is spread over many small step builders, which makes + the consistency between steps harder to see and requires unit tests per + builder. +- The price is also two concepts instead of one (the config and the step + builders), a `partial` is less discoverable than a keyword argument, the step + builders re-list the step-local fields of their step, and the config must + stay limited to shared settings or it turns into a grab-bag. + +## Alternatives Considered + +### Typed forwarding of per-step options + +Builders could take the step-local settings of each inner step as a +`TypedDict` and forward them (`translation={"enable_itir_transforms": False}`). +That keeps `factory-boy`'s one-call ergonomics, keeps the configuration flow +from parent to child visible, and is statically checked: `mypy` checks the +keys when a `TypedDict` is unpacked into the step constructor, so a stale key +is an error, and leaving the shared settings out of the `TypedDict`s keeps +them in sync. Like the step builders chosen here, a `TypedDict` only exposes +the step settings someone added to it. + +It was not chosen because replacing a whole step needs a second mechanism next +to the options, typically injecting a pre-built step, which brings back the +problems of the next alternative; and because every shared setting has to be +forwarded by hand through each builder layer. + +### Inject pre-built steps, used verbatim + +Builders could take shared settings as keyword arguments and accept a +pre-built step, used verbatim. Changing one setting of an inner step then +means building the whole step and repeating the shared settings in it, which +the builder already knew; a GPU toolchain with a CPU translation step is only +caught if the builder checks for it. + +### Stamp shared settings onto injected steps + +A `with_changes(step, **changes)` helper would stamp the shared settings onto +whichever step is present, applying only the fields the target declares. +Silently ignoring the fields a target does not declare reproduces the failure +mode that motivated this ADR. + +### Edit the built toolchain + +`dataclasses.replace` on a built toolchain keeps working, but the caller must +know the nesting of wrappers — the path the silently dropped override above +never reached — and the values the builder derived (cache folders, name, +allocator) are not recomputed. diff --git a/docs/development/ADRs/next/README.md b/docs/development/ADRs/next/README.md index 1bf9d21812..f19167ef9b 100644 --- a/docs/development/ADRs/next/README.md +++ b/docs/development/ADRs/next/README.md @@ -52,6 +52,7 @@ Writing a new ADR is simple: - [0016 - Multiple Backends and Build Systems](0016-Multiple-Backends-and-Build-Systems.md) - [0017 - Toolchain Configuration](0017-Toolchain-Configuration.md) - [0027 - External Workspace Memory for DaCe Transients](0027-External_Workspace_Memory.md) +- [0028 - Plain Builders Instead of factory-boy Factories](0028-Plain-Builders-Instead-of-Factories.md) ### Python Integration diff --git a/docs/user/next/advanced/HackTheToolchain.md b/docs/user/next/advanced/HackTheToolchain.md index 15e2e98ff1..773070130b 100644 --- a/docs/user/next/advanced/HackTheToolchain.md +++ b/docs/user/next/advanced/HackTheToolchain.md @@ -46,25 +46,44 @@ skip_linting_transforms = SkipLinting(**same_steps) skip_linting_transforms.step_order(DUMMY_FOP) ``` -## Alternative Factory +## Alternative Workflow + +A toolchain is built from one configuration, `GTFNConfig` (or `DaCeConfig`), +which holds the settings all steps must agree on: the device, the build type, +the cache lifetime, the data layout. Each step is created by a step builder +that receives that configuration. To change a single setting of one step, pass +a `functools.partial` of its default builder; the other settings still come +from the configuration. ```python -class MyCodeGen: ... +import functools +gtfn = gtx.program_processors.runners.gtfn -class Cpp2BindingsGen: ... +debug_gpu_no_transforms = gtfn.make_gtfn_toolchain( + gtfn.GTFNConfig(gpu=True, cmake_build_type=gtx.config.CMakeBuildType.DEBUG), + name_postfix="_debug_no_transforms", + translation=functools.partial(gtfn.make_gtfn_translation, enable_itir_transforms=False), +) +``` +To replace a step, pass any callable that takes the configuration and returns +the step. It is still wrapped in the translation cache. Configuring it +consistently with the configuration it receives (the device, for instance) is +up to the callable. + +```python +class MyCodeGen: ... -class PureCpp2WorkflowFactory(gtx.program_processors.runners.gtfn.GTFNCompileWorkflowFactory): - translation: workflow.Workflow[ - gtx.otf.stages.CompilableProgramDef, gtx.otf.artifacts.ProgramSource - ] = MyCodeGen() - bindings: workflow.Workflow[ - gtx.otf.artifacts.ProgramSource, gtx.otf.artifacts.ExtensionSource - ] = Cpp2BindingsGen() +class Cpp2BindingsGen: ... -PureCpp2WorkflowFactory(cmake_build_type=gtx.config.CMAKE_BUILD_TYPE.DEBUG) + +pure_cpp2_workflow = gtfn.make_gtfn_compile_workflow( + gtfn.GTFNConfig(cmake_build_type=gtx.config.CMakeBuildType.DEBUG, cached_translation=False), + translation=lambda cfg: MyCodeGen(), + bindings=lambda cfg: Cpp2BindingsGen(), +) ``` ## Invent new Workflow Types diff --git a/docs/user/next/advanced/WorkflowPatterns.md b/docs/user/next/advanced/WorkflowPatterns.md index 0e0abc4aea..b8ce74a9be 100644 --- a/docs/user/next/advanced/WorkflowPatterns.md +++ b/docs/user/next/advanced/WorkflowPatterns.md @@ -17,7 +17,6 @@ jupyter: import dataclasses import re -import factory import gt4py.next as gtx @@ -199,7 +198,7 @@ Let's say we want to make our calculation workflow compatible with string input. ```python editable=true slideshow={"slide_type": ""} # A plain conversion step turning a string into an int, chained into the -# workflow below and reused by `StrToIntFactory(cached=True)`. +# workflow below and reused by `make_str_to_int(cached=True)`. def to_int(inp: str) -> int: assert isinstance(inp, str), "Can not work with 'int'!" # yes, this is horribly contrived return int(inp) @@ -214,9 +213,9 @@ str_calc("1") -### Step with factory (builder) +### Step with a builder -If a step can be useful with different combinations of parameters and wrappers, it should have a factory. In this case we will add a neutral wrapper around it, so we can put any combination of wrappers into that: +If a step is useful with different combinations of parameters and wrappers, give it a **builder function**: a plain function taking the cross-cutting options and returning the assembled step. Steps are frozen dataclasses, so the builder is ordinary code — no factory framework involved, and the result is fully type-checked. @@ -229,32 +228,23 @@ class AnyStrToInt(gtx.otf.workflow.ChainableWorkflowMixin[str | int, int]): return self.inner_step(inp) -class StrToIntFactory(factory.Factory): - class Meta: - model = AnyStrToInt +def make_str_to_int( + *, cached: bool = False, step: gtx.otf.workflow.Workflow[str, int] = to_int +) -> AnyStrToInt: + if cached: + step = gtx.otf.workflow.CachedStep.in_memory(step=step, input_fingerprinter=str) + return AnyStrToInt(inner_step=step) - class Params: - default_step = to_int - cached = factory.Trait( - inner_step=factory.LazyAttribute( - lambda o: gtx.otf.workflow.CachedStep.in_memory( - step=o.default_step, input_fingerprinter=str - ) - ) - ) - inner_step = factory.LazyAttribute(lambda o: o.default_step) - - -cached = StrToIntFactory(cached=True) -uncached = StrToIntFactory() +cached = make_str_to_int(cached=True) +uncached = make_str_to_int() uncached.inner_step ``` ### Example in the Wild ```python -gtx.ffront.past_passes.linters.LinterFactory?? +gtx.ffront.past_passes.linters.linter_factory?? ``` @@ -413,5 +403,5 @@ gtx.program_processors.runners.gtfn.run_gtfn_gpu.executor.otf_workflow?? ``` ```python -gtx.program_processors.runners.gtfn.GTFNBackendFactory?? +gtx.program_processors.runners.gtfn.make_gtfn_toolchain?? ``` diff --git a/pyproject.toml b/pyproject.toml index e5717fe037..e131ebda0e 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -43,6 +43,7 @@ profiling = [ ] scripts = ["pyyaml>=6.0.1", "typer>=0.16.0", "packaging"] test = [ + 'factory-boy>=3.3.3', 'hypothesis>=6.0.0', 'nbmake>=1.4.6', 'nox>=2025.02.09', @@ -99,7 +100,6 @@ dependencies = [ 'dace==2.0.0a9', 'deepdiff>=8.1.0', 'devtools>=0.6', - 'factory-boy>=3.3.3', "filelock>=3.18.0", 'frozendict>=2.3', 'gridtools-cpp>=2.3.9,==2.*', @@ -231,12 +231,6 @@ module = 'gt4py.next.iterator.*' ignore_errors = true module = 'gt4py.next.iterator.runtime' -[[tool.mypy.overrides]] -ignore_missing_imports = true -implicit_reexport = true -# factory-boy is broken, see https://github.com/FactoryBoy/factory_boy/pull/1114 -module = "factory.*" - # -- pytest -- [tool.pytest] diff --git a/src/gt4py/next/backend.py b/src/gt4py/next/backend.py index eae12981f0..94f54172f4 100644 --- a/src/gt4py/next/backend.py +++ b/src/gt4py/next/backend.py @@ -12,7 +12,7 @@ from typing import Generic from gt4py._core import definitions as core_defs -from gt4py.next import custom_layout_allocators as next_allocators +from gt4py.next import config, custom_layout_allocators as next_allocators from gt4py.next.ffront import ( foast_to_gtir, foast_to_past, @@ -172,3 +172,51 @@ def __gt_allocator__( self, ) -> next_allocators.FieldBufferAllocatorProtocol[core_defs.DeviceTypeT]: return self.allocator + + +@dataclasses.dataclass(frozen=True) +class ToolchainConfig: + """ + Settings that describe what a compiled toolchain builds for. + + A toolchain builder creates every step from one config, so steps that must + agree on a setting (the target device, the build type, the cache lifetime, + the data layout) read it from the same place and cannot drift apart. + Settings that only tune how a single step does its job are not part of the + config: they are keyword arguments of that step's builder. + + Defaults are read from `gt4py.next.config` when the config is created, not + when this module is imported, so the toolchains built from a default config + follow the current user configuration. + """ + + #: Target a GPU (the one CuPy was built for) instead of the CPU. + gpu: bool = False + #: Wrap the translation step in a persistent cache. + cached_translation: bool = True + cmake_build_type: config.CMakeBuildType = dataclasses.field( + default_factory=lambda: config.CMAKE_BUILD_TYPE + ) + cache_lifetime: config.BuildCacheLifetime = dataclasses.field( + default_factory=lambda: config.BUILD_CACHE_LIFETIME + ) + #: Assume unit stride in the horizontal dimension of unstructured fields. + unstructured_horizontal_has_unit_stride: bool = dataclasses.field( + default_factory=lambda: config.UNSTRUCTURED_HORIZONTAL_HAS_UNIT_STRIDE + ) + + @property + def device_type(self) -> core_defs.DeviceType: + if self.gpu: + return core_defs.CUPY_DEVICE_TYPE or core_defs.DeviceType.CUDA + return core_defs.DeviceType.CPU + + @property + def device_name(self) -> str: + """Device part of the toolchain name.""" + return "gpu" if self.gpu else "cpu" + + def make_allocator(self) -> next_allocators.FieldBufferAllocatorProtocol: + if self.gpu: + return next_allocators.StandardGPUFieldBufferAllocator() + return next_allocators.StandardCPUFieldBufferAllocator() diff --git a/src/gt4py/next/otf/compilation/cache.py b/src/gt4py/next/otf/compilation/cache.py index eac9c07dc2..7950db129a 100644 --- a/src/gt4py/next/otf/compilation/cache.py +++ b/src/gt4py/next/otf/compilation/cache.py @@ -10,10 +10,11 @@ import pathlib import tempfile -from typing import Final +from typing import Final, TypeVar +from gt4py._core import filecache from gt4py.next import config, fingerprinting -from gt4py.next.otf import artifacts +from gt4py.next.otf import artifacts, workflow #: Regex describing the folder names produced by `get_cache_folder` (use @@ -44,9 +45,12 @@ TRANSLATION_CACHE_DIR_NAME: Final[str] = "translation_cache" #: Backends that persist the output of their translation step, i.e. those whose -#: workflow factory enables the `cached_translation` trait. +#: builders wrap it with `persistent_translation_cache`. TRANSLATION_CACHE_BACKENDS: Final[tuple[str, ...]] = ("dace", "gtfn") +StartT = TypeVar("StartT") +EndT = TypeVar("EndT") + _session_cache_dir = tempfile.TemporaryDirectory(prefix="gt4py_session_") _session_cache_dir_path = pathlib.Path(_session_cache_dir.name) @@ -62,6 +66,29 @@ def get_translation_cache_folder(cache_base: pathlib.Path, backend: str) -> path return cache_base / TRANSLATION_CACHE_DIR_NAME / backend +def persistent_translation_cache( + step: workflow.Workflow[StartT, EndT], backend: str, lifetime: config.BuildCacheLifetime +) -> workflow.CachedStep[StartT, EndT, str]: + """ + Wrap a translation step in the persistent translation cache of `backend`. + + Args: + step: The translation step to cache. + backend: Name of the backend family, which selects the cache folder. + lifetime: Build-cache lifetime, which selects the cache base path. + + Returns: + The step, cached in the translation cache folder of `backend`. + """ + return workflow.CachedStep[StartT, EndT, str].persistent( + step, + input_fingerprinter=fingerprinting.strict_fingerprinter, + cache=filecache.FileCache( + get_translation_cache_folder(get_cache_base_path(lifetime), backend) + ), + ) + + def get_cache_base_path(lifetime: config.BuildCacheLifetime) -> pathlib.Path: """Return the base directory for cached artifacts with the given lifetime.""" match lifetime: diff --git a/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py b/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py index dcf12cf509..f861d4b182 100644 --- a/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py +++ b/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py @@ -11,7 +11,6 @@ import dataclasses from typing import Any, Final, Optional -import factory import numpy as np from gt4py._core import definitions as core_defs @@ -48,7 +47,9 @@ class GTFNTranslationStep( code_spec: Optional[artifacts.HeaderAndSourceCodeSpec] = None # TODO replace by more general mechanism, see https://github.com/GridTools/gt4py/issues/1135 enable_itir_transforms: bool = True - device_type: core_defs.DeviceType = core_defs.DeviceType.CPU + # No default: the device must agree with the other steps of the pipeline, so + # forgetting to pass it must fail instead of silently targeting the CPU. + device_type: core_defs.DeviceType = dataclasses.field(kw_only=True) symbolic_domain_sizes: dict[str, itir.Expr] | None = None use_max_domain_range_on_unstructured_shift: bool | None = None @@ -265,13 +266,10 @@ def _not_implemented_for_device_type(self) -> NotImplementedError: ) -class GTFNTranslationStepFactory(factory.Factory[GTFNTranslationStep]): - class Meta: - model = GTFNTranslationStep - - -translate_program_cpu: Final[stages.TranslationStep] = GTFNTranslationStepFactory() # type: ignore[assignment] # factory-boy typing not precise enough +translate_program_cpu: Final[stages.TranslationStep] = GTFNTranslationStep( + device_type=core_defs.DeviceType.CPU +) -translate_program_gpu: Final[stages.TranslationStep] = GTFNTranslationStepFactory( # type: ignore[assignment] # factory-boy typing not precise enough +translate_program_gpu: Final[stages.TranslationStep] = GTFNTranslationStep( device_type=core_defs.DeviceType.CUDA ) diff --git a/src/gt4py/next/program_processors/formatters/gtfn.py b/src/gt4py/next/program_processors/formatters/gtfn.py index 75494a1759..a2ccab6168 100644 --- a/src/gt4py/next/program_processors/formatters/gtfn.py +++ b/src/gt4py/next/program_processors/formatters/gtfn.py @@ -10,14 +10,12 @@ from gt4py.next.iterator import ir as itir from gt4py.next.program_processors import program_formatter -from gt4py.next.program_processors.codegens.gtfn.gtfn_module import GTFNTranslationStep from gt4py.next.program_processors.runners import gtfn @program_formatter.program_formatter def format_cpp(program: itir.Program, *args: Any, **kwargs: Any) -> str: - gtfn_translation = gtfn.GTFNCompileWorkflowFactory(cached_translation=False).translation - assert isinstance(gtfn_translation, GTFNTranslationStep) + gtfn_translation = gtfn.make_gtfn_translation(gtfn.GTFNConfig()) return gtfn_translation.generate_stencil_source( program, offset_provider=kwargs.get("offset_provider", {}), diff --git a/src/gt4py/next/program_processors/runners/dace/__init__.py b/src/gt4py/next/program_processors/runners/dace/__init__.py index 0e560fa761..2e06493bd0 100644 --- a/src/gt4py/next/program_processors/runners/dace/__init__.py +++ b/src/gt4py/next/program_processors/runners/dace/__init__.py @@ -10,16 +10,36 @@ from gt4py.next.program_processors.runners.dace.sdfg_callable import get_sdfg_args from gt4py.next.program_processors.runners.dace.workflow.backend import ( make_dace_backend, + make_dace_toolchain, run_dace_cpu, run_dace_cpu_noopt, run_dace_gpu, run_dace_gpu_noopt, ) +from gt4py.next.program_processors.runners.dace.workflow.factory import ( + DaCeBindingsBuilder, + DaCeCompilationBuilder, + DaCeConfig, + DaCeTranslationBuilder, + make_dace_bindings, + make_dace_compile_workflow, + make_dace_compiler, + make_dace_translator, +) __all__ = [ + "DaCeBindingsBuilder", + "DaCeCompilationBuilder", + "DaCeConfig", + "DaCeTranslationBuilder", "get_sdfg_args", "make_dace_backend", + "make_dace_bindings", + "make_dace_compile_workflow", + "make_dace_compiler", + "make_dace_toolchain", + "make_dace_translator", "run_dace_cpu", "run_dace_cpu_noopt", "run_dace_gpu", diff --git a/src/gt4py/next/program_processors/runners/dace/workflow/__init__.py b/src/gt4py/next/program_processors/runners/dace/workflow/__init__.py index 4d825c0c9b..2159bcc0a3 100644 --- a/src/gt4py/next/program_processors/runners/dace/workflow/__init__.py +++ b/src/gt4py/next/program_processors/runners/dace/workflow/__init__.py @@ -15,6 +15,7 @@ - `compilation` for compiling the SDFG into a program - `decoration` to parse the program arguments and pass them to the program call -The GTIR-DaCe backend factory extends `CachedBackendFactory`, thus it provides -caching of the GTIR program. +The toolchain builders create every step from one `DaCeConfig`, and wrap the +translation step in a persistent `CachedStep` unless the config disables it, +thus they provide caching of the GTIR program. """ diff --git a/src/gt4py/next/program_processors/runners/dace/workflow/backend.py b/src/gt4py/next/program_processors/runners/dace/workflow/backend.py index c0ea33daf0..98149f20c4 100644 --- a/src/gt4py/next/program_processors/runners/dace/workflow/backend.py +++ b/src/gt4py/next/program_processors/runners/dace/workflow/backend.py @@ -9,16 +9,12 @@ from __future__ import annotations import dataclasses +import functools import warnings -from typing import Any, Final +from typing import Any -import factory - -import gt4py.next.custom_layout_allocators as next_allocators -from gt4py._core import definitions as core_defs -from gt4py.next import backend, common, config +from gt4py.next import backend, config from gt4py.next.otf import artifacts -from gt4py.next.program_processors.runners.dace import transformations as gtx_transformations from gt4py.next.program_processors.runners.dace.workflow import ( common as gtx_wfdcommon, decoration as gtx_wfddecoration, @@ -40,41 +36,41 @@ def load_artifact(self, artifact: artifacts.CompilationArtifact) -> artifacts.Ex return program -class DaCeBackendFactory(factory.Factory): +def make_dace_toolchain( + cfg: gtx_wfdfactory.DaCeConfig | None = None, + /, + *, + name_postfix: str = "", + translation: gtx_wfdfactory.DaCeTranslationBuilder = gtx_wfdfactory.make_dace_translator, + bindings: gtx_wfdfactory.DaCeBindingsBuilder = gtx_wfdfactory.make_dace_bindings, + compilation: gtx_wfdfactory.DaCeCompilationBuilder = gtx_wfdfactory.make_dace_compiler, +) -> DaCeBackend: """ - Workflow factory for the GTIR-DaCe backend. - - Several parameters are inherithed from `backend.Backend`, see below the specific ones. + Build a DaCe toolchain. Args: - auto_optimize: Enables the SDFG transformation pipeline. - """ + cfg: The toolchain configuration. Defaults to `DaCeConfig()`. + name_postfix: Appended to the toolchain name, which must stay unique. + translation: Builder of the translation step, see + `make_dace_compile_workflow`. + bindings: Builder of the bindings step. + compilation: Builder of the compilation step. - class Meta: - model = DaCeBackend - - class Params: - name_device = "cpu" - name_postfix = "" - gpu = factory.Trait( - allocator=next_allocators.StandardGPUFieldBufferAllocator(), - device_type=core_defs.CUPY_DEVICE_TYPE or core_defs.DeviceType.CUDA, - name_device="gpu", - ) - device_type = core_defs.DeviceType.CPU - otf_workflow = factory.SubFactory( - gtx_wfdfactory.DaCeWorkflowFactory, - cached_translation=True, - device_type=factory.SelfAttribute("..device_type"), - auto_optimize=factory.SelfAttribute("..auto_optimize"), - ) - auto_optimize = factory.Trait(name_postfix="_opt") - - name = factory.LazyAttribute(lambda o: f"run_dace_{o.name_device}{o.name_postfix}") - executor = factory.LazyAttribute(lambda o: o.otf_workflow) - allocator = next_allocators.StandardCPUFieldBufferAllocator() - transforms = backend.DEFAULT_TRANSFORMS - external_workspace = None + Returns: + The configured toolchain. + """ + if cfg is None: + cfg = gtx_wfdfactory.DaCeConfig() + + return DaCeBackend( + name=f"run_dace_{cfg.device_name}{'_opt' if cfg.auto_optimize else ''}{name_postfix}", + executor=gtx_wfdfactory.make_dace_compile_workflow( + cfg, translation=translation, bindings=bindings, compilation=compilation + ), + allocator=cfg.make_allocator(), + transforms=backend.DEFAULT_TRANSFORMS, + external_workspace=cfg.external_workspace, + ) def make_dace_backend( @@ -87,9 +83,14 @@ def make_dace_backend( use_metrics: bool = True, use_zero_origin: bool = False, use_max_domain_range_on_unstructured_shift: bool | None = None, -) -> backend.Backend: +) -> DaCeBackend: """Customize the dace backend with the given configuration parameters. + Deprecated: use `make_dace_toolchain` with a `DaCeConfig` for the shared + settings and `functools.partial(make_dace_translator, ...)` for the + translation settings. This flat-keyword front end builds exactly that from + its arguments. + Args: gpu: Enable GPU transformations and code generation. auto_optimize: Enable the SDFG auto-optimize pipeline. @@ -106,85 +107,45 @@ def make_dace_backend( use_zero_origin: Can be set to `True` when all fields passed as program arguments have zero-based origin. This setting will skip generation of range start-symbols `_range_0` since they can be assumed to be zero. + use_max_domain_range_on_unstructured_shift: See `DaCeTranslator`. Note that `gt_auto_optimize()` parameters that are derived from GT4Py configuration - cannot be overriden, and therefore cannot appear here. Thus, this function will - throw an exception if called with any argument included in `gt_optimization_args`. + cannot be overridden, and therefore cannot appear in `optimization_args`. Returns: A dace backend with custom configuration for the target device. - """ - # The `gt_optimization_args` set contains the parameters of `gt_auto_optimize()` - # that are derived from the gt4py configuration, and therefore cannot be customized. - gt_optimization_args: Final[set[str]] = {"gpu", "constant_symbols", "unit_strides_kind"} - - if optimization_args is None: - optimization_args = {} - elif optimization_args and not auto_optimize: - warnings.warn("Optimizations args given, but auto-optimize is disabled.", stacklevel=2) - elif intersect_args := gt_optimization_args.intersection(optimization_args.keys()): - raise ValueError( - f"The following optimization arguments cannot be overriden: {intersect_args}." - ) - - # Set `unit_strides_kind` based on the gt4py env configuration. - optimization_args = optimization_args | { - "unit_strides_kind": common.DimensionKind.HORIZONTAL - if unstructured_horizontal_has_unit_stride - else None - } - - if external_workspace is None: - if ( - optimization_args.get("transient_memory_mode") - is gtx_transformations.TransientMemoryMode.EXTERNAL - ): - raise ValueError( - "External memory workspace must be provided when 'transient_memory_mode' is 'EXTERNAL'." - ) - elif transient_memory_mode := optimization_args.get("transient_memory_mode"): - if transient_memory_mode is not gtx_transformations.TransientMemoryMode.EXTERNAL: - warnings.warn( - f"External memory workspace provided but 'transient_memory_mode' is '{transient_memory_mode}', it requires '{gtx_transformations.TransientMemoryMode.EXTERNAL}'.", - stacklevel=2, - ) - else: - optimization_args["transient_memory_mode"] = ( - gtx_transformations.TransientMemoryMode.EXTERNAL - ) - - return DaCeBackendFactory( # type: ignore[return-value] # factory-boy typing not precise enough - gpu=gpu, - auto_optimize=auto_optimize, - external_workspace=external_workspace, - otf_workflow__bare_translation__async_sdfg_call=(async_sdfg_call if gpu else False), - otf_workflow__bare_translation__auto_optimize_args=optimization_args, - otf_workflow__bare_translation__unstructured_horizontal_has_unit_stride=unstructured_horizontal_has_unit_stride, - otf_workflow__bare_translation__use_metrics=use_metrics, - otf_workflow__bare_translation__disable_field_origin_on_program_arguments=use_zero_origin, - otf_workflow__bare_translation__use_max_domain_range_on_unstructured_shift=use_max_domain_range_on_unstructured_shift, + Raises: + ValueError: If `optimization_args` sets a parameter derived from the + configuration, or requests the `EXTERNAL` transient memory mode + without an `external_workspace`. + """ + warnings.warn( + "'make_dace_backend' is deprecated, use 'make_dace_toolchain' with a 'DaCeConfig' and" + " 'functools.partial(make_dace_translator, ...)' instead.", + DeprecationWarning, + stacklevel=2, + ) + return make_dace_toolchain( + gtx_wfdfactory.DaCeConfig( + gpu=gpu, + auto_optimize=auto_optimize, + external_workspace=external_workspace, + unstructured_horizontal_has_unit_stride=unstructured_horizontal_has_unit_stride, + ), + translation=functools.partial( + gtx_wfdfactory.make_dace_translator, + auto_optimize_args=optimization_args, + async_sdfg_call=async_sdfg_call, + use_metrics=use_metrics, + disable_field_origin_on_program_arguments=use_zero_origin, + use_max_domain_range_on_unstructured_shift=use_max_domain_range_on_unstructured_shift, + ), ) -run_dace_cpu = make_dace_backend( - gpu=False, - auto_optimize=True, - async_sdfg_call=False, -) -run_dace_cpu_noopt = make_dace_backend( - gpu=False, - auto_optimize=False, - async_sdfg_call=False, -) +run_dace_cpu = make_dace_toolchain(gtx_wfdfactory.DaCeConfig(auto_optimize=True)) +run_dace_cpu_noopt = make_dace_toolchain(gtx_wfdfactory.DaCeConfig(auto_optimize=False)) -run_dace_gpu = make_dace_backend( - gpu=True, - auto_optimize=True, - async_sdfg_call=True, -) -run_dace_gpu_noopt = make_dace_backend( - gpu=True, - auto_optimize=False, - async_sdfg_call=True, -) +run_dace_gpu = make_dace_toolchain(gtx_wfdfactory.DaCeConfig(gpu=True, auto_optimize=True)) +run_dace_gpu_noopt = make_dace_toolchain(gtx_wfdfactory.DaCeConfig(gpu=True, auto_optimize=False)) diff --git a/src/gt4py/next/program_processors/runners/dace/workflow/compilation.py b/src/gt4py/next/program_processors/runners/dace/workflow/compilation.py index b202dea1cf..3a8bc9eac1 100644 --- a/src/gt4py/next/program_processors/runners/dace/workflow/compilation.py +++ b/src/gt4py/next/program_processors/runners/dace/workflow/compilation.py @@ -18,7 +18,6 @@ import dace import dace.codegen.compiler as dace_compiler -import factory from gt4py._core import definitions as core_defs, locking from gt4py.eve import xtyping @@ -367,8 +366,3 @@ def __call__(self, inp: SDFGExtensionSource) -> DaCeCompilationArtifact: bind_func_name=self.bind_func_name, device_type=self.device_type, ) - - -class DaCeCompilationStepFactory(factory.Factory): - class Meta: - model = DaCeCompiler diff --git a/src/gt4py/next/program_processors/runners/dace/workflow/factory.py b/src/gt4py/next/program_processors/runners/dace/workflow/factory.py index 2f37f90cd7..6b3d52dba6 100644 --- a/src/gt4py/next/program_processors/runners/dace/workflow/factory.py +++ b/src/gt4py/next/program_processors/runners/dace/workflow/factory.py @@ -8,69 +8,207 @@ from __future__ import annotations +import dataclasses import functools -from typing import Final +import pathlib +import warnings +from collections.abc import Callable +from typing import Any, ClassVar, Final, TypeAlias -import factory - -from gt4py._core import definitions as core_defs, filecache -from gt4py.next import config, fingerprinting -from gt4py.next.otf import recipes, workflow +import gt4py +from gt4py.next import backend as next_backend, config +from gt4py.next.otf import artifacts, recipes, stages, workflow from gt4py.next.otf.compilation import cache -from gt4py.next.program_processors.runners.dace.workflow import bindings as bindings_step -from gt4py.next.program_processors.runners.dace.workflow.compilation import ( - DaCeCompilationStepFactory, -) -from gt4py.next.program_processors.runners.dace.workflow.translation import ( - DaCeTranslationStepFactory, +from gt4py.next.program_processors.runners.dace import transformations as gtx_transformations +from gt4py.next.program_processors.runners.dace.workflow import ( + bindings as bindings_step, + common as gtx_wfdcommon, ) +from gt4py.next.program_processors.runners.dace.workflow.compilation import DaCeCompiler +from gt4py.next.program_processors.runners.dace.workflow.translation import DaCeTranslator -_GT_DACE_BINDING_FUNCTION_NAME: Final[str] = "update_sdfg_args" +#: Warnings about the builder arguments point at the first caller outside GT4Py. +_GT4PY_SOURCE_PREFIX: Final[str] = str(pathlib.Path(gt4py.__file__).parent) -class DaCeWorkflowFactory(factory.Factory): - class Meta: - model = recipes.OTFCompileWorkflow +@dataclasses.dataclass(frozen=True) +class DaCeConfig(next_backend.ToolchainConfig): + """Settings shared by the steps of a DaCe toolchain, see `ToolchainConfig`.""" - class Params: - auto_optimize: bool = False - device_type: core_defs.DeviceType = core_defs.DeviceType.CPU - cmake_build_type: config.CMakeBuildType = factory.LazyFunction( # type: ignore[assignment] # factory-boy typing not precise enough - lambda: config.CMAKE_BUILD_TYPE - ) + #: Run the SDFG auto-optimize pipeline. + auto_optimize: bool = True + #: Workspace memory allocated outside the SDFG. The toolchain injects it into + #: the loaded programs, and the translation step then defaults to the + #: `EXTERNAL` transient memory mode, which stores the transients in it. The + #: workspace is a dict, so a config holding one is not hashable. + external_workspace: gtx_wfdcommon.ExternalWorkspace | None = None - cached_translation = factory.Trait( - translation=factory.LazyAttribute( - lambda o: workflow.CachedStep.persistent( - o.bare_translation, - input_fingerprinter=fingerprinting.strict_fingerprinter, - cache=filecache.FileCache( - cache.get_translation_cache_folder( - cache.get_cache_base_path(config.BUILD_CACHE_LIFETIME), "dace" - ) - ), - ) - ), - ) + #: Name of the function that binds the SDFG arguments. The bindings and the + #: compilation steps must agree on it. + bind_func_name: ClassVar[str] = "update_sdfg_args" + + +def make_dace_translator( + cfg: DaCeConfig, + /, + *, + auto_optimize_args: dict[str, Any] | None = None, + async_sdfg_call: bool = True, + use_metrics: bool = True, + disable_itir_transforms: bool = False, + disable_field_origin_on_program_arguments: bool = False, + use_max_domain_range_on_unstructured_shift: bool | None = None, +) -> DaCeTranslator: + """ + Build the GTIR -> SDFG translation step. + + The keyword arguments are the step-local fields of `DaCeTranslator`. - bare_translation = factory.SubFactory( - DaCeTranslationStepFactory, - device_type=factory.SelfAttribute("..device_type"), - auto_optimize=factory.SelfAttribute("..auto_optimize"), + Args: + cfg: The toolchain configuration. + auto_optimize_args: Configuration for the SDFG auto-optimize pipeline, + see `gt_auto_optimize()`. The parameters `DaCeTranslator` derives + itself cannot be set here. With an external workspace in `cfg`, the + `transient_memory_mode` defaults to `EXTERNAL`. + async_sdfg_call: Make an asynchronous SDFG call, to overlap GPU kernel + execution with the Python driver code. Only effective on GPU. + use_metrics: Add SDFG instrumentation for stencil compute time. + disable_itir_transforms: Skip the GTIR transformation passes. + disable_field_origin_on_program_arguments: Assume that all fields passed + as program arguments have zero-based origin, which skips the range + start-symbols `_range_0`. + use_max_domain_range_on_unstructured_shift: See `DaCeTranslator`. + + Returns: + The translation step, targeting `cfg.device_type`. + + Raises: + ValueError: If `auto_optimize_args` sets a parameter `DaCeTranslator` + derives itself, or requests the `EXTERNAL` transient memory mode + without an external workspace in `cfg`. + """ + optimization_args = dict(auto_optimize_args or {}) + if optimization_args and not cfg.auto_optimize: + warnings.warn( + "Optimizations args given, but auto-optimize is disabled.", + skip_file_prefixes=(_GT4PY_SOURCE_PREFIX,), ) - translation = factory.LazyAttribute(lambda o: o.bare_translation) - bindings = factory.LazyAttribute( - lambda o: functools.partial( - bindings_step.bind_sdfg, - bind_func_name=_GT_DACE_BINDING_FUNCTION_NAME, + if cfg.external_workspace is None: + if ( + optimization_args.get("transient_memory_mode") + is gtx_transformations.TransientMemoryMode.EXTERNAL + ): + raise ValueError( + "External memory workspace must be provided when 'transient_memory_mode' is 'EXTERNAL'." + ) + elif transient_memory_mode := optimization_args.get("transient_memory_mode"): + if transient_memory_mode is not gtx_transformations.TransientMemoryMode.EXTERNAL: + warnings.warn( + f"External memory workspace provided but 'transient_memory_mode' is '{transient_memory_mode}', it requires '{gtx_transformations.TransientMemoryMode.EXTERNAL}'.", + skip_file_prefixes=(_GT4PY_SOURCE_PREFIX,), + ) + else: + optimization_args["transient_memory_mode"] = ( + gtx_transformations.TransientMemoryMode.EXTERNAL ) + + return DaCeTranslator( + device_type=cfg.device_type, + auto_optimize=cfg.auto_optimize, + auto_optimize_args=optimization_args, + async_sdfg_call=async_sdfg_call and cfg.gpu, + unstructured_horizontal_has_unit_stride=cfg.unstructured_horizontal_has_unit_stride, + use_metrics=use_metrics, + disable_itir_transforms=disable_itir_transforms, + disable_field_origin_on_program_arguments=disable_field_origin_on_program_arguments, + use_max_domain_range_on_unstructured_shift=use_max_domain_range_on_unstructured_shift, + ) + + +def make_dace_bindings( + cfg: DaCeConfig, / +) -> workflow.Workflow[artifacts.ProgramSource, artifacts.ExtensionSource]: + """Build the step generating the bindings of the translated SDFG.""" + return functools.partial(bindings_step.bind_sdfg, bind_func_name=cfg.bind_func_name) + + +def make_dace_compiler( + cfg: DaCeConfig, /, *, add_gpu_trace_markers: bool | None = None +) -> DaCeCompiler: + """ + Build the compilation step, targeting `cfg.device_type`. + + Args: + cfg: The toolchain configuration. + add_gpu_trace_markers: Add GPU trace markers to the generated code. + Defaults to the value in `config`. + + Returns: + The compilation step. + """ + if add_gpu_trace_markers is None: + add_gpu_trace_markers = config.ADD_GPU_TRACE_MARKERS + return DaCeCompiler( + bind_func_name=cfg.bind_func_name, + cache_lifetime=cfg.cache_lifetime, + device_type=cfg.device_type, + cmake_build_type=cfg.cmake_build_type, + add_gpu_trace_markers=add_gpu_trace_markers, ) - compilation = factory.SubFactory( - DaCeCompilationStepFactory, - bind_func_name=_GT_DACE_BINDING_FUNCTION_NAME, - cache_lifetime=factory.LazyFunction(lambda: config.BUILD_CACHE_LIFETIME), - device_type=factory.SelfAttribute("..device_type"), - cmake_build_type=factory.SelfAttribute("..cmake_build_type"), + + +#: Step builders: callables creating a step from the toolchain configuration. +DaCeTranslationBuilder: TypeAlias = Callable[[DaCeConfig], stages.TranslationStep] +DaCeBindingsBuilder: TypeAlias = Callable[ + [DaCeConfig], workflow.Workflow[artifacts.ProgramSource, artifacts.ExtensionSource] +] +DaCeCompilationBuilder: TypeAlias = Callable[ + [DaCeConfig], workflow.Workflow[artifacts.ExtensionSource, artifacts.CompilationArtifact] +] + + +def make_dace_compile_workflow( + cfg: DaCeConfig | None = None, + /, + *, + translation: DaCeTranslationBuilder = make_dace_translator, + bindings: DaCeBindingsBuilder = make_dace_bindings, + compilation: DaCeCompilationBuilder = make_dace_compiler, +) -> recipes.OTFCompileWorkflow: + """ + Build the DaCe translation -> bindings -> compilation workflow. + + Every step is created by a step builder that receives `cfg`, so all steps + agree on the settings in it. To customize a step, pass a different + builder: a `functools.partial` of the default one to change a step-local + setting, e.g. `translation=functools.partial(make_dace_translator, use_metrics=False)`, + or any callable taking the config to replace the step. The translation + step is wrapped in the cache here, after its builder ran, so a custom + translation step is cached like the default one. A custom step builder is + responsible for configuring its step from `cfg`. + + Args: + cfg: The toolchain configuration. Defaults to `DaCeConfig()`. + translation: Builder of the translation step. + bindings: Builder of the bindings step. + compilation: Builder of the compilation step. + + Returns: + The composed compile workflow. + """ + if cfg is None: + cfg = DaCeConfig() + + translation_step = translation(cfg) + compilation_step = compilation(cfg) + + if cfg.cached_translation: + translation_step = cache.persistent_translation_cache( + translation_step, "dace", cfg.cache_lifetime + ) + + return recipes.OTFCompileWorkflow( + translation=translation_step, bindings=bindings(cfg), compilation=compilation_step ) diff --git a/src/gt4py/next/program_processors/runners/dace/workflow/translation.py b/src/gt4py/next/program_processors/runners/dace/workflow/translation.py index f26e982eff..ee96fbeff2 100644 --- a/src/gt4py/next/program_processors/runners/dace/workflow/translation.py +++ b/src/gt4py/next/program_processors/runners/dace/workflow/translation.py @@ -9,10 +9,9 @@ from __future__ import annotations import dataclasses -from typing import Any, Optional +from typing import Any, Final, Optional import dace -import factory from gt4py._core import definitions as core_defs from gt4py.next import common @@ -340,6 +339,13 @@ def make_sdfg_call_sync(sdfg: dace.SDFG, gpu: bool) -> None: ) +#: The parameters of `gt_auto_optimize()` that the translation step derives from +#: its own configuration, and which therefore cannot be customized. +_DERIVED_OPTIMIZATION_ARGS: Final[frozenset[str]] = frozenset( + {"gpu", "constant_symbols", "unit_strides_kind"} +) + + @dataclasses.dataclass(frozen=True) class DaCeTranslator( workflow.ChainableWorkflowMixin[ @@ -362,6 +368,14 @@ class DaCeTranslator( disable_field_origin_on_program_arguments: bool = False use_max_domain_range_on_unstructured_shift: bool | None = None + def __post_init__(self) -> None: + if self.auto_optimize_args and ( + derived_args := self.auto_optimize_args.keys() & _DERIVED_OPTIMIZATION_ARGS + ): + raise ValueError( + f"The following optimization arguments cannot be overridden: {derived_args}." + ) + def generate_sdfg( self, *args: Any, @@ -402,6 +416,11 @@ def _generate_sdfg_without_configuring_dace( sdfg, gpu=on_gpu, constant_symbols=constant_symbols, + unit_strides_kind=( + common.DimensionKind.HORIZONTAL + if self.unstructured_horizontal_has_unit_stride + else None + ), **auto_optimize_args, ) elif on_gpu: @@ -467,8 +486,3 @@ def __call__( code_spec=artifacts.SDFGCodeSpec(), ) return module - - -class DaCeTranslationStepFactory(factory.Factory): - class Meta: - model = DaCeTranslator diff --git a/src/gt4py/next/program_processors/runners/gtfn.py b/src/gt4py/next/program_processors/runners/gtfn.py index 624ca269a3..f73f4e0cd6 100644 --- a/src/gt4py/next/program_processors/runners/gtfn.py +++ b/src/gt4py/next/program_processors/runners/gtfn.py @@ -7,19 +7,19 @@ # SPDX-License-Identifier: BSD-3-Clause import dataclasses +import functools import pathlib -from typing import Any +from collections.abc import Callable +from typing import Any, TypeAlias -import factory import numpy as np import gt4py._core.definitions as core_defs -import gt4py.next.custom_layout_allocators as next_allocators -from gt4py._core import filecache -from gt4py.next import backend, common, config, field_utils, fingerprinting +from gt4py.next import backend, common, field_utils from gt4py.next.embedded import nd_array_field from gt4py.next.instrumentation import metrics -from gt4py.next.otf import artifacts, recipes, workflow +from gt4py.next.iterator import ir as itir +from gt4py.next.otf import artifacts, recipes, stages, workflow from gt4py.next.otf.binding import nanobind from gt4py.next.otf.compilation import cache, compiler from gt4py.next.otf.compilation.build_systems import compiledb @@ -123,91 +123,203 @@ def _make_artifact( ) -class GTFNCompilerFactory(factory.Factory): - class Meta: - model = GTFNCompiler +@dataclasses.dataclass(frozen=True) +class GTFNConfig(backend.ToolchainConfig): + """Settings shared by the steps of a GTFN toolchain, see `ToolchainConfig`.""" + + +#: Step builders: callables creating a step from the toolchain configuration. +GTFNTranslationBuilder: TypeAlias = Callable[[GTFNConfig], stages.TranslationStep] +GTFNBindingsBuilder: TypeAlias = Callable[ + [GTFNConfig], workflow.Workflow[artifacts.ProgramSource, artifacts.ExtensionSource] +] +GTFNCompilationBuilder: TypeAlias = Callable[ + [GTFNConfig], workflow.Workflow[artifacts.ExtensionSource, artifacts.CompilationArtifact] +] + + +def make_gtfn_translation( + cfg: GTFNConfig, + /, + *, + enable_itir_transforms: bool = True, + symbolic_domain_sizes: dict[str, itir.Expr] | None = None, + use_max_domain_range_on_unstructured_shift: bool | None = None, +) -> gtfn_module.GTFNTranslationStep: + """ + Build the GTIR -> C++ translation step. + + Args: + cfg: The toolchain configuration. + enable_itir_transforms: Run the GTIR transformation passes before code + generation. + symbolic_domain_sizes: Symbolic sizes of the domains, by dimension name. + use_max_domain_range_on_unstructured_shift: See `GTFNTranslationStep`. + + Returns: + The translation step, targeting `cfg.device_type`. + """ + return gtfn_module.GTFNTranslationStep( + device_type=cfg.device_type, + enable_itir_transforms=enable_itir_transforms, + symbolic_domain_sizes=symbolic_domain_sizes, + use_max_domain_range_on_unstructured_shift=use_max_domain_range_on_unstructured_shift, + ) -class GTFNCompileWorkflowFactory(factory.Factory): - class Meta: - model = recipes.OTFCompileWorkflow +def make_gtfn_bindings( + cfg: GTFNConfig, / +) -> workflow.Workflow[artifacts.ProgramSource, artifacts.ExtensionSource]: + """Build the step generating the nanobind bindings of the translated program.""" + # `OTFCompileWorkflow` is not parameterized over the code spec, so its + # `bindings` field is typed for `ProgramSource[Any]` while + # `ExtensionGenerator` accepts only C++-like specs. Parameterizing the + # pipeline is the real fix and belongs with the pipeline rework. + return nanobind.ExtensionGenerator( # type: ignore[return-value] # see comment above + unstructured_horizontal_has_unit_stride=cfg.unstructured_horizontal_has_unit_stride + ) - class Params: - device_type: core_defs.DeviceType = core_defs.DeviceType.CPU - cmake_build_type: config.CMakeBuildType = factory.LazyFunction( # type: ignore[assignment] # factory-boy typing not precise enough - lambda: config.CMAKE_BUILD_TYPE - ) - unstructured_horizontal_has_unit_stride: bool = factory.LazyFunction( # type: ignore[assignment] # factory-boy typing not precise enough - lambda: config.UNSTRUCTURED_HORIZONTAL_HAS_UNIT_STRIDE - ) - builder_factory: compiler.BuildSystemProjectGenerator = factory.LazyAttribute( # type: ignore[assignment] # factory-boy typing not precise enough - lambda o: compiledb.CompiledbFactory(cmake_build_type=o.cmake_build_type) - ) - cached_translation = factory.Trait( - translation=factory.LazyAttribute( - lambda o: workflow.CachedStep.persistent( - o.bare_translation, - input_fingerprinter=fingerprinting.strict_fingerprinter, - cache=filecache.FileCache( - cache.get_translation_cache_folder( - cache.get_cache_base_path(config.BUILD_CACHE_LIFETIME), "gtfn" - ) - ), - ) - ), - ) +def make_gtfn_build_system( + cfg: GTFNConfig, + /, + *, + cmake_extra_flags: list[str] | None = None, + renew_compiledb: bool = False, +) -> compiledb.CompiledbFactory: + """ + Build the build-system project generator used by the GTFN compiler. + + Args: + cfg: The toolchain configuration. + cmake_extra_flags: Extra flags passed to CMake. + renew_compiledb: Regenerate the compilation database even if one exists. + + Returns: + A `CompiledbFactory` using `cfg.cmake_build_type`. + """ + return compiledb.CompiledbFactory( + cmake_build_type=cfg.cmake_build_type, + cmake_extra_flags=cmake_extra_flags or [], + renew_compiledb=renew_compiledb, + ) - bare_translation = factory.SubFactory( - gtfn_module.GTFNTranslationStepFactory, - device_type=factory.SelfAttribute("..device_type"), - ) - translation = factory.LazyAttribute(lambda o: o.bare_translation) - bindings: workflow.Workflow[artifacts.ProgramSource, artifacts.ExtensionSource] = ( - factory.LazyAttribute( # type: ignore[assignment] # factory-boy typing not precise enough - lambda o: nanobind.ExtensionGenerator( - unstructured_horizontal_has_unit_stride=o.unstructured_horizontal_has_unit_stride - ) - ) - ) - compilation = factory.SubFactory( - GTFNCompilerFactory, - cache_lifetime=factory.LazyFunction(lambda: config.BUILD_CACHE_LIFETIME), - builder_factory=factory.SelfAttribute("..builder_factory"), - device_type=factory.SelfAttribute("..device_type"), +def make_gtfn_compiler( + cfg: GTFNConfig, + /, + *, + build_system: Callable[[GTFNConfig], compiler.BuildSystemProjectGenerator] = ( + make_gtfn_build_system + ), + force_recompile: bool = False, +) -> GTFNCompiler: + """ + Build the compilation step. + + Args: + cfg: The toolchain configuration. + build_system: Builder of the build-system project generator. + force_recompile: Recompile even if a cached build exists. + + Returns: + The compilation step, targeting `cfg.device_type`. + """ + return GTFNCompiler( + cache_lifetime=cfg.cache_lifetime, + builder_factory=build_system(cfg), + device_type=cfg.device_type, + force_recompile=force_recompile, ) -class GTFNBackendFactory(factory.Factory): - class Meta: - model = backend.Backend - - class Params: - name_device = "cpu" - name_postfix = "" - gpu = factory.Trait( - allocator=next_allocators.StandardGPUFieldBufferAllocator(), - device_type=core_defs.CUPY_DEVICE_TYPE or core_defs.DeviceType.CUDA, - name_device="gpu", - ) - device_type = core_defs.DeviceType.CPU - otf_workflow = factory.SubFactory( - GTFNCompileWorkflowFactory, - cached_translation=True, - device_type=factory.SelfAttribute("..device_type"), +def make_gtfn_compile_workflow( + cfg: GTFNConfig | None = None, + /, + *, + translation: GTFNTranslationBuilder = make_gtfn_translation, + bindings: GTFNBindingsBuilder = make_gtfn_bindings, + compilation: GTFNCompilationBuilder = make_gtfn_compiler, +) -> recipes.OTFCompileWorkflow: + """ + Build the GTFN translation -> bindings -> compilation workflow. + + Every step is created by a step builder that receives `cfg`, so all steps + agree on the settings in it. To customize a step, pass a different + builder: a `functools.partial` of the default one to change a + step-local setting, e.g. + `translation=functools.partial(make_gtfn_translation, enable_itir_transforms=False)`, + or any callable taking the config to replace the step. The translation + step is wrapped in the cache here, after its builder ran, so a custom + translation step is cached like the default one. A custom step builder is + responsible for configuring its step from `cfg`. + + Args: + cfg: The toolchain configuration. Defaults to `GTFNConfig()`. + translation: Builder of the translation step. + bindings: Builder of the bindings step. + compilation: Builder of the compilation step. + + Returns: + The composed compile workflow. + """ + if cfg is None: + cfg = GTFNConfig() + + translation_step = translation(cfg) + compilation_step = compilation(cfg) + + if cfg.cached_translation: + translation_step = cache.persistent_translation_cache( + translation_step, "gtfn", cfg.cache_lifetime ) - name = factory.LazyAttribute(lambda o: f"run_gtfn_{o.name_device}{o.name_postfix}") - executor = factory.LazyAttribute(lambda o: o.otf_workflow) - allocator = next_allocators.StandardCPUFieldBufferAllocator() - transforms = backend.DEFAULT_TRANSFORMS + return recipes.OTFCompileWorkflow( + translation=translation_step, bindings=bindings(cfg), compilation=compilation_step + ) + + +def make_gtfn_toolchain( + cfg: GTFNConfig | None = None, + /, + *, + name_postfix: str = "", + translation: GTFNTranslationBuilder = make_gtfn_translation, + bindings: GTFNBindingsBuilder = make_gtfn_bindings, + compilation: GTFNCompilationBuilder = make_gtfn_compiler, +) -> backend.Backend: + """ + Build a GTFN toolchain. + + Args: + cfg: The toolchain configuration. Defaults to `GTFNConfig()`. + name_postfix: Appended to the toolchain name, which must stay unique. + translation: Builder of the translation step, see + `make_gtfn_compile_workflow`. + bindings: Builder of the bindings step. + compilation: Builder of the compilation step. + + Returns: + The configured toolchain. + """ + if cfg is None: + cfg = GTFNConfig() + + return backend.Backend( + name=f"run_gtfn_{cfg.device_name}{name_postfix}", + executor=make_gtfn_compile_workflow( + cfg, translation=translation, bindings=bindings, compilation=compilation + ), + allocator=cfg.make_allocator(), + transforms=backend.DEFAULT_TRANSFORMS, + ) -run_gtfn = GTFNBackendFactory() +run_gtfn = make_gtfn_toolchain() -run_gtfn_gpu = GTFNBackendFactory(gpu=True) +run_gtfn_gpu = make_gtfn_toolchain(GTFNConfig(gpu=True)) -run_gtfn_no_transforms = GTFNBackendFactory( - otf_workflow__bare_translation__enable_itir_transforms=False +run_gtfn_no_transforms = make_gtfn_toolchain( + name_postfix="_no_transforms", + translation=functools.partial(make_gtfn_translation, enable_itir_transforms=False), ) diff --git a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_temporaries_with_sizes.py b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_temporaries_with_sizes.py index 80963d83ae..8900154140 100644 --- a/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_temporaries_with_sizes.py +++ b/tests/next_tests/integration_tests/feature_tests/ffront_tests/test_temporaries_with_sizes.py @@ -5,14 +5,15 @@ # # Please, refer to the LICENSE file in the root directory. # SPDX-License-Identifier: BSD-3-Clause +import functools + import pytest from numpy import int32 from gt4py import next as gtx -from gt4py.next import backend, common +from gt4py.next import common from gt4py.next.iterator.transforms import apply_common_transforms from gt4py.next.program_processors.runners import gtfn -from gt4py.next import custom_layout_allocators as next_allocators from next_tests.integration_tests import cases from next_tests.integration_tests.cases import ( @@ -31,19 +32,17 @@ # see https://docs.pytest.org/en/latest/how-to/fixtures.html#override-a-fixture-on-a-test-module-level @pytest.fixture def exec_alloc_descriptor(): - return backend.Backend( - name="run_gtfn_with_temporaries_and_sizes", - transforms=backend.DEFAULT_TRANSFORMS, - executor=gtfn.GTFNCompileWorkflowFactory( - translation=gtfn.gtfn_module.GTFNTranslationStepFactory( - symbolic_domain_sizes={ - "Cell": "num_cells", - "Edge": "num_edges", - "Vertex": "num_vertices", - } - ) + return gtfn.make_gtfn_toolchain( + gtfn.GTFNConfig(cached_translation=False), + name_postfix="_with_temporaries_and_sizes", + translation=functools.partial( + gtfn.make_gtfn_translation, + symbolic_domain_sizes={ + "Cell": "num_cells", + "Edge": "num_edges", + "Vertex": "num_vertices", + }, ), - allocator=next_allocators.StandardCPUFieldBufferAllocator(), ) diff --git a/tests/next_tests/unit_tests/otf_tests/test_compiled_program.py b/tests/next_tests/unit_tests/otf_tests/test_compiled_program.py index 9b3f4bc2cb..3d9fcaa834 100644 --- a/tests/next_tests/unit_tests/otf_tests/test_compiled_program.py +++ b/tests/next_tests/unit_tests/otf_tests/test_compiled_program.py @@ -130,7 +130,9 @@ def pirate(program: workflow.ConcreteArtifact): hijacked_program = program return _NoOpArtifact() - hacked_gtfn_backend = gtfn.GTFNBackendFactory(name_postfix="_custom", executor=pirate) + hacked_gtfn_backend = dataclasses.replace( + gtfn.make_gtfn_toolchain(name_postfix="_custom"), executor=pirate + ) testee = testee_prog.with_backend(hacked_gtfn_backend).compile(cond=[True], offset_provider={}) testee( diff --git a/tests/next_tests/unit_tests/program_processor_tests/codegens_tests/gtfn_tests/test_gtfn_module.py b/tests/next_tests/unit_tests/program_processor_tests/codegens_tests/gtfn_tests/test_gtfn_module.py index fc077b3a90..af958b84d9 100644 --- a/tests/next_tests/unit_tests/program_processor_tests/codegens_tests/gtfn_tests/test_gtfn_module.py +++ b/tests/next_tests/unit_tests/program_processor_tests/codegens_tests/gtfn_tests/test_gtfn_module.py @@ -134,12 +134,12 @@ def test_gtfn_file_cache(program_example): data=fencil, args=arguments.CompileTimeArgs.from_concrete(*parameters, **{"offset_provider": {}}), ) - cached_gtfn_translation_step = gtfn.GTFNCompileWorkflowFactory( - cached_translation=True + cached_gtfn_translation_step = gtfn.make_gtfn_compile_workflow( + gtfn.GTFNConfig(cached_translation=True) ).translation - bare_gtfn_translation_step = gtfn.GTFNCompileWorkflowFactory( - cached_translation=False + bare_gtfn_translation_step = gtfn.make_gtfn_compile_workflow( + gtfn.GTFNConfig(cached_translation=False) ).translation cache_key = cached_gtfn_translation_step.cache_key(compilable_program) diff --git a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_backend.py b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_backend.py index 488364b360..13abf117de 100644 --- a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_backend.py +++ b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_backend.py @@ -9,6 +9,7 @@ """Test the bindings stage of the dace backend workflow.""" import dataclasses +import functools import re import unittest.mock as mock from typing import Any @@ -21,6 +22,7 @@ from gt4py._core import definitions as core_defs from gt4py.next import config from gt4py.next.otf import runners, stages +from gt4py.next.otf import workflow as gtx_workflow from gt4py.next.program_processors.runners.dace import transformations as gtx_transformations from gt4py.next.program_processors.runners.dace.transformations import ( auto_optimize as gtx_auto_optimize, @@ -29,6 +31,8 @@ backend as dace_wf_backend, common as dace_wf_common, decoration as dace_wf_decoration, + factory as dace_wf_factory, + translation as dace_wf_translation, ) from next_tests.integration_tests import cases, cases_utils @@ -101,13 +105,18 @@ def mocked_gpu_transformation(*args, **kwargs) -> dace.SDFG: monkeypatch.setattr(gtx_transformations, "gt_auto_optimize", mocked_auto_optimize) monkeypatch.setattr(gtx_transformations, "gt_gpu_transformation", mocked_gpu_transformation) - custom_backend = dace_wf_backend.make_dace_backend( - gpu=on_gpu, - auto_optimize=auto_optimize, - async_sdfg_call=True, - optimization_args=optimization_args, - unstructured_horizontal_has_unit_stride=on_gpu, - use_metrics=True, + custom_backend = dace_wf_backend.make_dace_toolchain( + dace_wf_factory.DaCeConfig( + gpu=on_gpu, + auto_optimize=auto_optimize, + unstructured_horizontal_has_unit_stride=on_gpu, + ), + translation=functools.partial( + dace_wf_factory.make_dace_translator, + auto_optimize_args=optimization_args, + async_sdfg_call=True, + use_metrics=True, + ), ) # The monkeypatched transformation functions exist only in this process, so # compilation must not be offloaded to a worker. @@ -184,14 +193,14 @@ class _RecordingWorkspace: def test_make_backend_accepts_external_workspace_with_external_mode(): workspace = _RecordingWorkspace() - backend = dace_wf_backend.make_dace_backend( - gpu=False, - auto_optimize=True, - async_sdfg_call=False, - optimization_args={ - "transient_memory_mode": gtx_transformations.TransientMemoryMode.EXTERNAL, - }, - external_workspace={core_defs.DeviceType.CPU: workspace}, + backend = dace_wf_backend.make_dace_toolchain( + dace_wf_factory.DaCeConfig(external_workspace={core_defs.DeviceType.CPU: workspace}), + translation=functools.partial( + dace_wf_factory.make_dace_translator, + auto_optimize_args={ + "transient_memory_mode": gtx_transformations.TransientMemoryMode.EXTERNAL, + }, + ), ) assert backend.external_workspace[core_defs.DeviceType.CPU] is workspace @@ -200,11 +209,8 @@ def test_make_backend_accepts_external_workspace_with_external_mode(): def test_make_backend_infers_external_mode_when_workspace_is_provided(): workspace = _RecordingWorkspace() - backend = dace_wf_backend.make_dace_backend( - gpu=False, - auto_optimize=True, - async_sdfg_call=False, - external_workspace={core_defs.DeviceType.CPU: workspace}, + backend = dace_wf_backend.make_dace_toolchain( + dace_wf_factory.DaCeConfig(external_workspace={core_defs.DeviceType.CPU: workspace}) ) assert ( @@ -218,14 +224,14 @@ def test_make_backend_warns_external_workspace_without_external_mode(): workspace = _RecordingWorkspace() with pytest.warns(UserWarning, match="External memory workspace provided"): - backend = dace_wf_backend.make_dace_backend( - gpu=False, - auto_optimize=True, - async_sdfg_call=False, - optimization_args={ - "transient_memory_mode": gtx_transformations.TransientMemoryMode.POOL, - }, - external_workspace={core_defs.DeviceType.CPU: workspace}, + backend = dace_wf_backend.make_dace_toolchain( + dace_wf_factory.DaCeConfig(external_workspace={core_defs.DeviceType.CPU: workspace}), + translation=functools.partial( + dace_wf_factory.make_dace_translator, + auto_optimize_args={ + "transient_memory_mode": gtx_transformations.TransientMemoryMode.POOL, + }, + ), ) # Explicit mode stays as requested by the caller; backend only warns. @@ -236,6 +242,130 @@ def test_make_backend_warns_external_workspace_without_external_mode(): assert backend.external_workspace[core_defs.DeviceType.CPU] is workspace +def test_make_toolchain_derives_workspace_and_memory_mode_from_one_config(): + """The toolchain and its translation step read the workspace from one config.""" + workspace = _RecordingWorkspace() + cfg = dace_wf_factory.DaCeConfig(external_workspace={core_defs.DeviceType.CPU: workspace}) + + backend = dace_wf_backend.make_dace_toolchain( + cfg, translation=functools.partial(dace_wf_factory.make_dace_translator, use_metrics=False) + ) + + translator = backend.executor.translation.step + assert translator.use_metrics is False + assert ( + translator.auto_optimize_args["transient_memory_mode"] + == gtx_transformations.TransientMemoryMode.EXTERNAL + ) + assert backend.external_workspace is cfg.external_workspace + assert backend.executor.compilation.bind_func_name == cfg.bind_func_name + + +def test_make_toolchain_rejects_derived_optimization_args(): + with pytest.raises(ValueError, match="cannot be overridden"): + dace_wf_backend.make_dace_toolchain( + translation=functools.partial( + dace_wf_factory.make_dace_translator, + auto_optimize_args={"unit_strides_kind": None}, + ) + ) + + +def test_make_toolchain_uncached_translation(): + backend = dace_wf_backend.make_dace_toolchain( + dace_wf_factory.DaCeConfig(cached_translation=False) + ) + + assert isinstance(backend.executor.translation, dace_wf_translation.DaCeTranslator) + + +def test_make_dace_backend_is_deprecated(): + workspace = {core_defs.DeviceType.CPU: _RecordingWorkspace()} + with pytest.warns(DeprecationWarning, match="make_dace_toolchain"): + deprecated = dace_wf_backend.make_dace_backend( + gpu=False, + async_sdfg_call=False, + optimization_args={"blocking_size": 10}, + external_workspace=workspace, + unstructured_horizontal_has_unit_stride=True, + use_metrics=False, + use_zero_origin=True, + use_max_domain_range_on_unstructured_shift=True, + ) + toolchain = dace_wf_backend.make_dace_toolchain( + dace_wf_factory.DaCeConfig( + external_workspace=workspace, unstructured_horizontal_has_unit_stride=True + ), + translation=functools.partial( + dace_wf_factory.make_dace_translator, + auto_optimize_args={"blocking_size": 10}, + async_sdfg_call=False, + use_metrics=False, + disable_field_origin_on_program_arguments=True, + use_max_domain_range_on_unstructured_shift=True, + ), + ) + + # The cached steps own distinct `FileCache` objects and the bindings are + # distinct `functools.partial` objects, so compare what they wrap. + assert deprecated.name == toolchain.name + assert deprecated.allocator == toolchain.allocator + assert deprecated.external_workspace is toolchain.external_workspace + assert deprecated.executor.translation.step == toolchain.executor.translation.step + assert deprecated.executor.translation.cache.path == toolchain.executor.translation.cache.path + assert deprecated.executor.bindings.func is toolchain.executor.bindings.func + assert deprecated.executor.bindings.keywords == toolchain.executor.bindings.keywords + assert deprecated.executor.compilation == toolchain.executor.compilation + + +def test_translator_rejects_derived_optimization_args_on_every_route(): + with pytest.raises(ValueError, match="cannot be overridden"): + dace_wf_translation.DaCeTranslator( + device_type=core_defs.DeviceType.CPU, + auto_optimize=False, + auto_optimize_args={"gpu": True}, + async_sdfg_call=False, + unstructured_horizontal_has_unit_stride=False, + use_metrics=False, + ) + translator = dace_wf_backend.run_dace_cpu.executor.translation.step + with pytest.raises(ValueError, match="cannot be overridden"): + dataclasses.replace(translator, auto_optimize_args={"constant_symbols": {}}) + + +def test_step_builders_reach_all_step_settings(): + toolchain = dace_wf_backend.make_dace_toolchain( + translation=functools.partial( + dace_wf_factory.make_dace_translator, disable_itir_transforms=True + ), + compilation=functools.partial( + dace_wf_factory.make_dace_compiler, add_gpu_trace_markers=True + ), + ) + + assert toolchain.executor.translation.step.disable_itir_transforms is True + assert toolchain.executor.compilation.add_gpu_trace_markers is True + + +def test_unused_optimization_args_warning_points_at_the_caller(): + with pytest.warns(UserWarning, match="auto-optimize is disabled") as record: + dace_wf_backend.make_dace_toolchain( + dace_wf_factory.DaCeConfig(auto_optimize=False), + translation=functools.partial( + dace_wf_factory.make_dace_translator, auto_optimize_args={"blocking_size": 10} + ), + ) + + assert record[0].filename == __file__ + + +def test_compile_workflow_without_config_caches_translation(): + workflow = dace_wf_factory.make_dace_compile_workflow() + + assert isinstance(workflow.translation, gtx_workflow.CachedStep) + assert isinstance(workflow.translation.step, dace_wf_translation.DaCeTranslator) + + def _parse_generated_code_from_sdfg(sdfg: dace.SDFG, gpu_api_prefix: str) -> str: # Helper function to ignore the GPU device initialization code in the generated # cuda code, which is not relevant to the test. @@ -289,14 +419,15 @@ def test_transient_memory_mode(device_type, transient_memory_mode, monkeypatch): else None ) - custom_backend = dace_wf_backend.make_dace_backend( - gpu=on_gpu, - auto_optimize=True, - async_sdfg_call=False, - optimization_args={ - "transient_memory_mode": transient_memory_mode, - }, - external_workspace=external_workspace, + custom_backend = dace_wf_backend.make_dace_toolchain( + dace_wf_factory.DaCeConfig(gpu=on_gpu, external_workspace=external_workspace), + translation=functools.partial( + dace_wf_factory.make_dace_translator, + auto_optimize_args={ + "transient_memory_mode": transient_memory_mode, + }, + async_sdfg_call=False, + ), ) @gtx.field_operator diff --git a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_bindings.py b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_bindings.py index e4ec853c23..df65b39c4c 100644 --- a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_bindings.py +++ b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_bindings.py @@ -294,10 +294,12 @@ def testee( ): testee_op(a, b, out=out, domain={IDim: (1, M - 1), JDim: (2, N - 2), KDim: (3, K - 3)}) - backend = dace_runner.make_dace_backend( - gpu=False, - use_metrics=use_metrics, - use_zero_origin=use_zero_origin, + backend = dace_runner.make_dace_toolchain( + translation=functools.partial( + dace_runner.make_dace_translator, + use_metrics=use_metrics, + disable_field_origin_on_program_arguments=use_zero_origin, + ) ) monkeypatch.setattr( dace_workflow.compilation.DaCeCompiler, @@ -348,10 +350,12 @@ def testee_op(a: cases.VField) -> cases.VField: def testee(a: cases.VField, b: cases.VField): testee_op(a, out=b) - backend = dace_runner.make_dace_backend( - gpu=False, - use_metrics=use_metrics, - use_zero_origin=use_zero_origin, + backend = dace_runner.make_dace_toolchain( + translation=functools.partial( + dace_runner.make_dace_translator, + use_metrics=use_metrics, + disable_field_origin_on_program_arguments=use_zero_origin, + ) ) monkeypatch.setattr( dace_workflow.compilation.DaCeCompiler, diff --git a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/test_gtfn.py b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/test_gtfn.py index cb72c80f42..70ef246bce 100644 --- a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/test_gtfn.py +++ b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/test_gtfn.py @@ -19,19 +19,23 @@ other variables are computed at import time based on them. """ +import functools import pathlib import unittest.mock +import pytest + import gt4py._core.definitions as core_defs from gt4py.next import config, custom_layout_allocators from gt4py.next.otf import workflow from gt4py.next.otf.compilation import build_data, cache, compiler, importer +from gt4py.next.program_processors.codegens.gtfn import gtfn_module from gt4py.next.program_processors.runners import gtfn -def test_backend_factory_trait_device(): - cpu_version = gtfn.GTFNBackendFactory(gpu=False) - gpu_version = gtfn.GTFNBackendFactory(gpu=True) +def test_make_gtfn_toolchain_device(): + cpu_version = gtfn.make_gtfn_toolchain(gtfn.GTFNConfig(gpu=False)) + gpu_version = gtfn.make_gtfn_toolchain(gtfn.GTFNConfig(gpu=True)) assert cpu_version.name == "run_gtfn_cpu" assert isinstance(cpu_version.executor.translation, workflow.CachedStep) @@ -52,11 +56,11 @@ def test_backend_factory_trait_device(): ) -def test_backend_factory_build_cache_config(monkeypatch): +def test_make_gtfn_toolchain_build_cache_config(monkeypatch): monkeypatch.setattr(config, "BUILD_CACHE_LIFETIME", config.BuildCacheLifetime.SESSION) - session_version = gtfn.GTFNBackendFactory() + session_version = gtfn.make_gtfn_toolchain() monkeypatch.setattr(config, "BUILD_CACHE_LIFETIME", config.BuildCacheLifetime.PERSISTENT) - persistent_version = gtfn.GTFNBackendFactory() + persistent_version = gtfn.make_gtfn_toolchain() assert session_version.executor.compilation.cache_lifetime is config.BuildCacheLifetime.SESSION assert ( @@ -65,11 +69,11 @@ def test_backend_factory_build_cache_config(monkeypatch): ) -def test_backend_factory_build_type_config(monkeypatch): +def test_make_gtfn_toolchain_build_type_config(monkeypatch): monkeypatch.setattr(config, "CMAKE_BUILD_TYPE", config.CMakeBuildType.RELEASE) - release_version = gtfn.GTFNBackendFactory() + release_version = gtfn.make_gtfn_toolchain() monkeypatch.setattr(config, "CMAKE_BUILD_TYPE", config.CMakeBuildType.MIN_SIZE_REL) - min_size_version = gtfn.GTFNBackendFactory() + min_size_version = gtfn.make_gtfn_toolchain() assert ( release_version.executor.compilation.builder_factory.cmake_build_type @@ -90,9 +94,9 @@ def test_cmake_build_type_changes_build_folder(monkeypatch, tmp_path): land in different cache folders. """ monkeypatch.setattr(config, "CMAKE_BUILD_TYPE", config.CMakeBuildType.RELEASE) - release_version = gtfn.GTFNBackendFactory() + release_version = gtfn.make_gtfn_toolchain() monkeypatch.setattr(config, "CMAKE_BUILD_TYPE", config.CMakeBuildType.DEBUG) - debug_version = gtfn.GTFNBackendFactory() + debug_version = gtfn.make_gtfn_toolchain() release_compiler = release_version.executor.compilation debug_compiler = debug_version.executor.compilation @@ -126,3 +130,77 @@ def fake_get_cache_folder( assert len(build_context_ids) == 2 assert build_context_ids[0] != build_context_ids[1] + + +def test_step_builder_partial_keeps_config_settings(): + """A partial of a default step builder changes a step-local setting only.""" + cfg = gtfn.GTFNConfig(gpu=True, cmake_build_type=config.CMakeBuildType.DEBUG) + + toolchain = gtfn.make_gtfn_toolchain( + cfg, + translation=functools.partial(gtfn.make_gtfn_translation, enable_itir_transforms=False), + compilation=functools.partial( + gtfn.make_gtfn_compiler, + build_system=functools.partial( + gtfn.make_gtfn_build_system, cmake_extra_flags=["-DEXTRA=ON"] + ), + ), + ) + + translation = toolchain.executor.translation + assert isinstance(translation, workflow.CachedStep) + assert translation.step.enable_itir_transforms is False + assert translation.step.device_type is cfg.device_type + compilation = toolchain.executor.compilation + assert compilation.device_type is cfg.device_type + assert compilation.builder_factory.cmake_extra_flags == ["-DEXTRA=ON"] + assert compilation.builder_factory.cmake_build_type is config.CMakeBuildType.DEBUG + + +def test_custom_step_builder_output_is_cached(): + """The cache wraps whatever the translation step builder returns.""" + custom_step = gtfn_module.GTFNTranslationStep( + device_type=core_defs.DeviceType.CPU, use_max_domain_range_on_unstructured_shift=True + ) + + toolchain = gtfn.make_gtfn_toolchain(translation=lambda cfg: custom_step) + + assert isinstance(toolchain.executor.translation, workflow.CachedStep) + assert toolchain.executor.translation.step is custom_step + + +def test_uncached_translation(): + toolchain = gtfn.make_gtfn_toolchain(gtfn.GTFNConfig(cached_translation=False)) + + assert isinstance(toolchain.executor.translation, gtfn_module.GTFNTranslationStep) + + +def test_step_builder_cannot_override_config_setting(): + with pytest.raises(TypeError, match="device_type"): + gtfn.make_gtfn_toolchain( + gtfn.GTFNConfig(gpu=True), + translation=functools.partial( + gtfn.make_gtfn_translation, device_type=core_defs.DeviceType.CPU + ), + ) + + +def test_prebuilt_toolchain_names_are_unique(): + names = [gtfn.run_gtfn.name, gtfn.run_gtfn_gpu.name, gtfn.run_gtfn_no_transforms.name] + + assert gtfn.run_gtfn_no_transforms.name == "run_gtfn_cpu_no_transforms" + assert len(set(names)) == len(names) + + +def test_translation_step_requires_device_type(): + # A step that silently defaulted to the CPU could disagree with the rest of + # a GPU pipeline, so the device has no default. + with pytest.raises(TypeError, match="device_type"): + gtfn_module.GTFNTranslationStep() + + +def test_compile_workflow_without_config_caches_translation(): + workflow_ = gtfn.make_gtfn_compile_workflow() + + assert isinstance(workflow_.translation, workflow.CachedStep) + assert isinstance(workflow_.translation.step, gtfn_module.GTFNTranslationStep) diff --git a/uv.lock b/uv.lock index d024752ae3..ef2d5ca914 100644 --- a/uv.lock +++ b/uv.lock @@ -1278,7 +1278,6 @@ dependencies = [ { name = "dace" }, { name = "deepdiff" }, { name = "devtools" }, - { name = "factory-boy" }, { name = "filelock" }, { name = "frozendict" }, { name = "gridtools-cpp" }, @@ -1357,6 +1356,7 @@ dev = [ { name = "atlas4py" }, { name = "cython" }, { name = "esbonio" }, + { name = "factory-boy" }, { name = "hypothesis" }, { name = "jupytext" }, { name = "matplotlib" }, @@ -1421,6 +1421,7 @@ scripts = [ { name = "typer" }, ] test = [ + { name = "factory-boy" }, { name = "hypothesis" }, { name = "nbmake" }, { name = "nox" }, @@ -1469,7 +1470,6 @@ requires-dist = [ { name = "dace", specifier = "==2.0.0a9" }, { name = "deepdiff", specifier = ">=8.1.0" }, { name = "devtools", specifier = ">=0.6" }, - { name = "factory-boy", specifier = ">=3.3.3" }, { name = "filelock", specifier = ">=3.18.0" }, { name = "frozendict", specifier = ">=2.3" }, { name = "gridtools-cpp", specifier = "==2.*,>=2.3.9" }, @@ -1514,6 +1514,7 @@ dev = [ { name = "atlas4py", specifier = ">=0.41", index = "https://test.pypi.org/simple" }, { name = "cython", specifier = ">=3.0.0" }, { name = "esbonio", specifier = ">=0.16.0" }, + { name = "factory-boy", specifier = ">=3.3.3" }, { name = "hypothesis", specifier = ">=6.0.0" }, { name = "jupytext", specifier = ">=1.14" }, { name = "matplotlib", specifier = ">=3.9.0" }, @@ -1576,6 +1577,7 @@ scripts = [ { name = "typer", specifier = ">=0.16.0" }, ] test = [ + { name = "factory-boy", specifier = ">=3.3.3" }, { name = "hypothesis", specifier = ">=6.0.0" }, { name = "nbmake", specifier = ">=1.4.6" }, { name = "nox", specifier = ">=2025.2.9" },