diff --git a/asl_xdsl/dialects/asl.py b/asl_xdsl/dialects/asl.py index c39ebb0..33cf006 100644 --- a/asl_xdsl/dialects/asl.py +++ b/asl_xdsl/dialects/asl.py @@ -20,12 +20,12 @@ from xdsl.irdl import ( BaseAttr, IRDLOperation, - ParameterDef, VarConstraint, irdl_attr_definition, irdl_op_definition, operand_def, opt_operand_def, + param_def, prop_def, region_def, result_def, @@ -58,10 +58,7 @@ class ConstraintExactAttr(ParametrizedAttribute): name = "asl.int_constraint_exact" - value_attr: ParameterDef[builtin.IntAttr] - - def __init__(self, value: int): - super().__init__([builtin.IntAttr(value)]) + value_attr: builtin.IntAttr = param_def(converter=builtin.IntAttr.get) @property def value(self) -> int: @@ -71,12 +68,12 @@ def value(self) -> int: def parse_parameter(cls, parser: AttrParser) -> int: """Parse the attribute parameter.""" with parser.in_angle_brackets(): - value = parser.parse_integer() - return value + return parser.parse_integer() def print_parameter(self, printer: Printer) -> None: """Print the attribute parameter.""" - printer.print("<", self.value, ">") + with printer.in_angle_brackets(): + return printer.print_int(self.value_attr.data) @irdl_attr_definition @@ -85,11 +82,8 @@ class ConstraintRangeAttr(ParametrizedAttribute): name = "asl.int_constraint_range" - min_value_attr: ParameterDef[builtin.IntAttr] - max_value_attr: ParameterDef[builtin.IntAttr] - - def __init__(self, min_value: int, max_value: int): - super().__init__([builtin.IntAttr(min_value), builtin.IntAttr(max_value)]) + min_value_attr: builtin.IntAttr = param_def(converter=builtin.IntAttr.get) + max_value_attr: builtin.IntAttr = param_def(converter=builtin.IntAttr.get) @property def min_value(self) -> int: @@ -110,7 +104,10 @@ def parse_parameters(cls, parser: AttrParser) -> Sequence[Attribute]: def print_parameters(self, printer: Printer) -> None: """Print the attribute parameters.""" - printer.print("<", self.min_value, ":", self.max_value, ">") + with printer.in_angle_brackets(): + printer.print_int(self.min_value) + printer.print_string(":") + printer.print_int(self.max_value) def _parse_integer_constraint( @@ -129,9 +126,11 @@ def _print_integer_constraint( ): """Print an integer constraint using a shorthand syntax.""" if isinstance(constraint, ConstraintExactAttr): - printer.print(constraint.value) + printer.print_int(constraint.value) else: - printer.print(constraint.min_value, ":", constraint.max_value) + printer.print_int(constraint.min_value) + printer.print_string(":") + printer.print_int(constraint.max_value) @irdl_attr_definition @@ -147,12 +146,10 @@ class IntegerType(ParametrizedAttribute, TypeAttribute): name = "asl.int" - constraints_attr: ParameterDef[ - builtin.ArrayAttr[ConstraintExactAttr | ConstraintRangeAttr] - ] + constraints_attr: builtin.ArrayAttr[ConstraintExactAttr | ConstraintRangeAttr] def __init__(self, constraints: Sequence[Attribute] = ()): - super().__init__([builtin.ArrayAttr(constraints)]) + super().__init__(builtin.ArrayAttr(constraints)) @property def constraints(self) -> Sequence[ConstraintExactAttr | ConstraintRangeAttr]: @@ -173,12 +170,11 @@ def parse_parameters(cls, parser: AttrParser) -> Sequence[Attribute]: def print_parameters(self, printer: Printer) -> None: if not self.constraints: return - printer.print("<") - printer.print_list( - self.constraints, - lambda constr: _print_integer_constraint(constr, printer), - ) - printer.print(">") + with printer.in_angle_brackets(): + printer.print_list( + self.constraints, + lambda constr: _print_integer_constraint(constr, printer), + ) @irdl_attr_definition @@ -187,10 +183,7 @@ class BitVectorType(ParametrizedAttribute, TypeAttribute): name = "asl.bits" - width: ParameterDef[builtin.IntAttr] - - def __init__(self, width: int): - super().__init__([builtin.IntAttr(width)]) + width: builtin.IntAttr = param_def(converter=builtin.IntAttr.get) @classmethod def parse_parameters(cls, parser: AttrParser) -> Sequence[Attribute]: @@ -202,9 +195,8 @@ def parse_parameters(cls, parser: AttrParser) -> Sequence[Attribute]: def print_parameters(self, printer: Printer) -> None: """Print the attribute parameters.""" - printer.print("<") - printer.print(self.width.data) - printer.print(">") + with printer.in_angle_brackets(): + printer.print_int(self.width.data) ArrayElementType: TypeAlias = BitVectorType @@ -221,15 +213,8 @@ class ArrayType( name = "asl.array" - shape: ParameterDef[builtin.ArrayAttr[builtin.IntAttr]] - element_type: ParameterDef[ArrayElementType] - - def __init__( - self, - shape: builtin.ArrayAttr[builtin.IntAttr], - element_type: ArrayElementType, - ): - super().__init__([shape, element_type]) + shape: builtin.ArrayAttr[builtin.IntAttr] + element_type: ArrayElementType def verify(self) -> None: if not self.shape.data: @@ -271,8 +256,12 @@ class BitVectorAttr(ParametrizedAttribute): name = "asl.bits_attr" - value: ParameterDef[builtin.IntAttr] - type: ParameterDef[BitVectorType] + value: builtin.IntAttr + type: BitVectorType + + def __init__(self, value: int, type: BitVectorType): + value = self.normalize_value(value, type.width.data) + super().__init__(builtin.IntAttr(value), type) def maximum_value(self) -> int: """Return the maximum value that can be represented.""" @@ -284,10 +273,6 @@ def normalize_value(value: int, width: int) -> int: max_value = 1 << width return ((value % max_value) + max_value) % max_value - def __init__(self, value: int, type: BitVectorType): - value = self.normalize_value(value, type.width.data) - super().__init__([builtin.IntAttr(value), type]) - def _verify(self) -> None: if self.value.data < 0 or self.value.data > self.maximum_value(): raise VerifyException( @@ -307,11 +292,10 @@ def parse_parameters(cls, parser: AttrParser) -> Sequence[Attribute]: def print_parameters(self, printer: Printer) -> None: """Print the attribute parameters.""" - printer.print("<") - printer.print(self.value.data) - printer.print(" : ") - printer.print(self.type.width.data) - printer.print(">") + with printer.in_angle_brackets(): + printer.print_string(str(self.value.data)) + printer.print_string(" : ") + printer.print_string(str(self.type.width.data)) class ConstantIntIntegerRangeTrait(IntegerRangeTrait): @@ -352,9 +336,10 @@ def parse(cls, parser: Parser) -> ConstantIntOp: def print(self, printer: Printer) -> None: """Print the operation.""" - printer.print(" ", self.value.data) + printer.print_string(" ") + printer.print_int(self.value.data) if self.attributes: - printer.print(" ") + printer.print_string(" ") printer.print_attr_dict(self.attributes) @@ -394,9 +379,12 @@ def parse(cls, parser: Parser) -> ConstantBitVectorOp: def print(self, printer: Printer) -> None: """Print the operation.""" - printer.print(" ", self.value.value.data, " : ", self.res.type) + printer.print_string(" ") + printer.print_int(self.value.value.data) + printer.print_string(" : ") + printer.print_attribute(self.res.type) if self.attributes: - printer.print(" ") + printer.print_string(" ") printer.print_attr_dict(self.attributes) diff --git a/asl_xdsl/dialects/asl_dep.py b/asl_xdsl/dialects/asl_dep.py index b95ca13..e440361 100644 --- a/asl_xdsl/dialects/asl_dep.py +++ b/asl_xdsl/dialects/asl_dep.py @@ -101,9 +101,10 @@ def parse(cls, parser: Parser) -> ConstantIntOp: def print(self, printer: Printer) -> None: """Print the operation.""" - printer.print(" ", self.value_attr.data) + printer.print_string(" ") + printer.print_int(self.value_attr.data) if self.attributes: - printer.print(" ") + printer.print_string(" ") printer.print_attr_dict(self.attributes) @@ -144,9 +145,13 @@ def parse(cls, parser: Parser) -> ConstantBitsOp: return op def print(self, printer: Printer) -> None: - printer.print(" ", self.value_attr.data, " : bits<", self.value_width, ">") + printer.print_string(" ") + printer.print_int(self.value_attr.data) + printer.print_string(" : bits<") + printer.print_ssa_value(self.value_width) + printer.print_string(">") if self.attributes: - printer.print(" ") + printer.print_string(" ") printer.print_attr_dict(self.attributes) diff --git a/asl_xdsl/interpreters/asl.py b/asl_xdsl/interpreters/asl.py index 7e6f8e7..0717fe5 100644 --- a/asl_xdsl/interpreters/asl.py +++ b/asl_xdsl/interpreters/asl.py @@ -15,6 +15,7 @@ ) from xdsl.ir import Operation from xdsl.utils.comparisons import to_signed, to_unsigned +from xdsl.utils.hints import isa from asl_xdsl.dialects import asl @@ -541,7 +542,7 @@ def asl_print_bits_hex( def asl_print_sintN_hex( self, interpreter: Interpreter, op: asl.PrintSIntNHexOp, args: PythonValues ) -> PythonValues: - assert isinstance(op.arg.type, builtin.IntegerType) + assert isa(op.arg.type, builtin.IntegerType) width = op.arg.type.width.data arg: int (arg,) = args @@ -556,7 +557,7 @@ def asl_print_sintN_hex( def asl_print_sintN_dec( self, interpreter: Interpreter, op: asl.PrintSIntNDecOp, args: PythonValues ) -> PythonValues: - assert isinstance(op.arg.type, builtin.IntegerType) + assert isa(op.arg.type, builtin.IntegerType) width = op.arg.type.width.data arg: int (arg,) = args diff --git a/pyproject.toml b/pyproject.toml index f0ebf45..4fbec9a 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,7 +1,7 @@ [project] name = "asl-xdsl" version = "0.0.0" -dependencies = ["xdsl==0.40.0"] +dependencies = ["xdsl==0.51.0"] requires-python = ">=3.10" [project.optional-dependencies] diff --git a/tests/filecheck/analysis/test-integer-range-analysis.mlir b/tests/filecheck/analysis/test-integer-range-analysis.mlir index 155fa1f..0127f5a 100644 --- a/tests/filecheck/analysis/test-integer-range-analysis.mlir +++ b/tests/filecheck/analysis/test-integer-range-analysis.mlir @@ -2,12 +2,12 @@ builtin.module { - // CHECK: %unknown = "test.op"() {__integer_ranges = {{[}}[#none, #none]]} : () -> !asl.int + // CHECK: %unknown = "test.op"() {__integer_ranges = {{[}}[none, none]]} : () -> !asl.int %unknown = "test.op"() : () -> !asl.int - + // CHECK-NEXT: %cst16 = asl.constant_int 16 {__integer_ranges = {{[}}[#builtin.int<16>, #builtin.int<16>]]} %cst16 = asl.constant_int 16 - + // CHECK-NEXT: %res = asl.mod_pow2_int %unknown, %cst16 : (!asl.int, !asl.int) -> !asl.int {__integer_ranges = {{[}}[#builtin.int<0>, #builtin.int<65535>]]} %res = asl.mod_pow2_int %unknown, %cst16 : (!asl.int, !asl.int) -> !asl.int } diff --git a/uv.lock b/uv.lock index ef4843f..7f1b129 100644 --- a/uv.lock +++ b/uv.lock @@ -33,7 +33,7 @@ requires-dist = [ { name = "psutil", marker = "extra == 'dev'", specifier = "==7.0.0" }, { name = "pyright", marker = "extra == 'dev'", specifier = "==1.1.405" }, { name = "pytest", marker = "extra == 'dev'", specifier = "<8.5" }, - { name = "xdsl", specifier = "==0.40.0" }, + { name = "xdsl", specifier = "==0.51.0" }, ] provides-extras = ["dev"] @@ -354,14 +354,14 @@ wheels = [ [[package]] name = "xdsl" -version = "0.40.0" +version = "0.51.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "immutabledict" }, { name = "ordered-set" }, { name = "typing-extensions" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/f7/c2/895c17df346274f0dba78fa34ef2f80e10b996c3059cdb4f368969d831e6/xdsl-0.40.0.tar.gz", hash = "sha256:ffd8a5b8c2ca27d7b736f8b9dfe3359cdfa24d1cae735ac506a4ccd5a28f3817", size = 1437311, upload-time = "2025-06-09T15:59:18.262Z" } +sdist = { url = "https://files.pythonhosted.org/packages/6e/11/bbadef39add07dbc6c2a7cf81b4c3f918b473ac2f4f6433090112dd3ad71/xdsl-0.51.0.tar.gz", hash = "sha256:b9b1b4271b9e5065329bf70419bff9c54897dd1d5bf6004b2f1353838ede30f1", size = 3774880, upload-time = "2025-09-11T15:24:29.183Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/46/5c/3bbefdc9001bdb7fc9a6ccf083121a6131d932d99a300135714954709edf/xdsl-0.40.0-py3-none-any.whl", hash = "sha256:3d5333c5bee18649201131c285ef28b6d65dfd094f88044c8a1df17c6b523ff6", size = 1655407, upload-time = "2025-06-09T15:59:16.682Z" }, + { url = "https://files.pythonhosted.org/packages/19/89/0848a39f78f4d6616304e3ccc972004e67e1488a500012e3b7285aeec9cc/xdsl-0.51.0-py3-none-any.whl", hash = "sha256:eeff28c51c9a21770d3e021532d93d38b012324959776eccbafd8badcd71e13b", size = 3992449, upload-time = "2025-09-11T15:24:27.016Z" }, ]