Darkweb007's picture
Upload folder using huggingface_hub
fa6719a verified
Raw
History Blame Contribute Delete
2.97 kB
"""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