Skip to content
Merged
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
12 changes: 6 additions & 6 deletions docs/user/next/advanced/HackTheToolchain.md
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,7 @@ import dataclasses
import typing

from gt4py import next as gtx
from gt4py.next.otf import toolchain, workflow
from gt4py.next.otf import workflow
from gt4py.next.ffront import field_operator_ast as foast, stages as ff_stages
from gt4py import eve
```
Expand All @@ -22,7 +22,7 @@ cached_lowering_toolchain = gtx.backend.DEFAULT_TRANSFORMS.replace(
## Skip Steps / Change Order

```python
DUMMY_FOP = toolchain.ConcreteArtifact(
DUMMY_FOP = workflow.ConcreteArtifact(
data=ff_stages.DSLFieldOperatorDef(definition=None), args=None
)
```
Expand Down Expand Up @@ -57,11 +57,11 @@ class Cpp2BindingsGen: ...

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


PureCpp2WorkflowFactory(cmake_build_type=gtx.config.CMAKE_BUILD_TYPE.DEBUG)
Expand Down
30 changes: 15 additions & 15 deletions src/gt4py/next/backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,7 +24,7 @@
)
from gt4py.next.ffront.past_passes import linters as past_linters
from gt4py.next.iterator import ir as itir
from gt4py.next.otf import arguments, definitions, stages, toolchain, workflow
from gt4py.next.otf import arguments, artifacts, stages, toolchain, workflow


def jit_to_aot_args(
Expand All @@ -34,8 +34,8 @@ def jit_to_aot_args(


def adapted_jit_to_aot_args_factory() -> workflow.Workflow[
definitions.ConcreteProgramDef[definitions.IRDefinitionT, arguments.JITArgs],
definitions.ConcreteProgramDef[definitions.IRDefinitionT, arguments.CompileTimeArgs],
stages.ConcreteProgramDef[stages.IRDefinitionT, arguments.JITArgs],
stages.ConcreteProgramDef[stages.IRDefinitionT, arguments.CompileTimeArgs],
]:
"""Wrap `jit_to_aot` into a workflow adapter to fit into backend transform workflows."""
return toolchain.ArgsOnlyAdapter(jit_to_aot_args)
Expand All @@ -44,8 +44,8 @@ def adapted_jit_to_aot_args_factory() -> workflow.Workflow[
@dataclasses.dataclass(frozen=True)
class Transforms(
workflow.MultiWorkflow[
definitions.ConcreteProgramDef[definitions.IRDefinitionT, definitions.ArgsDefinitionT],
definitions.CompilableProgramDef,
stages.ConcreteProgramDef[stages.IRDefinitionT, stages.ArgsDefinitionT],
stages.CompilableProgramDef,
]
):
"""
Expand All @@ -63,8 +63,8 @@ class Transforms(
"""

aotify_args: workflow.Workflow[
definitions.ConcreteProgramDef[definitions.IRDefinitionT, arguments.JITArgs],
definitions.ConcreteProgramDef[definitions.IRDefinitionT, arguments.CompileTimeArgs],
stages.ConcreteProgramDef[stages.IRDefinitionT, arguments.JITArgs],
stages.ConcreteProgramDef[stages.IRDefinitionT, arguments.CompileTimeArgs],
] = dataclasses.field(default_factory=adapted_jit_to_aot_args_factory)

func_to_foast: workflow.Workflow[
Expand Down Expand Up @@ -92,10 +92,10 @@ class Transforms(
] = dataclasses.field(default_factory=past_process_args.transform_program_args_factory)

past_to_itir: workflow.Workflow[
ffront_stages.ConcretePASTProgramDef, definitions.CompilableProgramDef
ffront_stages.ConcretePASTProgramDef, stages.CompilableProgramDef
] = dataclasses.field(default_factory=past_to_itir.past_to_gtir_factory)

def step_order(self, inp: definitions.ConcreteProgramDef) -> list[str]:
def step_order(self, inp: stages.ConcreteProgramDef) -> list[str]:
steps: list[str] = []
if isinstance(inp.args, arguments.JITArgs):
steps.append("aotify_args")
Expand Down Expand Up @@ -147,19 +147,19 @@ def step_order(self, inp: definitions.ConcreteProgramDef) -> list[str]:
@dataclasses.dataclass(frozen=True)
class Backend(Generic[core_defs.DeviceTypeT]):
name: str
executor: workflow.Workflow[definitions.CompilableProgramDef, stages.CompilationArtifact]
executor: workflow.Workflow[stages.CompilableProgramDef, artifacts.CompilationArtifact]
allocator: next_allocators.FieldBufferAllocatorProtocol[core_defs.DeviceTypeT]
transforms: workflow.Workflow[definitions.ConcreteProgramDef, definitions.CompilableProgramDef]
transforms: workflow.Workflow[stages.ConcreteProgramDef, stages.CompilableProgramDef]

def compile(
self, program: definitions.IRDefinitionT, compile_time_args: arguments.CompileTimeArgs
) -> stages.ExecutableProgram:
self, program: stages.IRDefinitionT, compile_time_args: arguments.CompileTimeArgs
) -> artifacts.ExecutableProgram:
artifact = self.executor(
self.transforms(definitions.ConcreteProgramDef(data=program, args=compile_time_args))
self.transforms(stages.ConcreteProgramDef(data=program, args=compile_time_args))
)
return self.load_artifact(artifact)

def load_artifact(self, artifact: stages.CompilationArtifact) -> stages.ExecutableProgram:
def load_artifact(self, artifact: artifacts.CompilationArtifact) -> artifacts.ExecutableProgram:
"""Load an artifact into an executable program.

Backends may override this method to inject backend-specific runtime data
Expand Down
12 changes: 5 additions & 7 deletions src/gt4py/next/ffront/decorator.py
Original file line number Diff line number Diff line change
Expand Up @@ -47,7 +47,7 @@
from gt4py.next.ffront.gtcallable import GTCallable
from gt4py.next.instrumentation import hook_machinery, metrics
from gt4py.next.iterator import ir as itir
from gt4py.next.otf import arguments, compiled_program, options, toolchain
from gt4py.next.otf import arguments, compiled_program, options, workflow
from gt4py.next.type_system import type_info, type_specifications as ts, type_translation


Expand Down Expand Up @@ -264,9 +264,7 @@ def __gt_type__(self) -> ts_ffront.ProgramType:

# TODO(ricoh): linting should become optional, up to the backend.
def __post_init__(self) -> None:
no_args_past = toolchain.ConcreteArtifact(
self.past_stage, arguments.CompileTimeArgs.empty()
)
no_args_past = workflow.ConcreteArtifact(self.past_stage, arguments.CompileTimeArgs.empty())
_ = self._frontend_transforms.past_lint(no_args_past).data

@property
Expand All @@ -289,7 +287,7 @@ def definition(self) -> types.FunctionType:
@functools.cached_property
def past_stage(self) -> ffront_stages.PASTProgramDef:
# backwards compatibility for backends that do not support the full toolchain
no_args_def = toolchain.ConcreteArtifact(
no_args_def = workflow.ConcreteArtifact(
self.definition_stage, arguments.CompileTimeArgs.empty()
)
return self._frontend_transforms.func_to_past(no_args_def).data
Expand All @@ -309,7 +307,7 @@ def _all_closure_vars(self) -> dict[str, Any]:

@functools.cached_property
def gtir(self) -> itir.Program:
no_args_past = toolchain.ConcreteArtifact(
no_args_past = workflow.ConcreteArtifact(
data=ffront_stages.PASTProgramDef(
past_node=self.past_stage.past_node,
closure_vars=self.past_stage.closure_vars,
Expand Down Expand Up @@ -609,7 +607,7 @@ def __post_init__(self) -> None:
@functools.cached_property
def foast_stage(self) -> ffront_stages.FOASTOperatorDef:
return self._frontend_transforms.func_to_foast(
toolchain.ConcreteArtifact(
workflow.ConcreteArtifact(
data=self.definition_stage, args=arguments.CompileTimeArgs.empty()
)
).data
Expand Down
8 changes: 4 additions & 4 deletions src/gt4py/next/ffront/foast_to_past.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,7 +20,7 @@
from gt4py.next.ffront.past_passes import closure_var_type_deduction, type_deduction
from gt4py.next.ffront.stages import ConcreteFOASTOperatorDef, ConcretePASTProgramDef
from gt4py.next.iterator import ir as itir
from gt4py.next.otf import toolchain, workflow
from gt4py.next.otf import workflow
from gt4py.next.type_system import type_info, type_specifications as ts


Expand Down Expand Up @@ -62,7 +62,7 @@ class OperatorToProgram(workflow.Workflow[ConcreteFOASTOperatorDef, ConcretePAST

Example:
>>> from gt4py import next as gtx
>>> from gt4py.next.otf import arguments, toolchain
>>> from gt4py.next.otf import arguments, workflow
>>> IDim = gtx.Dimension("I")

>>> @gtx.field_operator
Expand All @@ -83,7 +83,7 @@ class OperatorToProgram(workflow.Workflow[ConcreteFOASTOperatorDef, ConcretePAST
... )

>>> copy_program = op_to_prog(
... toolchain.ConcreteArtifact(copy.foast_stage, compile_time_args)
... workflow.ConcreteArtifact(copy.foast_stage, compile_time_args)
... )

>>> print(copy_program.data.past_node.id)
Expand Down Expand Up @@ -169,7 +169,7 @@ def __call__(self, inp: ConcreteFOASTOperatorDef) -> ConcretePASTProgramDef:
)
past_node = type_deduction.ProgramTypeDeduction.apply(untyped_past_node)

return toolchain.ConcreteArtifact(
return workflow.ConcreteArtifact(
data=ffront_stages.PASTProgramDef(
past_node=past_node,
closure_vars=fieldop_itir_closure_vars, # type: ignore[arg-type]
Expand Down
4 changes: 2 additions & 2 deletions src/gt4py/next/ffront/past_process_args.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,7 +14,7 @@
stages as ffront_stages,
type_specifications as ts_ffront,
)
from gt4py.next.otf import arguments, toolchain, workflow
from gt4py.next.otf import arguments, workflow
from gt4py.next.type_system import type_info, type_specifications as ts


Expand All @@ -24,7 +24,7 @@ def transform_program_args(
rewritten_args, rewritten_kwargs = _process_args(
past_node=inp.data.past_node, args=inp.args.args, kwargs=inp.args.kwargs
)
return toolchain.ConcreteArtifact(
return workflow.ConcreteArtifact(
data=inp.data,
args=arguments.CompileTimeArgs(
args=rewritten_args,
Expand Down
12 changes: 6 additions & 6 deletions src/gt4py/next/ffront/past_to_itir.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,19 +29,19 @@
from gt4py.next.iterator import ir as itir
from gt4py.next.iterator.ir_utils import ir_makers as im
from gt4py.next.iterator.transforms import remap_symbols, replace_get_domain_range_with_constants
from gt4py.next.otf import arguments, definitions, workflow
from gt4py.next.otf import arguments, stages, workflow
from gt4py.next.type_system import type_info, type_specifications as ts


# FIXME[#1582](tehrengruber): This should only depend on the program not the arguments. Remove
# dependency as soon as column axis can be deduced from ITIR in consumers of the CompilableProgram.
def past_to_gtir(inp: ConcretePASTProgramDef) -> definitions.CompilableProgramDef:
def past_to_gtir(inp: ConcretePASTProgramDef) -> stages.CompilableProgramDef:
"""
Lower a PAST program definition to Iterator IR.

Example:
>>> from gt4py import next as gtx
>>> from gt4py.next.otf import arguments, toolchain
>>> from gt4py.next.otf import arguments, workflow
>>> IDim = gtx.Dimension("I")

>>> @gtx.field_operator
Expand All @@ -63,7 +63,7 @@ def past_to_gtir(inp: ConcretePASTProgramDef) -> definitions.CompilableProgramDe
... )

>>> itir_copy = past_to_gtir(
... toolchain.ConcreteArtifact(copy_program.past_stage, compile_time_args)
... workflow.ConcreteArtifact(copy_program.past_stage, compile_time_args)
... )

>>> print(itir_copy.data.id)
Expand Down Expand Up @@ -144,12 +144,12 @@ def past_to_gtir(inp: ConcretePASTProgramDef) -> definitions.CompilableProgramDe
if config.DEBUG or inp.data.debug:
devtools.debug(itir_program)

return definitions.CompilableProgramDef(data=itir_program, args=compile_time_args)
return stages.CompilableProgramDef(data=itir_program, args=compile_time_args)


def past_to_gtir_factory(
cached: bool = True,
) -> workflow.Workflow[ConcretePASTProgramDef, definitions.CompilableProgramDef]:
) -> workflow.Workflow[ConcretePASTProgramDef, stages.CompilableProgramDef]:
wf = workflow.make_step(past_to_gtir)
if cached:
wf = workflow.CachedStep.in_memory(
Expand Down
10 changes: 5 additions & 5 deletions src/gt4py/next/ffront/stages.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,7 +29,7 @@

from gt4py.next import common, fingerprinting
from gt4py.next.ffront import field_operator_ast as foast, program_ast as past, source_utils
from gt4py.next.otf import arguments, toolchain
from gt4py.next.otf import arguments, workflow


@dataclasses.dataclass(frozen=True)
Expand Down Expand Up @@ -79,7 +79,7 @@ class DSLFieldOperatorDef(BaseStage):
debug: bool = False


ConcreteDSLFieldOperatorDef: typing.TypeAlias = toolchain.ConcreteArtifact[
ConcreteDSLFieldOperatorDef: typing.TypeAlias = workflow.ConcreteArtifact[
DSLFieldOperatorDef, arguments.CompileTimeArgs
]

Expand All @@ -93,7 +93,7 @@ class FOASTOperatorDef(BaseStage):
debug: bool = False


ConcreteFOASTOperatorDef: typing.TypeAlias = toolchain.ConcreteArtifact[
ConcreteFOASTOperatorDef: typing.TypeAlias = workflow.ConcreteArtifact[
FOASTOperatorDef, arguments.CompileTimeArgs
]

Expand All @@ -105,7 +105,7 @@ class DSLProgramDef(BaseStage):
debug: bool = False


ConcreteDSLProgramDef: typing.TypeAlias = toolchain.ConcreteArtifact[
ConcreteDSLProgramDef: typing.TypeAlias = workflow.ConcreteArtifact[
DSLProgramDef, arguments.CompileTimeArgs
]

Expand All @@ -118,7 +118,7 @@ class PASTProgramDef(BaseStage):
debug: bool = False


ConcretePASTProgramDef: typing.TypeAlias = toolchain.ConcreteArtifact[
ConcretePASTProgramDef: typing.TypeAlias = workflow.ConcreteArtifact[
PASTProgramDef, arguments.CompileTimeArgs
]

Expand Down
Loading