Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 7 additions & 0 deletions src/gt4py/next/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 = (
Expand Down
28 changes: 19 additions & 9 deletions src/gt4py/next/otf/artifacts.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand All @@ -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)
Expand All @@ -71,6 +76,8 @@ class SDFGCodeSpec(SourceCodeSpec):

source_language: str = "SDFG"
file_extension: str = "sdfg"
# 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)
Expand Down Expand Up @@ -118,12 +125,15 @@ class HIPCodeSpec(CPPLikeCodeSpec):


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 {})
)
"""Format `source` as configured in `source_code_spec` (no-op if `.format_source` is false)."""
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)
Expand Down
4 changes: 3 additions & 1 deletion src/gt4py/next/otf/compilation/build_systems/compiledb.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
106 changes: 53 additions & 53 deletions src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,13 +9,13 @@
from __future__ import annotations

import dataclasses
from collections.abc import Callable
from typing import Any, Final, Optional

import factory
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
Expand All @@ -29,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<generated::block_sizes_t>{}",
core_defs.DeviceType.ROCM: "gridtools::fn::backend::gpu<generated::block_sizes_t>{}",
}
_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_)
Expand All @@ -52,16 +77,23 @@ class GTFNTranslationStep(
symbolic_domain_sizes: dict[str, itir.Expr] | None = None
use_max_domain_range_on_unstructured_shift: bool | None = None

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()
def __post_init__(self) -> None:
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)):
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 _process_regular_arguments(
self,
Expand Down Expand Up @@ -171,14 +203,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)
Expand All @@ -195,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))
Expand All @@ -210,9 +243,9 @@ def __call__(
inp.args.column_axis,
)
source_code = artifacts.format_source(
self._code_spec(),
code_spec,
f"""
#include <{self._backend_header()}>
#include <{_BACKEND_HEADERS[self.device_type]}>
#include <gridtools/sid/dimension_to_tuple_like.hpp>
{stencil_src}
{decl_src}
Expand All @@ -222,48 +255,15 @@ 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=self._code_spec(),
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<generated::block_sizes_t>{}"
case core_defs.DeviceType.CPU:
return "gridtools::fn::backend::naive{}"
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:
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:
Expand Down
5 changes: 4 additions & 1 deletion src/gt4py/next/program_processors/formatters/gtfn.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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")
21 changes: 12 additions & 9 deletions src/gt4py/next/program_processors/runners/roundtrip.py
Original file line number Diff line number Diff line change
Expand Up @@ -110,13 +110,15 @@ 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(
ir: itir.Program,
debug: bool,
format_source: bool,
use_embedded: bool,
offset_provider: common.OffsetProvider,
transforms: itir_transforms.GTIRTransform,
Expand All @@ -128,7 +130,7 @@ def _generate_source(
(
ir,
transforms,
debug,
format_source,
use_embedded,
tuple(common.offset_provider_to_type(offset_provider).items()),
)
Expand All @@ -142,9 +144,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] = (
Expand Down Expand Up @@ -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.
Expand All @@ -205,7 +206,7 @@ def _load_module(source_code: str, debug: bool) -> types.ModuleType:
mod = types.ModuleType("roundtrip_module")
exec(compile(source_code, "<roundtrip>", "exec"), mod.__dict__)

_MODULE_CACHE[source_code] = mod
_MODULE_CACHE[cache_key] = mod
return mod


Expand Down Expand Up @@ -258,6 +259,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
Expand All @@ -266,6 +268,7 @@ def __call__(self, inp: stages.CompilableProgramDef) -> RoundtripArtifact:
inp.data,
offset_provider=inp.args.offset_provider,
debug=debug,
format_source=self.format_source,
Comment thread
egparedes marked this conversation as resolved.
use_embedded=self.use_embedded,
transforms=self.transforms,
)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -91,3 +92,28 @@ 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(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,
)

assert fingerprinting.strict_fingerprinter(
prototypes[0]
) == fingerprinting.strict_fingerprinter(prototypes[1])
Loading
Loading