diff --git a/src/tilefoundry/passes/transforms/hir_to_tir.py b/src/tilefoundry/passes/transforms/hir_to_tir.py index 85382781..b4381db9 100644 --- a/src/tilefoundry/passes/transforms/hir_to_tir.py +++ b/src/tilefoundry/passes/transforms/hir_to_tir.py @@ -90,6 +90,7 @@ from tilefoundry.ir.types.storage import StorageKind from tilefoundry.ir.visitor import ExprVisitor from tilefoundry.passes.pass_base import ModulePass +from tilefoundry.target import Target from tilefoundry.visitor_registry.registries import ( hir_lowering_registry, register_hir_lowering, @@ -1400,6 +1401,18 @@ def _lower_single_output( lo._items.append(_eval_call(Copy(), (src, out_var))) +def _target_kwarg(target: Target | None) -> dict[str, Target]: + """``{"target": target}``, or empty when the Module declared none. + + The target picks the backend and, for cuda, the SM arch codegen builds for, + so dropping it silently compiles for the fallback arch. Omitting the kwarg + leaves ``PrimFunction``'s own default in place; passing ``None`` would store + ``None`` and trip ``group_functions_by_target``. The rule: + tilefoundry spec core-ir target-inheritance. + """ + return {} if target is None else {"target": target} + + def _lower_function( fn: HirFunction, *, @@ -1411,6 +1424,7 @@ def _lower_function( override_name: str | None = None, dispatch_groups: "dict[str, tuple[HirFunction, ...]] | None" = None, mangled_registry: "dict[str, PrimFunction] | None" = None, + target: Target | None = None, ) -> PrimFunction: """Materialise the HIR `Function(params) -> tensor` as an explicit-output-param ``PrimFunction``. The function-end @@ -1535,12 +1549,15 @@ def _lower_function( params=final_params, body=body, output_count=len(out_vars), + **_target_kwarg(target), ) def _build_dispatch_entry( group: tuple[HirFunction, ...], mangled_pfs: list[PrimFunction], + *, + target: Target | None = None, ) -> PrimFunction: """Build the unmangled entry PrimFunction holding the DispatchCall. @@ -1633,6 +1650,9 @@ def _build_dispatch_entry( params=(*entry_params, *out_vars, shape_param), body=body, output_count=len(out_vars), + # Same target as the variants this entry dispatches to, so the entry and + # its callees stay on one backend / arch. + **_target_kwarg(target), ) @@ -1778,6 +1798,15 @@ def run(self, module: Module) -> Module: if isinstance(fn, HirFunction) and fn.variants: dispatch_view[fn.name] = fn.variants + # The Module — not the Function — owns the execution context, so the + # target is resolved once here through the owner chain and handed to + # every PrimFunction built below. A Module that declares no target + # anywhere leaves it unset, which keeps PrimFunction's own default. + try: + module_target = module.resolve_target() + except ValueError: + module_target = None + mangled_registry: dict[str, PrimFunction] = {} mangled_by_group: dict[str, list[PrimFunction]] = {} @@ -1811,6 +1840,7 @@ def run(self, module: Module) -> Module: override_name=mangled_name, dispatch_groups=dispatch_view, mangled_registry=mangled_registry, + target=module_target, ) mangled_registry[mangled_name] = pf lowered.append(pf) @@ -1841,7 +1871,9 @@ def run(self, module: Module) -> Module: new_fns.extend(mangled_for_group) new_fns.append( _build_dispatch_entry( - dispatch_view[group_name], mangled_for_group + dispatch_view[group_name], + mangled_for_group, + target=module_target, ) ) continue @@ -1860,6 +1892,7 @@ def run(self, module: Module) -> Module: thread_var_name=self.thread_var_name, dispatch_groups=dispatch_view, mangled_registry=mangled_registry, + target=module_target, ) ) diff --git a/tests/passes/test_hir_to_tir.py b/tests/passes/test_hir_to_tir.py index 09b9a547..926d0036 100644 --- a/tests/passes/test_hir_to_tir.py +++ b/tests/passes/test_hir_to_tir.py @@ -29,6 +29,7 @@ ) from tilefoundry.ir.core import Call, Constant, Var from tilefoundry.ir.core.module import Module +from tilefoundry.ir.core.pattern import DimVarRangePat from tilefoundry.ir.hir.function import Function as HirFunction from tilefoundry.ir.hir.grid_region import GridRegionExpr from tilefoundry.ir.hir.nn.relu import ReLU as HirReLU @@ -36,6 +37,7 @@ from tilefoundry.ir.tir.reduce import Reduce as TirReduce from tilefoundry.ir.tir.stmts import Evaluate from tilefoundry.ir.types import DType, TensorType +from tilefoundry.ir.types.dim import DimVar from tilefoundry.ir.types.shard.layout import Layout from tilefoundry.ir.types.shard.shard_layout import ShardLayout as SL from tilefoundry.ir.types.shard.shard_layout import Split @@ -45,6 +47,8 @@ _collect_hir_callee_names, _derive_meshes_from_body, ) +from tilefoundry.target import default_target +from tilefoundry.target.cuda.target import CudaTarget def test_umat_param_rejected_at_lowering() -> None: @@ -189,3 +193,96 @@ def test_the_hir_walks_reach_every_child_of_a_grid_region() -> None: ) assert _collect_hir_callee_names(in_yield) == {"callee_fn"} + + +def test_lowered_functions_carry_the_modules_declared_target() -> None: + """The Module owns the execution context, so its Target must reach every + PrimFunction the pass builds — the static body, each mangled dispatch + variant, and the dispatch entry alike. + + A dropped target is silent rather than fatal: ``PrimFunction`` defaults to + ``default_target()``, so codegen compiles for the fallback arch and the + driver PTX-JITs the result to whatever card is present. The kernel still + computes the right answer, so no runtime witness can catch this — only the + identity of the propagated value can. ``default_target()`` builds a fresh + equal value per call, so ``is`` separates "propagated" from "defaulted" + where ``==`` would not. + """ + declared = CudaTarget("nvidia.h200_sxm") + ty = TensorType(shape=(8,), dtype=DType.f32, layout=None, storage="gmem") + x = Var(type=ty, name="x") + fn = HirFunction.build(name="static_fn", params=(x,), body=x, return_type=ty) + + out = HirToTirPass().run( + Module(name="m", functions=(fn,), entry="static_fn", target=declared) + ) + assert out.functions[0].target is declared + + # Dispatch path: the mangled variants and the entry that selects between + # them are separate construction sites and must agree on one arch. + dim = DimVar(name="S", lo=1, hi=7) + dyn_ty = TensorType(shape=(dim,), dtype=DType.f32, layout=None, storage="gmem") + + def _variant(lo: int, hi: int) -> HirFunction: + v = Var(type=dyn_ty, name="x") + return HirFunction.build( + name="main", params=(v,), body=v, return_type=dyn_ty, + specializations=(DimVarRangePat("S", lo, hi),), + ) + + proto_param = Var(type=dyn_ty, name="x") + proto = HirFunction.build( + name="main", params=(proto_param,), body=None, return_type=dyn_ty + ) + for v in (_variant(1, 3), _variant(4, 7)): + proto.add_variant(v) + + out = HirToTirPass().run( + Module(name="m", functions=(proto,), entry="main", target=declared) + ) + assert sorted(f.name for f in out.functions) == [ + "main", "main$S$1_3", "main$S$4_7", + ] + for pf in out.functions: + assert pf.target is declared, f"{pf.name} lost the declared target" + + +def test_a_child_module_lowers_against_the_target_it_inherits() -> None: + """Only the root Module may declare a target; a child inherits it through + the owner chain. The pass therefore resolves the target rather than reading + the field, so lowering a child selected out of a tree does not fall back to + the default arch.""" + declared = CudaTarget("nvidia.h200_sxm") + ty = TensorType(shape=(8,), dtype=DType.f32, layout=None, storage="gmem") + x = Var(type=ty, name="x") + fn = HirFunction.build(name="inner_fn", params=(x,), body=x, return_type=ty) + + root = Module( + name="root", + functions=(), + entry=None, + modules=(Module(name="child", functions=(fn,), entry="inner_fn"),), + target=declared, + ) + child = root.modules[0] + assert child.target is None, "the child must not declare a target of its own" + + out = HirToTirPass().run(child) + assert out.functions[0].target is declared + + +def test_a_module_declaring_no_target_keeps_the_primfunction_default() -> None: + """A Module with no target anywhere in its owner chain must still lower. + Resolution fails there by contract, so the pass leaves the target unset and + ``PrimFunction``'s own default applies — the pre-existing behaviour for the + many tests and tools that build a Module without naming hardware.""" + ty = TensorType(shape=(8,), dtype=DType.f32, layout=None, storage="gmem") + x = Var(type=ty, name="x") + fn = HirFunction.build(name="static_fn", params=(x,), body=x, return_type=ty) + + module = Module(name="m", functions=(fn,), entry="static_fn") + with pytest.raises(ValueError, match="no target is declared"): + module.resolve_target() + + pf = HirToTirPass().run(module).functions[0] + assert pf.target == default_target()