From a18a1960b3e0340736aa16cf0e5aba5f29a80940 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Wed, 23 Sep 2026 16:27:11 +0200 Subject: [PATCH 1/6] refactor[next]: replace factory-boy factories with plain builders `factory-boy` is a test-data library, but `gt4py.next` used it in production to compose the GTFN and DaCe backends and their compile workflows. Every object it built is already a frozen dataclass, so the `Trait`/`SubFactory`/`SelfAttribute`/`LazyAttribute` machinery added a second construction language that no type checker can see, and its `__`-path overrides failed silently (`run_gtfn_imperative` was the declarative backend under another name for its whole life). Replace the eight factory classes with plain builder functions and move `factory-boy` to the `test` dependency group (the cartesian and eve IR test-data factories keep using it for its intended purpose). The builders follow three rules (ADR 0028): 1. **Shared settings live in one config.** `GTFNConfig` and `DaCeConfig` (both extending `backend.ToolchainConfig`) hold what the steps must agree on: device, build type, cache lifetime, data layout, translation caching, and for DaCe auto-optimize and the external workspace. Defaults are read from `gt4py.next.config` when the config is created. 2. **Every step is created by a step builder that receives the config** and takes only step-local settings as keyword arguments, so a shared setting cannot be set for one step alone. A step is customized by passing a different builder: a `functools.partial` of the default one, or any callable taking the config. 3. **The toolchain builder owns the composition.** It wraps the translation step in the cache after its builder ran, so a customization always lands on the bare step, and checks the device of whatever a custom builder returns (`workflow.check_device_agreement`: check, never mutate). ```python gtfn.make_gtfn_toolchain( gtfn.GTFNConfig(gpu=True), name_postfix="_no_transforms", translation=functools.partial(gtfn.make_gtfn_translation, enable_itir_transforms=False), ) ``` mypy rejects misspelled step-local settings, wrong value types, and attempts to set a shared setting through a step builder; the same mistakes raise `TypeError` when the toolchain is built. API: - `make_gtfn_toolchain`, `make_gtfn_compile_workflow`, `make_gtfn_translation`, `make_gtfn_bindings`, `make_gtfn_compiler`, `make_gtfn_build_system`. - `make_dace_toolchain`, `make_dace_compile_workflow`, `make_dace_translator`, `make_dace_bindings`, `make_dace_compiler`. - `make_dace_backend` keeps its flat keyword signature as a front end over `make_dace_toolchain`, so external callers (icon4py) are unaffected; it builds field-identical toolchains for the same arguments. - `GTFNTranslationStep.device_type` has no default any more: a builder that forgets to pass it fails instead of silently targeting the CPU. Rebase adaptation (2026-09): while this PR was pending, the gtfn imperative backend was removed from the codebase in #2877, resolving #2810. `run_gtfn_imperative` is therefore dropped instead of fixed, together with the two `#2810` xfails; the incident itself is documented in the Context of ADR 0028. Latent bug: `run_gtfn_no_transforms.name` was `run_gtfn_cpu`, colliding with `run_gtfn`. It is now `run_gtfn_cpu_no_transforms`. This rotates no cache -- the build cache keys on the entry-point name plus a fingerprint of the `ExtensionSource`, and the translation-cache directory is keyed on the literal backend family (`gtfn` / `dace`). `Backend.name` reaches only the metrics source key and one error message, so the collision's real cost was two distinct backends sharing one metrics identity. All other pre-built backends are unchanged, verified field-by-field against the previous construction. Removing the factories also removed the 8 `# type: ignore[assignment] # factory-boy typing not precise enough` suppressions in `src/`, which had been masking real typing problems. One remains as a scoped, documented `type: ignore` in `make_gtfn_bindings`: `OTFCompileWorkflow` is not parameterized over the code spec, while `ExtensionGenerator` accepts only C++-like specs. Parameterizing the pipeline belongs with the pipeline rework. Migration for downstream code: - `GTFNBackendFactory(gpu=on_gpu)` -> `make_gtfn_toolchain(GTFNConfig(gpu=on_gpu))` - `DaCeBackendFactory(..., otf_workflow__bare_translation__async_sdfg_call=False)` -> `make_dace_backend(..., async_sdfg_call=False)` See ADR 0028. --- ...028-Plain-Builders-Instead-of-Factories.md | 204 +++++++++++++ docs/development/ADRs/next/README.md | 1 + docs/user/next/advanced/HackTheToolchain.md | 40 ++- docs/user/next/advanced/WorkflowPatterns.md | 36 +-- pyproject.toml | 8 +- src/gt4py/next/backend.py | 50 +++- src/gt4py/next/otf/workflow.py | 36 ++- .../codegens/gtfn/gtfn_module.py | 16 +- .../program_processors/formatters/gtfn.py | 4 +- .../runners/dace/__init__.py | 12 + .../runners/dace/workflow/__init__.py | 5 +- .../runners/dace/workflow/backend.py | 192 +++++------- .../runners/dace/workflow/compilation.py | 6 - .../runners/dace/workflow/factory.py | 238 ++++++++++++--- .../runners/dace/workflow/translation.py | 6 - .../next/program_processors/runners/gtfn.py | 280 +++++++++++++----- .../test_temporaries_with_sizes.py | 27 +- .../otf_tests/test_compiled_program.py | 4 +- .../gtfn_tests/test_gtfn_module.py | 8 +- .../dace_tests/test_dace_backend.py | 31 ++ .../runners_tests/test_gtfn.py | 87 +++++- uv.lock | 6 +- 22 files changed, 958 insertions(+), 339 deletions(-) create mode 100644 docs/development/ADRs/next/0028-Plain-Builders-Instead-of-Factories.md 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..056a68112d --- /dev/null +++ b/docs/development/ADRs/next/0028-Plain-Builders-Instead-of-Factories.md @@ -0,0 +1,204 @@ +--- +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-23 + +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 cannot silently disagree on shared settings, one fewer +runtime dependency, and loud failures where the factories failed silently. + +## Context + +Every object these factories build — `Backend`, `OTFCompileWorkflow`, +`GTFNTranslationStep`, `DaCeTranslator`, the compilers — is already a frozen +dataclass. `factory-boy` added a second, parallel construction language on top: + +- **Untyped.** The declarations are class attributes of a `Params` block, so + `mypy` cannot check them. `src/` carried **8** + `# type: ignore[assignment] # factory-boy typing not precise enough` + suppressions solely to keep the factories quiet. + +- **Silently wrong.** Overrides are `__`-delimited strings resolved at + runtime. When a path does not resolve, nothing happens. This was not + hypothetical: `run_gtfn_imperative` was declared as + + ```python + run_gtfn_imperative = GTFNBackendFactory( + name_postfix="_imperative", + otf_workflow__translation__use_imperative_backend=True, + ) + ``` + + but the `cached_translation` trait replaces `translation` with a + `CachedStep`, so the path never reached the wrapped `GTFNTranslationStep`. + The backend had `use_imperative_backend=False` — it was the declarative + backend under another name, and the `GTFN_CPU_IMPERATIVE` entry of the test + matrix had therefore never exercised imperative code generation. + `run_gtfn_no_transforms` was likewise named `run_gtfn_cpu`, colliding with + `run_gtfn`. (The imperative backend was removed in #2877 while this change + was pending, so `run_gtfn_imperative` is dropped rather than fixed; the + incident remains the motivating evidence.) + +- **A runtime dependency for a test-time concern.** `factory-boy` sat in + `[project] dependencies`, shipped to every user, to compose four 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` dependency group, where the `cartesian` and `eve` 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, `GTFNConfig` and `DaCeConfig`, both 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, for DaCe, the settings that couple + the toolchain to its translation step (auto-optimize, the external + workspace). 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 — the concrete device type from `gpu`, the + allocator, the DaCe transient memory mode — are derived in exactly one + place. + +2. **Every step is created by a step builder that receives the config.** The + step builders (`make_gtfn_translation`, `make_gtfn_bindings`, + `make_gtfn_compiler`, `make_dace_translator`, …) take the config as their + only positional argument and **step-local settings only** as keyword + arguments. A shared setting therefore cannot be set for one step alone. A + step is customized by passing a different builder to `make_*_toolchain` or + `make_*_compile_workflow`: 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: the GTFN build system is a + step builder argument of `make_gtfn_compiler`. + +3. **The toolchain builder owns the composition and checks custom steps.** It + calls the step builders, then wraps the translation step in the cache, so a + customization always lands on the bare step — the path that + `run_gtfn_imperative` never reached. A step builder may be arbitrary user + code that ignores the config, so the builder checks with + `workflow.check_device_agreement` that a step recording a device agrees + with the config. It checks, never mutates. + +```python +gtfn.make_gtfn_toolchain( + gtfn.GTFNConfig(gpu=True), + name_postfix="_no_transforms", + translation=functools.partial(gtfn.make_gtfn_translation, enable_itir_transforms=False), +) +``` + +`make_dace_backend` keeps its flat keyword signature as a front end over +`make_dace_toolchain`, so existing callers are unaffected. + +## Consequences + +- Composition is ordinary, statically checked Python. `mypy` rejects a + misspelled step-local setting, a value of the wrong type, and an attempt to + set a shared setting through a step builder + (`partial(make_gtfn_translation, device_type=...)`); the same mistakes raise + `TypeError` when the toolchain is built. The 8 factory-related + `type: ignore` suppressions are gone. One remains, scoped and documented, in + `make_gtfn_bindings`: `OTFCompileWorkflow` is not parameterized over the + code spec, while `ExtensionGenerator` accepts only C++-like specs. +- Default and partially customized steps agree on the shared settings by + construction, because they read them from the same config. A fully custom + step builder can still ignore the config; the device check catches that + case for the device only, and only for steps that record it as + `device_type`. Invariants checked by the assembled pipeline itself, which + would also cover `dataclasses.replace` on a built toolchain, are left to the + pipeline rework. +- Step fields that must agree with other steps get no default, so a builder + that forgets to pass one fails instead of silently using the default: + `GTFNTranslationStep.device_type` no longer defaults to the CPU. +- A new shared setting is one config field, read where it is needed, instead + of a keyword argument threaded through every builder layer. +- The price is 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. +- Step builders run at build time and are not stored, so a `lambda` step + builder does not make the toolchain unpicklable (offloading compilation to + worker processes needs a picklable executor). +- One config means one default: a compile workflow built on its own now + caches its translation step, like the toolchains always did. + `GTFNConfig(cached_translation=False)` opts out. +- **`run_gtfn_no_transforms` is renamed** from `run_gtfn_cpu` to + `run_gtfn_cpu_no_transforms`, removing the collision with `run_gtfn`. No + cache is affected: the build cache keys on the entry-point name plus a + fingerprint of the `ExtensionSource`, and the translation-cache directory + is keyed on the literal backend family (`gtfn` / `dace`). `Backend.name` + reaches only the metrics source key and one error message, so what the + collision actually cost was two distinct backends sharing one metrics + identity. +- All other pre-built toolchains are unchanged, verified field by field + against the previous construction, and `make_dace_backend` builds + field-identical toolchains for the same arguments. +- Downstream code migrates as `GTFNBackendFactory(gpu=on_gpu)` → + `make_gtfn_toolchain(GTFNConfig(gpu=on_gpu))`, and + `DaCeBackendFactory(..., otf_workflow__bare_translation__async_sdfg_call=False)` + → `make_dace_backend(..., async_sdfg_call=False)`. + +## Alternatives Considered + +### Inject pre-built steps, used verbatim + +A first version of this change had builders take shared settings as keyword +arguments and accept a pre-built step, used verbatim and checked for device +agreement. Changing one setting of an inner step then meant building the +whole step and repeating the shared settings in it: +`make_gtfn_backend(gpu=True, translation=GTFNTranslationStep(enable_itir_transforms=False))` +raised, because the injected step defaulted to the CPU, and the caller had to +re-derive the GPU device type (`CUPY_DEVICE_TYPE or CUDA`) the builder already +knew. Only the device was checked, so other shared settings could still +disagree silently, and each shared setting had to be threaded through every +builder layer by hand. + +### 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 and is statically checked, and +leaving the shared settings out of the `TypedDict`s keeps them in sync. But +the `TypedDict`s mirror the step fields and drift from them, they cannot +replace a step (which needs a second, instance-injection mechanism with the +problems above), and the builders still forward every shared setting by hand +through each layer, where a forgotten forward silently falls back to the +step's default. + +### 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 exactly the +failure mode that motivated this ADR — the `run_gtfn_imperative` bug is what a +silent no-op looks like after a year. Checking is the same amount of +introspection with the opposite failure mode. + +### Edit the built toolchain + +`dataclasses.replace` on a built toolchain keeps working, but the caller must +know the nesting of wrappers (`executor.translation.step`) — the path the +`run_gtfn_imperative` override 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..ac149e357e 100644 --- a/docs/user/next/advanced/HackTheToolchain.md +++ b/docs/user/next/advanced/HackTheToolchain.md @@ -46,25 +46,43 @@ 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, and a step that +records a device other than the configured one is rejected. + +```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/workflow.py b/src/gt4py/next/otf/workflow.py index fa57a9ac7f..b30845e112 100644 --- a/src/gt4py/next/otf/workflow.py +++ b/src/gt4py/next/otf/workflow.py @@ -15,7 +15,7 @@ import typing from typing import Any, Callable, Generic, Protocol, Self, TypeVar -from gt4py._core import filecache +from gt4py._core import definitions as core_defs, filecache from gt4py.eve.xtyping import OpaqueMutableMapping from gt4py.next import config, fingerprinting, utils @@ -359,3 +359,37 @@ def __call__(self, inp: StartT) -> EndT: def cache_key(self, inp: StartT) -> str: return self.step_fingerprinter((self._step_fingerprint, self.input_fingerprinter(inp))) + + +@typing.runtime_checkable +class DeviceConfigurable(Protocol): + """A step that records the device it was configured for.""" + + device_type: core_defs.DeviceType + + +def check_device_agreement(step: Any, device_type: core_defs.DeviceType, what: str) -> None: + """ + Raise if a step is configured for a different device than its pipeline. + + Toolchain builders create every step from one configuration, but a step + builder can be replaced by arbitrary user code, which may ignore the + configured device. Without this check a mismatch would silently produce a + pipeline whose steps disagree about the target device, which surfaces much + later as a confusing compilation or runtime failure. The check never + modifies the step. + + Args: + step: The step to check. Steps that do not record a device are accepted. + device_type: The device the surrounding pipeline is built for. + what: Name of the step, used in the error message. + + Raises: + ValueError: If `step` records a device other than `device_type`. + """ + if isinstance(step, DeviceConfigurable) and step.device_type is not device_type: + raise ValueError( + f"The {what} is configured for device '{step.device_type.name}', but the" + f" toolchain is being built for '{device_type.name}'. A custom step builder" + " must configure the step with the 'device_type' of the config it receives." + ) 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..faa842d122 100644 --- a/src/gt4py/next/program_processors/runners/dace/__init__.py +++ b/src/gt4py/next/program_processors/runners/dace/__init__.py @@ -10,16 +10,28 @@ 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 ( + DaCeConfig, + make_dace_bindings, + make_dace_compiler, + make_dace_translator, +) __all__ = [ + "DaCeConfig", "get_sdfg_args", "make_dace_backend", + "make_dace_bindings", + "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..7ccbda52c5 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 warnings -from typing import Any, Final +import functools +from collections.abc import Callable +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.otf import artifacts -from gt4py.next.program_processors.runners.dace import transformations as gtx_transformations +from gt4py.next import backend, config +from gt4py.next.otf import artifacts, stages, workflow from gt4py.next.program_processors.runners.dace.workflow import ( common as gtx_wfdcommon, decoration as gtx_wfddecoration, @@ -40,41 +36,53 @@ 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: Callable[ + [gtx_wfdfactory.DaCeConfig], stages.TranslationStep + ] = gtx_wfdfactory.make_dace_translator, + bindings: Callable[ + [gtx_wfdfactory.DaCeConfig], + workflow.Workflow[artifacts.ProgramSource, artifacts.ExtensionSource], + ] = gtx_wfdfactory.make_dace_bindings, + compilation: Callable[ + [gtx_wfdfactory.DaCeConfig], + workflow.Workflow[artifacts.ExtensionSource, artifacts.CompilationArtifact], + ] = 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. + + Raises: + ValueError: If a step builder returns a step configured for a device + other than `cfg.device_type`. + """ + 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 +95,13 @@ 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. + A flat-keyword front end for `make_dace_toolchain`, kept for existing + callers: it builds the `DaCeConfig` and the translation step builder from + its arguments. + Args: gpu: Enable GPU transformations and code generation. auto_optimize: Enable the SDFG auto-optimize pipeline. @@ -106,85 +118,39 @@ 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 overriden, 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`. + """ + 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, + optimization_args=optimization_args, + async_sdfg_call=async_sdfg_call, + use_metrics=use_metrics, + use_zero_origin=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..33181b84a2 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,209 @@ 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 -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._core import filecache +from gt4py.next import backend as next_backend, common, fingerprinting +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 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.translation import ( - DaCeTranslationStepFactory, +from gt4py.next.program_processors.runners.dace.workflow.compilation import DaCeCompiler +from gt4py.next.program_processors.runners.dace.workflow.translation import DaCeTranslator + + +#: The parameters of `gt_auto_optimize()` that are derived from the toolchain +#: configuration, and therefore cannot be customized. +_DERIVED_OPTIMIZATION_ARGS: Final[frozenset[str]] = frozenset( + {"gpu", "constant_symbols", "unit_strides_kind"} ) +#: Warnings about the builder arguments point at the first caller outside GT4Py. +_GT4PY_SOURCE_PREFIX: Final[str] = str(pathlib.Path(gt4py.__file__).parent) -_GT_DACE_BINDING_FUNCTION_NAME: Final[str] = "update_sdfg_args" +@dataclasses.dataclass(frozen=True) +class DaCeConfig(next_backend.ToolchainConfig): + """Settings shared by the steps of a DaCe toolchain, see `ToolchainConfig`.""" -class DaCeWorkflowFactory(factory.Factory): - class Meta: - model = recipes.OTFCompileWorkflow + #: 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. + external_workspace: gtx_wfdcommon.ExternalWorkspace | None = None - 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 - ) + #: 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" - 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" - ) - ), - ) - ), - ) - bare_translation = factory.SubFactory( - DaCeTranslationStepFactory, - device_type=factory.SelfAttribute("..device_type"), - auto_optimize=factory.SelfAttribute("..auto_optimize"), +def make_dace_translator( + cfg: DaCeConfig, + /, + *, + optimization_args: dict[str, Any] | None = None, + async_sdfg_call: bool = True, + use_metrics: bool = True, + use_zero_origin: bool = False, + use_max_domain_range_on_unstructured_shift: bool | None = None, +) -> DaCeTranslator: + """ + Build the GTIR -> SDFG translation step. + + Args: + cfg: The toolchain configuration. + optimization_args: Configuration for the SDFG auto-optimize pipeline, see + `gt_auto_optimize()`. The parameters derived from `cfg` cannot be + set here. + 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. + use_zero_origin: 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 `optimization_args` sets a parameter derived from `cfg`, + or requests the `EXTERNAL` transient memory mode without an + external workspace in `cfg`. + """ + if optimization_args is None: + optimization_args = {} + elif optimization_args and not cfg.auto_optimize: + warnings.warn( + "Optimizations args given, but auto-optimize is disabled.", + skip_file_prefixes=(_GT4PY_SOURCE_PREFIX,), ) + elif intersect_args := optimization_args.keys() & _DERIVED_OPTIMIZATION_ARGS: + raise ValueError( + f"The following optimization arguments cannot be overriden: {intersect_args}." + ) + + optimization_args = optimization_args | { + "unit_strides_kind": common.DimensionKind.HORIZONTAL + if cfg.unstructured_horizontal_has_unit_stride + else None + } - 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_field_origin_on_program_arguments=use_zero_origin, + 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, /) -> DaCeCompiler: + """Build the compilation step, targeting `cfg.device_type`.""" + 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, ) - 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"), + + +def make_dace_compile_workflow( + cfg: DaCeConfig | None = None, + /, + *, + translation: Callable[[DaCeConfig], stages.TranslationStep] = make_dace_translator, + bindings: Callable[ + [DaCeConfig], workflow.Workflow[artifacts.ProgramSource, artifacts.ExtensionSource] + ] = make_dace_bindings, + compilation: Callable[ + [DaCeConfig], workflow.Workflow[artifacts.ExtensionSource, artifacts.CompilationArtifact] + ] = 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. + + 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. + + Raises: + ValueError: If a step builder returns a step configured for a device + other than `cfg.device_type`. + """ + if cfg is None: + cfg = DaCeConfig() + + translation_step = translation(cfg) + workflow.check_device_agreement(translation_step, cfg.device_type, "DaCe translation step") + compilation_step = compilation(cfg) + workflow.check_device_agreement(compilation_step, cfg.device_type, "DaCe compilation step") + + if cfg.cached_translation: + translation_step = workflow.CachedStep[ + stages.CompilableProgramDef, artifacts.ProgramSource, str + ].persistent( + translation_step, + input_fingerprinter=fingerprinting.strict_fingerprinter, + cache=filecache.FileCache( + cache.get_translation_cache_folder( + cache.get_cache_base_path(cfg.cache_lifetime), "dace" + ) + ), + ) + + 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..e45f3840b3 100644 --- a/src/gt4py/next/program_processors/runners/dace/workflow/translation.py +++ b/src/gt4py/next/program_processors/runners/dace/workflow/translation.py @@ -12,7 +12,6 @@ from typing import Any, Optional import dace -import factory from gt4py._core import definitions as core_defs from gt4py.next import common @@ -467,8 +466,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..b7ed317642 100644 --- a/src/gt4py/next/program_processors/runners/gtfn.py +++ b/src/gt4py/next/program_processors/runners/gtfn.py @@ -7,19 +7,20 @@ # SPDX-License-Identifier: BSD-3-Clause import dataclasses +import functools import pathlib +from collections.abc import Callable from typing import Any -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, fingerprinting 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 +124,218 @@ 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`.""" + + +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: Callable[[GTFNConfig], stages.TranslationStep] = make_gtfn_translation, + bindings: Callable[ + [GTFNConfig], workflow.Workflow[artifacts.ProgramSource, artifacts.ExtensionSource] + ] = make_gtfn_bindings, + compilation: Callable[ + [GTFNConfig], workflow.Workflow[artifacts.ExtensionSource, artifacts.CompilationArtifact] + ] = 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. + + 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. + + Raises: + ValueError: If a step builder returns a step configured for a device + other than `cfg.device_type`. + """ + if cfg is None: + cfg = GTFNConfig() + + translation_step = translation(cfg) + workflow.check_device_agreement(translation_step, cfg.device_type, "GTFN translation step") + compilation_step = compilation(cfg) + workflow.check_device_agreement(compilation_step, cfg.device_type, "GTFN compilation step") + + if cfg.cached_translation: + translation_step = workflow.CachedStep[ + stages.CompilableProgramDef, artifacts.ProgramSource, str + ].persistent( + translation_step, + input_fingerprinter=fingerprinting.strict_fingerprinter, + cache=filecache.FileCache( + cache.get_translation_cache_folder( + cache.get_cache_base_path(cfg.cache_lifetime), "gtfn" + ) + ), ) - 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: Callable[[GTFNConfig], stages.TranslationStep] = make_gtfn_translation, + bindings: Callable[ + [GTFNConfig], workflow.Workflow[artifacts.ProgramSource, artifacts.ExtensionSource] + ] = make_gtfn_bindings, + compilation: Callable[ + [GTFNConfig], workflow.Workflow[artifacts.ExtensionSource, artifacts.CompilationArtifact] + ] = 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. + + Raises: + ValueError: If a step builder returns a step configured for a device + other than `cfg.device_type`. + """ + 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..e9717efbd8 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 @@ -29,6 +30,7 @@ backend as dace_wf_backend, common as dace_wf_common, decoration as dace_wf_decoration, + factory as dace_wf_factory, ) from next_tests.integration_tests import cases, cases_utils @@ -236,6 +238,35 @@ 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 overriden"): + dace_wf_backend.make_dace_toolchain( + translation=functools.partial( + dace_wf_factory.make_dace_translator, + optimization_args={"unit_strides_kind": None}, + ) + ) + + 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. 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..d24abc1cb6 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,64 @@ 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_ignoring_config_device_raises(): + def cpu_only_translation(cfg: gtfn.GTFNConfig) -> gtfn_module.GTFNTranslationStep: + return gtfn_module.GTFNTranslationStep(device_type=core_defs.DeviceType.CPU) + + with pytest.raises(ValueError, match="toolchain is being built for 'CUDA'"): + gtfn.make_gtfn_toolchain(gtfn.GTFNConfig(gpu=True), translation=cpu_only_translation) + + +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 + ), + ) 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" }, From e3940c75fb7e4345e46af64423dc685332d68306 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Thu, 24 Sep 2026 10:35:28 +0200 Subject: [PATCH 2/6] refactor[next]: re-export make_dace_compile_workflow and test the DaCe device check Address review: the dace package facade exposed every new builder except make_dace_compile_workflow, and only the GTFN side of the device-agreement check of custom step builders was tested. Also pin the corrected run_gtfn_no_transforms name. --- .../runners/dace/__init__.py | 2 ++ .../dace_tests/test_dace_backend.py | 22 +++++++++++++++++++ .../runners_tests/test_gtfn.py | 7 ++++++ 3 files changed, 31 insertions(+) diff --git a/src/gt4py/next/program_processors/runners/dace/__init__.py b/src/gt4py/next/program_processors/runners/dace/__init__.py index faa842d122..c8e96ad515 100644 --- a/src/gt4py/next/program_processors/runners/dace/__init__.py +++ b/src/gt4py/next/program_processors/runners/dace/__init__.py @@ -19,6 +19,7 @@ from gt4py.next.program_processors.runners.dace.workflow.factory import ( DaCeConfig, make_dace_bindings, + make_dace_compile_workflow, make_dace_compiler, make_dace_translator, ) @@ -29,6 +30,7 @@ "get_sdfg_args", "make_dace_backend", "make_dace_bindings", + "make_dace_compile_workflow", "make_dace_compiler", "make_dace_toolchain", "make_dace_translator", 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 e9717efbd8..ba8a32daa7 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 @@ -267,6 +267,28 @@ def test_make_toolchain_rejects_derived_optimization_args(): ) +@pytest.mark.parametrize( + "step_builders", + [ + # Each builder ignores the GPU config it receives and targets the CPU. + { + "translation": lambda cfg: dace_wf_factory.make_dace_translator( + dace_wf_factory.DaCeConfig() + ) + }, + { + "compilation": lambda cfg: dace_wf_factory.make_dace_compiler( + dace_wf_factory.DaCeConfig() + ) + }, + ], + ids=["translation", "compilation"], +) +def test_make_toolchain_rejects_step_builder_ignoring_config_device(step_builders): + with pytest.raises(ValueError, match="toolchain is being built for"): + dace_wf_backend.make_dace_toolchain(dace_wf_factory.DaCeConfig(gpu=True), **step_builders) + + 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. 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 d24abc1cb6..c1abc3fdd1 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 @@ -191,3 +191,10 @@ def test_step_builder_cannot_override_config_setting(): 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) From b4f19bb35a86e6c7775b571e4a41059a75f9dd26 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Fri, 25 Sep 2026 08:45:36 +0200 Subject: [PATCH 3/6] test[next]: cover the DaCe translation-cache opt-out Address review: only the GTFN side of 'cached_translation=False' was tested. --- .../runners_tests/dace_tests/test_dace_backend.py | 9 +++++++++ 1 file changed, 9 insertions(+) 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 ba8a32daa7..2441e21f26 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 @@ -31,6 +31,7 @@ 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 @@ -289,6 +290,14 @@ def test_make_toolchain_rejects_step_builder_ignoring_config_device(step_builder dace_wf_backend.make_dace_toolchain(dace_wf_factory.DaCeConfig(gpu=True), **step_builders) +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 _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. From 0e472ed75deafed785fb22396f83c37ef58d6e47 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Wed, 30 Sep 2026 08:11:39 +0200 Subject: [PATCH 4/6] refactor[next]: address review of the plain builders - Remove the builder-level device check (`check_device_agreement` and the `DeviceConfigurable` protocol): default and partially customized steps are in sync by construction, and configuring a fully custom step builder consistently is the caller's responsibility. - Deprecate `make_dace_backend` in favor of `make_dace_toolchain`, and move the internal callers to the new builder. - Trim ADR 0028 to the decision: statements about specific code move to the PR description, the `TypedDict` alternative is described accurately, and the structural costs raised in review are recorded as consequences. --- ...028-Plain-Builders-Instead-of-Factories.md | 225 +++++++----------- docs/user/next/advanced/HackTheToolchain.md | 5 +- src/gt4py/next/otf/workflow.py | 36 +-- .../runners/dace/workflow/backend.py | 16 +- .../runners/dace/workflow/factory.py | 9 +- .../next/program_processors/runners/gtfn.py | 13 +- .../dace_tests/test_dace_backend.py | 105 ++++---- .../dace_tests/test_dace_bindings.py | 20 +- .../runners_tests/test_gtfn.py | 8 - 9 files changed, 166 insertions(+), 271 deletions(-) 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 index 056a68112d..860449b917 100644 --- a/docs/development/ADRs/next/0028-Plain-Builders-Instead-of-Factories.md +++ b/docs/development/ADRs/next/0028-Plain-Builders-Instead-of-Factories.md @@ -7,7 +7,7 @@ tags: [backend, otf, toolchain, workflows, dependencies] - **Status**: valid - **Authors**: Enrique González Paredes (@egparedes) - **Created**: 2026-08-20 -- **Updated**: 2026-09-23 +- **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 @@ -16,43 +16,25 @@ 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 cannot silently disagree on shared settings, one fewer -runtime dependency, and loud failures where the factories failed silently. +composition, steps that agree on shared settings by construction, and one +fewer runtime dependency. ## Context -Every object these factories build — `Backend`, `OTFCompileWorkflow`, -`GTFNTranslationStep`, `DaCeTranslator`, the compilers — is already a frozen -dataclass. `factory-boy` added a second, parallel construction language on top: +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. `src/` carried **8** - `# type: ignore[assignment] # factory-boy typing not precise enough` - suppressions solely to keep the factories quiet. - + `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: `run_gtfn_imperative` was declared as - - ```python - run_gtfn_imperative = GTFNBackendFactory( - name_postfix="_imperative", - otf_workflow__translation__use_imperative_backend=True, - ) - ``` - - but the `cached_translation` trait replaces `translation` with a - `CachedStep`, so the path never reached the wrapped `GTFNTranslationStep`. - The backend had `use_imperative_backend=False` — it was the declarative - backend under another name, and the `GTFN_CPU_IMPERATIVE` entry of the test - matrix had therefore never exercised imperative code generation. - `run_gtfn_no_transforms` was likewise named `run_gtfn_cpu`, colliding with - `run_gtfn`. (The imperative backend was removed in #2877 while this change - was pending, so `run_gtfn_imperative` is dropped rather than fixed; the - incident remains the motivating evidence.) - -- **A runtime dependency for a test-time concern.** `factory-boy` sat in - `[project] dependencies`, shipped to every user, to compose four backends. + 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 @@ -65,140 +47,111 @@ the first, `SubFactory` plus `__` paths for the second. ## Decision Factory classes are replaced by **plain builder functions**; `factory-boy` -moves to the `test` dependency group, where the `cartesian` and `eve` IR -test-data factories keep using it for what it is designed for. The builders -follow three rules. +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, `GTFNConfig` and `DaCeConfig`, both 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, for DaCe, the settings that couple - the toolchain to its translation step (auto-optimize, the external - workspace). 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 — the concrete device type from `gpu`, the - allocator, the DaCe transient memory mode — are derived in exactly one - place. - -2. **Every step is created by a step builder that receives the config.** The - step builders (`make_gtfn_translation`, `make_gtfn_bindings`, - `make_gtfn_compiler`, `make_dace_translator`, …) take the config as their - only positional argument and **step-local settings only** as keyword - arguments. A shared setting therefore cannot be set for one step alone. A - step is customized by passing a different builder to `make_*_toolchain` or - `make_*_compile_workflow`: 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: the GTFN build system is a - step builder argument of `make_gtfn_compiler`. - -3. **The toolchain builder owns the composition and checks custom steps.** It - calls the step builders, then wraps the translation step in the cache, so a - customization always lands on the bare step — the path that - `run_gtfn_imperative` never reached. A step builder may be arbitrary user - code that ignores the config, so the builder checks with - `workflow.check_device_agreement` that a step recording a device agrees - with the config. It checks, never mutates. + 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 -gtfn.make_gtfn_toolchain( - gtfn.GTFNConfig(gpu=True), +make_gtfn_toolchain( + GTFNConfig(gpu=True), name_postfix="_no_transforms", - translation=functools.partial(gtfn.make_gtfn_translation, enable_itir_transforms=False), + translation=functools.partial(make_gtfn_translation, enable_itir_transforms=False), ) ``` -`make_dace_backend` keeps its flat keyword signature as a front end over -`make_dace_toolchain`, so existing callers are unaffected. +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. `mypy` rejects a - misspelled step-local setting, a value of the wrong type, and an attempt to - set a shared setting through a step builder - (`partial(make_gtfn_translation, device_type=...)`); the same mistakes raise - `TypeError` when the toolchain is built. The 8 factory-related - `type: ignore` suppressions are gone. One remains, scoped and documented, in - `make_gtfn_bindings`: `OTFCompileWorkflow` is not parameterized over the - code spec, while `ExtensionGenerator` accepts only C++-like specs. +- 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. A fully custom - step builder can still ignore the config; the device check catches that - case for the device only, and only for steps that record it as - `device_type`. Invariants checked by the assembled pipeline itself, which - would also cover `dataclasses.replace` on a built toolchain, are left to the - pipeline rework. + construction, because they read them from the same config. Steps from fully + custom step builders are not checked. - Step fields that must agree with other steps get no default, so a builder - that forgets to pass one fails instead of silently using the default: - `GTFNTranslationStep.device_type` no longer defaults to the CPU. + that forgets to pass one fails instead of silently using the default. - A new shared setting is one config field, read where it is needed, instead of a keyword argument threaded through every builder layer. -- The price is 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. - Step builders run at build time and are not stored, so a `lambda` step - builder does not make the toolchain unpicklable (offloading compilation to - worker processes needs a picklable executor). -- One config means one default: a compile workflow built on its own now - caches its translation step, like the toolchains always did. - `GTFNConfig(cached_translation=False)` opts out. -- **`run_gtfn_no_transforms` is renamed** from `run_gtfn_cpu` to - `run_gtfn_cpu_no_transforms`, removing the collision with `run_gtfn`. No - cache is affected: the build cache keys on the entry-point name plus a - fingerprint of the `ExtensionSource`, and the translation-cache directory - is keyed on the literal backend family (`gtfn` / `dace`). `Backend.name` - reaches only the metrics source key and one error message, so what the - collision actually cost was two distinct backends sharing one metrics - identity. -- All other pre-built toolchains are unchanged, verified field by field - against the previous construction, and `make_dace_backend` builds - field-identical toolchains for the same arguments. -- Downstream code migrates as `GTFNBackendFactory(gpu=on_gpu)` → - `make_gtfn_toolchain(GTFNConfig(gpu=on_gpu))`, and - `DaCeBackendFactory(..., otf_workflow__bare_translation__async_sdfg_call=False)` - → `make_dace_backend(..., async_sdfg_call=False)`. + 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 -### Inject pre-built steps, used verbatim - -A first version of this change had builders take shared settings as keyword -arguments and accept a pre-built step, used verbatim and checked for device -agreement. Changing one setting of an inner step then meant building the -whole step and repeating the shared settings in it: -`make_gtfn_backend(gpu=True, translation=GTFNTranslationStep(enable_itir_transforms=False))` -raised, because the injected step defaulted to the CPU, and the caller had to -re-derive the GPU device type (`CUPY_DEVICE_TYPE or CUDA`) the builder already -knew. Only the device was checked, so other shared settings could still -disagree silently, and each shared setting had to be threaded through every -builder layer by hand. - ### 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 and is statically checked, and -leaving the shared settings out of the `TypedDict`s keeps them in sync. But -the `TypedDict`s mirror the step fields and drift from them, they cannot -replace a step (which needs a second, instance-injection mechanism with the -problems above), and the builders still forward every shared setting by hand -through each layer, where a forgotten forward silently falls back to the -step's default. +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 exactly the -failure mode that motivated this ADR — the `run_gtfn_imperative` bug is what a -silent no-op looks like after a year. Checking is the same amount of -introspection with the opposite failure mode. +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 (`executor.translation.step`) — the path the -`run_gtfn_imperative` override never reached — and the values the builder -derived (cache folders, name, allocator) are not recomputed. +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/user/next/advanced/HackTheToolchain.md b/docs/user/next/advanced/HackTheToolchain.md index ac149e357e..773070130b 100644 --- a/docs/user/next/advanced/HackTheToolchain.md +++ b/docs/user/next/advanced/HackTheToolchain.md @@ -68,8 +68,9 @@ debug_gpu_no_transforms = gtfn.make_gtfn_toolchain( ``` To replace a step, pass any callable that takes the configuration and returns -the step. It is still wrapped in the translation cache, and a step that -records a device other than the configured one is rejected. +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: ... diff --git a/src/gt4py/next/otf/workflow.py b/src/gt4py/next/otf/workflow.py index b30845e112..fa57a9ac7f 100644 --- a/src/gt4py/next/otf/workflow.py +++ b/src/gt4py/next/otf/workflow.py @@ -15,7 +15,7 @@ import typing from typing import Any, Callable, Generic, Protocol, Self, TypeVar -from gt4py._core import definitions as core_defs, filecache +from gt4py._core import filecache from gt4py.eve.xtyping import OpaqueMutableMapping from gt4py.next import config, fingerprinting, utils @@ -359,37 +359,3 @@ def __call__(self, inp: StartT) -> EndT: def cache_key(self, inp: StartT) -> str: return self.step_fingerprinter((self._step_fingerprint, self.input_fingerprinter(inp))) - - -@typing.runtime_checkable -class DeviceConfigurable(Protocol): - """A step that records the device it was configured for.""" - - device_type: core_defs.DeviceType - - -def check_device_agreement(step: Any, device_type: core_defs.DeviceType, what: str) -> None: - """ - Raise if a step is configured for a different device than its pipeline. - - Toolchain builders create every step from one configuration, but a step - builder can be replaced by arbitrary user code, which may ignore the - configured device. Without this check a mismatch would silently produce a - pipeline whose steps disagree about the target device, which surfaces much - later as a confusing compilation or runtime failure. The check never - modifies the step. - - Args: - step: The step to check. Steps that do not record a device are accepted. - device_type: The device the surrounding pipeline is built for. - what: Name of the step, used in the error message. - - Raises: - ValueError: If `step` records a device other than `device_type`. - """ - if isinstance(step, DeviceConfigurable) and step.device_type is not device_type: - raise ValueError( - f"The {what} is configured for device '{step.device_type.name}', but the" - f" toolchain is being built for '{device_type.name}'. A custom step builder" - " must configure the step with the 'device_type' of the config it receives." - ) 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 7ccbda52c5..55a8d9b5d5 100644 --- a/src/gt4py/next/program_processors/runners/dace/workflow/backend.py +++ b/src/gt4py/next/program_processors/runners/dace/workflow/backend.py @@ -10,6 +10,7 @@ import dataclasses import functools +import warnings from collections.abc import Callable from typing import Any @@ -66,10 +67,6 @@ def make_dace_toolchain( Returns: The configured toolchain. - - Raises: - ValueError: If a step builder returns a step configured for a device - other than `cfg.device_type`. """ if cfg is None: cfg = gtx_wfdfactory.DaCeConfig() @@ -98,8 +95,9 @@ def make_dace_backend( ) -> DaCeBackend: """Customize the dace backend with the given configuration parameters. - A flat-keyword front end for `make_dace_toolchain`, kept for existing - callers: it builds the `DaCeConfig` and the translation step builder from + 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: @@ -131,6 +129,12 @@ def make_dace_backend( 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, 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 33181b84a2..ecb42a47d5 100644 --- a/src/gt4py/next/program_processors/runners/dace/workflow/factory.py +++ b/src/gt4py/next/program_processors/runners/dace/workflow/factory.py @@ -175,7 +175,8 @@ def make_dace_compile_workflow( 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. + 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()`. @@ -185,18 +186,12 @@ def make_dace_compile_workflow( Returns: The composed compile workflow. - - Raises: - ValueError: If a step builder returns a step configured for a device - other than `cfg.device_type`. """ if cfg is None: cfg = DaCeConfig() translation_step = translation(cfg) - workflow.check_device_agreement(translation_step, cfg.device_type, "DaCe translation step") compilation_step = compilation(cfg) - workflow.check_device_agreement(compilation_step, cfg.device_type, "DaCe compilation step") if cfg.cached_translation: translation_step = workflow.CachedStep[ diff --git a/src/gt4py/next/program_processors/runners/gtfn.py b/src/gt4py/next/program_processors/runners/gtfn.py index b7ed317642..0e1f4a95ec 100644 --- a/src/gt4py/next/program_processors/runners/gtfn.py +++ b/src/gt4py/next/program_processors/runners/gtfn.py @@ -246,7 +246,8 @@ def make_gtfn_compile_workflow( `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. + 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()`. @@ -256,18 +257,12 @@ def make_gtfn_compile_workflow( Returns: The composed compile workflow. - - Raises: - ValueError: If a step builder returns a step configured for a device - other than `cfg.device_type`. """ if cfg is None: cfg = GTFNConfig() translation_step = translation(cfg) - workflow.check_device_agreement(translation_step, cfg.device_type, "GTFN translation step") compilation_step = compilation(cfg) - workflow.check_device_agreement(compilation_step, cfg.device_type, "GTFN compilation step") if cfg.cached_translation: translation_step = workflow.CachedStep[ @@ -313,10 +308,6 @@ def make_gtfn_toolchain( Returns: The configured toolchain. - - Raises: - ValueError: If a step builder returns a step configured for a device - other than `cfg.device_type`. """ if cfg is None: cfg = GTFNConfig() 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 2441e21f26..d962daa62b 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 @@ -104,13 +104,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, + optimization_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. @@ -187,14 +192,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, + optimization_args={ + "transient_memory_mode": gtx_transformations.TransientMemoryMode.EXTERNAL, + }, + ), ) assert backend.external_workspace[core_defs.DeviceType.CPU] is workspace @@ -203,11 +208,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 ( @@ -221,14 +223,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, + optimization_args={ + "transient_memory_mode": gtx_transformations.TransientMemoryMode.POOL, + }, + ), ) # Explicit mode stays as requested by the caller; backend only warns. @@ -268,28 +270,6 @@ def test_make_toolchain_rejects_derived_optimization_args(): ) -@pytest.mark.parametrize( - "step_builders", - [ - # Each builder ignores the GPU config it receives and targets the CPU. - { - "translation": lambda cfg: dace_wf_factory.make_dace_translator( - dace_wf_factory.DaCeConfig() - ) - }, - { - "compilation": lambda cfg: dace_wf_factory.make_dace_compiler( - dace_wf_factory.DaCeConfig() - ) - }, - ], - ids=["translation", "compilation"], -) -def test_make_toolchain_rejects_step_builder_ignoring_config_device(step_builders): - with pytest.raises(ValueError, match="toolchain is being built for"): - dace_wf_backend.make_dace_toolchain(dace_wf_factory.DaCeConfig(gpu=True), **step_builders) - - def test_make_toolchain_uncached_translation(): backend = dace_wf_backend.make_dace_toolchain( dace_wf_factory.DaCeConfig(cached_translation=False) @@ -298,6 +278,14 @@ def test_make_toolchain_uncached_translation(): assert isinstance(backend.executor.translation, dace_wf_translation.DaCeTranslator) +def test_make_dace_backend_is_deprecated(): + with pytest.warns(DeprecationWarning, match="make_dace_toolchain"): + backend = dace_wf_backend.make_dace_backend(gpu=False, use_metrics=False) + + assert backend.name == "run_dace_cpu_opt" + assert backend.executor.translation.step.use_metrics is False + + 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. @@ -351,14 +339,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, + optimization_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..ddf78afe13 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, + use_zero_origin=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, + use_zero_origin=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 c1abc3fdd1..1b3599f5a5 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 @@ -175,14 +175,6 @@ def test_uncached_translation(): assert isinstance(toolchain.executor.translation, gtfn_module.GTFNTranslationStep) -def test_step_builder_ignoring_config_device_raises(): - def cpu_only_translation(cfg: gtfn.GTFNConfig) -> gtfn_module.GTFNTranslationStep: - return gtfn_module.GTFNTranslationStep(device_type=core_defs.DeviceType.CPU) - - with pytest.raises(ValueError, match="toolchain is being built for 'CUDA'"): - gtfn.make_gtfn_toolchain(gtfn.GTFNConfig(gpu=True), translation=cpu_only_translation) - - def test_step_builder_cannot_override_config_setting(): with pytest.raises(TypeError, match="device_type"): gtfn.make_gtfn_toolchain( From 265837321587aae0f6d620980110205e8fab52cc Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Thu, 1 Oct 2026 15:31:45 +0200 Subject: [PATCH 5/6] refactor[next]: address independent review of the plain builders - `DaCeTranslator` derives `unit_strides_kind` itself, next to `gpu` and `constant_symbols`, and rejects these three keys in `auto_optimize_args` on construction, so the check covers every way a translator is built. The warning about unused optimization arguments stays in the step builder, where it is attributed to the caller. - `make_dace_translator` takes the step-local fields of `DaCeTranslator` under their own names (`auto_optimize_args`, `disable_field_origin_on_program_arguments`) and now also exposes `disable_itir_transforms`; `make_dace_compiler` exposes `add_gpu_trace_markers`. The deprecated `make_dace_backend` maps its legacy names onto them. - Share the translation-cache wrapping in `cache.persistent_translation_cache`, and name the step-builder types (`GTFN*Builder`, `DaCe*Builder`). - The deprecation test checks that `make_dace_backend` builds the same toolchain as the equivalent `make_dace_toolchain` call. - ADR 0028 states the no-default rule accurately; fix a stale comment in `otf/compilation/cache.py`; note that a `DaCeConfig` holding a workspace is not hashable. The DaCe translators no longer store `unit_strides_kind` in `auto_optimize_args`, so DaCe translation-cache keys rotate once; the value passed to `gt_auto_optimize` is unchanged. --- ...028-Plain-Builders-Instead-of-Factories.md | 7 +- src/gt4py/next/otf/compilation/cache.py | 33 +++++- .../runners/dace/workflow/backend.py | 21 +--- .../runners/dace/workflow/factory.py | 109 +++++++++--------- .../runners/dace/workflow/translation.py | 22 +++- .../next/program_processors/runners/gtfn.py | 47 ++++---- .../dace_tests/test_dace_backend.py | 88 ++++++++++++-- .../dace_tests/test_dace_bindings.py | 4 +- 8 files changed, 220 insertions(+), 111 deletions(-) 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 index 860449b917..e5f4fbd0d0 100644 --- a/docs/development/ADRs/next/0028-Plain-Builders-Instead-of-Factories.md +++ b/docs/development/ADRs/next/0028-Plain-Builders-Instead-of-Factories.md @@ -95,8 +95,11 @@ front end over the config-based builder, for existing callers. - 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. -- Step fields that must agree with other steps get no default, so a builder - that forgets to pass one fails instead of silently using the default. +- 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 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/runners/dace/workflow/backend.py b/src/gt4py/next/program_processors/runners/dace/workflow/backend.py index 55a8d9b5d5..e356523aee 100644 --- a/src/gt4py/next/program_processors/runners/dace/workflow/backend.py +++ b/src/gt4py/next/program_processors/runners/dace/workflow/backend.py @@ -11,11 +11,10 @@ import dataclasses import functools import warnings -from collections.abc import Callable from typing import Any from gt4py.next import backend, config -from gt4py.next.otf import artifacts, stages, workflow +from gt4py.next.otf import artifacts from gt4py.next.program_processors.runners.dace.workflow import ( common as gtx_wfdcommon, decoration as gtx_wfddecoration, @@ -42,17 +41,9 @@ def make_dace_toolchain( /, *, name_postfix: str = "", - translation: Callable[ - [gtx_wfdfactory.DaCeConfig], stages.TranslationStep - ] = gtx_wfdfactory.make_dace_translator, - bindings: Callable[ - [gtx_wfdfactory.DaCeConfig], - workflow.Workflow[artifacts.ProgramSource, artifacts.ExtensionSource], - ] = gtx_wfdfactory.make_dace_bindings, - compilation: Callable[ - [gtx_wfdfactory.DaCeConfig], - workflow.Workflow[artifacts.ExtensionSource, artifacts.CompilationArtifact], - ] = gtx_wfdfactory.make_dace_compiler, + 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: """ Build a DaCe toolchain. @@ -144,10 +135,10 @@ def make_dace_backend( ), translation=functools.partial( gtx_wfdfactory.make_dace_translator, - optimization_args=optimization_args, + auto_optimize_args=optimization_args, async_sdfg_call=async_sdfg_call, use_metrics=use_metrics, - use_zero_origin=use_zero_origin, + disable_field_origin_on_program_arguments=use_zero_origin, use_max_domain_range_on_unstructured_shift=use_max_domain_range_on_unstructured_shift, ), ) 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 ecb42a47d5..6b3d52dba6 100644 --- a/src/gt4py/next/program_processors/runners/dace/workflow/factory.py +++ b/src/gt4py/next/program_processors/runners/dace/workflow/factory.py @@ -13,11 +13,10 @@ import pathlib import warnings from collections.abc import Callable -from typing import Any, ClassVar, Final +from typing import Any, ClassVar, Final, TypeAlias import gt4py -from gt4py._core import filecache -from gt4py.next import backend as next_backend, common, fingerprinting +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 import transformations as gtx_transformations @@ -29,12 +28,6 @@ from gt4py.next.program_processors.runners.dace.workflow.translation import DaCeTranslator -#: The parameters of `gt_auto_optimize()` that are derived from the toolchain -#: configuration, and therefore cannot be customized. -_DERIVED_OPTIMIZATION_ARGS: Final[frozenset[str]] = frozenset( - {"gpu", "constant_symbols", "unit_strides_kind"} -) - #: Warnings about the builder arguments point at the first caller outside GT4Py. _GT4PY_SOURCE_PREFIX: Final[str] = str(pathlib.Path(gt4py.__file__).parent) @@ -47,7 +40,8 @@ class DaCeConfig(next_backend.ToolchainConfig): 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. + #: `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 #: Name of the function that binds the SDFG arguments. The bindings and the @@ -59,52 +53,47 @@ def make_dace_translator( cfg: DaCeConfig, /, *, - optimization_args: dict[str, Any] | None = None, + auto_optimize_args: dict[str, Any] | None = None, async_sdfg_call: bool = True, use_metrics: bool = True, - use_zero_origin: bool = False, + 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`. + Args: cfg: The toolchain configuration. - optimization_args: Configuration for the SDFG auto-optimize pipeline, see - `gt_auto_optimize()`. The parameters derived from `cfg` cannot be - set here. + 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. - use_zero_origin: Assume that all fields passed as program arguments have - zero-based origin, which skips the range start-symbols `_range_0`. + 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 `optimization_args` sets a parameter derived from `cfg`, - or requests the `EXTERNAL` transient memory mode without an - external workspace in `cfg`. + ValueError: If `auto_optimize_args` sets a parameter `DaCeTranslator` + derives itself, or requests the `EXTERNAL` transient memory mode + without an external workspace in `cfg`. """ - if optimization_args is None: - optimization_args = {} - elif optimization_args and not cfg.auto_optimize: + 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,), ) - elif intersect_args := optimization_args.keys() & _DERIVED_OPTIMIZATION_ARGS: - raise ValueError( - f"The following optimization arguments cannot be overriden: {intersect_args}." - ) - - optimization_args = optimization_args | { - "unit_strides_kind": common.DimensionKind.HORIZONTAL - if cfg.unstructured_horizontal_has_unit_stride - else None - } if cfg.external_workspace is None: if ( @@ -132,7 +121,8 @@ def make_dace_translator( 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_field_origin_on_program_arguments=use_zero_origin, + 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, ) @@ -144,27 +134,48 @@ def make_dace_bindings( return functools.partial(bindings_step.bind_sdfg, bind_func_name=cfg.bind_func_name) -def make_dace_compiler(cfg: DaCeConfig, /) -> DaCeCompiler: - """Build the compilation step, targeting `cfg.device_type`.""" +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, ) +#: 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: Callable[[DaCeConfig], stages.TranslationStep] = make_dace_translator, - bindings: Callable[ - [DaCeConfig], workflow.Workflow[artifacts.ProgramSource, artifacts.ExtensionSource] - ] = make_dace_bindings, - compilation: Callable[ - [DaCeConfig], workflow.Workflow[artifacts.ExtensionSource, artifacts.CompilationArtifact] - ] = make_dace_compiler, + translation: DaCeTranslationBuilder = make_dace_translator, + bindings: DaCeBindingsBuilder = make_dace_bindings, + compilation: DaCeCompilationBuilder = make_dace_compiler, ) -> recipes.OTFCompileWorkflow: """ Build the DaCe translation -> bindings -> compilation workflow. @@ -194,16 +205,8 @@ def make_dace_compile_workflow( compilation_step = compilation(cfg) if cfg.cached_translation: - translation_step = workflow.CachedStep[ - stages.CompilableProgramDef, artifacts.ProgramSource, str - ].persistent( - translation_step, - input_fingerprinter=fingerprinting.strict_fingerprinter, - cache=filecache.FileCache( - cache.get_translation_cache_folder( - cache.get_cache_base_path(cfg.cache_lifetime), "dace" - ) - ), + translation_step = cache.persistent_translation_cache( + translation_step, "dace", cfg.cache_lifetime ) return recipes.OTFCompileWorkflow( 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 e45f3840b3..b3b2203eb9 100644 --- a/src/gt4py/next/program_processors/runners/dace/workflow/translation.py +++ b/src/gt4py/next/program_processors/runners/dace/workflow/translation.py @@ -9,7 +9,7 @@ from __future__ import annotations import dataclasses -from typing import Any, Optional +from typing import Any, Final, Optional import dace @@ -339,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[ @@ -361,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 overriden: {derived_args}." + ) + def generate_sdfg( self, *args: Any, @@ -401,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: diff --git a/src/gt4py/next/program_processors/runners/gtfn.py b/src/gt4py/next/program_processors/runners/gtfn.py index 0e1f4a95ec..f73f4e0cd6 100644 --- a/src/gt4py/next/program_processors/runners/gtfn.py +++ b/src/gt4py/next/program_processors/runners/gtfn.py @@ -10,13 +10,12 @@ import functools import pathlib from collections.abc import Callable -from typing import Any +from typing import Any, TypeAlias import numpy as np import gt4py._core.definitions as core_defs -from gt4py._core import filecache -from gt4py.next import backend, common, 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.iterator import ir as itir @@ -129,6 +128,16 @@ 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, /, @@ -228,13 +237,9 @@ def make_gtfn_compile_workflow( cfg: GTFNConfig | None = None, /, *, - translation: Callable[[GTFNConfig], stages.TranslationStep] = make_gtfn_translation, - bindings: Callable[ - [GTFNConfig], workflow.Workflow[artifacts.ProgramSource, artifacts.ExtensionSource] - ] = make_gtfn_bindings, - compilation: Callable[ - [GTFNConfig], workflow.Workflow[artifacts.ExtensionSource, artifacts.CompilationArtifact] - ] = make_gtfn_compiler, + translation: GTFNTranslationBuilder = make_gtfn_translation, + bindings: GTFNBindingsBuilder = make_gtfn_bindings, + compilation: GTFNCompilationBuilder = make_gtfn_compiler, ) -> recipes.OTFCompileWorkflow: """ Build the GTFN translation -> bindings -> compilation workflow. @@ -265,16 +270,8 @@ def make_gtfn_compile_workflow( compilation_step = compilation(cfg) if cfg.cached_translation: - translation_step = workflow.CachedStep[ - stages.CompilableProgramDef, artifacts.ProgramSource, str - ].persistent( - translation_step, - input_fingerprinter=fingerprinting.strict_fingerprinter, - cache=filecache.FileCache( - cache.get_translation_cache_folder( - cache.get_cache_base_path(cfg.cache_lifetime), "gtfn" - ) - ), + translation_step = cache.persistent_translation_cache( + translation_step, "gtfn", cfg.cache_lifetime ) return recipes.OTFCompileWorkflow( @@ -287,13 +284,9 @@ def make_gtfn_toolchain( /, *, name_postfix: str = "", - translation: Callable[[GTFNConfig], stages.TranslationStep] = make_gtfn_translation, - bindings: Callable[ - [GTFNConfig], workflow.Workflow[artifacts.ProgramSource, artifacts.ExtensionSource] - ] = make_gtfn_bindings, - compilation: Callable[ - [GTFNConfig], workflow.Workflow[artifacts.ExtensionSource, artifacts.CompilationArtifact] - ] = make_gtfn_compiler, + translation: GTFNTranslationBuilder = make_gtfn_translation, + bindings: GTFNBindingsBuilder = make_gtfn_bindings, + compilation: GTFNCompilationBuilder = make_gtfn_compiler, ) -> backend.Backend: """ Build a GTFN toolchain. 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 d962daa62b..ceb7245d88 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 @@ -112,7 +112,7 @@ def mocked_gpu_transformation(*args, **kwargs) -> dace.SDFG: ), translation=functools.partial( dace_wf_factory.make_dace_translator, - optimization_args=optimization_args, + auto_optimize_args=optimization_args, async_sdfg_call=True, use_metrics=True, ), @@ -196,7 +196,7 @@ def test_make_backend_accepts_external_workspace_with_external_mode(): dace_wf_factory.DaCeConfig(external_workspace={core_defs.DeviceType.CPU: workspace}), translation=functools.partial( dace_wf_factory.make_dace_translator, - optimization_args={ + auto_optimize_args={ "transient_memory_mode": gtx_transformations.TransientMemoryMode.EXTERNAL, }, ), @@ -227,7 +227,7 @@ def test_make_backend_warns_external_workspace_without_external_mode(): dace_wf_factory.DaCeConfig(external_workspace={core_defs.DeviceType.CPU: workspace}), translation=functools.partial( dace_wf_factory.make_dace_translator, - optimization_args={ + auto_optimize_args={ "transient_memory_mode": gtx_transformations.TransientMemoryMode.POOL, }, ), @@ -265,7 +265,7 @@ def test_make_toolchain_rejects_derived_optimization_args(): dace_wf_backend.make_dace_toolchain( translation=functools.partial( dace_wf_factory.make_dace_translator, - optimization_args={"unit_strides_kind": None}, + auto_optimize_args={"unit_strides_kind": None}, ) ) @@ -279,11 +279,83 @@ def test_make_toolchain_uncached_translation(): def test_make_dace_backend_is_deprecated(): + workspace = {core_defs.DeviceType.CPU: _RecordingWorkspace()} with pytest.warns(DeprecationWarning, match="make_dace_toolchain"): - backend = dace_wf_backend.make_dace_backend(gpu=False, use_metrics=False) + 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 overriden"): + 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 overriden"): + 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 backend.name == "run_dace_cpu_opt" - assert backend.executor.translation.step.use_metrics is False + assert record[0].filename == __file__ def _parse_generated_code_from_sdfg(sdfg: dace.SDFG, gpu_api_prefix: str) -> str: @@ -343,7 +415,7 @@ def test_transient_memory_mode(device_type, transient_memory_mode, monkeypatch): dace_wf_factory.DaCeConfig(gpu=on_gpu, external_workspace=external_workspace), translation=functools.partial( dace_wf_factory.make_dace_translator, - optimization_args={ + auto_optimize_args={ "transient_memory_mode": transient_memory_mode, }, async_sdfg_call=False, 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 ddf78afe13..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 @@ -298,7 +298,7 @@ def testee( translation=functools.partial( dace_runner.make_dace_translator, use_metrics=use_metrics, - use_zero_origin=use_zero_origin, + disable_field_origin_on_program_arguments=use_zero_origin, ) ) monkeypatch.setattr( @@ -354,7 +354,7 @@ def testee(a: cases.VField, b: cases.VField): translation=functools.partial( dace_runner.make_dace_translator, use_metrics=use_metrics, - use_zero_origin=use_zero_origin, + disable_field_origin_on_program_arguments=use_zero_origin, ) ) monkeypatch.setattr( From 65e9f39606201dd83d4e69e8cf6bdf30d92186bd Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Thu, 1 Oct 2026 16:55:31 +0200 Subject: [PATCH 6/6] refactor[next]: export the DaCe step-builder types and cover builder defaults Address review: - Export `DaCeTranslationBuilder`, `DaCeBindingsBuilder` and `DaCeCompilationBuilder` from `runners.dace`, next to the other builders. - Fix "overriden" in the error for derived optimization arguments. - Test that `GTFNTranslationStep` requires `device_type`, and that `make_gtfn_compile_workflow()` and `make_dace_compile_workflow()` without a config cache their translation step. --- .../program_processors/runners/dace/__init__.py | 6 ++++++ .../runners/dace/workflow/backend.py | 2 +- .../runners/dace/workflow/translation.py | 2 +- .../runners_tests/dace_tests/test_dace_backend.py | 14 +++++++++++--- .../runners_tests/test_gtfn.py | 14 ++++++++++++++ 5 files changed, 33 insertions(+), 5 deletions(-) diff --git a/src/gt4py/next/program_processors/runners/dace/__init__.py b/src/gt4py/next/program_processors/runners/dace/__init__.py index c8e96ad515..2e06493bd0 100644 --- a/src/gt4py/next/program_processors/runners/dace/__init__.py +++ b/src/gt4py/next/program_processors/runners/dace/__init__.py @@ -17,7 +17,10 @@ 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, @@ -26,7 +29,10 @@ __all__ = [ + "DaCeBindingsBuilder", + "DaCeCompilationBuilder", "DaCeConfig", + "DaCeTranslationBuilder", "get_sdfg_args", "make_dace_backend", "make_dace_bindings", 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 e356523aee..98149f20c4 100644 --- a/src/gt4py/next/program_processors/runners/dace/workflow/backend.py +++ b/src/gt4py/next/program_processors/runners/dace/workflow/backend.py @@ -110,7 +110,7 @@ def make_dace_backend( 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 in `optimization_args`. + cannot be overridden, and therefore cannot appear in `optimization_args`. Returns: A dace backend with custom configuration for the target device. 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 b3b2203eb9..ee96fbeff2 100644 --- a/src/gt4py/next/program_processors/runners/dace/workflow/translation.py +++ b/src/gt4py/next/program_processors/runners/dace/workflow/translation.py @@ -373,7 +373,7 @@ def __post_init__(self) -> None: derived_args := self.auto_optimize_args.keys() & _DERIVED_OPTIMIZATION_ARGS ): raise ValueError( - f"The following optimization arguments cannot be overriden: {derived_args}." + f"The following optimization arguments cannot be overridden: {derived_args}." ) def generate_sdfg( 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 ceb7245d88..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 @@ -22,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, @@ -261,7 +262,7 @@ def test_make_toolchain_derives_workspace_and_memory_mode_from_one_config(): def test_make_toolchain_rejects_derived_optimization_args(): - with pytest.raises(ValueError, match="cannot be overriden"): + with pytest.raises(ValueError, match="cannot be overridden"): dace_wf_backend.make_dace_toolchain( translation=functools.partial( dace_wf_factory.make_dace_translator, @@ -318,7 +319,7 @@ def test_make_dace_backend_is_deprecated(): def test_translator_rejects_derived_optimization_args_on_every_route(): - with pytest.raises(ValueError, match="cannot be overriden"): + with pytest.raises(ValueError, match="cannot be overridden"): dace_wf_translation.DaCeTranslator( device_type=core_defs.DeviceType.CPU, auto_optimize=False, @@ -328,7 +329,7 @@ def test_translator_rejects_derived_optimization_args_on_every_route(): use_metrics=False, ) translator = dace_wf_backend.run_dace_cpu.executor.translation.step - with pytest.raises(ValueError, match="cannot be overriden"): + with pytest.raises(ValueError, match="cannot be overridden"): dataclasses.replace(translator, auto_optimize_args={"constant_symbols": {}}) @@ -358,6 +359,13 @@ def test_unused_optimization_args_warning_points_at_the_caller(): 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. 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 1b3599f5a5..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 @@ -190,3 +190,17 @@ def test_prebuilt_toolchain_names_are_unique(): 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)