From 94557933c38e93cf0f87276bf7ae6817069d0ada Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Thu, 24 Sep 2026 16:05:50 +0200 Subject: [PATCH 1/6] feat[next]: add FORMAT_SOURCES config option to control formatting of generated code Add `config.FORMAT_SOURCES` (env var `GT4PY_FORMAT_SOURCES`, defaulting to `GT4PY_DEBUG`) to control whether generated sources are run through black/clang-format. The value is captured at construction time in `SourceCodeSpec.format_source` and `Roundtrip.format_source`, so it is part of the fingerprints and cache keys of specs and workflow steps. - `GTFNTranslationStep` resolves its default code spec in `__post_init__` so the setting is part of the translation cache key. - Drop the redundant clang-format pass in `generate_stencil_source`; `format_cpp` always formats explicitly. - `SDFGCodeSpec` pins `format_source=False` and the compiledb prototype normalizes it, so neither DaCe builds nor the compiledb cache split on it. --- src/gt4py/next/config.py | 7 ++ src/gt4py/next/otf/artifacts.py | 16 ++++- .../compilation/build_systems/compiledb.py | 4 +- .../codegens/gtfn/gtfn_module.py | 26 ++++--- .../program_processors/formatters/gtfn.py | 5 +- .../program_processors/runners/roundtrip.py | 10 +-- .../build_systems_tests/test_compiledb.py | 17 +++++ .../unit_tests/otf_tests/test_languages.py | 36 ++++++++++ .../gtfn_tests/test_gtfn_module.py | 23 +++++- .../runners_tests/test_roundtrip.py | 71 +++++++++++++++++++ 10 files changed, 197 insertions(+), 18 deletions(-) create mode 100644 tests/next_tests/unit_tests/program_processor_tests/runners_tests/test_roundtrip.py diff --git a/src/gt4py/next/config.py b/src/gt4py/next/config.py index 56eec687a0..308193e43c 100644 --- a/src/gt4py/next/config.py +++ b/src/gt4py/next/config.py @@ -106,6 +106,13 @@ def env_flag_to_int(name: str, default: int) -> int: ) +#: Run source formatters (e.g. black, clang-format) on generated code. +#: Only affects the readability of the generated code, never its semantics. +#: The value is captured when a code spec or workflow step is created, so +#: changing it later does not affect already existing backends. +FORMAT_SOURCES: bool = env_flag_to_bool("GT4PY_FORMAT_SOURCES", default=DEBUG) + + #: Where generated code projects should be persisted. #: Only active if BUILD_CACHE_LIFETIME is set to PERSISTENT BUILD_CACHE_DIR: pathlib.Path = ( diff --git a/src/gt4py/next/otf/artifacts.py b/src/gt4py/next/otf/artifacts.py index 204e2e75d0..06fcbcde8f 100644 --- a/src/gt4py/next/otf/artifacts.py +++ b/src/gt4py/next/otf/artifacts.py @@ -30,6 +30,7 @@ from typing import Any, Generic, Optional, Protocol, TypeAlias, TypeVar, runtime_checkable from gt4py.eve import codegen +from gt4py.next import config from gt4py.next.otf.binding import interface @@ -38,15 +39,19 @@ class SourceCodeSpec: """ Basic settings for any source programming language. - Formatting will happen through ``eve.codegen.format_source``. - For available formatting options, check the options of the - specific formatter used depending on ``.formatter_key``. + Formatting will happen through `eve.codegen.format_source`, only if + `.format_source` is true. For available formatting options, check the + options of the specific formatter used depending on `.formatter_key`. + + `.format_source` defaults to the value of `config.FORMAT_SOURCES` at + creation time, and it is part of the spec (and thus of its fingerprint). """ source_language: str file_extension: str formatter_key: str | None = None formatter_options: Mapping[str, Any] | None = None + format_source: bool = dataclasses.field(default_factory=lambda: config.FORMAT_SOURCES) @dataclasses.dataclass(frozen=True, kw_only=True) @@ -71,6 +76,8 @@ class SDFGCodeSpec(SourceCodeSpec): source_language: str = "SDFG" file_extension: str = "sdfg" + # There is no SDFG formatter: keep the spec independent of `config.FORMAT_SOURCES` + format_source: bool = False @dataclasses.dataclass(frozen=True, kw_only=True) @@ -118,6 +125,9 @@ class HIPCodeSpec(CPPLikeCodeSpec): def format_source(source_code_spec: SourceCodeSpec, source: str) -> str: + """Format `source` as configured in `source_code_spec` (no-op if `.format_source` is false).""" + if not source_code_spec.format_source: + return source assert source_code_spec.formatter_key is not None, ( "No formatter key specified in source code specification." ) diff --git a/src/gt4py/next/otf/compilation/build_systems/compiledb.py b/src/gt4py/next/otf/compilation/build_systems/compiledb.py index 92a9facb08..c131605951 100644 --- a/src/gt4py/next/otf/compilation/build_systems/compiledb.py +++ b/src/gt4py/next/otf/compilation/build_systems/compiledb.py @@ -262,7 +262,9 @@ def _cc_prototype_program_source( entry_point=interface.Function(name=name, parameters=()), source_code="", library_deps=deps, - code_spec=code_spec, + # The compiledb does not depend on source formatting: normalize it so the + # cache folder (keyed by the prototype source) is shared across settings. + code_spec=dataclasses.replace(code_spec, format_source=False), ) 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..42822cede5 100644 --- a/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py +++ b/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py @@ -15,7 +15,6 @@ import numpy as np from gt4py._core import definitions as core_defs -from gt4py.eve import codegen from gt4py.next import common from gt4py.next.ffront import fbuiltins from gt4py.next.iterator import ir as itir @@ -52,6 +51,19 @@ class GTFNTranslationStep( symbolic_domain_sizes: dict[str, itir.Expr] | None = None use_max_domain_range_on_unstructured_shift: bool | None = None + def __post_init__(self) -> None: + # Resolve the default code spec eagerly, so its settings (e.g. `format_source`, + # which follows `config.FORMAT_SOURCES`) are part of the step and its fingerprint. + default_code_spec = self._default_code_spec() + if self.code_spec is None: + object.__setattr__(self, "code_spec", default_code_spec) + elif not isinstance(self.code_spec, type(default_code_spec)): + raise ValueError( + f"Code spec '{type(self.code_spec).__name__}' does not match device type " + f"'{self.device_type.name}' (expected '{type(default_code_spec).__name__}'). " + "When replacing the device type, pass 'code_spec=None' to use the default spec." + ) + def _default_code_spec(self) -> artifacts.HeaderAndSourceCodeSpec: match self.device_type: case core_defs.DeviceType.CUDA: @@ -171,14 +183,15 @@ def generate_stencil_source( column_axis=column_axis, ) - generated_code = GTFNCodegen.apply(gtfn_ir) - return codegen.format_source("cpp", generated_code, style="LLVM") + return GTFNCodegen.apply(gtfn_ir) def __call__( self, inp: stages.CompilableProgramDef ) -> artifacts.ProgramSource[artifacts.HeaderAndSourceCodeSpec]: """Generate GTFN C++ code from the ITIR definition.""" program: itir.Program = inp.data + code_spec = self.code_spec + assert code_spec is not None # resolved in `__post_init__` # handle regular parameters and arguments of the program (i.e. what the user defined in # the program) @@ -210,7 +223,7 @@ def __call__( inp.args.column_axis, ) source_code = artifacts.format_source( - self._code_spec(), + code_spec, f""" #include <{self._backend_header()}> #include @@ -224,7 +237,7 @@ def __call__( entry_point=function, library_deps=(interface.LibraryDependency(self._library_name(), "master"),), source_code=source_code, - code_spec=self._code_spec(), + code_spec=code_spec, ) ) return module @@ -247,9 +260,6 @@ def _backend_type(self) -> str: case _: raise self._not_implemented_for_device_type() - def _code_spec(self) -> artifacts.HeaderAndSourceCodeSpec: - return self.code_spec if self.code_spec is not None else self._default_code_spec() - def _library_name(self) -> str: match self.device_type: case core_defs.DeviceType.CUDA | core_defs.DeviceType.ROCM: diff --git a/src/gt4py/next/program_processors/formatters/gtfn.py b/src/gt4py/next/program_processors/formatters/gtfn.py index 75494a1759..2932a742b7 100644 --- a/src/gt4py/next/program_processors/formatters/gtfn.py +++ b/src/gt4py/next/program_processors/formatters/gtfn.py @@ -8,6 +8,7 @@ from typing import Any +from gt4py.eve import codegen 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 @@ -18,8 +19,10 @@ def format_cpp(program: itir.Program, *args: Any, **kwargs: Any) -> str: gtfn_translation = gtfn.GTFNCompileWorkflowFactory(cached_translation=False).translation assert isinstance(gtfn_translation, GTFNTranslationStep) - return gtfn_translation.generate_stencil_source( + generated_code = gtfn_translation.generate_stencil_source( program, offset_provider=kwargs.get("offset_provider", {}), column_axis=kwargs.get("column_axis", None), ) + # The purpose of this formatter is producing human-readable code, so always format it + return codegen.format_source("cpp", generated_code, style="LLVM") diff --git a/src/gt4py/next/program_processors/runners/roundtrip.py b/src/gt4py/next/program_processors/runners/roundtrip.py index 09f173d3f9..2ef7122f49 100644 --- a/src/gt4py/next/program_processors/runners/roundtrip.py +++ b/src/gt4py/next/program_processors/runners/roundtrip.py @@ -117,6 +117,7 @@ def visit_Temporary(self, node: itir.Temporary, **kwargs: Any) -> str: def _generate_source( ir: itir.Program, debug: bool, + format_source: bool, use_embedded: bool, offset_provider: common.OffsetProvider, transforms: itir_transforms.GTIRTransform, @@ -128,7 +129,7 @@ def _generate_source( ( ir, transforms, - debug, + format_source, use_embedded, tuple(common.offset_provider_to_type(offset_provider).items()), ) @@ -142,9 +143,8 @@ def _generate_source( program = EmbeddedDSL.apply(ir) - # format output in debug mode for better debuggability - # (e.g. line numbers, overview in the debugger). - if debug: + # format output for better debuggability (e.g. line numbers, overview in the debugger). + if format_source: program = codegen.format_python_source(program) offset_literals: Iterable[str] = ( @@ -258,6 +258,7 @@ class Roundtrip(workflow.Workflow[stages.CompilableProgramDef, RoundtripArtifact use_embedded: bool = True dispatch_backend: Optional[next_backend.Backend] = None transforms: itir_transforms.GTIRTransform = itir_transforms.apply_common_transforms # type: ignore[assignment] # TODO(havogt): cleanup interface of `apply_common_transforms` + format_source: bool = dataclasses.field(default_factory=lambda: config.FORMAT_SOURCES) def __call__(self, inp: stages.CompilableProgramDef) -> RoundtripArtifact: debug = config.DEBUG if self.debug is None else self.debug @@ -266,6 +267,7 @@ def __call__(self, inp: stages.CompilableProgramDef) -> RoundtripArtifact: inp.data, offset_provider=inp.args.offset_provider, debug=debug, + format_source=self.format_source, use_embedded=self.use_embedded, transforms=self.transforms, ) diff --git a/tests/next_tests/unit_tests/otf_tests/compilation_tests/build_systems_tests/test_compiledb.py b/tests/next_tests/unit_tests/otf_tests/compilation_tests/build_systems_tests/test_compiledb.py index b208b7cd6f..7df35c1889 100644 --- a/tests/next_tests/unit_tests/otf_tests/compilation_tests/build_systems_tests/test_compiledb.py +++ b/tests/next_tests/unit_tests/otf_tests/compilation_tests/build_systems_tests/test_compiledb.py @@ -13,6 +13,7 @@ import pytest from gt4py.next import config, fingerprinting +from gt4py.next.otf import artifacts from gt4py.next.otf.compilation import build_data, cache, importer from gt4py.next.otf.compilation.build_systems import compiledb @@ -91,3 +92,19 @@ def test_compiledb_project_is_relocatable(extension_source_example, clean_compil assert hasattr( importer.import_from_path(relocated_dir / new_data.module), new_data.entry_point_name ) + + +def test_compiledb_prototype_ignores_format_source(): + prototypes = [ + compiledb._cc_prototype_program_source( + deps=(), + build_type=config.CMakeBuildType.RELEASE, + cmake_flags=[], + code_spec=artifacts.CPPCodeSpec(format_source=format_source), + ) + for format_source in (True, False) + ] + + assert fingerprinting.strict_fingerprinter( + prototypes[0] + ) == fingerprinting.strict_fingerprinter(prototypes[1]) diff --git a/tests/next_tests/unit_tests/otf_tests/test_languages.py b/tests/next_tests/unit_tests/otf_tests/test_languages.py index 08a738d017..0815a41deb 100644 --- a/tests/next_tests/unit_tests/otf_tests/test_languages.py +++ b/tests/next_tests/unit_tests/otf_tests/test_languages.py @@ -8,6 +8,7 @@ import pytest +from gt4py.next import config, fingerprinting from gt4py.next.otf import artifacts from gt4py.next.otf.binding import interface @@ -19,3 +20,38 @@ def test_header_files_settings_with_cpp_accepted(): library_deps=(), code_spec=artifacts.CPPCodeSpec(), ) + + +@pytest.mark.parametrize("flag", [True, False]) +def test_format_source_default_captures_config_at_creation(monkeypatch, flag): + monkeypatch.setattr(config, "FORMAT_SOURCES", flag) + code_spec = artifacts.CPPCodeSpec() + assert code_spec.format_source is flag + + # changing the global option afterwards does not affect existing specs + monkeypatch.setattr(config, "FORMAT_SOURCES", not flag) + assert code_spec.format_source is flag + assert artifacts.CPPCodeSpec().format_source is not flag + + +@pytest.mark.parametrize("flag", [True, False]) +def test_format_source_is_part_of_fingerprint(monkeypatch, flag): + monkeypatch.setattr(config, "FORMAT_SOURCES", flag) + assert fingerprinting.strict_fingerprinter( + artifacts.CPPCodeSpec() + ) != fingerprinting.strict_fingerprinter(artifacts.CPPCodeSpec(format_source=not flag)) + + +@pytest.mark.parametrize("flag", [True, False]) +def test_sdfg_code_spec_ignores_format_sources_config(monkeypatch, flag): + monkeypatch.setattr(config, "FORMAT_SOURCES", flag) + assert artifacts.SDFGCodeSpec().format_source is False + + +def test_format_source_follows_code_spec(): + source = "x=( 1 )" + + assert artifacts.format_source(artifacts.PythonCodeSpec(format_source=False), source) == source + assert ( + artifacts.format_source(artifacts.PythonCodeSpec(format_source=True), source) == "x = 1\n" + ) 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..f4b473f362 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 @@ -15,7 +15,7 @@ import gt4py.next as gtx from gt4py.next.iterator import builtins, ir as itir from gt4py.next.iterator.ir_utils import ir_makers as im -from gt4py.next import fingerprinting +from gt4py.next import config, fingerprinting from gt4py.next.otf import arguments, artifacts, stages from gt4py.next.program_processors.codegens.gtfn import gtfn_module from gt4py.next.program_processors.runners import gtfn @@ -86,6 +86,27 @@ def test_codegen(program_example): assert isinstance(module.code_spec, artifacts.CPPCodeSpec) +@pytest.mark.parametrize("flag", [True, False]) +def test_code_spec_is_resolved_at_construction(monkeypatch, flag): + monkeypatch.setattr(config, "FORMAT_SOURCES", flag) + step = gtfn_module.GTFNTranslationStep() + monkeypatch.setattr(config, "FORMAT_SOURCES", not flag) + + assert isinstance(step.code_spec, artifacts.CPPCodeSpec) + assert step.code_spec.format_source is flag + # the formatting setting is part of the step fingerprint, and thus of translation cache keys + assert fingerprinting.strict_fingerprinter(step) != fingerprinting.strict_fingerprinter( + gtfn_module.GTFNTranslationStep() + ) + + +def test_code_spec_device_type_mismatch(): + with pytest.raises(ValueError, match="does not match device type"): + gtfn_module.GTFNTranslationStep( + device_type=gtx.DeviceType.CUDA, code_spec=artifacts.CPPCodeSpec() + ) + + def test_hash_and_diskcache(program_example, tmp_path): fencil, parameters = program_example compilable_program = stages.CompilableProgramDef( diff --git a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/test_roundtrip.py b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/test_roundtrip.py new file mode 100644 index 0000000000..5b70bbbf54 --- /dev/null +++ b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/test_roundtrip.py @@ -0,0 +1,71 @@ +# GT4Py - GridTools Framework +# +# Copyright (c) 2014-2024, ETH Zurich +# All rights reserved. +# +# Please, refer to the LICENSE file in the root directory. +# SPDX-License-Identifier: BSD-3-Clause + +import pytest + +from gt4py.eve import codegen +from gt4py.next import config +from gt4py.next.iterator import ir as itir +from gt4py.next.otf import arguments, stages +from gt4py.next.program_processors.runners import roundtrip + + +@pytest.fixture +def empty_program(): + return itir.Program(id="empty", function_definitions=[], params=[], declarations=[], body=[]) + + +@pytest.fixture +def formatter_spy(monkeypatch): + formatted_sources = [] + + def spy_formatter(source, **kwargs): + formatted_sources.append(source) + return source + + monkeypatch.setattr(codegen, "format_python_source", spy_formatter) + monkeypatch.setattr(roundtrip, "_SOURCE_CACHE", {}) + return formatted_sources + + +@pytest.mark.parametrize("flag", [True, False]) +def test_format_source_default_captures_config_at_creation(monkeypatch, flag): + monkeypatch.setattr(config, "FORMAT_SOURCES", flag) + step = roundtrip.Roundtrip() + monkeypatch.setattr(config, "FORMAT_SOURCES", not flag) + + assert step.format_source is flag + + +@pytest.mark.parametrize("format_source", [True, False]) +def test_generate_source_formatting(empty_program, formatter_spy, format_source): + roundtrip._generate_source( + empty_program, + debug=False, + format_source=format_source, + use_embedded=True, + offset_provider={}, + transforms=lambda ir, offset_provider: ir, + ) + + assert len(formatter_spy) == int(format_source) + + +@pytest.mark.parametrize("format_source", [True, False]) +def test_roundtrip_step_formatting(empty_program, formatter_spy, format_source): + step = roundtrip.Roundtrip( + transforms=lambda ir, offset_provider: ir, format_source=format_source + ) + step( + stages.CompilableProgramDef( + data=empty_program, + args=arguments.CompileTimeArgs.from_concrete(offset_provider={}), + ) + ) + + assert len(formatter_spy) == int(format_source) From a607432da72df30782425c3bbd0faa23fa3dc8f4 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Thu, 24 Sep 2026 17:06:45 +0200 Subject: [PATCH 2/6] Address review comments - Restructure artifacts.format_source for readability. - Normalize format_source for the compiledb prototype at the call site. - Compute all device-dependent settings of GTFNTranslationStep in __post_init__. --- src/gt4py/next/otf/artifacts.py | 16 ++-- .../compilation/build_systems/compiledb.py | 8 +- .../codegens/gtfn/gtfn_module.py | 80 ++++++++----------- .../build_systems_tests/test_compiledb.py | 29 ++++--- 4 files changed, 64 insertions(+), 69 deletions(-) diff --git a/src/gt4py/next/otf/artifacts.py b/src/gt4py/next/otf/artifacts.py index 06fcbcde8f..ab6f35594d 100644 --- a/src/gt4py/next/otf/artifacts.py +++ b/src/gt4py/next/otf/artifacts.py @@ -126,14 +126,14 @@ class HIPCodeSpec(CPPLikeCodeSpec): def format_source(source_code_spec: SourceCodeSpec, source: str) -> str: """Format `source` as configured in `source_code_spec` (no-op if `.format_source` is false).""" - if not source_code_spec.format_source: - return source - assert source_code_spec.formatter_key is not None, ( - "No formatter key specified in source code specification." - ) - return codegen.format_source( - source_code_spec.formatter_key, source, **(source_code_spec.formatter_options or {}) - ) + if source_code_spec.format_source: + assert source_code_spec.formatter_key is not None, ( + "No formatter key specified in source code specification." + ) + source = codegen.format_source( + source_code_spec.formatter_key, source, **(source_code_spec.formatter_options or {}) + ) + return source CodeSpecT = TypeVar("CodeSpecT", bound=SourceCodeSpec) diff --git a/src/gt4py/next/otf/compilation/build_systems/compiledb.py b/src/gt4py/next/otf/compilation/build_systems/compiledb.py index c131605951..a9060075e9 100644 --- a/src/gt4py/next/otf/compilation/build_systems/compiledb.py +++ b/src/gt4py/next/otf/compilation/build_systems/compiledb.py @@ -65,7 +65,9 @@ def __call__( deps=source.library_deps, build_type=self.cmake_build_type, cmake_flags=self.cmake_extra_flags or [], - code_spec=source.program_source.code_spec, + # The compiledb does not depend on source formatting: normalize it so the + # cache folder (keyed by the prototype source) is shared across settings. + code_spec=dataclasses.replace(source.program_source.code_spec, format_source=False), ) compiledb_template = _cc_get_compiledb( @@ -262,9 +264,7 @@ def _cc_prototype_program_source( entry_point=interface.Function(name=name, parameters=()), source_code="", library_deps=deps, - # The compiledb does not depend on source formatting: normalize it so the - # cache folder (keyed by the prototype source) is shared across settings. - code_spec=dataclasses.replace(code_spec, format_source=False), + code_spec=code_spec, ) 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 42822cede5..4e3c1dc814 100644 --- a/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py +++ b/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py @@ -51,10 +51,36 @@ class GTFNTranslationStep( symbolic_domain_sizes: dict[str, itir.Expr] | None = None use_max_domain_range_on_unstructured_shift: bool | None = None + # Device-dependent settings, initialized in `__post_init__` + _backend_header: str = dataclasses.field(init=False, repr=False) + _backend_type: str = dataclasses.field(init=False, repr=False) + _library_name: str = dataclasses.field(init=False, repr=False) + def __post_init__(self) -> None: + default_code_spec: artifacts.HeaderAndSourceCodeSpec + match self.device_type: + case core_defs.DeviceType.CUDA | core_defs.DeviceType.ROCM: + default_code_spec = ( + artifacts.CUDACodeSpec() + if self.device_type == core_defs.DeviceType.CUDA + else artifacts.HIPCodeSpec() + ) + backend_header = "gridtools/fn/backend/gpu.hpp" + backend_type = "gridtools::fn::backend::gpu{}" + library_name = "gridtools_gpu" + case core_defs.DeviceType.CPU: + default_code_spec = artifacts.CPPCodeSpec() + backend_header = "gridtools/fn/backend/naive.hpp" + backend_type = "gridtools::fn::backend::naive{}" + library_name = "gridtools_cpu" + case _: + raise NotImplementedError( + f"{self.__class__.__name__} is not implemented for device type " + f"{self.device_type.name}" + ) + # Resolve the default code spec eagerly, so its settings (e.g. `format_source`, # which follows `config.FORMAT_SOURCES`) are part of the step and its fingerprint. - default_code_spec = self._default_code_spec() if self.code_spec is None: object.__setattr__(self, "code_spec", default_code_spec) elif not isinstance(self.code_spec, type(default_code_spec)): @@ -63,17 +89,9 @@ def __post_init__(self) -> None: f"'{self.device_type.name}' (expected '{type(default_code_spec).__name__}'). " "When replacing the device type, pass 'code_spec=None' to use the default spec." ) - - def _default_code_spec(self) -> artifacts.HeaderAndSourceCodeSpec: - match self.device_type: - case core_defs.DeviceType.CUDA: - return artifacts.CUDACodeSpec() - case core_defs.DeviceType.ROCM: - return artifacts.HIPCodeSpec() - case core_defs.DeviceType.CPU: - return artifacts.CPPCodeSpec() - case _: - raise self._not_implemented_for_device_type() + object.__setattr__(self, "_backend_header", backend_header) + object.__setattr__(self, "_backend_type", backend_type) + object.__setattr__(self, "_library_name", library_name) def _process_regular_arguments( self, @@ -208,7 +226,7 @@ def __call__( # combine into a format that is aligned with what the backend expects parameters: list[interface.Parameter] = regular_parameters + connectivity_parameters - backend_arg = self._backend_type() + backend_arg = self._backend_type args_expr: list[str] = [backend_arg, *regular_args_expr] function = interface.Function(program.id, tuple(parameters)) @@ -225,7 +243,7 @@ def __call__( source_code = artifacts.format_source( code_spec, f""" - #include <{self._backend_header()}> + #include <{self._backend_header}> #include {stencil_src} {decl_src} @@ -235,45 +253,13 @@ def __call__( module: artifacts.ProgramSource[artifacts.HeaderAndSourceCodeSpec] = ( artifacts.ProgramSource( entry_point=function, - library_deps=(interface.LibraryDependency(self._library_name(), "master"),), + library_deps=(interface.LibraryDependency(self._library_name, "master"),), source_code=source_code, code_spec=code_spec, ) ) return module - def _backend_header(self) -> str: - match self.device_type: - case core_defs.DeviceType.CUDA | core_defs.DeviceType.ROCM: - return "gridtools/fn/backend/gpu.hpp" - case core_defs.DeviceType.CPU: - return "gridtools/fn/backend/naive.hpp" - case _: - raise self._not_implemented_for_device_type() - - def _backend_type(self) -> str: - match self.device_type: - case core_defs.DeviceType.CUDA | core_defs.DeviceType.ROCM: - return "gridtools::fn::backend::gpu{}" - case core_defs.DeviceType.CPU: - return "gridtools::fn::backend::naive{}" - case _: - raise self._not_implemented_for_device_type() - - def _library_name(self) -> str: - match self.device_type: - case core_defs.DeviceType.CUDA | core_defs.DeviceType.ROCM: - return "gridtools_gpu" - case core_defs.DeviceType.CPU: - return "gridtools_cpu" - case _: - raise self._not_implemented_for_device_type() - - def _not_implemented_for_device_type(self) -> NotImplementedError: - return NotImplementedError( - f"{self.__class__.__name__} is not implemented for device type {self.device_type.name}" - ) - class GTFNTranslationStepFactory(factory.Factory[GTFNTranslationStep]): class Meta: diff --git a/tests/next_tests/unit_tests/otf_tests/compilation_tests/build_systems_tests/test_compiledb.py b/tests/next_tests/unit_tests/otf_tests/compilation_tests/build_systems_tests/test_compiledb.py index 7df35c1889..9edc6c0a59 100644 --- a/tests/next_tests/unit_tests/otf_tests/compilation_tests/build_systems_tests/test_compiledb.py +++ b/tests/next_tests/unit_tests/otf_tests/compilation_tests/build_systems_tests/test_compiledb.py @@ -6,6 +6,7 @@ # Please, refer to the LICENSE file in the root directory. # SPDX-License-Identifier: BSD-3-Clause +import dataclasses import pathlib import shutil import tempfile @@ -13,7 +14,6 @@ import pytest from gt4py.next import config, fingerprinting -from gt4py.next.otf import artifacts from gt4py.next.otf.compilation import build_data, cache, importer from gt4py.next.otf.compilation.build_systems import compiledb @@ -94,16 +94,25 @@ def test_compiledb_project_is_relocatable(extension_source_example, clean_compil ) -def test_compiledb_prototype_ignores_format_source(): - prototypes = [ - compiledb._cc_prototype_program_source( - deps=(), - build_type=config.CMakeBuildType.RELEASE, - cmake_flags=[], - code_spec=artifacts.CPPCodeSpec(format_source=format_source), +def test_compiledb_prototype_ignores_format_source(monkeypatch, extension_source_example): + prototypes = [] + + def fake_get_compiledb(renew_compiledb, prototype_program_source, **kwargs): + prototypes.append(prototype_program_source) + return pathlib.Path("compile_commands.json") + + monkeypatch.setattr(compiledb, "_cc_get_compiledb", fake_get_compiledb) + + program_source = extension_source_example.program_source + for format_source in (True, False): + code_spec = dataclasses.replace(program_source.code_spec, format_source=format_source) + compiledb.CompiledbFactory()( + dataclasses.replace( + extension_source_example, + program_source=dataclasses.replace(program_source, code_spec=code_spec), + ), + cache_lifetime=config.BuildCacheLifetime.SESSION, ) - for format_source in (True, False) - ] assert fingerprinting.strict_fingerprinter( prototypes[0] From e70db7ce9a934d64150278f20f30fdd5de52dd64 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Thu, 24 Sep 2026 21:39:27 +0200 Subject: [PATCH 3/6] Move device-dependent gtfn settings to module-level mappings --- .../codegens/gtfn/gtfn_module.py | 70 ++++++++++--------- 1 file changed, 37 insertions(+), 33 deletions(-) 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 4e3c1dc814..3eba92252f 100644 --- a/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py +++ b/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py @@ -9,6 +9,7 @@ from __future__ import annotations import dataclasses +from collections.abc import Callable from typing import Any, Final, Optional import factory @@ -28,6 +29,31 @@ GENERATED_CONNECTIVITY_PARAM_PREFIX = "gt_conn_" +# Device-dependent settings of the generated code. Code specs are stored as factories +# (called per translation step) since their defaults depend on `config`. +_DEFAULT_CODE_SPEC_FACTORIES: Final[ + dict[core_defs.DeviceType, Callable[[], artifacts.HeaderAndSourceCodeSpec]] +] = { + core_defs.DeviceType.CPU: artifacts.CPPCodeSpec, + core_defs.DeviceType.CUDA: artifacts.CUDACodeSpec, + core_defs.DeviceType.ROCM: artifacts.HIPCodeSpec, +} +_BACKEND_HEADERS: Final[dict[core_defs.DeviceType, str]] = { + core_defs.DeviceType.CPU: "gridtools/fn/backend/naive.hpp", + core_defs.DeviceType.CUDA: "gridtools/fn/backend/gpu.hpp", + core_defs.DeviceType.ROCM: "gridtools/fn/backend/gpu.hpp", +} +_BACKEND_TYPES: Final[dict[core_defs.DeviceType, str]] = { + core_defs.DeviceType.CPU: "gridtools::fn::backend::naive{}", + core_defs.DeviceType.CUDA: "gridtools::fn::backend::gpu{}", + core_defs.DeviceType.ROCM: "gridtools::fn::backend::gpu{}", +} +_LIBRARY_NAMES: Final[dict[core_defs.DeviceType, str]] = { + core_defs.DeviceType.CPU: "gridtools_cpu", + core_defs.DeviceType.CUDA: "gridtools_gpu", + core_defs.DeviceType.ROCM: "gridtools_gpu", +} + def get_param_description(name: str, type_: Any) -> interface.Parameter: return interface.Parameter(name, type_) @@ -51,36 +77,15 @@ class GTFNTranslationStep( symbolic_domain_sizes: dict[str, itir.Expr] | None = None use_max_domain_range_on_unstructured_shift: bool | None = None - # Device-dependent settings, initialized in `__post_init__` - _backend_header: str = dataclasses.field(init=False, repr=False) - _backend_type: str = dataclasses.field(init=False, repr=False) - _library_name: str = dataclasses.field(init=False, repr=False) - def __post_init__(self) -> None: - default_code_spec: artifacts.HeaderAndSourceCodeSpec - match self.device_type: - case core_defs.DeviceType.CUDA | core_defs.DeviceType.ROCM: - default_code_spec = ( - artifacts.CUDACodeSpec() - if self.device_type == core_defs.DeviceType.CUDA - else artifacts.HIPCodeSpec() - ) - backend_header = "gridtools/fn/backend/gpu.hpp" - backend_type = "gridtools::fn::backend::gpu{}" - library_name = "gridtools_gpu" - case core_defs.DeviceType.CPU: - default_code_spec = artifacts.CPPCodeSpec() - backend_header = "gridtools/fn/backend/naive.hpp" - backend_type = "gridtools::fn::backend::naive{}" - library_name = "gridtools_cpu" - case _: - raise NotImplementedError( - f"{self.__class__.__name__} is not implemented for device type " - f"{self.device_type.name}" - ) - + if (code_spec_factory := _DEFAULT_CODE_SPEC_FACTORIES.get(self.device_type)) is None: + raise NotImplementedError( + f"{self.__class__.__name__} is not implemented for device type " + f"{self.device_type.name}" + ) # Resolve the default code spec eagerly, so its settings (e.g. `format_source`, # which follows `config.FORMAT_SOURCES`) are part of the step and its fingerprint. + default_code_spec = code_spec_factory() if self.code_spec is None: object.__setattr__(self, "code_spec", default_code_spec) elif not isinstance(self.code_spec, type(default_code_spec)): @@ -89,9 +94,6 @@ def __post_init__(self) -> None: f"'{self.device_type.name}' (expected '{type(default_code_spec).__name__}'). " "When replacing the device type, pass 'code_spec=None' to use the default spec." ) - object.__setattr__(self, "_backend_header", backend_header) - object.__setattr__(self, "_backend_type", backend_type) - object.__setattr__(self, "_library_name", library_name) def _process_regular_arguments( self, @@ -226,7 +228,7 @@ def __call__( # combine into a format that is aligned with what the backend expects parameters: list[interface.Parameter] = regular_parameters + connectivity_parameters - backend_arg = self._backend_type + backend_arg = _BACKEND_TYPES[self.device_type] args_expr: list[str] = [backend_arg, *regular_args_expr] function = interface.Function(program.id, tuple(parameters)) @@ -243,7 +245,7 @@ def __call__( source_code = artifacts.format_source( code_spec, f""" - #include <{self._backend_header}> + #include <{_BACKEND_HEADERS[self.device_type]}> #include {stencil_src} {decl_src} @@ -253,7 +255,9 @@ def __call__( module: artifacts.ProgramSource[artifacts.HeaderAndSourceCodeSpec] = ( artifacts.ProgramSource( entry_point=function, - library_deps=(interface.LibraryDependency(self._library_name, "master"),), + library_deps=( + interface.LibraryDependency(_LIBRARY_NAMES[self.device_type], "master"), + ), source_code=source_code, code_spec=code_spec, ) From b76226deddc4b544fcb6b50b82e60dfca8d37dbe Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Thu, 24 Sep 2026 21:47:20 +0200 Subject: [PATCH 4/6] Key the roundtrip module cache by source and debug mode --- .../next/program_processors/runners/roundtrip.py | 11 ++++++----- .../runners_tests/test_roundtrip.py | 14 ++++++++++++++ 2 files changed, 20 insertions(+), 5 deletions(-) diff --git a/src/gt4py/next/program_processors/runners/roundtrip.py b/src/gt4py/next/program_processors/runners/roundtrip.py index 2ef7122f49..9e1eeb42ac 100644 --- a/src/gt4py/next/program_processors/runners/roundtrip.py +++ b/src/gt4py/next/program_processors/runners/roundtrip.py @@ -110,8 +110,9 @@ def visit_Temporary(self, node: itir.Temporary, **kwargs: Any) -> str: # Caches the generated source by IR hash so re-codegen is skipped within a process. _SOURCE_CACHE: dict[int, tuple[str, str]] = {} -# Caches the loaded module by source string so re-exec is skipped within a process. -_MODULE_CACHE: dict[str, types.ModuleType] = {} +# Caches the loaded module by source string and debug mode (debug modules are loaded +# from a temporary file) so re-exec is skipped within a process. +_MODULE_CACHE: dict[tuple[str, bool], types.ModuleType] = {} def _generate_source( @@ -187,8 +188,8 @@ def _generate_source( def _load_module(source_code: str, debug: bool) -> types.ModuleType: - if source_code in _MODULE_CACHE: - return _MODULE_CACHE[source_code] + if (cache_key := (source_code, debug)) in _MODULE_CACHE: + return _MODULE_CACHE[cache_key] if debug: # Write to a real .py so debuggers/tracebacks have file/line info. @@ -205,7 +206,7 @@ def _load_module(source_code: str, debug: bool) -> types.ModuleType: mod = types.ModuleType("roundtrip_module") exec(compile(source_code, "", "exec"), mod.__dict__) - _MODULE_CACHE[source_code] = mod + _MODULE_CACHE[cache_key] = mod return mod diff --git a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/test_roundtrip.py b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/test_roundtrip.py index 5b70bbbf54..2f0b15ff29 100644 --- a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/test_roundtrip.py +++ b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/test_roundtrip.py @@ -69,3 +69,17 @@ def test_roundtrip_step_formatting(empty_program, formatter_spy, format_source): ) assert len(formatter_spy) == int(format_source) + + +def test_load_module_cache_respects_debug_mode(monkeypatch, tmp_path): + monkeypatch.setattr(roundtrip, "_MODULE_CACHE", {}) + monkeypatch.setattr(roundtrip.tempfile, "tempdir", str(tmp_path)) + source_code = "x = 1\n" + + roundtrip._load_module(source_code, debug=False) + assert not list(tmp_path.glob("*.py")) + + # loading the same source in debug mode must still write the temporary `.py` file + debug_module = roundtrip._load_module(source_code, debug=True) + assert list(tmp_path.glob("*.py")) + assert debug_module.__file__.startswith(str(tmp_path)) From 845f05f8f64b22a1305ad2cfe24fa86aff790d74 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Fri, 25 Sep 2026 08:45:34 +0200 Subject: [PATCH 5/6] Pin SDFGCodeSpec.format_source and test format_cpp formatting --- src/gt4py/next/otf/artifacts.py | 4 ++-- .../unit_tests/otf_tests/test_languages.py | 5 +++++ .../gtfn_tests/test_gtfn_module.py | 19 +++++++++++++++++++ 3 files changed, 26 insertions(+), 2 deletions(-) diff --git a/src/gt4py/next/otf/artifacts.py b/src/gt4py/next/otf/artifacts.py index ab6f35594d..ba59ed1191 100644 --- a/src/gt4py/next/otf/artifacts.py +++ b/src/gt4py/next/otf/artifacts.py @@ -76,8 +76,8 @@ class SDFGCodeSpec(SourceCodeSpec): source_language: str = "SDFG" file_extension: str = "sdfg" - # There is no SDFG formatter: keep the spec independent of `config.FORMAT_SOURCES` - format_source: bool = False + # There is no SDFG formatter: pin to `False`, independent of `config.FORMAT_SOURCES` + format_source: bool = dataclasses.field(default=False, init=False) @dataclasses.dataclass(frozen=True, kw_only=True) diff --git a/tests/next_tests/unit_tests/otf_tests/test_languages.py b/tests/next_tests/unit_tests/otf_tests/test_languages.py index 0815a41deb..7c69c73805 100644 --- a/tests/next_tests/unit_tests/otf_tests/test_languages.py +++ b/tests/next_tests/unit_tests/otf_tests/test_languages.py @@ -48,6 +48,11 @@ def test_sdfg_code_spec_ignores_format_sources_config(monkeypatch, flag): assert artifacts.SDFGCodeSpec().format_source is False +def test_sdfg_code_spec_rejects_format_source(): + with pytest.raises(TypeError): + artifacts.SDFGCodeSpec(format_source=True) + + def test_format_source_follows_code_spec(): source = "x=( 1 )" 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 f4b473f362..5ac71b3af2 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 @@ -13,11 +13,13 @@ import pytest import gt4py.next as gtx +from gt4py.eve import codegen from gt4py.next.iterator import builtins, ir as itir from gt4py.next.iterator.ir_utils import ir_makers as im from gt4py.next import config, fingerprinting from gt4py.next.otf import arguments, artifacts, stages from gt4py.next.program_processors.codegens.gtfn import gtfn_module +from gt4py.next.program_processors.formatters import gtfn as gtfn_formatters from gt4py.next.program_processors.runners import gtfn from gt4py.next.type_system import type_translation from gt4py.next import custom_layout_allocators as next_allocators @@ -100,6 +102,23 @@ def test_code_spec_is_resolved_at_construction(monkeypatch, flag): ) +@pytest.mark.parametrize("flag", [True, False]) +def test_format_cpp_always_formats_once(monkeypatch, program_example, flag): + monkeypatch.setattr(config, "FORMAT_SOURCES", flag) + formatted = [] + + def spy_format_source(language, source, **kwargs): + formatted.append((language, kwargs)) + return source + + monkeypatch.setattr(codegen, "format_source", spy_format_source) + + program, _ = program_example + gtfn_formatters.format_cpp(program, offset_provider={}) + + assert formatted == [("cpp", {"style": "LLVM"})] + + def test_code_spec_device_type_mismatch(): with pytest.raises(ValueError, match="does not match device type"): gtfn_module.GTFNTranslationStep( From 07d1973963d8d374089ea28f9430e51304e3bb17 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Fri, 25 Sep 2026 08:53:08 +0200 Subject: [PATCH 6/6] Test GT4PY_FORMAT_SOURCES precedence over GT4PY_DEBUG --- tests/next_tests/unit_tests/test_config.py | 30 ++++++++++++++++++++++ 1 file changed, 30 insertions(+) diff --git a/tests/next_tests/unit_tests/test_config.py b/tests/next_tests/unit_tests/test_config.py index a33bd5734a..2f23a6fdf3 100644 --- a/tests/next_tests/unit_tests/test_config.py +++ b/tests/next_tests/unit_tests/test_config.py @@ -6,6 +6,7 @@ # Please, refer to the LICENSE file in the root directory. # SPDX-License-Identifier: BSD-3-Clause +import importlib.util import os import pytest @@ -46,3 +47,32 @@ def test_env_flag_to_bool_invalid(env_var): def test_env_flag_to_bool_unset(env_var): _ = os.environ.pop(env_var, None) assert config.env_flag_to_bool(env_var, default=False) is False + + +def _load_fresh_config(): + """Execute a fresh copy of the `config` module, leaving `gt4py.next.config` untouched.""" + spec = importlib.util.spec_from_file_location("_fresh_gt4py_next_config", config.__file__) + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return module + + +@pytest.mark.parametrize( + "debug, format_sources, expected", + [ + (None, None, False), + ("1", None, True), + ("0", "1", True), + ("1", "0", False), + ], +) +def test_format_sources_precedence(monkeypatch, debug, format_sources, expected): + # avoid adding global warning filters when executing the module + monkeypatch.setenv("GT4PY_SKIP_DACE_WARNINGS", "0") + for name, value in (("GT4PY_DEBUG", debug), ("GT4PY_FORMAT_SOURCES", format_sources)): + if value is None: + monkeypatch.delenv(name, raising=False) + else: + monkeypatch.setenv(name, value) + + assert _load_fresh_config().FORMAT_SOURCES is expected