bloqade.squin.analysis.schedule.StmtDag
classStmtDag¶source
bloqade.squin.analysis.schedule.StmtDag
Bases: graph.Graph[ir.Statement]
Signature
class StmtDag(id_table: idtable.IdTable[ir.Statement] = (lambda: idtable.IdTable())(), stmts: Dict[str, ir.Statement] = OrderedDict(), out_edges: Dict[str, Set[str]] = OrderedDict(), inc_edges: Dict[str, Set[str]] = OrderedDict(), stmt_index: Dict[ir.Statement, int] = OrderedDict())Parameters
| Name | Type | Default | Description |
|---|---|---|---|
id_table | idtable.IdTable[ir.Statement] | (lambda: idtable.IdTable())() | |
stmts | Dict[str, ir.Statement] | OrderedDict() | |
out_edges | Dict[str, Set[str]] | OrderedDict() | |
inc_edges | Dict[str, Set[str]] | OrderedDict() | |
stmt_index | Dict[ir.Statement, int] | OrderedDict() |
Attributes
| Name | Type | Default | Description |
|---|---|---|---|
id_table | idtable.IdTable[ir.Statement] | field(default_factory=(lambda: idtable.IdTable())) | |
stmts | Dict[str, ir.Statement] | field(default_factory=OrderedDict) | |
out_edges | Dict[str, Set[str]] | field(default_factory=OrderedDict) | |
inc_edges | Dict[str, Set[str]] | field(default_factory=OrderedDict) | |
stmt_index | Dict[ir.Statement, int] | field(default_factory=OrderedDict) |
methodupdate_index¶source
bloqade.squin.analysis.schedule.StmtDag.update_index
methodadd_node¶source
bloqade.squin.analysis.schedule.StmtDag.add_node
methodadd_edge¶source
bloqade.squin.analysis.schedule.StmtDag.add_edge
def add_edge(src: ir.Statement, dst: ir.Statement)Parameters
| Name | Type | Description |
|---|---|---|
src | ir.Statement | |
dst | ir.Statement |
methodget_parents¶source
bloqade.squin.analysis.schedule.StmtDag.get_parents
def get_parents(node: ir.Statement) -> Iterable[ir.Statement]Parameters
| Name | Type | Description |
|---|---|---|
node | ir.Statement |
Returns
Iterable[ir.Statement]
methodget_children¶source
bloqade.squin.analysis.schedule.StmtDag.get_children
def get_children(node: ir.Statement) -> Iterable[ir.Statement]Parameters
| Name | Type | Description |
|---|---|---|
node | ir.Statement |
Returns
Iterable[ir.Statement]
methodget_neighbors¶source
bloqade.squin.analysis.schedule.StmtDag.get_neighbors
def get_neighbors(node: ir.Statement) -> Iterable[ir.Statement]Parameters
| Name | Type | Description |
|---|---|---|
node | ir.Statement |
Returns
Iterable[ir.Statement]
methodget_nodes¶source
bloqade.squin.analysis.schedule.StmtDag.get_nodes
methodget_edges¶source
bloqade.squin.analysis.schedule.StmtDag.get_edges
def get_edges() -> Iterable[tuple[ir.Statement, ir.Statement]]Returns
Iterable[tuple[ir.Statement, ir.Statement]]
methodprint¶source
bloqade.squin.analysis.schedule.StmtDag.print
Signature
def print(printer: Optional[Printer] = None, analysis: dict[ir.SSAValue, Any] | None = None) -> NoneParameters
| Name | Type | Default | Description |
|---|---|---|---|
printer | Optional[Printer] | None | |
analysis | dict[ir.SSAValue, Any] | None | None |
methodtopological_groups¶source
bloqade.squin.analysis.schedule.StmtDag.topological_groups
def topological_groups()Split the dag into topological groups where each group contains nodes that have no dependencies on each other, but have dependencies on nodes in one or more previous groups.
The idea is to yield all nodes with no dependencies, then remove those nodes from the graph repeating until no nodes are left or we reach some upper limit. Worse case is a linear dag, so we can use len(dag.stmts) as the upper limit
If we reach the limit and there are still nodes left, then we have a cyclic dependency.
Raises
| Type | Description |
|---|---|
ValueError | If a cyclic dependency is detected |
Yields
List[str]: A list of node ids in a topological group
source