diff --git a/lighthouse/dialects/transform/smt_ext/ops/constrain_params.py b/lighthouse/dialects/transform/smt_ext/ops/constrain_params.py index 148d7fec..17e5e526 100644 --- a/lighthouse/dialects/transform/smt_ext/ops/constrain_params.py +++ b/lighthouse/dialects/transform/smt_ext/ops/constrain_params.py @@ -78,11 +78,13 @@ def allow_repeated_handle_operands(_op: "ConstrainParamsOp") -> bool: class ConstrainParamsMemoryEffectsOpInterfaceModel(ir.MemoryEffectsOpInterface): @staticmethod - def get_effects(op: "ConstrainParamsOp", effects): + def get_effects(op: "ConstrainParamsOp"): + effects = [] if op.op_operands: - transform.only_reads_handle(op.op_operands, effects) - transform.produces_handle(op.results, effects) - transform.only_reads_payload(effects) + effects += transform.only_reads_handle(op.op_operands) + effects += transform.produces_handle(op.results) + effects += transform.only_reads_payload() + return effects class MixedResultConstrainParamsOp(ConstrainParamsOp): diff --git a/lighthouse/dialects/transform/transform_ext/ops/assign_tile_sizes.py b/lighthouse/dialects/transform/transform_ext/ops/assign_tile_sizes.py index 0add6ed0..5c19430f 100644 --- a/lighthouse/dialects/transform/transform_ext/ops/assign_tile_sizes.py +++ b/lighthouse/dialects/transform/transform_ext/ops/assign_tile_sizes.py @@ -73,10 +73,12 @@ def allow_repeated_handle_operands(_op: "AssignTileSizesOp") -> bool: class MemoryEffectsOpInterfaceModel(ir.MemoryEffectsOpInterface): @staticmethod - def get_effects(op: ir.Operation, effects): - transform.only_reads_handle(op.op_operands, effects) - transform.produces_handle(op.results, effects) - transform.modifies_payload(effects) + def get_effects(op: ir.Operation): + return ( + transform.only_reads_handle(op.op_operands) + + transform.produces_handle(op.results) + + transform.modifies_payload() + ) def assign_tile_sizes( diff --git a/lighthouse/dialects/transform/transform_ext/ops/clear_tile_and_fuse_annotations.py b/lighthouse/dialects/transform/transform_ext/ops/clear_tile_and_fuse_annotations.py index 35289035..0fb4b2f7 100644 --- a/lighthouse/dialects/transform/transform_ext/ops/clear_tile_and_fuse_annotations.py +++ b/lighthouse/dialects/transform/transform_ext/ops/clear_tile_and_fuse_annotations.py @@ -66,10 +66,12 @@ def allow_repeated_handle_operands( class MemoryEffectsOpInterfaceModel(ir.MemoryEffectsOpInterface): @staticmethod - def get_effects(op: ir.Operation, effects): - transform.only_reads_handle(op.op_operands, effects) - transform.produces_handle(op.results, effects) - transform.modifies_payload(effects) + def get_effects(op: ir.Operation): + return ( + transform.only_reads_handle(op.op_operands) + + transform.produces_handle(op.results) + + transform.modifies_payload() + ) def clear_tile_and_fuse_annotations( diff --git a/lighthouse/dialects/transform/transform_ext/ops/convert_func_results_to_args.py b/lighthouse/dialects/transform/transform_ext/ops/convert_func_results_to_args.py index 2fc43a5d..021c62df 100644 --- a/lighthouse/dialects/transform/transform_ext/ops/convert_func_results_to_args.py +++ b/lighthouse/dialects/transform/transform_ext/ops/convert_func_results_to_args.py @@ -128,10 +128,12 @@ def allow_repeated_handle_operands(_op: "ConvertFuncResultsToArgsOp") -> bool: class MemoryEffectsOpInterfaceModel(ir.MemoryEffectsOpInterface): @staticmethod - def get_effects(op: "ConvertFuncResultsToArgsOp", effects): - transform.consumes_handle(op.op_operands, effects) - transform.produces_handle(op.results, effects) - transform.modifies_payload(effects) + def get_effects(op: "ConvertFuncResultsToArgsOp"): + return ( + transform.consumes_handle(op.op_operands) + + transform.produces_handle(op.results) + + transform.modifies_payload() + ) def convert_func_results_to_args( diff --git a/lighthouse/dialects/transform/transform_ext/ops/extract_handle.py b/lighthouse/dialects/transform/transform_ext/ops/extract_handle.py index 09157ed5..649671c5 100644 --- a/lighthouse/dialects/transform/transform_ext/ops/extract_handle.py +++ b/lighthouse/dialects/transform/transform_ext/ops/extract_handle.py @@ -55,10 +55,12 @@ def allow_repeated_handle_operands(_op: "ExtractHandleOp") -> bool: class MemoryEffectsOpInterfaceModel(ir.MemoryEffectsOpInterface): @staticmethod - def get_effects(op: ir.Operation, effects): - transform.only_reads_handle(op.op_operands, effects) - transform.produces_handle(op.results, effects) - transform.only_reads_payload(effects) + def get_effects(op: ir.Operation): + return ( + transform.only_reads_handle(op.op_operands) + + transform.produces_handle(op.results) + + transform.only_reads_payload() + ) def extract_handle( diff --git a/lighthouse/dialects/transform/transform_ext/ops/fold_singleton_extract_slice.py b/lighthouse/dialects/transform/transform_ext/ops/fold_singleton_extract_slice.py index ce094494..e211512c 100644 --- a/lighthouse/dialects/transform/transform_ext/ops/fold_singleton_extract_slice.py +++ b/lighthouse/dialects/transform/transform_ext/ops/fold_singleton_extract_slice.py @@ -184,10 +184,12 @@ def allow_repeated_handle_operands(_op: "FoldSingletonExtractSliceOp") -> bool: class MemoryEffectsOpInterfaceModel(ir.MemoryEffectsOpInterface): @staticmethod - def get_effects(op: ir.Operation, effects): - transform.only_reads_handle(op.op_operands, effects) - transform.produces_handle(op.results, effects) - transform.modifies_payload(effects) + def get_effects(op: ir.Operation): + return ( + transform.only_reads_handle(op.op_operands) + + transform.produces_handle(op.results) + + transform.modifies_payload() + ) def fold_singleton_extract_slice( diff --git a/lighthouse/dialects/transform/transform_ext/ops/get_fusion_roots.py b/lighthouse/dialects/transform/transform_ext/ops/get_fusion_roots.py index dd5f48f8..6958f95a 100644 --- a/lighthouse/dialects/transform/transform_ext/ops/get_fusion_roots.py +++ b/lighthouse/dialects/transform/transform_ext/ops/get_fusion_roots.py @@ -149,10 +149,12 @@ def allow_repeated_handle_operands(_op: "GetFusionRootsOp") -> bool: class MemoryEffectsOpInterfaceModel(ir.MemoryEffectsOpInterface): @staticmethod - def get_effects(op: ir.Operation, effects): - transform.only_reads_handle(op.op_operands, effects) - transform.produces_handle(op.results, effects) - transform.only_reads_payload(effects) + def get_effects(op: ir.Operation): + return ( + transform.only_reads_handle(op.op_operands) + + transform.produces_handle(op.results) + + transform.only_reads_payload() + ) def get_fusion_roots( diff --git a/lighthouse/dialects/transform/transform_ext/ops/get_leading_unit_tile_sizes.py b/lighthouse/dialects/transform/transform_ext/ops/get_leading_unit_tile_sizes.py index 19ac1809..6765566b 100644 --- a/lighthouse/dialects/transform/transform_ext/ops/get_leading_unit_tile_sizes.py +++ b/lighthouse/dialects/transform/transform_ext/ops/get_leading_unit_tile_sizes.py @@ -57,10 +57,12 @@ def allow_repeated_handle_operands(_op: "GetLeadingUnitTileSizesOp") -> bool: class MemoryEffectsOpInterfaceModel(ir.MemoryEffectsOpInterface): @staticmethod - def get_effects(op: ir.Operation, effects): - transform.only_reads_handle(op.op_operands, effects) - transform.produces_handle(op.results, effects) - transform.only_reads_payload(effects) + def get_effects(op: ir.Operation): + return ( + transform.only_reads_handle(op.op_operands) + + transform.produces_handle(op.results) + + transform.only_reads_payload() + ) def get_leading_unit_tile_sizes( diff --git a/lighthouse/dialects/transform/transform_ext/ops/get_named_attribute.py b/lighthouse/dialects/transform/transform_ext/ops/get_named_attribute.py index 567e7698..310ae734 100644 --- a/lighthouse/dialects/transform/transform_ext/ops/get_named_attribute.py +++ b/lighthouse/dialects/transform/transform_ext/ops/get_named_attribute.py @@ -49,10 +49,12 @@ def allow_repeated_handle_operands(_op: "GetNamedAttributeOp") -> bool: class MemoryEffectsOpInterfaceModel(ir.MemoryEffectsOpInterface): @staticmethod - def get_effects(op: ir.Operation, effects): - transform.only_reads_handle(op.op_operands, effects) - transform.produces_handle(op.results, effects) - transform.only_reads_payload(effects) + def get_effects(op: ir.Operation): + return ( + transform.only_reads_handle(op.op_operands) + + transform.produces_handle(op.results) + + transform.only_reads_payload() + ) def get_named_attribute( diff --git a/lighthouse/dialects/transform/transform_ext/ops/get_tile_sizes.py b/lighthouse/dialects/transform/transform_ext/ops/get_tile_sizes.py index c58e6a02..c91c89f2 100644 --- a/lighthouse/dialects/transform/transform_ext/ops/get_tile_sizes.py +++ b/lighthouse/dialects/transform/transform_ext/ops/get_tile_sizes.py @@ -59,10 +59,12 @@ def allow_repeated_handle_operands(_op: "GetTileSizesOp") -> bool: class MemoryEffectsOpInterfaceModel(ir.MemoryEffectsOpInterface): @staticmethod - def get_effects(op: ir.Operation, effects): - transform.only_reads_handle(op.op_operands, effects) - transform.produces_handle(op.results, effects) - transform.only_reads_payload(effects) + def get_effects(op: ir.Operation): + return ( + transform.only_reads_handle(op.op_operands) + + transform.produces_handle(op.results) + + transform.only_reads_payload() + ) def get_tile_sizes( diff --git a/lighthouse/dialects/transform/transform_ext/ops/get_tileable_consumers.py b/lighthouse/dialects/transform/transform_ext/ops/get_tileable_consumers.py index 2e40e4a2..483bee63 100644 --- a/lighthouse/dialects/transform/transform_ext/ops/get_tileable_consumers.py +++ b/lighthouse/dialects/transform/transform_ext/ops/get_tileable_consumers.py @@ -91,10 +91,12 @@ def allow_repeated_handle_operands(_op: "GetTileableConsumersOp") -> bool: class MemoryEffectsOpInterfaceModel(ir.MemoryEffectsOpInterface): @staticmethod - def get_effects(op: ir.Operation, effects): - transform.only_reads_handle(op.op_operands, effects) - transform.produces_handle(op.results, effects) - transform.only_reads_payload(effects) + def get_effects(op: ir.Operation): + return ( + transform.only_reads_handle(op.op_operands) + + transform.produces_handle(op.results) + + transform.only_reads_payload() + ) def get_tileable_consumers( diff --git a/lighthouse/dialects/transform/transform_ext/ops/get_tiling_sizes.py b/lighthouse/dialects/transform/transform_ext/ops/get_tiling_sizes.py index 1e80f899..9088bbe8 100644 --- a/lighthouse/dialects/transform/transform_ext/ops/get_tiling_sizes.py +++ b/lighthouse/dialects/transform/transform_ext/ops/get_tiling_sizes.py @@ -146,10 +146,12 @@ def allow_repeated_handle_operands(_op: "GetTilingSizesOp") -> bool: class MemoryEffectsOpInterfaceModel(ir.MemoryEffectsOpInterface): @staticmethod - def get_effects(op: ir.Operation, effects): - transform.only_reads_handle(op.op_operands, effects) - transform.produces_handle(op.results, effects) - transform.only_reads_payload(effects) + def get_effects(op: ir.Operation): + return ( + transform.only_reads_handle(op.op_operands) + + transform.produces_handle(op.results) + + transform.only_reads_payload() + ) def get_tiling_sizes( diff --git a/lighthouse/dialects/transform/transform_ext/ops/move_offsets_to_subview.py b/lighthouse/dialects/transform/transform_ext/ops/move_offsets_to_subview.py index 0b785f24..4408a0d8 100644 --- a/lighthouse/dialects/transform/transform_ext/ops/move_offsets_to_subview.py +++ b/lighthouse/dialects/transform/transform_ext/ops/move_offsets_to_subview.py @@ -168,10 +168,12 @@ def allow_repeated_handle_operands(_op: "MoveOffsetsToSubviewOp") -> bool: class MemoryEffectsOpInterfaceModel(ir.MemoryEffectsOpInterface): @staticmethod - def get_effects(op: "MoveOffsetsToSubviewOp", effects): - transform.consumes_handle(op.op_operands, effects) - transform.produces_handle(op.results, effects) - transform.modifies_payload(effects) + def get_effects(op: "MoveOffsetsToSubviewOp"): + return ( + transform.consumes_handle(op.op_operands) + + transform.produces_handle(op.results) + + transform.modifies_payload() + ) def move_offsets_to_subview( diff --git a/lighthouse/dialects/transform/transform_ext/ops/param_cmp_eq.py b/lighthouse/dialects/transform/transform_ext/ops/param_cmp_eq.py index 4b95f818..99ddf52d 100644 --- a/lighthouse/dialects/transform/transform_ext/ops/param_cmp_eq.py +++ b/lighthouse/dialects/transform/transform_ext/ops/param_cmp_eq.py @@ -44,9 +44,11 @@ def allow_repeated_handle_operands(_op: "ParamCmpEqOp") -> bool: class MemoryEffectsOpInterfaceModel(ir.MemoryEffectsOpInterface): @staticmethod - def get_effects(op: "ParamCmpEqOp", effects): - transform.only_reads_handle(op.op_operands, effects) - transform.only_reads_payload(effects) + def get_effects(op: "ParamCmpEqOp"): + return ( + transform.only_reads_handle(op.op_operands) + + transform.only_reads_payload() + ) def param_cmp_eq(lhs: ir.Value, rhs: ir.Value): diff --git a/lighthouse/dialects/transform/transform_ext/ops/propagate_tile_sizes.py b/lighthouse/dialects/transform/transform_ext/ops/propagate_tile_sizes.py index 3eda1f84..c2d6a116 100644 --- a/lighthouse/dialects/transform/transform_ext/ops/propagate_tile_sizes.py +++ b/lighthouse/dialects/transform/transform_ext/ops/propagate_tile_sizes.py @@ -135,10 +135,12 @@ def allow_repeated_handle_operands(_op: "PropagateTileSizesOp") -> bool: class MemoryEffectsOpInterfaceModel(ir.MemoryEffectsOpInterface): @staticmethod - def get_effects(op: ir.Operation, effects): - transform.only_reads_handle(op.op_operands, effects) - transform.produces_handle(op.results, effects) - transform.modifies_payload(effects) + def get_effects(op: ir.Operation): + return ( + transform.only_reads_handle(op.op_operands) + + transform.produces_handle(op.results) + + transform.modifies_payload() + ) def propagate_tile_sizes( diff --git a/lighthouse/dialects/transform/transform_ext/ops/replace.py b/lighthouse/dialects/transform/transform_ext/ops/replace.py index 76dc023f..8a6a10aa 100644 --- a/lighthouse/dialects/transform/transform_ext/ops/replace.py +++ b/lighthouse/dialects/transform/transform_ext/ops/replace.py @@ -101,12 +101,15 @@ def allow_repeated_handle_operands(_op: "ReplaceOp") -> bool: class MemoryEffectsOpInterfaceModel(ir.MemoryEffectsOpInterface): @staticmethod - def get_effects(op: ir.Operation, effects): - transform.consumes_handle(op.op_operands[:1], effects) + def get_effects(op: ir.Operation): + effects = transform.consumes_handle(op.op_operands[:1]) if new_operands_handles := op.op_operands[1:]: - transform.only_reads_handle(new_operands_handles, effects) - transform.produces_handle(op.results, effects) - transform.modifies_payload(effects) + effects += transform.only_reads_handle(new_operands_handles) + return ( + effects + + transform.produces_handle(op.results) + + transform.modifies_payload() + ) def replace( diff --git a/lighthouse/dialects/transform/transform_ext/ops/replace_with_fused_attention.py b/lighthouse/dialects/transform/transform_ext/ops/replace_with_fused_attention.py index 4e22eaf7..c6e7fcab 100644 --- a/lighthouse/dialects/transform/transform_ext/ops/replace_with_fused_attention.py +++ b/lighthouse/dialects/transform/transform_ext/ops/replace_with_fused_attention.py @@ -470,15 +470,17 @@ def allow_repeated_handle_operands(_op: "ReplaceWithFusedAttentionOp") -> bool: class MemoryEffectsOpInterfaceModel(ir.MemoryEffectsOpInterface): @staticmethod - def get_effects(op: ir.Operation, effects): - # Read Q, K, scale, V slices - transform.only_reads_handle(op.op_operands[:4], effects) - # Consume and replace output - transform.consumes_handle(op.op_operands[4:5], effects) - # Produce new output handle - transform.produces_handle(op.results, effects) - # Modify the payload - transform.modifies_payload(effects) + def get_effects(op: ir.Operation): + return ( + # Read Q, K, scale, V slices + transform.only_reads_handle(op.op_operands[:4]) + # Consume and replace output + + transform.consumes_handle(op.op_operands[4:5]) + # Produce new output handle + + transform.produces_handle(op.results) + # Modify the payload + + transform.modifies_payload() + ) def replace_with_fused_attention( diff --git a/lighthouse/dialects/transform/transform_ext/ops/reverse_handles.py b/lighthouse/dialects/transform/transform_ext/ops/reverse_handles.py index 819a9ded..fcd51c5b 100644 --- a/lighthouse/dialects/transform/transform_ext/ops/reverse_handles.py +++ b/lighthouse/dialects/transform/transform_ext/ops/reverse_handles.py @@ -41,10 +41,12 @@ def allow_repeated_handle_operands(_op: "ReverseHandlesOp") -> bool: class MemoryEffectsOpInterfaceModel(ir.MemoryEffectsOpInterface): @staticmethod - def get_effects(op: ir.Operation, effects): - transform.only_reads_handle(op.op_operands, effects) - transform.produces_handle(op.results, effects) - transform.only_reads_payload(effects) + def get_effects(op: ir.Operation): + return ( + transform.only_reads_handle(op.op_operands) + + transform.produces_handle(op.results) + + transform.only_reads_payload() + ) def reverse_handles( diff --git a/lighthouse/dialects/transform/transform_ext/ops/trace_producers.py b/lighthouse/dialects/transform/transform_ext/ops/trace_producers.py index 28b46488..a87b659d 100644 --- a/lighthouse/dialects/transform/transform_ext/ops/trace_producers.py +++ b/lighthouse/dialects/transform/transform_ext/ops/trace_producers.py @@ -69,10 +69,12 @@ def allow_repeated_handle_operands(_op: "TraceProducersOp") -> bool: class MemoryEffectsOpInterfaceModel(ir.MemoryEffectsOpInterface): @staticmethod - def get_effects(op: ir.Operation, effects): - transform.only_reads_handle(op.op_operands, effects) - transform.produces_handle(op.results, effects) - transform.only_reads_payload(effects) + def get_effects(op: ir.Operation): + return ( + transform.only_reads_handle(op.op_operands) + + transform.produces_handle(op.results) + + transform.only_reads_payload() + ) def trace_producers( diff --git a/lighthouse/dialects/transform/transform_ext/ops/update_address_space.py b/lighthouse/dialects/transform/transform_ext/ops/update_address_space.py index c0088553..3074a7fc 100644 --- a/lighthouse/dialects/transform/transform_ext/ops/update_address_space.py +++ b/lighthouse/dialects/transform/transform_ext/ops/update_address_space.py @@ -87,10 +87,12 @@ def allow_repeated_handle_operands(_op: "UpdateAddressSpaceOp") -> bool: class MemoryEffectsOpInterfaceModel(ir.MemoryEffectsOpInterface): @staticmethod - def get_effects(op: ir.Operation, effects): - transform.consumes_handle(op.op_operands[:1], effects) - transform.produces_handle(op.results, effects) - transform.modifies_payload(effects) + def get_effects(op: ir.Operation): + return ( + transform.consumes_handle(op.op_operands[:1]) + + transform.produces_handle(op.results) + + transform.modifies_payload() + ) def update_address_space( diff --git a/lighthouse/dialects/transform/transform_ext/ops/wrap_in_benching_func.py b/lighthouse/dialects/transform/transform_ext/ops/wrap_in_benching_func.py index 5ea7273a..82b7b849 100644 --- a/lighthouse/dialects/transform/transform_ext/ops/wrap_in_benching_func.py +++ b/lighthouse/dialects/transform/transform_ext/ops/wrap_in_benching_func.py @@ -101,10 +101,12 @@ def allow_repeated_handle_operands(_op: "WrapInBenchingFuncOp") -> bool: class MemoryEffectsOpInterfaceModel(ir.MemoryEffectsOpInterface): @staticmethod - def get_effects(op: "WrapInBenchingFuncOp", effects): - transform.only_reads_handle(op.op_operands, effects) - transform.produces_handle(op.results, effects) - transform.modifies_payload(effects) + def get_effects(op: "WrapInBenchingFuncOp"): + return ( + transform.only_reads_handle(op.op_operands) + + transform.produces_handle(op.results) + + transform.modifies_payload() + ) def wrap_in_benching_func( diff --git a/lighthouse/dialects/transform/transform_ext/utils/make_filter_handles_op.py b/lighthouse/dialects/transform/transform_ext/utils/make_filter_handles_op.py index 49e1e6de..5a67521a 100644 --- a/lighthouse/dialects/transform/transform_ext/utils/make_filter_handles_op.py +++ b/lighthouse/dialects/transform/transform_ext/utils/make_filter_handles_op.py @@ -71,10 +71,12 @@ def allow_repeated_handle_operands(_op: "_FilterHandlesOp") -> bool: class MemoryEffectsOpInterfaceModel(ir.MemoryEffectsOpInterface): @staticmethod - def get_effects(op: ir.Operation, effects): - transform.only_reads_handle(op.op_operands, effects) - transform.produces_handle(op.results, effects) - transform.only_reads_payload(effects) + def get_effects(op: ir.Operation): + return ( + transform.only_reads_handle(op.op_operands) + + transform.produces_handle(op.results) + + transform.only_reads_payload() + ) else: @@ -116,10 +118,12 @@ def allow_repeated_handle_operands(_op: "_FilterHandlesOp") -> bool: class MemoryEffectsOpInterfaceModel(ir.MemoryEffectsOpInterface): @staticmethod - def get_effects(op: ir.Operation, effects): - transform.only_reads_handle(op.op_operands, effects) - transform.produces_handle(op.results, effects) - transform.only_reads_payload(effects) + def get_effects(op: ir.Operation): + return ( + transform.only_reads_handle(op.op_operands) + + transform.produces_handle(op.results) + + transform.only_reads_payload() + ) _FilterHandlesOp.__name__ = op_name _FilterHandlesOp.__qualname__ = op_name diff --git a/pyproject.toml b/pyproject.toml index c98b02b7..b50dc983 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -3,7 +3,7 @@ name = "lighthouse" dynamic = ["version"] requires-python = ">=3.10,<3.13" # Bounds are due to torch-mlir's packaging dependencies = [ - "mlir-python-bindings==20260806+af2a0e9c8", + "mlir-python-bindings==20260820+002905df0", "pyyaml>=6.0", ]