[Minor] Enhance JIT kernel and add dev docs (#14570)

This commit is contained in:
DarkSharpness
2025-12-23 22:34:59 +08:00
committed by GitHub
parent d7301c89ba
commit 291f11ae39
13 changed files with 817 additions and 290 deletions

View File

@@ -0,0 +1,29 @@
from __future__ import annotations
import functools
from typing import TYPE_CHECKING
import torch
from sglang.jit_kernel.utils import load_jit, make_cpp_args
if TYPE_CHECKING:
from tvm_ffi.module import Module
@functools.cache
def _jit_add_constant_module(constant: int) -> Module:
args = make_cpp_args(constant) # pass all the template argument
return load_jit(
"add_constant",
*args,
cuda_files=["add_constant.cuh"],
cuda_wrappers=[("add_constant", f"add_constant<{args}>")],
)
def add_constant(src: torch.Tensor, constant: int) -> torch.Tensor:
dst = torch.empty_like(src)
module = _jit_add_constant_module(constant)
module.add_constant(dst, src)
return dst