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
37 changes: 26 additions & 11 deletions docs/user/next/advanced/HackTheToolchain.md
Original file line number Diff line number Diff line change
Expand Up @@ -46,25 +46,40 @@ skip_linting_transforms = SkipLinting(**same_steps)
skip_linting_transforms.step_order(DUMMY_FOP)
```

## Alternative Factory
## Alternative Workflow

The builders take the settings shared by several steps (the device, the build
type, translation caching) as keyword arguments, and the settings of a single
step as a dict. The keys of that dict are the step's own fields, so a typo is a
type error.

```python
class MyCodeGen: ...
gtfn = gtx.program_processors.runners.gtfn

debug_gpu_no_transforms = gtfn.make_gtfn_backend(
gpu=True,
cmake_build_type=gtx.config.CMakeBuildType.DEBUG,
name_postfix="_debug_no_transforms",
translation={"enable_itir_transforms": False},
)
```

class Cpp2BindingsGen: ...
Compile workflows are plain frozen dataclasses, so a whole step is replaced
with `dataclasses.replace` on the workflow a builder returned. The replacement
is used as given.

```python
class MyCodeGen: ...

class PureCpp2WorkflowFactory(gtx.program_processors.runners.gtfn.GTFNCompileWorkflowFactory):
translation: workflow.Workflow[
gtx.otf.stages.CompilableProgramDef, gtx.otf.artifacts.ProgramSource
] = MyCodeGen()
bindings: workflow.Workflow[
gtx.otf.artifacts.ProgramSource, gtx.otf.artifacts.ExtensionSource
] = Cpp2BindingsGen()

class Cpp2BindingsGen: ...

PureCpp2WorkflowFactory(cmake_build_type=gtx.config.CMAKE_BUILD_TYPE.DEBUG)

pure_cpp2_workflow = dataclasses.replace(
gtfn.make_gtfn_compile_workflow(cmake_build_type=gtx.config.CMakeBuildType.DEBUG),
translation=MyCodeGen(),
bindings=Cpp2BindingsGen(),
)
```

## Invent new Workflow Types
Expand Down
37 changes: 13 additions & 24 deletions docs/user/next/advanced/WorkflowPatterns.md
Original file line number Diff line number Diff line change
Expand Up @@ -17,8 +17,6 @@ jupyter:
import dataclasses
import re

import factory

import gt4py.next as gtx

import devtools
Expand Down Expand Up @@ -199,7 +197,7 @@ Let's say we want to make our calculation workflow compatible with string input.

```python editable=true slideshow={"slide_type": ""}
# A plain conversion step turning a string into an int, chained into the
# workflow below and reused by `StrToIntFactory(cached=True)`.
# workflow below and reused by `make_str_to_int(cached=True)`.
def to_int(inp: str) -> int:
assert isinstance(inp, str), "Can not work with 'int'!" # yes, this is horribly contrived
return int(inp)
Expand All @@ -214,9 +212,9 @@ str_calc("1")

<!-- #region editable=true slideshow={"slide_type": ""} -->

### Step with factory (builder)
### Step with a builder

If a step can be useful with different combinations of parameters and wrappers, it should have a factory. In this case we will add a neutral wrapper around it, so we can put any combination of wrappers into that:
If a step is useful with different combinations of parameters and wrappers, give it a **builder function**: a plain function taking the cross-cutting options and returning the assembled step. Steps are frozen dataclasses, so the builder is ordinary code — no factory framework involved, and the result is fully type-checked.

<!-- #endregion -->

Expand All @@ -229,32 +227,23 @@ class AnyStrToInt(gtx.otf.workflow.ChainableWorkflowMixin[str | int, int]):
return self.inner_step(inp)


class StrToIntFactory(factory.Factory):
class Meta:
model = AnyStrToInt

class Params:
default_step = to_int
cached = factory.Trait(
inner_step=factory.LazyAttribute(
lambda o: gtx.otf.workflow.CachedStep.in_memory(
step=o.default_step, input_fingerprinter=str
)
)
)

inner_step = factory.LazyAttribute(lambda o: o.default_step)
def make_str_to_int(
*, cached: bool = False, step: gtx.otf.workflow.Workflow[str, int] = to_int
) -> AnyStrToInt:
if cached:
step = gtx.otf.workflow.CachedStep.in_memory(step=step, input_fingerprinter=str)
return AnyStrToInt(inner_step=step)


cached = StrToIntFactory(cached=True)
uncached = StrToIntFactory()
cached = make_str_to_int(cached=True)
uncached = make_str_to_int()
uncached.inner_step
```

### Example in the Wild

```python
gtx.ffront.past_passes.linters.LinterFactory??
gtx.ffront.past_passes.linters.linter_factory??
```

<!-- #region editable=true slideshow={"slide_type": ""} tags=["skip-execution"] -->
Expand Down Expand Up @@ -413,5 +402,5 @@ gtx.program_processors.runners.gtfn.run_gtfn_gpu.executor.otf_workflow??
```

```python
gtx.program_processors.runners.gtfn.GTFNBackendFactory??
gtx.program_processors.runners.gtfn.make_gtfn_backend??
```
8 changes: 1 addition & 7 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -43,6 +43,7 @@ profiling = [
]
scripts = ["pyyaml>=6.0.1", "typer>=0.16.0", "packaging"]
test = [
'factory-boy>=3.3.3',
'hypothesis>=6.0.0',
'nbmake>=1.4.6',
'nox>=2025.02.09',
Expand Down Expand Up @@ -99,7 +100,6 @@ dependencies = [
'dace==2.0.0a9',
'deepdiff>=8.1.0',
'devtools>=0.6',
'factory-boy>=3.3.3',
"filelock>=3.18.0",
'frozendict>=2.3',
'gridtools-cpp>=2.3.9,==2.*',
Expand Down Expand Up @@ -231,12 +231,6 @@ module = 'gt4py.next.iterator.*'
ignore_errors = true
module = 'gt4py.next.iterator.runtime'

[[tool.mypy.overrides]]
ignore_missing_imports = true
implicit_reexport = true
# factory-boy is broken, see https://github.com/FactoryBoy/factory_boy/pull/1114
module = "factory.*"

# -- pytest --
[tool.pytest]

Expand Down
16 changes: 16 additions & 0 deletions src/gt4py/next/backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -140,6 +140,22 @@ def step_order(self, inp: stages.ConcreteProgramDef) -> list[str]:
DEFAULT_TRANSFORMS: Transforms = Transforms()


def select_device(
gpu: bool,
) -> tuple[core_defs.DeviceType, next_allocators.FieldBufferAllocatorProtocol]:
"""
Return the device type and default field allocator of a CPU or GPU backend.

The GPU is the one CuPy was built for, or CUDA if CuPy is not available.
"""
if gpu:
return (
core_defs.CUPY_DEVICE_TYPE or core_defs.DeviceType.CUDA,
next_allocators.StandardGPUFieldBufferAllocator(),
)
Comment thread
Copilot marked this conversation as resolved.
return core_defs.DeviceType.CPU, next_allocators.StandardCPUFieldBufferAllocator()


# TODO(tehrengruber): Rename class and `executor` & `transforms` attribute. Maybe:
# `Backend` -> `Toolchain`
# `transforms` -> `frontend_transforms`
Expand Down
31 changes: 29 additions & 2 deletions src/gt4py/next/otf/compilation/cache.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,10 +10,11 @@

import pathlib
import tempfile
from typing import Final
from typing import Final, TypeVar

from gt4py._core import filecache
from gt4py.next import config, fingerprinting
from gt4py.next.otf import artifacts
from gt4py.next.otf import artifacts, workflow


#: Regex describing the folder names produced by `get_cache_folder` (use
Expand Down Expand Up @@ -47,6 +48,9 @@
#: workflow factory enables the `cached_translation` trait.
TRANSLATION_CACHE_BACKENDS: Final[tuple[str, ...]] = ("dace", "gtfn")

StartT = TypeVar("StartT")
EndT = TypeVar("EndT")

_session_cache_dir = tempfile.TemporaryDirectory(prefix="gt4py_session_")

_session_cache_dir_path = pathlib.Path(_session_cache_dir.name)
Expand All @@ -62,6 +66,29 @@ def get_translation_cache_folder(cache_base: pathlib.Path, backend: str) -> path
return cache_base / TRANSLATION_CACHE_DIR_NAME / backend


def persistent_translation_cache(
step: workflow.Workflow[StartT, EndT], backend: str
) -> workflow.CachedStep[StartT, EndT, str]:
"""
Wrap a translation step in the persistent translation cache of `backend`.

Args:
step: The translation step to cache.
backend: Name of the backend family, which selects the cache folder.

Returns:
The step, cached in the translation cache folder of `backend` under the
cache base path of the configured build-cache lifetime.
"""
return workflow.CachedStep[StartT, EndT, str].persistent(
step,
input_fingerprinter=fingerprinting.strict_fingerprinter,
cache=filecache.FileCache(
get_translation_cache_folder(get_cache_base_path(config.BUILD_CACHE_LIFETIME), backend)
),
)


def get_cache_base_path(lifetime: config.BuildCacheLifetime) -> pathlib.Path:
"""Return the base directory for cached artifacts with the given lifetime."""
match lifetime:
Expand Down
10 changes: 2 additions & 8 deletions src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,6 @@
import dataclasses
from typing import Any, Final, Optional

import factory
import numpy as np

from gt4py._core import definitions as core_defs
Expand Down Expand Up @@ -265,13 +264,8 @@ def _not_implemented_for_device_type(self) -> NotImplementedError:
)


class GTFNTranslationStepFactory(factory.Factory[GTFNTranslationStep]):
class Meta:
model = GTFNTranslationStep
translate_program_cpu: Final[stages.TranslationStep] = GTFNTranslationStep()


translate_program_cpu: Final[stages.TranslationStep] = GTFNTranslationStepFactory() # type: ignore[assignment] # factory-boy typing not precise enough

translate_program_gpu: Final[stages.TranslationStep] = GTFNTranslationStepFactory( # type: ignore[assignment] # factory-boy typing not precise enough
translate_program_gpu: Final[stages.TranslationStep] = GTFNTranslationStep(
device_type=core_defs.DeviceType.CUDA
)
2 changes: 1 addition & 1 deletion src/gt4py/next/program_processors/formatters/gtfn.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,7 @@

@program_formatter.program_formatter
def format_cpp(program: itir.Program, *args: Any, **kwargs: Any) -> str:
gtfn_translation = gtfn.GTFNCompileWorkflowFactory(cached_translation=False).translation
gtfn_translation = gtfn.make_gtfn_compile_workflow().translation
assert isinstance(gtfn_translation, GTFNTranslationStep)
return gtfn_translation.generate_stencil_source(
program,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,6 @@
- `compilation` for compiling the SDFG into a program
- `decoration` to parse the program arguments and pass them to the program call

The GTIR-DaCe backend factory extends `CachedBackendFactory`, thus it provides
caching of the GTIR program.
The backend builder wraps the translation step in a persistent `CachedStep`,
thus it provides caching of the GTIR program.
"""
Loading
Loading