diff --git a/docs/development/ADRs/next/0012-GridTools_Cpp_OTF_Steps.md b/docs/development/ADRs/next/0012-GridTools_Cpp_OTF_Steps.md index 4c7a132d24..77b9dabdc0 100644 --- a/docs/development/ADRs/next/0012-GridTools_Cpp_OTF_Steps.md +++ b/docs/development/ADRs/next/0012-GridTools_Cpp_OTF_Steps.md @@ -10,7 +10,7 @@ tags: [backend, cpp, gridtools, bindings, otf] - **Updated**: 2026-09-15 > [!NOTE] -> The entry points named below have since been renamed or moved: `program_processors.formatters.gtfn.format_sourcecode` is now `format_cpp`; `program_processors.codegens.gtfn_modules.translate_program` is now the `GTFNTranslationStep` class in `program_processors.codegens.gtfn.gtfn_module`; `otf.binding.pybind.bind_source` is now `otf.binding.nanobind.create_bindings` (nanobind replaced pybind11); and `program_processors.runners.gtfn_cpu.run_gtfn` is now `program_processors.runners.gtfn.run_gtfn`. `otf.step_types` was merged into `otf.stages`, `otf.workflow.Step` is now the `otf.workflow.Workflow` protocol, and `processor_interface` was split up — `ProgramFormatter` lives in `program_processors.program_formatter`. The decision — building the GTFN backend out of composable OTF steps — is unchanged. +> The entry points named below have since been renamed or moved: `program_processors.formatters.gtfn.format_sourcecode` is now `format_cpp`; `program_processors.codegens.gtfn_modules.translate_program` is now the `GTFNTranslationStep` class in `program_processors.codegens.gtfn.gtfn_module`; `otf.binding.pybind.bind_source` is now `otf.binding.nanobind.create_bindings` (nanobind replaced pybind11); and `program_processors.runners.gtfn_cpu.run_gtfn` is now `program_processors.runners.gtfn.run_gtfn`. `otf.step_types` was merged into `otf.stages`, `otf.workflow.Step` is now the `otf.workflow.Workflow` protocol, and `processor_interface` was split up — `ProgramFormatter` lives in `program_processors.program_formatter`. The decision — building the GTFN backend out of composable OTF steps — is unchanged. `SourceCodeSpec` no longer carries formatting options, and generated code is no longer auto-formatted in `gt4py.next`. ## Context diff --git a/pyproject.toml b/pyproject.toml index e131ebda0e..8eecbcd3e1 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -43,6 +43,8 @@ profiling = [ ] scripts = ["pyyaml>=6.0.1", "typer>=0.16.0", "packaging"] test = [ + 'black>=25.11', + 'clang-format>=18.1', 'factory-boy>=3.3.3', 'hypothesis>=6.0.0', 'nbmake>=1.4.6', @@ -91,7 +93,6 @@ classifiers = [ dependencies = [ 'attrs>=21.3', 'array-api-compat>=1.13', - 'black>=25.11', 'boltons>=20.1', 'cached-property>=1.5.1', 'click>=8.0.0', @@ -142,7 +143,7 @@ readme = 'README.md' requires-python = '>=3.12, <3.15, !=3.13.10, !=3.14.1' [project.optional-dependencies] -cartesian = ['gt4py[jax,standard,testing]'] +cartesian = ['black>=25.11', 'clang-format>=18.1', 'gt4py[jax,standard,testing]'] cuda12 = ['cupy-cuda12x>=12.0'] cuda13 = ['cupy-cuda13x>=14.0'] jax = [ @@ -154,7 +155,7 @@ jax-cuda13 = ['jax[cuda13-local]>=0.7.0', 'gt4py[cuda13]'] next = ['gt4py[jax,standard,testing]'] rocm6 = ['cupy>=13.4.1,<14.0'] rocm7 = ['cupy-rocm-7-0>=14.0'] -standard = ['clang-format>=18.1', 'scipy>=1.16.1'] +standard = ['scipy>=1.16.1'] testing = ['hypothesis>=6.93', 'pytest>=7.0'] [project.urls] diff --git a/src/gt4py/cartesian/backend/dace_backend.py b/src/gt4py/cartesian/backend/dace_backend.py index 3bd76bcfeb..bf3e09e592 100644 --- a/src/gt4py/cartesian/backend/dace_backend.py +++ b/src/gt4py/cartesian/backend/dace_backend.py @@ -45,7 +45,7 @@ from gt4py.cartesian.gtc.passes.oir_optimizations.utils import compute_fields_extents from gt4py.cartesian.gtc.passes.oir_pipeline import DefaultPipeline from gt4py.cartesian.utils import shash -from gt4py.eve import codegen +from gt4py.eve import formatting from gt4py.eve.codegen import MakoTemplate as as_mako from gt4py.storage.cartesian import layout, layout_registry @@ -679,7 +679,7 @@ def apply(cls, builder: StencilBuilder, sdfg: SDFG) -> str: """ if builder.options.format_source: - generated_code = codegen.format_source("cpp", generated_code, style="LLVM") + generated_code = formatting.format_cpp_source(generated_code) return generated_code @@ -863,7 +863,7 @@ def generate_sdfg_bindings(self, sdfg: SDFG, module_name: str) -> str: def apply(cls, sdfg: SDFG, module_name: str, *, backend: BaseDaceBackend) -> str: generated_code = cls(backend).generate_sdfg_bindings(sdfg, module_name) if backend.builder.options.format_source: - generated_code = codegen.format_source("cpp", generated_code, style="LLVM") + generated_code = formatting.format_cpp_source(generated_code) return generated_code diff --git a/src/gt4py/cartesian/backend/debug_backend.py b/src/gt4py/cartesian/backend/debug_backend.py index 08d139d3f5..5892cabfad 100644 --- a/src/gt4py/cartesian/backend/debug_backend.py +++ b/src/gt4py/cartesian/backend/debug_backend.py @@ -16,7 +16,7 @@ from gt4py.cartesian.gtc.debug.debug_codegen import DebugCodeGen from gt4py.cartesian.gtc.gtir_to_oir import GTIRToOIR from gt4py.cartesian.gtc.passes import oir_optimizations -from gt4py.eve import codegen +from gt4py.eve import formatting from gt4py.storage import layout from gt4py.storage.cartesian import layout_registry @@ -45,7 +45,7 @@ def _generate_computation(self) -> dict[str, str | dict]: source_code = DebugCodeGen().visit(oir) if self.builder.options.format_source: - source_code = codegen.format_source("python", source_code) + source_code = formatting.format_python_source(source_code) caching = self.builder.caching computation_name = f"{caching.module_prefix}computation{caching.module_postfix}.py" diff --git a/src/gt4py/cartesian/backend/gtcpp_backend.py b/src/gt4py/cartesian/backend/gtcpp_backend.py index d77c2fb37f..92305a3314 100644 --- a/src/gt4py/cartesian/backend/gtcpp_backend.py +++ b/src/gt4py/cartesian/backend/gtcpp_backend.py @@ -24,7 +24,7 @@ from gt4py.cartesian.gtc.gtcpp.oir_to_gtcpp import OIRToGTCpp from gt4py.cartesian.gtc.gtir_to_oir import GTIRToOIR from gt4py.cartesian.gtc.passes.oir_pipeline import DefaultPipeline -from gt4py.eve import codegen +from gt4py.eve import codegen, formatting from gt4py.storage.cartesian import layout, layout_registry @@ -117,7 +117,7 @@ def visit_Program(self, node: gtcpp.Program, **kwargs): def apply(cls, root, *, module_name="stencil", **kwargs) -> str: generated_code = cls(kwargs.get("backend")).visit(root, module_name=module_name, **kwargs) if kwargs.get("format_source", True): - generated_code = codegen.format_source("cpp", generated_code, style="LLVM") + generated_code = formatting.format_cpp_source(generated_code) return generated_code diff --git a/src/gt4py/cartesian/backend/module_generator.py b/src/gt4py/cartesian/backend/module_generator.py index 4ce5279feb..58f5c4f3c5 100644 --- a/src/gt4py/cartesian/backend/module_generator.py +++ b/src/gt4py/cartesian/backend/module_generator.py @@ -25,7 +25,7 @@ from gt4py.cartesian.gtc.passes.oir_access_kinds import compute_access_kinds from gt4py.cartesian.gtc.passes.oir_optimizations.utils import compute_fields_extents from gt4py.cartesian.gtc.utils import dimension_flags_to_names -from gt4py.eve import codegen +from gt4py.eve import formatting if TYPE_CHECKING: @@ -150,8 +150,8 @@ def __call__(self, args_data: ModuleData) -> str: implementation=self.generate_implementation(), ) if self.builder.options.as_dict()["format_source"]: - module_source = codegen.format_source( - "python", module_source, line_length=self.SOURCE_LINE_LENGTH + module_source = formatting.format_python_source( + module_source, line_length=self.SOURCE_LINE_LENGTH ) return module_source @@ -208,7 +208,7 @@ def generate_sources(self) -> dict[str, str]: """ if self.builder.gtir.sources is not None: return { - key: codegen.format_source("python", value, line_length=self.SOURCE_LINE_LENGTH) + key: formatting.format_python_source(value, line_length=self.SOURCE_LINE_LENGTH) for key, value in self.builder.gtir.sources.items() } return {} diff --git a/src/gt4py/cartesian/backend/numpy_backend.py b/src/gt4py/cartesian/backend/numpy_backend.py index c1c1c9e0b4..335feccf04 100644 --- a/src/gt4py/cartesian/backend/numpy_backend.py +++ b/src/gt4py/cartesian/backend/numpy_backend.py @@ -16,7 +16,7 @@ from gt4py.cartesian.gtc.gtir_to_oir import GTIRToOIR from gt4py.cartesian.gtc.numpy import npir from gt4py.cartesian.gtc.passes import oir_optimizations as oir_opt -from gt4py.eve import codegen +from gt4py.eve import formatting from gt4py.storage import layout from gt4py.storage.cartesian import layout_registry @@ -44,7 +44,7 @@ def generate_computation(self) -> dict[str, str | dict]: source = numpy.NpirCodegen.apply(self.npir, ignore_np_errstate=ignore_np_errstate) if self.builder.options.format_source: - source = codegen.format_source("python", source) + source = formatting.format_python_source(source) caching = self.builder.caching computation_name = f"{caching.module_prefix}computation{caching.module_postfix}.py" diff --git a/src/gt4py/cartesian/gtc/gtcpp/gtcpp_codegen.py b/src/gt4py/cartesian/gtc/gtcpp/gtcpp_codegen.py index 62eae238a9..5f0f30468d 100644 --- a/src/gt4py/cartesian/gtc/gtcpp/gtcpp_codegen.py +++ b/src/gt4py/cartesian/gtc/gtcpp/gtcpp_codegen.py @@ -20,7 +20,7 @@ UnaryOperator, ) from gt4py.cartesian.gtc.gtcpp import gtcpp -from gt4py.eve import codegen +from gt4py.eve import codegen, formatting from gt4py.eve.codegen import FormatTemplate as as_fmt, MakoTemplate as as_mako from gt4py.eve.concepts import LeafNode @@ -326,6 +326,6 @@ def apply(cls, root: LeafNode, **kwargs: Any) -> str: raise TypeError("apply() missing 1 required keyword-only argument: 'gt_backend_t'") generated_code = super().apply(root, offset_limit=_offset_limit(root), **kwargs) if kwargs.get("format_source", True): - generated_code = codegen.format_source("cpp", generated_code, style="LLVM") + generated_code = formatting.format_cpp_source(generated_code) return generated_code diff --git a/src/gt4py/eve/__init__.py b/src/gt4py/eve/__init__.py index 33295cefc6..70a4c999ad 100644 --- a/src/gt4py/eve/__init__.py +++ b/src/gt4py/eve/__init__.py @@ -11,7 +11,7 @@ The internal dependencies between modules are the following (each module depends on some of the previous ones): - 0. xtyping + 0. formatting, xtyping 1. exceptions, pattern_matching, type_definitions 2. utils 3. type_validation diff --git a/src/gt4py/eve/codegen.py b/src/gt4py/eve/codegen.py index e6e64adb55..e3ba40302e 100644 --- a/src/gt4py/eve/codegen.py +++ b/src/gt4py/eve/codegen.py @@ -14,17 +14,14 @@ import collections.abc import contextlib import inspect -import os import re import string -import subprocess import sys import textwrap import types -from collections.abc import Callable, Collection, Iterator, Mapping, Sequence +from collections.abc import Collection, Iterator, Mapping, Sequence from typing import Any, ClassVar, Optional, TypeVar, Union, overload -import black import jinja2 from mako import template as mako_tpl from typing_extensions import Protocol, runtime_checkable @@ -34,24 +31,6 @@ from .visitors import NodeVisitor -SourceFormatter = Callable[[str], str] - -SOURCE_FORMATTERS: dict[str, SourceFormatter] = {} -"""Global dict storing registered formatters.""" - - -class FormatterNameError(exceptions.EveRuntimeError): - """Run-time error registering a new source code formatter.""" - - ... - - -class FormattingError(exceptions.EveRuntimeError): - """Run-time error applying a source code formatter.""" - - ... - - class TemplateDefinitionError(exceptions.EveTypeError): """Template definition error.""" @@ -64,111 +43,6 @@ class TemplateRenderingError(exceptions.EveRuntimeError): ... -def register_formatter(language: str) -> Callable[[SourceFormatter], SourceFormatter]: - """Register source code formatters for specific languages (decorator).""" - - def _decorator(formatter: SourceFormatter) -> SourceFormatter: - if language in SOURCE_FORMATTERS: - raise FormatterNameError(f"Another formatter for language '{language}' already exists") - - assert callable(formatter) - SOURCE_FORMATTERS[language] = formatter - - return formatter - - return _decorator - - -@register_formatter("python") -def format_python_source( - source: str, - *, - line_length: int = 100, - python_versions: Optional[set[str]] = None, - string_normalization: bool = True, -) -> str: - """Format Python source code using black formatter.""" - python_versions = python_versions or {f"{sys.version_info.major}{sys.version_info.minor}"} - target_versions = set(black.TargetVersion[f"PY{v.replace('.', '')}"] for v in python_versions) # type: ignore[attr-defined] # .TargetVersion implicitly exported - - formatted_source = black.format_str( - source, - mode=black.FileMode( - line_length=line_length, - target_versions=target_versions, - string_normalization=string_normalization, - ), - ) - assert isinstance(formatted_source, str) - - return formatted_source - - -def _get_clang_format() -> Optional[str]: - """Return the clang-format executable, or None if not available.""" - executable = os.getenv("CLANG_FORMAT_EXECUTABLE", "clang-format") - try: - assert isinstance(executable, str) - if subprocess.run([executable, "--version"], capture_output=True).returncode != 0: - return None - except Exception: - return None - - return executable - - -_CLANG_FORMAT_EXECUTABLE = _get_clang_format() - - -if _CLANG_FORMAT_EXECUTABLE is not None: - - @register_formatter("cpp") - def format_cpp_source( - source: str, - *, - style: Optional[str] = None, - fallback_style: Optional[str] = None, - sort_includes: bool = False, - ) -> str: - """Format C++ source code using clang-format.""" - assert isinstance(_CLANG_FORMAT_EXECUTABLE, str) - args = [_CLANG_FORMAT_EXECUTABLE, "--assume-filename=_gt4py_generated_file.cpp"] - if style: - args.append(f"--style={style}") - if fallback_style: - args.append(f"--fallback-style={style}") - if sort_includes: - args.append("--sort-includes") - - try: - # use a timeout as clang-format used to deadlock on some sources - formatted_source = subprocess.run( - args, check=True, input=source, capture_output=True, text=True, timeout=3 - ).stdout - except subprocess.TimeoutExpired: - return source - - assert isinstance(formatted_source, str) - return formatted_source - - -def format_source(language: str, source: str, *, skip_errors: bool = True, **kwargs: Any) -> str: - """Format source code if a formatter exists for the specific language.""" - formatter = SOURCE_FORMATTERS.get(language, None) - try: - if formatter: - return formatter(source, **kwargs) - else: - raise FormattingError(f"Missing formatter for '{language}' language") - except Exception as e: - if skip_errors: - return source - else: - raise FormattingError( - f"Something went wrong when trying to format '{language}' source code" - ) from e - - class Name: """Text formatter with different case styles for symbol names in source code.""" diff --git a/src/gt4py/eve/formatting.py b/src/gt4py/eve/formatting.py new file mode 100644 index 0000000000..4ba9612696 --- /dev/null +++ b/src/gt4py/eve/formatting.py @@ -0,0 +1,94 @@ +# 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 + +"""Best-effort formatting of generated source code for human readers.""" + +from __future__ import annotations + +import functools +import os +import subprocess +import sys + + +def format_python_source(source: str, *, line_length: int = 100) -> str: + """Format Python source code with `black`. + + The target Python version is pinned to the running interpreter. + + Args: + source: Python source code. + line_length: Maximum line length of the formatted code. + + Returns: + The formatted source code, or `source` unchanged if `black` is not + installed, does not support the running interpreter, or formatting + fails. + """ + try: + # Lazy import: `black` is an optional dependency and slow to import. + import black + except ImportError: + return source + + try: + target_version = black.TargetVersion[ # type: ignore[attr-defined] # .TargetVersion implicitly exported + f"PY{sys.version_info.major}{sys.version_info.minor}" + ] + except KeyError: # `black` is too old to know the running interpreter + return source + + try: + return black.format_str( + source, mode=black.Mode(line_length=line_length, target_versions={target_version}) + ) + except Exception: # e.g. `black.InvalidInput` for unparsable source, or internal `black` errors + return source + + +def format_cpp_source(source: str) -> str: + """Format C++ source code with `clang-format` using the LLVM style. + + The executable can be overridden with the `CLANG_FORMAT_EXECUTABLE` + environment variable. Its availability is checked once, on first use. + + Args: + source: C++ source code. + + Returns: + The formatted source code, or `source` unchanged if `clang-format` is + not available or formatting fails. + """ + executable = _get_clang_format() + if executable is None: + return source + + args = [executable, "--style=LLVM", "--assume-filename=_gt4py_generated_file.cpp"] + try: + # use a timeout as clang-format used to deadlock on some sources + return subprocess.run( + args, input=source, capture_output=True, encoding="utf-8", check=True, timeout=3 + ).stdout + except (OSError, subprocess.SubprocessError, UnicodeError): + return source + + +@functools.cache +def _get_clang_format() -> str | None: + """Return the `clang-format` executable, or `None` if it is not available. + + The result is cached, so the executable is probed only once. + """ + executable = os.getenv("CLANG_FORMAT_EXECUTABLE", "clang-format") + try: + if subprocess.run([executable, "--version"], capture_output=True).returncode != 0: + return None + except Exception: + return None + + return executable diff --git a/src/gt4py/next/otf/artifacts.py b/src/gt4py/next/otf/artifacts.py index 204e2e75d0..fa3226153b 100644 --- a/src/gt4py/next/otf/artifacts.py +++ b/src/gt4py/next/otf/artifacts.py @@ -25,28 +25,18 @@ from __future__ import annotations import dataclasses -import functools -from collections.abc import Callable, Mapping -from typing import Any, Generic, Optional, Protocol, TypeAlias, TypeVar, runtime_checkable +from collections.abc import Callable +from typing import Generic, Optional, Protocol, TypeAlias, TypeVar, runtime_checkable -from gt4py.eve import codegen from gt4py.next.otf.binding import interface @dataclasses.dataclass(frozen=True, kw_only=True) 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``. - """ + """Basic settings for any source programming language.""" source_language: str file_extension: str - formatter_key: str | None = None - formatter_options: Mapping[str, Any] | None = None @dataclasses.dataclass(frozen=True, kw_only=True) @@ -62,7 +52,6 @@ class PythonCodeSpec(SourceCodeSpec): source_language: str = "python" file_extension: str = "py" - formatter_key: str = "python" @dataclasses.dataclass(frozen=True, kw_only=True) @@ -85,10 +74,6 @@ class CPPCodeSpec(CPPLikeCodeSpec): source_language: str = "CXX" file_extension: str = "cpp" header_extension: str = "hpp" - formatter_key: str = "cpp" - formatter_options: Mapping[str, Any] = dataclasses.field( - default_factory=functools.partial(dict, style="LLVM") - ) @dataclasses.dataclass(frozen=True, kw_only=True) @@ -98,10 +83,6 @@ class CUDACodeSpec(CPPLikeCodeSpec): source_language: str = "CUDA" file_extension: str = "cu" header_extension: str = "cuh" - formatter_key: str = "cpp" - formatter_options: Mapping[str, Any] = dataclasses.field( - default_factory=functools.partial(dict, style="LLVM") - ) @dataclasses.dataclass(frozen=True, kw_only=True) @@ -111,19 +92,6 @@ class HIPCodeSpec(CPPLikeCodeSpec): source_language: str = "HIP" file_extension: str = "hip" header_extension: str = "h" - formatter_key: str = "cpp" - formatter_options: Mapping[str, Any] = dataclasses.field( - default_factory=functools.partial(dict, style="LLVM") - ) - - -def format_source(source_code_spec: SourceCodeSpec, source: str) -> str: - 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 {}) - ) CodeSpecT = TypeVar("CodeSpecT", bound=SourceCodeSpec) diff --git a/src/gt4py/next/otf/binding/nanobind.py b/src/gt4py/next/otf/binding/nanobind.py index ab4004d44e..f67b3383a6 100644 --- a/src/gt4py/next/otf/binding/nanobind.py +++ b/src/gt4py/next/otf/binding/nanobind.py @@ -309,12 +309,11 @@ def create_bindings( ), ) - src = artifacts.format_source( - program_source.code_spec, BindingCodeGenerator.apply(file_binding) + return artifacts.BindingSource( + BindingCodeGenerator.apply(file_binding), + (interface.LibraryDependency("nanobind", "2.0.0"),), ) - return artifacts.BindingSource(src, (interface.LibraryDependency("nanobind", "2.0.0"),)) - @dataclasses.dataclass(frozen=True) class ExtensionGenerator: 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 f861d4b182..f1c6978791 100644 --- a/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py +++ b/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py @@ -14,7 +14,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 @@ -172,8 +171,7 @@ 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 @@ -210,14 +208,13 @@ def __call__( inp.args.offset_provider, inp.args.column_axis, ) - source_code = artifacts.format_source( - self._code_spec(), - f""" - #include <{self._backend_header()}> - #include - {stencil_src} - {decl_src} - """.strip(), + source_code = "\n".join( + [ + f"#include <{self._backend_header()}>", + "#include ", + stencil_src, + decl_src, + ] ) module: artifacts.ProgramSource[artifacts.HeaderAndSourceCodeSpec] = ( diff --git a/src/gt4py/next/program_processors/runners/roundtrip.py b/src/gt4py/next/program_processors/runners/roundtrip.py index 09f173d3f9..29afc3535a 100644 --- a/src/gt4py/next/program_processors/runners/roundtrip.py +++ b/src/gt4py/next/program_processors/runners/roundtrip.py @@ -110,8 +110,8 @@ 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 flag so re-exec is skipped within a process. +_MODULE_CACHE: dict[tuple[str, bool], types.ModuleType] = {} def _generate_source( @@ -128,7 +128,6 @@ def _generate_source( ( ir, transforms, - debug, use_embedded, tuple(common.offset_provider_to_type(offset_provider).items()), ) @@ -142,11 +141,6 @@ 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: - program = codegen.format_python_source(program) - offset_literals: Iterable[str] = ( ir.pre_walk_values() .if_isinstance(itir.OffsetLiteral) @@ -187,8 +181,9 @@ def _generate_source( def _load_module(source_code: str, debug: bool) -> types.ModuleType: - if source_code in _MODULE_CACHE: - return _MODULE_CACHE[source_code] + cache_key = (source_code, debug) + if cache_key in _MODULE_CACHE: + return _MODULE_CACHE[cache_key] if debug: # Write to a real .py so debuggers/tracebacks have file/line info. @@ -205,7 +200,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/tach.toml b/tach.toml index 78541c5dff..50abf97025 100644 --- a/tach.toml +++ b/tach.toml @@ -18,6 +18,7 @@ path = "gt4py.cartesian" depends_on = [ { path = "gt4py._core" }, { path = "gt4py.eve" }, + { path = "gt4py.eve.formatting" }, { path = "gt4py.storage" }, ] @@ -25,6 +26,12 @@ depends_on = [ path = "gt4py.eve" depends_on = [] +# Separate module so that only `gt4py.cartesian` can use the best-effort +# generated-code formatting (optional `black`/`clang-format` dependencies). +[[modules]] +path = "gt4py.eve.formatting" +depends_on = [] + [[modules]] path = "gt4py.next" depends_on = [ diff --git a/tests/eve_tests/unit_tests/test_formatting.py b/tests/eve_tests/unit_tests/test_formatting.py new file mode 100644 index 0000000000..2fbe0970e7 --- /dev/null +++ b/tests/eve_tests/unit_tests/test_formatting.py @@ -0,0 +1,104 @@ +# 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 shutil +import sys + +import pytest + +from gt4py.eve import formatting + + +UNFORMATTED_PYTHON = "def f( a,b ):\n return a+b\n" +UNFORMATTED_CPP = "int f( int a,int b ){return a+b;}\n" + + +# -- Python tests -- +def test_format_python_source(): + pytest.importorskip("black") + assert formatting.format_python_source(UNFORMATTED_PYTHON) == ( + "def f(a, b):\n return a + b\n" + ) + + +def test_format_python_source_line_length(): + pytest.importorskip("black") + source = "result = function_name(first_argument, second_argument)\n" + assert formatting.format_python_source(source) == source + assert formatting.format_python_source(source, line_length=40) == ( + "result = function_name(\n first_argument, second_argument\n)\n" + ) + + +def test_format_python_source_invalid_input(): + pytest.importorskip("black") + source = "def f(:\n" + assert formatting.format_python_source(source) == source + + +def test_format_python_source_black_failure(monkeypatch): + black = pytest.importorskip("black") + + def failing_format_str(*args, **kwargs): + raise RuntimeError("internal black error") + + monkeypatch.setattr(black, "format_str", failing_format_str) + assert formatting.format_python_source(UNFORMATTED_PYTHON) == UNFORMATTED_PYTHON + + +def test_format_python_source_unsupported_interpreter(monkeypatch): + black = pytest.importorskip("black") + # Simulate a `black` version that does not know the running interpreter. + monkeypatch.setattr(black, "TargetVersion", {}) + assert formatting.format_python_source(UNFORMATTED_PYTHON) == UNFORMATTED_PYTHON + + +def test_format_python_source_without_black(monkeypatch): + monkeypatch.setitem(sys.modules, "black", None) + assert formatting.format_python_source(UNFORMATTED_PYTHON) == UNFORMATTED_PYTHON + + +# -- C++ tests -- +@pytest.fixture(autouse=True) +def clear_clang_format_cache(): + formatting._get_clang_format.cache_clear() + yield + formatting._get_clang_format.cache_clear() + + +@pytest.mark.skipif(shutil.which("clang-format") is None, reason="clang-format not available") +def test_format_cpp_source(monkeypatch): + monkeypatch.delenv("CLANG_FORMAT_EXECUTABLE", raising=False) + assert ( + formatting.format_cpp_source(UNFORMATTED_CPP) == "int f(int a, int b) { return a + b; }\n" + ) + + +def test_format_cpp_source_missing_executable(monkeypatch): + monkeypatch.setenv("CLANG_FORMAT_EXECUTABLE", "/nonexistent") + assert formatting.format_cpp_source(UNFORMATTED_CPP) == UNFORMATTED_CPP + + +@pytest.mark.skipif(shutil.which("false") is None, reason="`false` executable not available") +def test_format_cpp_source_failing_executable(monkeypatch): + monkeypatch.setenv("CLANG_FORMAT_EXECUTABLE", "false") + assert formatting.format_cpp_source(UNFORMATTED_CPP) == UNFORMATTED_CPP + + +@pytest.mark.skipif(shutil.which("false") is None, reason="`false` executable not available") +def test_format_cpp_source_formatting_failure(monkeypatch): + # The executable passes the availability probe but fails when formatting. + monkeypatch.setattr(formatting, "_get_clang_format", lambda: "false") + assert formatting.format_cpp_source(UNFORMATTED_CPP) == UNFORMATTED_CPP + + +def test_format_cpp_source_probes_executable_once(monkeypatch): + monkeypatch.setenv("CLANG_FORMAT_EXECUTABLE", "/nonexistent") + formatting.format_cpp_source(UNFORMATTED_CPP) + formatting.format_cpp_source(UNFORMATTED_CPP) + assert formatting._get_clang_format.cache_info().misses == 1 diff --git a/tests/next_tests/integration_tests/cases_utils.py b/tests/next_tests/integration_tests/cases_utils.py index 18617a6f91..ea7c898f4d 100644 --- a/tests/next_tests/integration_tests/cases_utils.py +++ b/tests/next_tests/integration_tests/cases_utils.py @@ -143,10 +143,10 @@ def debug_itir(tree): """Compare tree snippets while debugging.""" from devtools import debug - from gt4py.eve.codegen import format_python_source + from gt4py.eve import formatting from gt4py.next.program_processors import EmbeddedDSL - debug(format_python_source(EmbeddedDSL.apply(tree))) + debug(formatting.format_python_source(EmbeddedDSL.apply(tree))) DimsType = TypeVar("DimsType") diff --git a/tests/next_tests/unit_tests/otf_tests/binding_tests/test_cpp_interface.py b/tests/next_tests/unit_tests/otf_tests/binding_tests/test_cpp_interface.py index a25732649a..8e5e8dcca4 100644 --- a/tests/next_tests/unit_tests/otf_tests/binding_tests/test_cpp_interface.py +++ b/tests/next_tests/unit_tests/otf_tests/binding_tests/test_cpp_interface.py @@ -11,7 +11,7 @@ import gt4py.next as gtx import gt4py.next.otf.binding.cpp_interface as cpp import gt4py.next.type_system.type_specifications as ts -from gt4py.eve.codegen import format_source +from gt4py.eve import formatting from gt4py.next.otf.binding import interface @@ -27,28 +27,22 @@ def function_scalar_example(): def test_render_function_declaration_scalar(function_scalar_example): - rendered = format_source( - "cpp", cpp.render_function_declaration(function_scalar_example, "return;"), style="LLVM" + rendered = formatting.format_cpp_source( + cpp.render_function_declaration(function_scalar_example, "return;") ) - expected = format_source( - "cpp", - """\ + expected = formatting.format_cpp_source("""\ decltype(auto) example(double a, std::int64_t b) { return; }\ -""", - style="LLVM", - ) +""") assert rendered == expected def test_render_function_call_scalar(function_scalar_example): - rendered = format_source( - "cpp", - cpp.render_function_call(function_scalar_example, args=["13.6", "get_arg()"]), - style="LLVM", + rendered = formatting.format_cpp_source( + cpp.render_function_call(function_scalar_example, args=["13.6", "get_arg()"]) ) - expected = format_source("cpp", """example(13.6, get_arg())""", style="LLVM") + expected = formatting.format_cpp_source("""example(13.6, get_arg())""") assert rendered == expected @@ -75,29 +69,23 @@ def function_buffer_example(): def test_render_function_declaration_buffer(function_buffer_example): - rendered = format_source( - "cpp", cpp.render_function_declaration(function_buffer_example, "return;"), style="LLVM" + rendered = formatting.format_cpp_source( + cpp.render_function_declaration(function_buffer_example, "return;") ) - expected = format_source( - "cpp", - """\ + expected = formatting.format_cpp_source("""\ template decltype(auto) example(ArgT0&& a_buf, ArgT1&& b_buf) { return; }\ -""", - style="LLVM", - ) +""") assert rendered == expected def test_render_function_call_buffer(function_buffer_example): - rendered = format_source( - "cpp", - cpp.render_function_call(function_buffer_example, args=["get_arg_1()", "get_arg_2()"]), - style="LLVM", + rendered = formatting.format_cpp_source( + cpp.render_function_call(function_buffer_example, args=["get_arg_1()", "get_arg_2()"]) ) - expected = format_source("cpp", """example(get_arg_1(), get_arg_2())""", style="LLVM") + expected = formatting.format_cpp_source("""example(get_arg_1(), get_arg_2())""") assert rendered == expected @@ -126,25 +114,21 @@ def function_tuple_example(): def test_render_function_declaration_tuple(function_tuple_example): - rendered = format_source( - "cpp", cpp.render_function_declaration(function_tuple_example, "return;"), style="LLVM" + rendered = formatting.format_cpp_source( + cpp.render_function_declaration(function_tuple_example, "return;") ) - expected = format_source( - "cpp", - """\ + expected = formatting.format_cpp_source("""\ template decltype(auto) example(ArgT0&& a_buf) { return; }\ -""", - style="LLVM", - ) +""") assert rendered == expected def test_render_function_call_tuple(function_tuple_example): - rendered = format_source( - "cpp", cpp.render_function_call(function_tuple_example, args=["get_arg_1()"]), style="LLVM" + rendered = formatting.format_cpp_source( + cpp.render_function_call(function_tuple_example, args=["get_arg_1()"]) ) - expected = format_source("cpp", """example(get_arg_1())""", style="LLVM") + expected = formatting.format_cpp_source("""example(get_arg_1())""") assert rendered == expected 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 df65b39c4c..2eb9d87172 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 @@ -12,7 +12,7 @@ import dace import numpy as np import pytest -from gt4py.eve import codegen +from gt4py.eve import formatting from gt4py import next as gtx from gt4py.next import common as gtx_common, int32 @@ -238,7 +238,7 @@ def mocked_compile_call( for line in inp.binding_source.source_code.splitlines() if not line.lstrip().startswith("assert") ) - assert codegen.format_python_source(binding_source_pruned) == binding_source_ref + assert formatting.format_python_source(binding_source_pruned) == binding_source_ref return _dace_compile_call(self, inp) 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..510aa1ef1d --- /dev/null +++ b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/test_roundtrip.py @@ -0,0 +1,72 @@ +# 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 os + +import pytest + +from gt4py.next.iterator import ir as itir +from gt4py.next.iterator.ir_utils import ir_makers as im +from gt4py.next.program_processors.runners import roundtrip + + +@pytest.mark.parametrize("debug_order", [(False, True), (True, False)]) +def test_generate_source_ignores_debug_flag(debug_order, monkeypatch): + monkeypatch.setattr(roundtrip, "_SOURCE_CACHE", {}) + domain = im.call("cartesian_domain")(im.named_range(itir.AxisLiteral(value="D"), 0, 1)) + program = itir.Program( + id="testee", + function_definitions=[], + params=[itir.Sym(id="out")], + declarations=[], + body=[ + itir.SetAt(expr=im.as_fieldop("deref")(), domain=domain, target=itir.SymRef(id="out")) + ], + ) + transform_calls = [] + + def transforms(ir, *, offset_provider): + transform_calls.append(ir) + return ir + + sources = [ + roundtrip._generate_source( + program, + debug=debug, + use_embedded=True, + offset_provider={}, + transforms=transforms, + ) + for debug in debug_order + ] + + assert sources[0] == sources[1] + assert sources[0][1] == "testee" + # The source is generated once and shared by both debug modes. + assert len(transform_calls) == 1 + assert len(roundtrip._SOURCE_CACHE) == 1 + + +@pytest.mark.parametrize("debug_order", [(False, True), (True, False)]) +def test_load_module_caches_by_debug_flag(debug_order, monkeypatch): + monkeypatch.setattr(roundtrip, "_MODULE_CACHE", {}) + source_code = "VALUE = 42\n" + + modules = {debug: roundtrip._load_module(source_code, debug) for debug in debug_order} + try: + assert modules[False] is not modules[True] + assert modules[False].VALUE == modules[True].VALUE == 42 + # Only the debug module is backed by a real '.py' file. + assert not hasattr(modules[False], "__file__") + assert modules[True].__file__.endswith(".py") + assert os.path.isfile(modules[True].__file__) + + for debug in debug_order: + assert roundtrip._load_module(source_code, debug) is modules[debug] + finally: + os.remove(modules[True].__file__) diff --git a/uv.lock b/uv.lock index ef2d5ca914..d897f527ec 100644 --- a/uv.lock +++ b/uv.lock @@ -1269,7 +1269,6 @@ source = { editable = "." } dependencies = [ { name = "array-api-compat" }, { name = "attrs" }, - { name = "black" }, { name = "boltons" }, { name = "cached-property" }, { name = "click" }, @@ -1300,6 +1299,7 @@ dependencies = [ [package.optional-dependencies] cartesian = [ + { name = "black" }, { name = "clang-format" }, { name = "hypothesis" }, { name = "jax" }, @@ -1324,7 +1324,6 @@ jax-cuda13 = [ { name = "jax", extra = ["cuda13-local"], marker = "extra == 'extra-5-gt4py-jax-cuda13' or (extra == 'extra-5-gt4py-cuda12' and extra == 'extra-5-gt4py-rocm6') or (extra == 'extra-5-gt4py-cuda12' and extra == 'extra-5-gt4py-rocm7') or (extra == 'extra-5-gt4py-rocm6' and extra == 'extra-5-gt4py-rocm7')" }, ] next = [ - { name = "clang-format" }, { name = "hypothesis" }, { name = "jax" }, { name = "pytest" }, @@ -1337,7 +1336,6 @@ rocm7 = [ { name = "cupy-rocm-7-0" }, ] standard = [ - { name = "clang-format" }, { name = "scipy" }, ] testing = [ @@ -1354,6 +1352,8 @@ build = [ ] dev = [ { name = "atlas4py" }, + { name = "black" }, + { name = "clang-format" }, { name = "cython" }, { name = "esbonio" }, { name = "factory-boy" }, @@ -1421,6 +1421,8 @@ scripts = [ { name = "typer" }, ] test = [ + { name = "black" }, + { name = "clang-format" }, { name = "factory-boy" }, { name = "hypothesis" }, { name = "nbmake" }, @@ -1456,10 +1458,10 @@ typing-exports = [ requires-dist = [ { name = "array-api-compat", specifier = ">=1.13" }, { name = "attrs", specifier = ">=21.3" }, - { name = "black", specifier = ">=25.11" }, + { name = "black", marker = "extra == 'cartesian'", specifier = ">=25.11" }, { name = "boltons", specifier = ">=20.1" }, { name = "cached-property", specifier = ">=1.5.1" }, - { name = "clang-format", marker = "extra == 'standard'", specifier = ">=18.1" }, + { name = "clang-format", marker = "extra == 'cartesian'", specifier = ">=18.1" }, { name = "click", specifier = ">=8.0.0" }, { name = "cmake", specifier = ">=3.22" }, { name = "cupy", marker = "extra == 'rocm6'", specifier = ">=13.4.1,<14.0" }, @@ -1512,6 +1514,8 @@ build = [ ] dev = [ { name = "atlas4py", specifier = ">=0.41", index = "https://test.pypi.org/simple" }, + { name = "black", specifier = ">=25.11" }, + { name = "clang-format", specifier = ">=18.1" }, { name = "cython", specifier = ">=3.0.0" }, { name = "esbonio", specifier = ">=0.16.0" }, { name = "factory-boy", specifier = ">=3.3.3" }, @@ -1577,6 +1581,8 @@ scripts = [ { name = "typer", specifier = ">=0.16.0" }, ] test = [ + { name = "black", specifier = ">=25.11" }, + { name = "clang-format", specifier = ">=18.1" }, { name = "factory-boy", specifier = ">=3.3.3" }, { name = "hypothesis", specifier = ">=6.0.0" }, { name = "nbmake", specifier = ">=1.4.6" },