[2 / 2] apply sgl-kernel weak_ref_tensor (#12978)
Co-authored-by: Yuan Luo <yuan.luo@hotmail.com>
This commit is contained in:
@@ -8,26 +8,14 @@ from unittest.mock import patch
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.fx as fx
|
import torch.fx as fx
|
||||||
|
from sgl_kernel import weak_ref_tensor
|
||||||
|
|
||||||
import sglang.srt.compilation.weak_ref_tensor_jit # noqa: F401
|
|
||||||
from sglang.srt.compilation.compilation_config import CompilationConfig
|
from sglang.srt.compilation.compilation_config import CompilationConfig
|
||||||
from sglang.srt.compilation.compilation_counter import compilation_counter
|
from sglang.srt.compilation.compilation_counter import compilation_counter
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
def weak_ref_tensor(tensor: Any) -> Any:
|
|
||||||
"""
|
|
||||||
Create a weak reference to a tensor.
|
|
||||||
The new tensor will share the same data as the original tensor,
|
|
||||||
but will not keep the original tensor alive.
|
|
||||||
"""
|
|
||||||
if isinstance(tensor, torch.Tensor):
|
|
||||||
# TODO(yuwei): introduce weak_ref_tensor from sgl_kernel
|
|
||||||
return torch.ops.jit_weak_ref_tensor.weak_ref_tensor(tensor)
|
|
||||||
return tensor
|
|
||||||
|
|
||||||
|
|
||||||
def weak_ref_tensors(
|
def weak_ref_tensors(
|
||||||
tensors: Union[torch.Tensor, list[torch.Tensor], tuple[torch.Tensor]]
|
tensors: Union[torch.Tensor, list[torch.Tensor], tuple[torch.Tensor]]
|
||||||
) -> Union[torch.Tensor, list[Any], tuple[Any], Any]:
|
) -> Union[torch.Tensor, list[Any], tuple[Any], Any]:
|
||||||
|
|||||||
@@ -1,28 +0,0 @@
|
|||||||
// Adapted from: https://github.com/vllm-project/vllm/blob/main/csrc/ops.h
|
|
||||||
|
|
||||||
#include <torch/extension.h>
|
|
||||||
#include <vector>
|
|
||||||
|
|
||||||
static at::Tensor weak_ref_tensor(at::Tensor &tensor) {
|
|
||||||
TORCH_CHECK(tensor.is_cuda(), "weak_ref_tensor expects a CUDA tensor");
|
|
||||||
|
|
||||||
void *data_ptr = tensor.data_ptr();
|
|
||||||
std::vector<int64_t> sizes = tensor.sizes().vec();
|
|
||||||
std::vector<int64_t> strides = tensor.strides().vec();
|
|
||||||
|
|
||||||
auto options = tensor.options();
|
|
||||||
|
|
||||||
auto new_tensor = torch::from_blob(data_ptr, sizes, strides, options);
|
|
||||||
|
|
||||||
return new_tensor;
|
|
||||||
}
|
|
||||||
|
|
||||||
TORCH_LIBRARY(jit_weak_ref_tensor, ops) {
|
|
||||||
ops.def("weak_ref_tensor(Tensor input) -> Tensor");
|
|
||||||
}
|
|
||||||
|
|
||||||
TORCH_LIBRARY_IMPL(jit_weak_ref_tensor, CUDA, ops) {
|
|
||||||
ops.impl("weak_ref_tensor", weak_ref_tensor);
|
|
||||||
}
|
|
||||||
|
|
||||||
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {}
|
|
||||||
@@ -1,11 +0,0 @@
|
|||||||
import os
|
|
||||||
|
|
||||||
from torch.utils.cpp_extension import load
|
|
||||||
|
|
||||||
_abs_path = os.path.dirname(os.path.abspath(__file__))
|
|
||||||
|
|
||||||
load(
|
|
||||||
name="weak_ref_tensor_ext",
|
|
||||||
sources=[f"{_abs_path}/weak_ref_tensor.cpp"],
|
|
||||||
extra_cflags=["-O3"],
|
|
||||||
)
|
|
||||||
Reference in New Issue
Block a user