File size: 2,969 Bytes
fa6719a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
"""Generate CUDA source for a single FusionGroup.

Each fusion group becomes one `__global__` kernel with a grid-stride loop.
External inputs are read once from global memory into a register per
element; every intermediate value in the group lives in a local (register)
variable; only the group's single external output is written back to global
memory. That's the entire point of fusion: N ops -> 1 kernel launch, 1 read
per external input, 1 write.
"""

from __future__ import annotations

from typing import Dict, List

from .fuser import FusionGroup
from .ops import OP_REGISTRY


def _c_identifier(name: str) -> str:
    """Sanitize a graph value name into a valid C identifier."""
    safe = "".join(c if (c.isalnum() or c == "_") else "_" for c in name)
    if safe and safe[0].isdigit():
        safe = "_" + safe
    return safe


def generate_cuda_source(group: FusionGroup, kernel_name: str = "fused_kernel") -> str:
    inputs = group.inputs
    output = group.output
    produced_in_group = {n.output for n in group.nodes}

    def value_expr(name: str) -> str:
        """How to refer to `name` inside the loop body: a local var if it's
        an intermediate produced within this group, else an indexed global
        memory read."""
        if name in produced_in_group:
            return _c_identifier(name)
        return f"{_c_identifier(name)}[i]"

    lines: List[str] = []
    for node in group.nodes:
        spec = OP_REGISTRY[node.op]
        fmt_args: Dict[str, str] = {}
        arg_letters = ["a", "b", "c", "d"]
        for letter, input_name in zip(arg_letters, node.inputs):
            fmt_args[letter] = value_expr(input_name)
        for scalar_name, scalar_val in node.scalar_args.items():
            fmt_args[scalar_name] = str(scalar_val)

        expr = spec.template.format(**fmt_args)
        var_name = _c_identifier(node.output)
        lines.append(f"        float {var_name} = {expr};")

    body = "\n".join(lines)
    output_var = _c_identifier(output)

    params = ", ".join(f"const float* __restrict__ {_c_identifier(name)}" for name in inputs)
    if params:
        params += ", "

    source = f"""\
// Auto-generated by fusion_compiler.codegen -- do not hand-edit.
// Fuses {len(group.nodes)} elementwise op(s) into a single kernel:
//   {' ; '.join(repr(n) for n in group.nodes)}
extern "C" __global__ void {kernel_name}(
    {params}float* __restrict__ {output_var}_out, int n)
{{
    for (int i = blockIdx.x * blockDim.x + threadIdx.x; i < n; i += blockDim.x * gridDim.x) {{
{body}
        {output_var}_out[i] = {output_var};
    }}
}}
"""
    return source


def generate_all(groups: List[FusionGroup], name_prefix: str = "fused_kernel") -> Dict[str, str]:
    """Generate CUDA source for every group, keyed by kernel name."""
    result = {}
    for idx, group in enumerate(groups):
        name = f"{name_prefix}_{idx}"
        result[name] = generate_cuda_source(group, kernel_name=name)
    return result