Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
31 changes: 31 additions & 0 deletions api/src/inference/decoder_fusion.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,31 @@
"""Opt-in decoder fusion: freeze weight-norm, fuse Snake via CUDA kernel.

Enable with KOKORO_DECODER_FUSION=1 after the model is on CUDA.
Fails closed when requirements are not met; never enables a CPU fallback.
Validated on Quadro K620 (sm_50): ~8.5% wall-time reduction with exact audio.
"""

from __future__ import annotations

import torch

from . import fused_snake


def enable(model):
if model.training or model.device.type != "cuda":
raise RuntimeError("Decoder fusion requires an evaluation model on CUDA")
if any(
p.dtype != torch.float32 or p.device != model.device for p in model.parameters()
):
raise RuntimeError("Decoder fusion requires FP32 parameters on one CUDA device")

with torch.no_grad():
kernel = fused_snake.Snake()
frozen = 0
for layer in model.modules():
if hasattr(layer, "weight_g") and hasattr(layer, "weight_v"):
torch.nn.utils.remove_weight_norm(layer)
frozen += 1
blocks = fused_snake.install(model, kernel)
return {"frozen_weights": frozen, "fused_blocks": len(blocks)}
189 changes: 189 additions & 0 deletions api/src/inference/fused_snake.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,189 @@
"""Opt-in FP32 Snake activation fused into one CUDA kernel.

Validated on Quadro K620 (sm_50) with exact waveform match vs stock Snake.
Requires CUDA + NVRTC; no CPU fallback.
"""

from __future__ import annotations

import ctypes as C
import glob
import os
from pathlib import Path
import types

import torch


def _find_nvrtc() -> str:
env = os.environ.get("KOKORO_NVRTC_LIB")
if env and Path(env).is_file():
return env
patterns = [
"/usr/local/cuda/lib64/libnvrtc.so*",
"/usr/lib/libnvrtc.so*",
"/usr/lib64/libnvrtc.so*",
str(Path(torch.__file__).resolve().parent / "lib" / "libnvrtc.so*"),
str(
Path.home()
/ ".local/lib/python*/site-packages/nvidia/cuda_nvrtc/lib/libnvrtc.so*"
),
"/usr/lib/python*/site-packages/nvidia/cuda_nvrtc/lib/libnvrtc.so*",
]
hits: list[str] = []
for pattern in patterns:
hits.extend(glob.glob(pattern))
# Prefer versioned .so.N over bare .so when both exist
hits = sorted(set(hits), key=lambda p: (0 if ".so." in Path(p).name else 1, p))
if not hits:
raise RuntimeError(
"libnvrtc not found; set KOKORO_NVRTC_LIB or install CUDA/NVRTC"
)
return hits[0]


class Snake:
"""Compile and launch a contiguous FP32 Snake kernel on the current stream."""

def __init__(self, arch: str | None = None):
major, minor = torch.cuda.get_device_capability()
arch = arch or f"compute_{major}{minor}"
nv = C.CDLL(_find_nvrtc())
self.driver = C.CDLL("libcuda.so.1")

def bind(lib, name, args):
fn = getattr(lib, name)
fn.argtypes = args
fn.restype = C.c_int
return fn

void = C.c_void_p
ptr = C.POINTER(void)
create = bind(
nv,
"nvrtcCreateProgram",
[ptr, C.c_char_p, C.c_char_p, C.c_int, C.POINTER(C.c_char_p), C.POINTER(C.c_char_p)],
)
compile_ = bind(nv, "nvrtcCompileProgram", [void, C.c_int, C.POINTER(C.c_char_p)])
logsize = bind(nv, "nvrtcGetProgramLogSize", [void, C.POINTER(C.c_size_t)])
getlog = bind(nv, "nvrtcGetProgramLog", [void, void])
getsize = bind(nv, "nvrtcGetPTXSize", [void, C.POINTER(C.c_size_t)])
getptx = bind(nv, "nvrtcGetPTX", [void, void])
destroy = bind(nv, "nvrtcDestroyProgram", [ptr])

source = b"""extern "C" __global__ void snake(
const float* x, const float* a, float* y, int n, int channels, int length) {
int i = blockIdx.x * blockDim.x + threadIdx.x;
if (i >= n) return;
float alpha = a[(i / length) % channels];
float v = x[i];
float s = sinf(__fmul_rn(alpha, v));
y[i] = __fadd_rn(v, __fmul_rn(__frcp_rn(alpha), __fmul_rn(s, s)));
}"""
prog = void()
self.check(create(C.byref(prog), source, b"snake.cu", 0, None, None))
options = (C.c_char_p * 2)(f"--gpu-architecture={arch}".encode(), b"--std=c++11")
err = compile_(prog, 2, options)
size = C.c_size_t()
self.check(logsize(prog, C.byref(size)))
log = C.create_string_buffer(size.value)
self.check(getlog(prog, log))
if err:
raise RuntimeError(log.value.decode())
self.check(getsize(prog, C.byref(size)))
ptx = C.create_string_buffer(size.value)
self.check(getptx(prog, ptx))
self.check(destroy(C.byref(prog)))

torch.cuda.current_stream()
self.module = void()
self.function = void()
self.check(
bind(self.driver, "cuModuleLoadData", [ptr, void])(
C.byref(self.module), ptx
)
)
self.check(
bind(self.driver, "cuModuleGetFunction", [ptr, void, C.c_char_p])(
C.byref(self.function), self.module, b"snake"
)
)
self.launch = bind(
self.driver,
"cuLaunchKernel",
[
void,
C.c_uint,
C.c_uint,
C.c_uint,
C.c_uint,
C.c_uint,
C.c_uint,
C.c_uint,
void,
ptr,
void,
],
)

@staticmethod
def check(err):
if err:
raise RuntimeError(f"CUDA/NVRTC error {err}")

def __call__(self, x, a):
assert x.is_cuda and a.is_cuda and x.device == a.device
assert x.dtype == a.dtype == torch.float32
assert x.is_contiguous() and a.is_contiguous()
assert x.ndim == 3 and a.shape == (1, x.shape[1], 1)
assert not torch.is_grad_enabled() and x.numel() < 2147483647
y = torch.empty_like(x)
values = [
C.c_uint64(x.data_ptr()),
C.c_uint64(a.data_ptr()),
C.c_uint64(y.data_ptr()),
C.c_int(x.numel()),
C.c_int(x.shape[1]),
C.c_int(x.shape[2]),
]
args = (C.c_void_p * len(values))(
*[C.cast(C.byref(v), C.c_void_p) for v in values]
)
self.check(
self.launch(
self.function,
(x.numel() + 255) // 256,
1,
1,
256,
1,
1,
0,
torch.cuda.current_stream().cuda_stream,
args,
None,
)
)
return y


def install(model, kernel):
"""Replace AdaINResBlock1.forward with fused Snake activations."""

def forward(self, x, s):
for c1, c2, n1, n2, a1, a2 in zip(
self.convs1, self.convs2, self.adain1, self.adain2, self.alpha1, self.alpha2
):
xt = kernel(n1(x, s), a1)
xt = c1(xt)
xt = kernel(n2(xt, s), a2)
xt = c2(xt)
x = xt + x
return x

names = []
for name, module in model.named_modules():
if type(module).__name__ == "AdaINResBlock1":
module.forward = types.MethodType(forward, module)
names.append(name)
return names
5 changes: 5 additions & 0 deletions api/src/inference/kokoro_v1.py
Original file line number Diff line number Diff line change
Expand Up @@ -177,6 +177,11 @@ async def load_model(self, path: str) -> None:
# ROCm reports device "cuda", so this is the ROCm path too.
_configure_rocm_backend()
self._model = self._model.cuda()
if os.environ.get("KOKORO_DECODER_FUSION") == "1":
from .decoder_fusion import enable

details = enable(self._model)
logger.info(f"Experimental decoder fusion enabled: {details}")
else:
self._model = self._model.cpu()

Expand Down