CoolFace
Apppublic

Darkweb007/cuda-kernels

sourceHugging Faceupdated 2mo agoView on Hugging Face
0likes
graph.py78 linesDownload Raw Back to fusion_compiler
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