Darkweb007/cuda-kernels
0
1"""Minimal computation-graph IR.2 3Deliberately tiny: a Graph is just an ordered list of Nodes. Each Node is one4op applied to named inputs, producing one named output. Graph inputs are any5name that's never produced by a node; graph outputs are names never consumed6(or explicitly marked).7"""8 9from __future__ import annotations10 11from dataclasses import dataclass, field12from typing import List, Optional13 14 15@dataclass16class Node:17 op: str # op name, must exist in ops.OP_REGISTRY18 inputs: List[str] # names of input values (graph inputs or other nodes' outputs)19 output: str # name of the value this node produces20 scalar_args: dict = field(default_factory=dict) # e.g. {"alpha": 0.5} for scalar_mul21 22 def __repr__(self) -> str:23 args = ", ".join(self.inputs)24 extra = f", {self.scalar_args}" if self.scalar_args else ""25 return f"{self.output} = {self.op}({args}{extra})"26 27 28class Graph:29 """Ordered list of Nodes forming a DAG (single static assignment: each30 output name is produced exactly once)."""31 32 def __init__(self):33 self.nodes: List[Node] = []34 self._outputs = set()35 36 def add(self, op: str, inputs: List[str], output: str, **scalar_args) -> Node:37 if output in self._outputs:38 raise ValueError(f"output name '{output}' already produced by another node (not SSA)")39 node = Node(op=op, inputs=list(inputs), output=output, scalar_args=scalar_args)40 self.nodes.append(node)41 self._outputs.add(output)42 return node43 44 def producer_of(self, name: str) -> Optional[Node]:45 """Return the Node that produces `name`, or None if `name` is a graph input."""46 for node in self.nodes:47 if node.output == name:48 return node49 return None50 51 def consumers_of(self, name: str) -> List[Node]:52 """All nodes that read `name` as an input."""53 return [n for n in self.nodes if name in n.inputs]54 55 def graph_inputs(self) -> List[str]:56 """Names consumed but never produced within this graph -- external inputs."""57 produced = {n.output for n in self.nodes}58 seen = []59 for n in self.nodes:60 for i in n.inputs:61 if i not in produced and i not in seen:62 seen.append(i)63 return seen64 65 def final_outputs(self) -> List[str]:66 """Names produced but never consumed within this graph -- graph outputs."""67 consumed = {i for n in self.nodes for i in n.inputs}68 return [n.output for n in self.nodes if n.output not in consumed]69 70 def __iter__(self):71 return iter(self.nodes)72 73 def __len__(self):74 return len(self.nodes)75 76 def __repr__(self) -> str:77 return "Graph(\n " + "\n ".join(repr(n) for n in self.nodes) + "\n)"78 