17 lines
521 B
Python
17 lines
521 B
Python
# SPDX-License-Identifier: Apache-2.0
|
|
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
|
from collections.abc import Sequence
|
|
|
|
import torch
|
|
|
|
from vllm.utils.platform_utils import is_uva_available
|
|
|
|
|
|
class UvaBuffer:
|
|
def __init__(self, size: int | Sequence[int], dtype: torch.dtype):
|
|
if not is_uva_available():
|
|
raise RuntimeError("UVA is not available")
|
|
self.cpu = torch.zeros(size, dtype=dtype, device="cpu")
|
|
self.np = self.cpu.numpy()
|
|
self.uva = self.cpu
|