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
102 changes: 45 additions & 57 deletions asl_xdsl/dialects/asl.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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:
Expand All @@ -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
Expand All @@ -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:
Expand All @@ -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(
Expand All @@ -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
Expand All @@ -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]:
Expand All @@ -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
Expand All @@ -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]:
Expand All @@ -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
Expand All @@ -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:
Expand Down Expand Up @@ -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."""
Expand All @@ -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(
Expand All @@ -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):
Expand Down Expand Up @@ -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)


Expand Down Expand Up @@ -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)


Expand Down
13 changes: 9 additions & 4 deletions asl_xdsl/dialects/asl_dep.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)


Expand Down Expand Up @@ -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)


Expand Down
5 changes: 3 additions & 2 deletions asl_xdsl/interpreters/asl.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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
Expand All @@ -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
Expand Down
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
@@ -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]
Expand Down
6 changes: 3 additions & 3 deletions tests/filecheck/analysis/test-integer-range-analysis.mlir
Original file line number Diff line number Diff line change
Expand Up @@ -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
}
8 changes: 4 additions & 4 deletions uv.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.