forked from Karylab-cklius/vllm
Merge branch 'main' into woosuk/ds-exp
This commit is contained in:
@@ -53,7 +53,7 @@ class TileGemm82 {
|
||||
const int64_t ldb, const int64_t ldc,
|
||||
const int32_t block_size, const int32_t dynamic_k_size,
|
||||
const bool accum_c) {
|
||||
static_assert(0 < M <= 8);
|
||||
static_assert(0 < M && M <= 8);
|
||||
using load_vec_t = typename VecTypeTrait<kv_cache_t>::vec_t;
|
||||
|
||||
kv_cache_t* __restrict__ curr_b_0 = b_tile;
|
||||
|
||||
@@ -68,7 +68,7 @@ class TileGemm161 {
|
||||
const int64_t ldb, const int64_t ldc,
|
||||
const int32_t block_size, const int32_t dynamic_k_size,
|
||||
const bool accum_c) {
|
||||
static_assert(0 < M <= 16);
|
||||
static_assert(0 < M && M <= 16);
|
||||
using load_vec_t = typename VecTypeTrait<kv_cache_t>::vec_t;
|
||||
|
||||
kv_cache_t* __restrict__ curr_b_0 = b_tile;
|
||||
|
||||
@@ -39,7 +39,7 @@ class TileGemm82 {
|
||||
|
||||
template <int32_t M>
|
||||
static void gemm_micro(DEFINE_CPU_MICRO_GEMM_PARAMS) {
|
||||
static_assert(0 < M <= 8);
|
||||
static_assert(0 < M && M <= 8);
|
||||
using load_vec_t = typename cpu_utils::VecTypeTrait<scalar_t>::vec_t;
|
||||
|
||||
scalar_t* __restrict__ curr_b_0 = b_ptr;
|
||||
|
||||
@@ -57,8 +57,8 @@ Modular kernels are supported by the following `FusedMoEMethodBase` classes.
|
||||
|
||||
- [`ModelOptFp8MoEMethod`][vllm.model_executor.layers.quantization.modelopt.ModelOptFp8MoEMethod]
|
||||
- [`Fp8MoEMethod`][vllm.model_executor.layers.quantization.fp8.Fp8MoEMethod]
|
||||
- [`CompressedTensorsW4A4Nvfp4MoEMethod`][vllm.model_executor.layers.quantization.compressed_tensors.compressed_tensors_moe.CompressedTensorsW4A4Nvfp4MoEMethod]
|
||||
- [`CompressedTensorsW8A8Fp8MoEMethod`][vllm.model_executor.layers.quantization.compressed_tensors.compressed_tensors_moe.CompressedTensorsW8A8Fp8MoEMethod]
|
||||
- [`CompressedTensorsW4A4Nvfp4MoEMethod`][vllm.model_executor.layers.quantization.compressed_tensors.compressed_tensors_moe.compressed_tensors_moe_w4a4_nvfp4.CompressedTensorsW4A4Nvfp4MoEMethod]
|
||||
- [`CompressedTensorsW8A8Fp8MoEMethod`][vllm.model_executor.layers.quantization.compressed_tensors.compressed_tensors_moe.compressed_tensors_moe_w8a8_fp8.CompressedTensorsW8A8Fp8MoEMethod]
|
||||
- [`Mxfp4MoEMethod`][vllm.model_executor.layers.quantization.mxfp4.Mxfp4MoEMethod]
|
||||
- [`UnquantizedFusedMoEMethod`][vllm.model_executor.layers.fused_moe.layer.UnquantizedFusedMoEMethod]
|
||||
|
||||
|
||||
@@ -0,0 +1,167 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""
|
||||
MkDocs hook to automatically convert inline code references to API doc links.
|
||||
|
||||
For example, `WeightTransferConfig` becomes
|
||||
[`WeightTransferConfig`][vllm.config.WeightTransferConfig]
|
||||
|
||||
This works with the `autorefs` plugin to create clickable cross-references
|
||||
to API documentation pages generated by `mkdocstrings`.
|
||||
|
||||
The hook builds an index of all documented public Python names (classes and
|
||||
functions with docstrings) from the vllm package at startup using AST parsing,
|
||||
then substitutes matching inline code spans on each page. Names without
|
||||
docstrings are excluded because mkdocstrings will not generate a page for them.
|
||||
"""
|
||||
|
||||
import ast
|
||||
import logging
|
||||
from pathlib import Path
|
||||
|
||||
import regex as re
|
||||
from mkdocs.config.defaults import MkDocsConfig
|
||||
from mkdocs.structure.files import Files
|
||||
from mkdocs.structure.pages import Page
|
||||
|
||||
logger = logging.getLogger("mkdocs")
|
||||
|
||||
ROOT_DIR = Path(__file__).parent.parent.parent.parent.resolve()
|
||||
VLLM_DIR = ROOT_DIR / "vllm"
|
||||
|
||||
# Maps short name -> qualified name (e.g. "ModelConfig" -> "vllm.config.ModelConfig")
|
||||
_name_index: dict[str, str] = {}
|
||||
|
||||
# Fenced code block pattern (``` or ~~~, with optional language specifier).
|
||||
_FENCED_BLOCK = re.compile(
|
||||
r"(?:^|\n)(?P<fence>`{3,}|~{3,})[^\n]*\n.*?(?:\n(?P=fence))", re.DOTALL
|
||||
)
|
||||
|
||||
# Inline code that is NOT already part of a markdown link.
|
||||
# Matches `Name` but not [`Name`] and not [`Name`][...] or [`Name`](...).
|
||||
_INLINE_CODE = re.compile(
|
||||
r"(?<!\[)" # not preceded by [
|
||||
r"`(?P<name>[A-Za-z0-9_]*)`" # `UpperCamelCase` or `UPPER_SNAKE`
|
||||
r"(?!\])" # not followed by ]
|
||||
)
|
||||
|
||||
|
||||
def _has_docstring(node: ast.AST) -> bool:
|
||||
"""Check if a class or function node has a docstring."""
|
||||
if not isinstance(node, ast.ClassDef | ast.FunctionDef | ast.AsyncFunctionDef):
|
||||
return False
|
||||
return ast.get_docstring(node, clean=False) is not None
|
||||
|
||||
|
||||
def _module_path(filepath: Path) -> str:
|
||||
"""Convert a filesystem path to a dotted module path."""
|
||||
rel = filepath.relative_to(ROOT_DIR)
|
||||
parts = list(rel.with_suffix("").parts)
|
||||
if parts[-1] == "__init__":
|
||||
parts = parts[:-1]
|
||||
return ".".join(parts)
|
||||
|
||||
|
||||
def _index_file(filepath: Path) -> dict[str, str]:
|
||||
"""Extract documented public names from a Python file using AST parsing.
|
||||
|
||||
Only classes and functions with docstrings are included, since
|
||||
mkdocstrings won't generate a page for undocumented symbols.
|
||||
"""
|
||||
names: dict[str, str] = {}
|
||||
try:
|
||||
source = filepath.read_text(encoding="utf-8")
|
||||
tree = ast.parse(source, filename=str(filepath))
|
||||
except (SyntaxError, UnicodeDecodeError):
|
||||
return names
|
||||
|
||||
module = _module_path(filepath)
|
||||
|
||||
for node in ast.iter_child_nodes(tree):
|
||||
if (
|
||||
# Class definitions (with docstring)
|
||||
isinstance(node, ast.ClassDef)
|
||||
and not node.name.startswith("_")
|
||||
and _has_docstring(node)
|
||||
) or (
|
||||
# Function definitions (with docstring, only uppercase/CamelCase)
|
||||
isinstance(node, ast.FunctionDef | ast.AsyncFunctionDef)
|
||||
and not node.name.startswith("_")
|
||||
and node.name[0].isupper()
|
||||
and _has_docstring(node)
|
||||
):
|
||||
names[node.name] = f"{module}.{node.name}"
|
||||
|
||||
return names
|
||||
|
||||
|
||||
def _build_index() -> dict[str, str]:
|
||||
"""Walk the vllm package and build a name -> qualified path index."""
|
||||
index: dict[str, str] = {}
|
||||
# Track conflicts: if multiple modules define the same name,
|
||||
# prefer shallower modules (more likely to be the public API).
|
||||
depth: dict[str, int] = {}
|
||||
|
||||
for filepath in sorted(VLLM_DIR.rglob("*.py")):
|
||||
# Skip internal/private modules
|
||||
if any(part.startswith("_") and part != "__init__" for part in filepath.parts):
|
||||
continue
|
||||
# Skip third-party vendored code
|
||||
rel = filepath.relative_to(VLLM_DIR)
|
||||
if rel.parts and rel.parts[0] in ("third_party", "vllm_flash_attn"):
|
||||
continue
|
||||
|
||||
module_depth = len(filepath.relative_to(ROOT_DIR).parts)
|
||||
file_names = _index_file(filepath)
|
||||
|
||||
for name, qualified in file_names.items():
|
||||
if name not in index or module_depth < depth[name]:
|
||||
index[name] = qualified
|
||||
depth[name] = module_depth
|
||||
|
||||
return index
|
||||
|
||||
|
||||
def on_startup(*, command: str, dirty: bool) -> None:
|
||||
"""Build the name index once at startup."""
|
||||
global _name_index
|
||||
_name_index = _build_index()
|
||||
logger.info("autoref_code: indexed %d names from vllm/", len(_name_index))
|
||||
|
||||
|
||||
def on_page_markdown(
|
||||
markdown: str, *, page: Page, config: MkDocsConfig, files: Files
|
||||
) -> str:
|
||||
"""Replace inline code references with autoref links."""
|
||||
if not _name_index:
|
||||
return markdown
|
||||
|
||||
# Skip API reference pages to avoid circular/redundant links.
|
||||
if page.file.src_path.startswith("api/"):
|
||||
return markdown
|
||||
|
||||
# Step 1: Mask fenced code blocks so we don't touch code inside them.
|
||||
masks: list[str] = []
|
||||
|
||||
def _mask_block(match: re.Match) -> str:
|
||||
masks.append(match.group(0))
|
||||
return f"\ue000CODEBLOCK{len(masks) - 1}\ue000"
|
||||
|
||||
masked = _FENCED_BLOCK.sub(_mask_block, markdown)
|
||||
|
||||
# Step 2: Replace inline code references.
|
||||
def _replace(match: re.Match) -> str:
|
||||
name = match.group("name")
|
||||
qualified = _name_index.get(name)
|
||||
if qualified is None:
|
||||
return match.group(0)
|
||||
logger.debug("autoref_code: linking `%s` to [%s]", name, qualified)
|
||||
return f"[`{name}`][{qualified}]"
|
||||
|
||||
result = _INLINE_CODE.sub(_replace, masked)
|
||||
|
||||
# Step 3: Restore masked code blocks.
|
||||
result = re.sub(
|
||||
r"\ue000CODEBLOCK(\d+)\ue000", lambda m: masks[int(m.group(1))], result
|
||||
)
|
||||
return result
|
||||
@@ -54,6 +54,7 @@ hooks:
|
||||
- docs/mkdocs/hooks/generate_argparse.py
|
||||
- docs/mkdocs/hooks/generate_metrics.py
|
||||
- docs/mkdocs/hooks/url_schemes.py
|
||||
- docs/mkdocs/hooks/autoref_code.py
|
||||
|
||||
plugins:
|
||||
- meta
|
||||
|
||||
@@ -1013,6 +1013,7 @@ package_data = {
|
||||
"model_executor/layers/quantization/utils/configs/*.json",
|
||||
"entrypoints/serve/instrumentator/static/*.js",
|
||||
"entrypoints/serve/instrumentator/static/*.css",
|
||||
"distributed/kv_transfer/kv_connector/v1/hf3fs/utils/*.cpp",
|
||||
]
|
||||
}
|
||||
|
||||
|
||||
@@ -3,4 +3,4 @@
|
||||
model_name: openai/gpt-oss-20b
|
||||
metric_threshold: 0.568
|
||||
reasoning_effort: low
|
||||
server_args: "--attention-backend ROCM_AITER_UNIFIED_ATTN"
|
||||
server_args: "--attention-backend ROCM_AITER_UNIFIED_ATTN --tensor-parallel-size 2"
|
||||
@@ -3,6 +3,6 @@
|
||||
model_name: amd/gpt-oss-20b-w-mxfp4-a-bf16
|
||||
metric_threshold: 0.568
|
||||
reasoning_effort: low
|
||||
server_args: "--attention-backend ROCM_AITER_UNIFIED_ATTN --moe-backend aiter"
|
||||
server_args: "--attention-backend ROCM_AITER_UNIFIED_ATTN --moe-backend aiter --tokenizer openai/gpt-oss-20b --tensor-parallel-size 2"
|
||||
env:
|
||||
VLLM_ROCM_USE_AITER: "1"
|
||||
VLLM_ROCM_USE_AITER: "1"
|
||||
@@ -3,4 +3,4 @@
|
||||
model_name: amd/gpt-oss-20b-w-mxfp4-a-bf16
|
||||
metric_threshold: 0.568
|
||||
reasoning_effort: low
|
||||
server_args: "--attention-backend ROCM_AITER_UNIFIED_ATTN --moe-backend triton"
|
||||
server_args: "--attention-backend ROCM_AITER_UNIFIED_ATTN --moe-backend triton --tokenizer openai/gpt-oss-20b --tensor-parallel-size 2"
|
||||
@@ -3,6 +3,6 @@
|
||||
model_name: amd/gpt-oss-20b-MoE-Quant-W-MXFP4-A-FP8-KV-FP8
|
||||
metric_threshold: 0.568
|
||||
reasoning_effort: low
|
||||
server_args: "--attention-backend ROCM_AITER_UNIFIED_ATTN"
|
||||
server_args: "--attention-backend ROCM_AITER_UNIFIED_ATTN --tensor-parallel-size 2"
|
||||
env:
|
||||
VLLM_ROCM_USE_AITER: "1"
|
||||
@@ -6,7 +6,7 @@ import torch
|
||||
|
||||
from tests.kernels.quant_utils import FP8_DTYPE
|
||||
from tests.kernels.utils import opcheck
|
||||
from vllm.model_executor.layers.layernorm import RMSNorm
|
||||
from vllm.model_executor.layers.layernorm import GemmaRMSNorm, RMSNorm
|
||||
from vllm.platforms import current_platform
|
||||
from vllm.utils.torch_utils import set_random_seed
|
||||
|
||||
@@ -162,3 +162,31 @@ def test_fused_rms_norm_quant(
|
||||
atol=1e-3,
|
||||
rtol=1e-3,
|
||||
)
|
||||
|
||||
|
||||
@torch.inference_mode()
|
||||
def test_gemma_rms_norm_mixed_input_weight_dtype(default_vllm_config) -> None:
|
||||
if not torch.cuda.is_available():
|
||||
pytest.skip("CUDA required")
|
||||
|
||||
device = CUDA_DEVICES[0]
|
||||
torch.set_default_device(device)
|
||||
|
||||
num_tokens, hidden_size = 32, 1024
|
||||
x = torch.randn(num_tokens, hidden_size, dtype=torch.bfloat16, device=device)
|
||||
layer = GemmaRMSNorm(hidden_size, eps=1e-6).to(device=device)
|
||||
layer.weight.data.normal_(mean=0.0, std=0.1)
|
||||
|
||||
# Gemma uses fp32 weight parameter while activations can be bf16.
|
||||
assert layer.weight.dtype == torch.float32
|
||||
out = layer(x)
|
||||
|
||||
x_fp32 = x.float()
|
||||
weight_fp32 = layer.weight.data.float() + 1.0
|
||||
variance = x_fp32.pow(2).mean(dim=-1, keepdim=True)
|
||||
ref = (x_fp32 * torch.rsqrt(variance + layer.variance_epsilon) * weight_fp32).to(
|
||||
x.dtype
|
||||
)
|
||||
|
||||
assert out.dtype == x.dtype
|
||||
torch.testing.assert_close(out, ref, atol=1e-2, rtol=1e-2)
|
||||
|
||||
@@ -23,16 +23,12 @@ from triton_kernels.numerics_details.mxfp import downcast_to_mxfp, upcast_from_m
|
||||
from triton_kernels.tensor import FP4, convert_layout, wrap_torch_tensor
|
||||
from triton_kernels.tensor_details import layout
|
||||
from triton_kernels.testing import assert_close
|
||||
from triton_kernels.topk import topk as topk_fn
|
||||
|
||||
from vllm.model_executor.layers.fused_moe.config import mxfp4_w4a16_moe_quant_config
|
||||
from vllm.model_executor.layers.fused_moe.gpt_oss_triton_kernels_moe import (
|
||||
legacy_routing,
|
||||
make_routing_data,
|
||||
triton_kernel_moe_forward,
|
||||
)
|
||||
from vllm.utils.math_utils import round_up
|
||||
from vllm.utils.torch_utils import set_random_seed
|
||||
|
||||
from .utils import shuffle_weight
|
||||
|
||||
@@ -97,10 +93,18 @@ def init_compute_data(M, K, N, E, a_dtype: str, w_dtype: str, num_warps: int):
|
||||
if w_dtype != "mx4":
|
||||
pytest.skip("NYI")
|
||||
else: # quantize to mx4
|
||||
# careful on the padding here, the activation padding need to be
|
||||
# multiple of 64, the actual engine is not implemented
|
||||
w1_bottom_pad = round_up(w1_tri.shape[1], 64) - w1_tri.shape[1]
|
||||
w1_right_pad = round_up(w1_tri.shape[2], 128) - w1_tri.shape[2]
|
||||
# Padding alignment depends on the platform. On CDNA4 the scale
|
||||
# swizzle requires SCALE_K % 8 == 0 (K % 256) and
|
||||
# SCALE_N % 32 == 0 (2*N % 512), matching the production
|
||||
# alignment in mxfp4_round_up_hidden_size_and_intermediate_size.
|
||||
# On CUDA (Hopper) the scale layout pads internally, so the
|
||||
# original 64/128 alignment is sufficient.
|
||||
if current_platform.is_rocm():
|
||||
k_align, n2_align = 256, 512
|
||||
else:
|
||||
k_align, n2_align = 64, 128
|
||||
w1_bottom_pad = round_up(w1_tri.shape[1], k_align) - w1_tri.shape[1]
|
||||
w1_right_pad = round_up(w1_tri.shape[2], n2_align) - w1_tri.shape[2]
|
||||
|
||||
w2_bottom_pad = w1_right_pad // 2
|
||||
w2_right_pad = w1_bottom_pad
|
||||
@@ -367,52 +371,3 @@ def test_unit_shuffle():
|
||||
)
|
||||
|
||||
assert_close(ref=out_ref, tri=out)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("num_tokens", [2, 8, 64])
|
||||
@pytest.mark.parametrize("num_experts", [32, 128])
|
||||
@pytest.mark.parametrize("topk", [1, 4])
|
||||
@pytest.mark.parametrize("renormalize", [True, False])
|
||||
@pytest.mark.parametrize("dtype", [torch.bfloat16])
|
||||
def test_legacy_routing(
|
||||
num_tokens: int, num_experts: int, topk: int, renormalize: bool, dtype: torch.dtype
|
||||
):
|
||||
set_random_seed(0)
|
||||
gating_output = torch.randn(num_tokens, num_experts, device="cuda", dtype=dtype)
|
||||
|
||||
sm_first = not renormalize
|
||||
logits = gating_output
|
||||
if sm_first:
|
||||
logits = torch.softmax(logits, dim=-1)
|
||||
topk_result = topk_fn(logits, topk, apply_softmax=not sm_first)
|
||||
# topk_fn returns SparseMatrix on NVIDIA, plain tuple on ROCm.
|
||||
if isinstance(topk_result, tuple):
|
||||
topk_weights, topk_ids_raw, bitmatrix = topk_result
|
||||
from triton_kernels.routing import routing_from_bitmatrix
|
||||
|
||||
routing_data_ref, gather_indx_ref, scatter_indx_ref = routing_from_bitmatrix(
|
||||
bitmatrix, topk_weights, topk_ids_raw, num_experts, topk
|
||||
)
|
||||
else:
|
||||
topk_ids = topk_result.indx.to(torch.long)
|
||||
topk_weights = topk_result.vals
|
||||
routing_data_ref, gather_indx_ref, scatter_indx_ref = make_routing_data(
|
||||
topk_ids, topk_weights, num_experts
|
||||
)
|
||||
|
||||
routing_data, gather_indx, scatter_indx = legacy_routing(
|
||||
gating_output, topk, sm_first=sm_first
|
||||
)
|
||||
|
||||
assert_close(
|
||||
ref=gather_indx_ref.src_indx, tri=gather_indx.src_indx, maxtol=0, rmstol=0
|
||||
)
|
||||
assert_close(
|
||||
ref=gather_indx_ref.dst_indx, tri=gather_indx.dst_indx, maxtol=0, rmstol=0
|
||||
)
|
||||
assert_close(
|
||||
ref=scatter_indx_ref.src_indx, tri=scatter_indx.src_indx, maxtol=0, rmstol=0
|
||||
)
|
||||
assert_close(
|
||||
ref=scatter_indx_ref.dst_indx, tri=scatter_indx.dst_indx, maxtol=0, rmstol=0
|
||||
)
|
||||
|
||||
@@ -4,12 +4,9 @@
|
||||
Tests that triton_kernel_moe_forward correctly applies expert_map
|
||||
remapping when expert parallelism (EP) is enabled.
|
||||
|
||||
Previously, legacy_routing was always used and it produced routing data
|
||||
with global expert IDs that didn't correspond to local weight indices,
|
||||
causing illegal memory access with EP. The fix splits routing: when
|
||||
expert_map is provided, topk selection is performed first, expert_map is
|
||||
applied to remap global→local IDs, and make_routing_data builds routing
|
||||
structures from the local IDs.
|
||||
Both EP and non-EP paths use topk + make_routing_data. When expert_map
|
||||
is provided, global expert IDs are remapped to local IDs before building
|
||||
routing structures.
|
||||
"""
|
||||
|
||||
from unittest.mock import MagicMock, patch
|
||||
@@ -24,21 +21,15 @@ class TestTritonMoeForwardExpertMap:
|
||||
|
||||
@pytest.mark.parametrize("expert_map_present", [False, True])
|
||||
def test_routing_path_selection(self, expert_map_present):
|
||||
"""Verify that the EP-aware routing path is taken when expert_map
|
||||
is present, and the legacy_routing path is taken otherwise."""
|
||||
"""Verify that both EP and non-EP paths use topk + make_routing_data,
|
||||
and that expert_map remapping is applied when present."""
|
||||
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
# This is a structural test: we mock the routing functions to
|
||||
# verify the correct path is exercised.
|
||||
mock_expert_map = (
|
||||
torch.tensor([0, -1, 1, -1], device=device) if expert_map_present else None
|
||||
)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"vllm.model_executor.layers.fused_moe."
|
||||
"gpt_oss_triton_kernels_moe.legacy_routing"
|
||||
) as mock_legacy,
|
||||
patch("triton_kernels.topk.topk") as mock_topk,
|
||||
patch(
|
||||
"vllm.model_executor.layers.fused_moe."
|
||||
@@ -53,27 +44,19 @@ class TestTritonMoeForwardExpertMap:
|
||||
triton_kernel_moe_forward,
|
||||
)
|
||||
|
||||
# Set up return values
|
||||
mock_routing_data = MagicMock()
|
||||
mock_gather = MagicMock()
|
||||
mock_scatter = MagicMock()
|
||||
|
||||
if expert_map_present:
|
||||
sparse_result = MagicMock()
|
||||
sparse_result.indx = torch.tensor([[0, 2]], dtype=torch.int32)
|
||||
sparse_result.vals = torch.tensor([[0.6, 0.4]])
|
||||
mock_topk.return_value = sparse_result
|
||||
mock_make_routing.return_value = (
|
||||
mock_routing_data,
|
||||
mock_gather,
|
||||
mock_scatter,
|
||||
)
|
||||
else:
|
||||
mock_legacy.return_value = (
|
||||
mock_routing_data,
|
||||
mock_gather,
|
||||
mock_scatter,
|
||||
)
|
||||
sparse_result = MagicMock()
|
||||
sparse_result.indx = torch.tensor([[0, 2]], dtype=torch.int32)
|
||||
sparse_result.vals = torch.tensor([[0.6, 0.4]])
|
||||
mock_topk.return_value = sparse_result
|
||||
mock_make_routing.return_value = (
|
||||
mock_routing_data,
|
||||
mock_gather,
|
||||
mock_scatter,
|
||||
)
|
||||
|
||||
mock_fused_experts.return_value = torch.zeros((1, 8), device=device)
|
||||
|
||||
@@ -92,20 +75,14 @@ class TestTritonMoeForwardExpertMap:
|
||||
expert_map=mock_expert_map,
|
||||
)
|
||||
|
||||
# Both paths use topk + make_routing_data
|
||||
mock_topk.assert_called_once()
|
||||
mock_make_routing.assert_called_once()
|
||||
|
||||
if expert_map_present:
|
||||
# EP path: should use topk + make_routing_data, NOT
|
||||
# legacy_routing
|
||||
mock_topk.assert_called_once()
|
||||
mock_make_routing.assert_called_once()
|
||||
mock_legacy.assert_not_called()
|
||||
# expert_map should be None in the fused_experts call
|
||||
# (already applied)
|
||||
call_kwargs = mock_fused_experts.call_args
|
||||
assert call_kwargs[1].get("expert_map") is None or (
|
||||
len(call_kwargs[0]) > 0
|
||||
)
|
||||
else:
|
||||
# Non-EP path: should use legacy_routing
|
||||
mock_legacy.assert_called_once()
|
||||
mock_topk.assert_not_called()
|
||||
mock_make_routing.assert_not_called()
|
||||
|
||||
@@ -109,6 +109,14 @@ def _load_hf_model(model_name: str, hf_spec: dict, device: torch.device):
|
||||
**extra,
|
||||
).to(device)
|
||||
model.eval()
|
||||
|
||||
# Transformers 5.0 weight materialization can clear non-persistent
|
||||
# buffers (e.g. rotary inv_freq) that were registered with
|
||||
# persistent=False. Re-compute them so the model produces valid output.
|
||||
for mod in model.modules():
|
||||
if hasattr(mod, "_compute_inv_freq") and hasattr(mod, "inv_freq"):
|
||||
mod.inv_freq = mod._compute_inv_freq(device=device)
|
||||
|
||||
return model
|
||||
|
||||
|
||||
|
||||
@@ -89,22 +89,33 @@ def test_models(example_prompts, model_name) -> None:
|
||||
|
||||
EAGER = [True, False]
|
||||
|
||||
SM_100_NVFP4_BACKENDS = [
|
||||
"flashinfer-cudnn",
|
||||
"flashinfer-trtllm",
|
||||
"flashinfer-cutlass",
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
not current_platform.has_device_capability(100),
|
||||
reason="modelopt_fp4 is not supported on this GPU type.",
|
||||
)
|
||||
@pytest.mark.parametrize("model", ["nvidia/Llama-3.1-8B-Instruct-NVFP4"])
|
||||
@pytest.mark.parametrize("eager", EAGER)
|
||||
@pytest.mark.parametrize(
|
||||
"backend",
|
||||
[
|
||||
"emulation",
|
||||
"flashinfer-cudnn",
|
||||
"flashinfer-trtllm", # the small seq_len ensures trtllm_8x4_layout backend is used
|
||||
"flashinfer-cutlass",
|
||||
],
|
||||
)
|
||||
def test_nvfp4(vllm_runner, model, eager, backend, monkeypatch):
|
||||
if (
|
||||
not current_platform.has_device_capability(100)
|
||||
and backend in SM_100_NVFP4_BACKENDS
|
||||
):
|
||||
pytest.skip(
|
||||
f"The backend {backend} is not supported with current_platform.has_device_capability(100) == False"
|
||||
)
|
||||
|
||||
monkeypatch.setenv("VLLM_NVFP4_GEMM_BACKEND", backend)
|
||||
with vllm_runner(model, enforce_eager=eager) as llm:
|
||||
output = llm.generate_greedy(["1 2 3 4 5"], max_tokens=2)
|
||||
|
||||
@@ -366,9 +366,6 @@ def test_compressed_tensors_kv_cache_fp8_per_attn_head(vllm_runner):
|
||||
assert output
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
not current_platform.is_cuda(), reason="This test is skipped on non-CUDA platform."
|
||||
)
|
||||
@pytest.mark.parametrize(
|
||||
"args",
|
||||
[
|
||||
@@ -398,7 +395,7 @@ def test_compressed_tensors_nvfp4(vllm_runner, args):
|
||||
assert qkv_proj.scheme.group_size == 16
|
||||
|
||||
llm.apply_model(check_model)
|
||||
output = llm.generate_greedy("Hello my name is", max_tokens=4)
|
||||
output = llm.generate_greedy(["Hello my name is"], max_tokens=4)
|
||||
print(output)
|
||||
assert output
|
||||
|
||||
|
||||
@@ -15,8 +15,9 @@
|
||||
|
||||
# set -xe
|
||||
|
||||
# Find the git repository root directory
|
||||
GIT_ROOT=$(git rev-parse --show-toplevel)
|
||||
# Resolve the repository root from the script location instead of `.git`.
|
||||
SCRIPT_DIR="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd -P)"
|
||||
GIT_ROOT="${GIT_ROOT:-$(cd -- "${SCRIPT_DIR}/../../../.." && pwd -P)}"
|
||||
|
||||
# Model to test
|
||||
MODEL="${MODEL:-Qwen/Qwen2.5-VL-3B-Instruct}"
|
||||
|
||||
@@ -85,8 +85,11 @@ DECODE_BLOCK_SIZE=${DECODE_BLOCK_SIZE:-128}
|
||||
# Comma-separated extra args for vllm serve (e.g. --max-model-len,2048)
|
||||
VLLM_SERVE_EXTRA_ARGS=${VLLM_SERVE_EXTRA_ARGS:-}
|
||||
|
||||
# Find the git repository root directory
|
||||
GIT_ROOT=$(git rev-parse --show-toplevel)
|
||||
# Resolve the repository root from the script location instead of `.git`.
|
||||
# The ROCm CI image copies `/vllm-workspace` without the Git metadata, so
|
||||
# `git rev-parse --show-toplevel` is not reliable at runtime.
|
||||
SCRIPT_DIR="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd -P)"
|
||||
GIT_ROOT="${GIT_ROOT:-$(cd -- "${SCRIPT_DIR}/../../../.." && pwd -P)}"
|
||||
|
||||
SMI_BIN=$(which nvidia-smi || which rocm-smi || echo "")
|
||||
|
||||
|
||||
@@ -33,8 +33,9 @@ MODELS=(
|
||||
"Qwen/Qwen3-0.6B"
|
||||
)
|
||||
|
||||
# Find the git repository root directory
|
||||
GIT_ROOT=$(git rev-parse --show-toplevel)
|
||||
# Resolve the repository root from the script location instead of `.git`.
|
||||
SCRIPT_DIR="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd -P)"
|
||||
GIT_ROOT="${GIT_ROOT:-$(cd -- "${SCRIPT_DIR}/../../../.." && pwd -P)}"
|
||||
|
||||
# Trap the SIGINT signal (triggered by Ctrl+C)
|
||||
trap 'kill $(jobs -pr)' SIGINT SIGTERM EXIT
|
||||
|
||||
@@ -20,7 +20,8 @@ BLOCK_SIZE=${BLOCK_SIZE:-32}
|
||||
|
||||
|
||||
# execution env
|
||||
GIT_ROOT=$(git rev-parse --show-toplevel)
|
||||
SCRIPT_DIR="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd -P)"
|
||||
GIT_ROOT="${GIT_ROOT:-$(cd -- "${SCRIPT_DIR}/../../../.." && pwd -P)}"
|
||||
EXP_ROOT="${GIT_ROOT}/tests/v1/kv_connector/nixl_integration"
|
||||
CONDA_PATH=${CONDA_PATH:-"/home/${USER}/anaconda3"}
|
||||
CONDA_ENV_NAME=${CONDA_ENV_NAME:-"nixl"}
|
||||
@@ -153,4 +154,4 @@ echo "-----P/D success----"
|
||||
rm "${OUTPUT_FILE}"
|
||||
cleanup
|
||||
|
||||
exit 0
|
||||
exit 0
|
||||
|
||||
@@ -20,7 +20,8 @@ BLOCK_SIZE=${BLOCK_SIZE:-32}
|
||||
|
||||
|
||||
# execution env
|
||||
GIT_ROOT=$(git rev-parse --show-toplevel)
|
||||
SCRIPT_DIR="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd -P)"
|
||||
GIT_ROOT="${GIT_ROOT:-$(cd -- "${SCRIPT_DIR}/../../../.." && pwd -P)}"
|
||||
EXP_ROOT="${GIT_ROOT}/tests/v1/kv_connector/nixl_integration"
|
||||
CONDA_PATH=${CONDA_PATH:-"/home/${USER}/anaconda3"}
|
||||
CONDA_ENV_NAME=${CONDA_ENV_NAME:-"nixl"}
|
||||
|
||||
@@ -44,7 +44,8 @@ DECODER_ZE_AFFINITY_MASK=${DECODER_ZE_AFFINITY_MASK:-$(generate_affinity_mask "$
|
||||
|
||||
|
||||
# execution env
|
||||
GIT_ROOT=$(git rev-parse --show-toplevel)
|
||||
SCRIPT_DIR="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd -P)"
|
||||
GIT_ROOT="${GIT_ROOT:-$(cd -- "${SCRIPT_DIR}/../../../.." && pwd -P)}"
|
||||
EXP_ROOT="${GIT_ROOT}/tests/v1/kv_connector/nixl_integration"
|
||||
|
||||
OUTPUT_FILE=${OUTPUT_FILE:-"${EXP_ROOT}/.xpu_accuracy_test_outputs.txt"}
|
||||
|
||||
@@ -52,7 +52,9 @@ DECODER_TP_SIZE=${DECODER_TP_SIZE:-1}
|
||||
GPU_MEMORY_UTILIZATION=${GPU_MEMORY_UTILIZATION:-0.7}
|
||||
BLOCK_SIZE=${BLOCK_SIZE:-16}
|
||||
|
||||
GIT_ROOT=$(git rev-parse --show-toplevel)
|
||||
# Resolve the repository root from the script location instead of `.git`.
|
||||
SCRIPT_DIR="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd -P)"
|
||||
GIT_ROOT="${GIT_ROOT:-$(cd -- "${SCRIPT_DIR}/../../../.." && pwd -P)}"
|
||||
|
||||
SMI_BIN=$(which nvidia-smi || which rocm-smi || echo "")
|
||||
|
||||
|
||||
@@ -0,0 +1,284 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""
|
||||
Tests for resource management in hf3fs_client.py: constructor failure cleanup
|
||||
and idempotent close(). Tests use mock to replace real I/O operations
|
||||
(hf3fs_fuse.io, SharedMemory, os, CUDA).
|
||||
Requires hf3fs_fuse.io to be installed; skipped otherwise.
|
||||
"""
|
||||
|
||||
from typing import Any
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
HF3FS_AVAILABLE = True
|
||||
try:
|
||||
from hf3fs_fuse.io import ( # noqa: F401
|
||||
deregister_fd,
|
||||
extract_mount_point,
|
||||
make_ioring,
|
||||
make_iovec,
|
||||
register_fd,
|
||||
)
|
||||
|
||||
from vllm.distributed.kv_transfer.kv_connector.v1.hf3fs.hf3fs_client import (
|
||||
Hf3fsClient,
|
||||
)
|
||||
except Exception:
|
||||
HF3FS_AVAILABLE = False
|
||||
|
||||
requires_hf3fs = pytest.mark.skipif(
|
||||
not HF3FS_AVAILABLE,
|
||||
reason="hf3fs_fuse.io is not available on this machine",
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class _FakeShm:
|
||||
"""Shared-memory stub matching the multiprocessing.shared_memory.SharedMemory
|
||||
interface used by Hf3fsClient:
|
||||
|
||||
Attributes accessed by the constructor:
|
||||
.buf – memoryview / buffer-protocol object consumed by torch.frombuffer
|
||||
Methods called during normal lifetime:
|
||||
.unlink() – called right after the iovec is set up
|
||||
.close() – called in _release_resources()
|
||||
"""
|
||||
|
||||
def __init__(self, size: int = 1024):
|
||||
self._data = bytearray(size)
|
||||
self.buf = memoryview(self._data)
|
||||
self.closed = False
|
||||
self.close_call_count = 0
|
||||
self.unlink_call_count = 0
|
||||
|
||||
def close(self):
|
||||
self.closed = True
|
||||
self.close_call_count += 1
|
||||
|
||||
def unlink(self):
|
||||
self.unlink_call_count += 1
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# TestHf3fsClientResourceManagement
|
||||
# ===========================================================================
|
||||
|
||||
|
||||
@requires_hf3fs
|
||||
class TestHf3fsClientResourceManagement:
|
||||
"""Tests for constructor failure cleanup and idempotent close()."""
|
||||
|
||||
_MOD = "vllm.distributed.kv_transfer.kv_connector.v1.hf3fs.hf3fs_client"
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Helper: build a minimal Hf3fsClient bypassing all real I/O so that
|
||||
# we can fully control its internal state.
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def _make_client(self, tmp_path):
|
||||
"""Return a fully-mocked Hf3fsClient with controllable internals."""
|
||||
fake_shm_r = _FakeShm()
|
||||
fake_shm_w = _FakeShm()
|
||||
|
||||
patcher_list: list[Any] = [
|
||||
patch(f"{self._MOD}.HF3FS_AVAILABLE", True),
|
||||
patch(f"{self._MOD}.register_fd"),
|
||||
patch(f"{self._MOD}.deregister_fd"),
|
||||
patch(f"{self._MOD}.extract_mount_point", return_value="/mnt/hf3fs"),
|
||||
patch(f"{self._MOD}.make_ioring", return_value=MagicMock()),
|
||||
patch(f"{self._MOD}.make_iovec", return_value=MagicMock()),
|
||||
patch(
|
||||
"multiprocessing.shared_memory.SharedMemory",
|
||||
side_effect=[fake_shm_r, fake_shm_w],
|
||||
),
|
||||
patch("os.open", return_value=99),
|
||||
patch("os.ftruncate"),
|
||||
patch("os.close"),
|
||||
patch("os.fsync"),
|
||||
patch("torch.cuda.Stream", return_value=MagicMock()),
|
||||
patch("torch.frombuffer", return_value=MagicMock()),
|
||||
patch("torch.empty", return_value=MagicMock()),
|
||||
]
|
||||
for p in patcher_list:
|
||||
p.start()
|
||||
|
||||
try:
|
||||
client = Hf3fsClient(
|
||||
path=str(tmp_path / "test.bin"),
|
||||
size=1024,
|
||||
bytes_per_page=256,
|
||||
entries=4,
|
||||
)
|
||||
finally:
|
||||
for p in patcher_list:
|
||||
p.stop()
|
||||
|
||||
# Manually point internal handles to our controllable fakes so that
|
||||
# assertions after close() can inspect them directly.
|
||||
client.shm_r = fake_shm_r
|
||||
client.shm_w = fake_shm_w
|
||||
client.file = 99
|
||||
return client, fake_shm_r, fake_shm_w
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# close() idempotency
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def test_close_idempotent_and_handles_cleared(self, tmp_path):
|
||||
"""Multiple close() calls must not raise; deregister_fd called exactly
|
||||
once, all handles set to None, shm.close() invoked."""
|
||||
client, shm_r, shm_w = self._make_client(tmp_path)
|
||||
|
||||
with (
|
||||
patch(f"{self._MOD}.deregister_fd") as mock_dereg,
|
||||
patch("os.close"),
|
||||
):
|
||||
client.close() # first close
|
||||
client.close() # second close — must be no-op
|
||||
client.close() # third close — must be no-op
|
||||
|
||||
assert client._closed is True
|
||||
assert mock_dereg.call_count == 1, (
|
||||
f"deregister_fd called {mock_dereg.call_count} times; expected 1"
|
||||
)
|
||||
for attr in ("iov_r", "iov_w", "ior_r", "ior_w", "shm_r", "shm_w", "file"):
|
||||
assert getattr(client, attr) is None, f"{attr} should be None after close()"
|
||||
assert shm_r.closed is True
|
||||
assert shm_w.closed is True
|
||||
|
||||
def test_flush_after_close_is_noop(self, tmp_path):
|
||||
"""flush() after close() must silently do nothing (no fsync call)."""
|
||||
client, _, _ = self._make_client(tmp_path)
|
||||
|
||||
with (
|
||||
patch(f"{self._MOD}.deregister_fd"),
|
||||
patch("os.close"),
|
||||
patch("os.fsync") as mock_fsync,
|
||||
):
|
||||
client.close()
|
||||
client.flush()
|
||||
|
||||
mock_fsync.assert_not_called()
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Constructor failure leaves no leaked resources
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def test_constructor_failure_after_file_open_cleans_file(self, tmp_path):
|
||||
"""If the constructor raises after os.open(), the fd must be closed."""
|
||||
with (
|
||||
patch(f"{self._MOD}.HF3FS_AVAILABLE", True),
|
||||
patch(f"{self._MOD}.register_fd"),
|
||||
patch(f"{self._MOD}.deregister_fd"),
|
||||
patch(
|
||||
f"{self._MOD}.extract_mount_point",
|
||||
side_effect=RuntimeError("mount point not found"),
|
||||
),
|
||||
patch("os.open", return_value=55),
|
||||
patch("os.ftruncate"),
|
||||
patch("os.close") as mock_os_close,
|
||||
patch("torch.cuda.Stream", return_value=MagicMock()),
|
||||
pytest.raises(RuntimeError, match="mount point not found"),
|
||||
):
|
||||
Hf3fsClient(
|
||||
path=str(tmp_path / "fail.bin"),
|
||||
size=1024,
|
||||
bytes_per_page=256,
|
||||
entries=4,
|
||||
)
|
||||
|
||||
mock_os_close.assert_called_once_with(55)
|
||||
|
||||
def test_constructor_failure_after_shm_alloc_closes_shm(self, tmp_path):
|
||||
"""Constructor raises after SharedMemory creation → both shm objects closed."""
|
||||
fake_shm_r = _FakeShm()
|
||||
fake_shm_w = _FakeShm()
|
||||
|
||||
with (
|
||||
patch(f"{self._MOD}.HF3FS_AVAILABLE", True),
|
||||
patch(f"{self._MOD}.register_fd"),
|
||||
patch(f"{self._MOD}.deregister_fd"),
|
||||
patch(f"{self._MOD}.extract_mount_point", return_value="/mnt/hf3fs"),
|
||||
patch(
|
||||
"multiprocessing.shared_memory.SharedMemory",
|
||||
side_effect=[fake_shm_r, fake_shm_w],
|
||||
),
|
||||
patch("os.open", return_value=66),
|
||||
patch("os.ftruncate"),
|
||||
patch("os.close"),
|
||||
patch("torch.frombuffer", return_value=MagicMock()),
|
||||
patch("torch.empty", return_value=MagicMock()),
|
||||
patch(
|
||||
f"{self._MOD}.make_ioring",
|
||||
side_effect=RuntimeError("ioring init failed"),
|
||||
),
|
||||
patch(f"{self._MOD}.make_iovec", return_value=MagicMock()),
|
||||
patch("torch.cuda.Stream", return_value=MagicMock()),
|
||||
pytest.raises(RuntimeError, match="ioring init failed"),
|
||||
):
|
||||
Hf3fsClient(
|
||||
path=str(tmp_path / "fail2.bin"),
|
||||
size=1024,
|
||||
bytes_per_page=256,
|
||||
entries=4,
|
||||
)
|
||||
|
||||
assert fake_shm_r.closed is True, (
|
||||
"shm_r was not closed after constructor failure"
|
||||
)
|
||||
assert fake_shm_w.closed is True, (
|
||||
"shm_w was not closed after constructor failure"
|
||||
)
|
||||
|
||||
def test_constructor_failure_does_not_close_unallocated_shm(self, tmp_path):
|
||||
"""Failure before SharedMemory is created must not raise AttributeError
|
||||
or TypeError from cleanup."""
|
||||
with (
|
||||
patch(f"{self._MOD}.HF3FS_AVAILABLE", True),
|
||||
patch(f"{self._MOD}.register_fd"),
|
||||
patch(f"{self._MOD}.deregister_fd"),
|
||||
patch(
|
||||
f"{self._MOD}.extract_mount_point",
|
||||
side_effect=RuntimeError("early failure"),
|
||||
),
|
||||
patch("os.open", return_value=77),
|
||||
patch("os.ftruncate"),
|
||||
patch("os.close"),
|
||||
patch("torch.cuda.Stream", return_value=MagicMock()),
|
||||
pytest.raises(RuntimeError, match="early failure"),
|
||||
):
|
||||
Hf3fsClient(
|
||||
path=str(tmp_path / "early_fail.bin"),
|
||||
size=1024,
|
||||
bytes_per_page=256,
|
||||
entries=4,
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# _release_resources on already-cleared state must be a no-op
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def test_release_resources_on_empty_state_is_safe(self, tmp_path):
|
||||
"""_release_resources() on a fully-cleared client must not raise."""
|
||||
client, _, _ = self._make_client(tmp_path)
|
||||
|
||||
with (
|
||||
patch(f"{self._MOD}.deregister_fd"),
|
||||
patch("os.close"),
|
||||
):
|
||||
client.close() # clears all handles
|
||||
|
||||
with (
|
||||
patch(f"{self._MOD}.deregister_fd") as mock_dereg2,
|
||||
patch("os.close") as mock_os_close2,
|
||||
):
|
||||
client._release_resources() # must not raise
|
||||
|
||||
mock_dereg2.assert_not_called()
|
||||
mock_os_close2.assert_not_called()
|
||||
@@ -0,0 +1,230 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""
|
||||
Tests for HF3FS KV Connector high-level components:
|
||||
- TestHf3fsMockClient : file-backed mock client I/O correctness
|
||||
- TestHF3FSKVConnectorStats: metric collection, aggregation, serialisation
|
||||
"""
|
||||
|
||||
import os
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from vllm.distributed.kv_transfer.kv_connector.v1.hf3fs.hf3fs_connector import (
|
||||
HF3FSKVConnectorStats,
|
||||
)
|
||||
from vllm.distributed.kv_transfer.kv_connector.v1.hf3fs.utils.hf3fs_mock_client import (
|
||||
Hf3fsClient as MockHf3fsClient,
|
||||
)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def hf3fs_stats():
|
||||
"""Fresh HF3FSKVConnectorStats instance."""
|
||||
return HF3FSKVConnectorStats()
|
||||
|
||||
|
||||
def _make_cuda_event():
|
||||
"""Return a real CUDA event when available, otherwise a MagicMock."""
|
||||
if torch.cuda.is_available():
|
||||
return torch.cuda.Event()
|
||||
return MagicMock()
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# TestHf3fsMockClient
|
||||
# ===========================================================================
|
||||
|
||||
|
||||
class TestHf3fsMockClient:
|
||||
"""Tests for hf3fs_mock_client.Hf3fsClient (file-backend mock)."""
|
||||
|
||||
def test_init_creates_file(self, tmp_path):
|
||||
"""Initializing the client should create the backing file."""
|
||||
path = str(tmp_path / "test_file")
|
||||
client = MockHf3fsClient(path=path, size=4096, bytes_per_page=512, entries=4)
|
||||
assert os.path.exists(path), "Backing file should be created on init"
|
||||
assert os.path.getsize(path) == 4096
|
||||
client.close()
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"dtype, bytes_per_page",
|
||||
[
|
||||
(torch.float32, 512),
|
||||
(torch.float16, 256),
|
||||
(torch.bfloat16, 256),
|
||||
],
|
||||
ids=["float32", "float16", "bfloat16"],
|
||||
)
|
||||
def test_batch_write_and_read_dtype(self, tmp_path, dtype, bytes_per_page):
|
||||
"""Write a tensor of the given dtype and verify round-trip correctness."""
|
||||
path = str(tmp_path / f"rw_{dtype}")
|
||||
client = MockHf3fsClient(
|
||||
path=path, size=bytes_per_page * 8, bytes_per_page=bytes_per_page, entries=4
|
||||
)
|
||||
elem_size = torch.tensor([], dtype=dtype).element_size()
|
||||
numel = bytes_per_page // elem_size
|
||||
tensor_write = torch.arange(numel, dtype=dtype)
|
||||
event = _make_cuda_event()
|
||||
|
||||
results = client.batch_write([0], [tensor_write], event)
|
||||
assert results == [bytes_per_page], f"Write should succeed, got {results}"
|
||||
|
||||
tensor_read = torch.zeros(numel, dtype=dtype)
|
||||
results = client.batch_read([0], [tensor_read])
|
||||
assert results == [bytes_per_page], f"Read should succeed, got {results}"
|
||||
assert torch.equal(tensor_write, tensor_read), (
|
||||
"Read tensor should match written tensor"
|
||||
)
|
||||
client.close()
|
||||
|
||||
def test_batch_read_empty_file_returns_error(self, tmp_path):
|
||||
"""Reading out-of-bounds offset should return -1."""
|
||||
bytes_per_page = 128
|
||||
size = bytes_per_page * 4
|
||||
path = str(tmp_path / "empty_read")
|
||||
client = MockHf3fsClient(
|
||||
path=path, size=size, bytes_per_page=bytes_per_page, entries=4
|
||||
)
|
||||
numel = bytes_per_page // 4
|
||||
tensor_read = torch.zeros(numel, dtype=torch.float32)
|
||||
results = client.batch_read([size], [tensor_read]) # offset == size => OOB
|
||||
assert results[0] == -1, "Out-of-bounds read should return -1"
|
||||
client.close()
|
||||
|
||||
def test_batch_write_out_of_bounds_returns_error(self, tmp_path):
|
||||
"""Writing at an offset beyond file size should return -1."""
|
||||
bytes_per_page = 128
|
||||
size = bytes_per_page * 4
|
||||
path = str(tmp_path / "oob_write")
|
||||
client = MockHf3fsClient(
|
||||
path=path, size=size, bytes_per_page=bytes_per_page, entries=4
|
||||
)
|
||||
numel = bytes_per_page // 4
|
||||
tensor = torch.ones(numel, dtype=torch.float32)
|
||||
event = _make_cuda_event()
|
||||
results = client.batch_write([size], [tensor], event) # OOB offset
|
||||
assert results[0] == -1, "Out-of-bounds write should return -1"
|
||||
client.close()
|
||||
|
||||
def test_multiple_tensors_rw(self, tmp_path):
|
||||
"""Write multiple tensors at different offsets, then read all back."""
|
||||
bytes_per_page = 128
|
||||
n = 4
|
||||
path = str(tmp_path / "multi_rw")
|
||||
client = MockHf3fsClient(
|
||||
path=path,
|
||||
size=bytes_per_page * n * 2,
|
||||
bytes_per_page=bytes_per_page,
|
||||
entries=8,
|
||||
)
|
||||
tensors_write = [
|
||||
torch.full((bytes_per_page // 4,), float(i), dtype=torch.float32)
|
||||
for i in range(n)
|
||||
]
|
||||
offsets = [i * bytes_per_page for i in range(n)]
|
||||
event = _make_cuda_event()
|
||||
|
||||
results = client.batch_write(offsets, tensors_write, event)
|
||||
assert all(r == bytes_per_page for r in results)
|
||||
|
||||
tensors_read = [
|
||||
torch.zeros(bytes_per_page // 4, dtype=torch.float32) for _ in range(n)
|
||||
]
|
||||
results = client.batch_read(offsets, tensors_read)
|
||||
assert all(r == bytes_per_page for r in results)
|
||||
|
||||
for i, (tw, tr) in enumerate(zip(tensors_write, tensors_read)):
|
||||
assert torch.allclose(tw, tr), f"Tensor {i} mismatch after round-trip"
|
||||
client.close()
|
||||
|
||||
def test_flush_and_close_no_error(self, tmp_path):
|
||||
"""flush() and close() should not raise exceptions."""
|
||||
path = str(tmp_path / "flush_close")
|
||||
client = MockHf3fsClient(path=path, size=1024, bytes_per_page=128, entries=4)
|
||||
client.flush()
|
||||
client.close()
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# TestHF3FSKVConnectorStats
|
||||
# ===========================================================================
|
||||
|
||||
|
||||
class TestHF3FSKVConnectorStats:
|
||||
"""Tests for HF3FSKVConnectorStats metric collection and aggregation."""
|
||||
|
||||
def test_initial_is_empty(self, hf3fs_stats):
|
||||
"""Fresh stats object should report is_empty() == True."""
|
||||
assert hf3fs_stats.is_empty() is True
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"task_type, duration_key",
|
||||
[
|
||||
("Saved", "save_duration"),
|
||||
("Loaded", "load_duration"),
|
||||
],
|
||||
ids=["save", "load"],
|
||||
)
|
||||
def test_record_success_duration(self, hf3fs_stats, task_type, duration_key):
|
||||
"""Recording a successful task should update duration list and total count."""
|
||||
hf3fs_stats.record_success_task_duration(task_type, 0.5)
|
||||
assert not hf3fs_stats.is_empty()
|
||||
assert len(hf3fs_stats.data[duration_key]) == 1
|
||||
assert hf3fs_stats.data[duration_key][0] == pytest.approx(0.5)
|
||||
assert hf3fs_stats.data["num_transfer_task"] == 1
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"task_type, failed_key",
|
||||
[
|
||||
("Saved", "num_failed_save"),
|
||||
("Loaded", "num_failed_load"),
|
||||
],
|
||||
ids=["save", "load"],
|
||||
)
|
||||
def test_record_failed_task(self, hf3fs_stats, task_type, failed_key):
|
||||
"""Recording a failed task should increment the corresponding counter."""
|
||||
hf3fs_stats.record_failed_task_count(task_type)
|
||||
assert hf3fs_stats.data[failed_key] == 1
|
||||
assert hf3fs_stats.data["num_transfer_task"] == 1
|
||||
|
||||
def test_aggregate_two_stats(self):
|
||||
"""aggregate() should merge save/load duration lists and sum counters."""
|
||||
stats1 = HF3FSKVConnectorStats()
|
||||
stats1.record_success_task_duration("Saved", 0.1)
|
||||
stats1.record_success_task_duration("Loaded", 0.2)
|
||||
|
||||
stats2 = HF3FSKVConnectorStats()
|
||||
stats2.record_success_task_duration("Saved", 0.3)
|
||||
stats2.record_failed_task_count("Loaded")
|
||||
|
||||
stats1.aggregate(stats2)
|
||||
assert stats1.data["save_duration"] == pytest.approx([0.1, 0.3])
|
||||
assert stats1.data["load_duration"] == pytest.approx([0.2])
|
||||
assert stats1.data["num_failed_load"] == 1
|
||||
assert stats1.data["num_transfer_task"] == 4
|
||||
|
||||
def test_reduce_with_data(self):
|
||||
"""reduce() computes correct averages when data is present."""
|
||||
stats = HF3FSKVConnectorStats()
|
||||
stats.record_success_task_duration("Saved", 1.0)
|
||||
stats.record_success_task_duration("Saved", 3.0)
|
||||
result = stats.reduce()
|
||||
assert result["Num save task success"] == pytest.approx(2.0, rel=0.01)
|
||||
assert result["Num save task failed"] == pytest.approx(0.0, rel=0.01)
|
||||
assert result["Avg save duration (ms)"] == pytest.approx(2000.0, rel=0.01)
|
||||
|
||||
def test_clone_and_reset(self, hf3fs_stats):
|
||||
"""clone_and_reset() returns a copy with data and resets the original."""
|
||||
hf3fs_stats.record_success_task_duration("Saved", 0.7)
|
||||
hf3fs_stats.record_success_task_duration("Loaded", 0.4)
|
||||
|
||||
clone = hf3fs_stats.clone_and_reset()
|
||||
assert clone.data["num_transfer_task"] == 2
|
||||
assert hf3fs_stats.is_empty()
|
||||
@@ -0,0 +1,193 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""
|
||||
Tests for HF3FS metadata server data structures and allocation logic:
|
||||
- RankFileMetadata : page allocation / release primitives
|
||||
- KeyMetadata : per-key rank-page tracking and completion detection
|
||||
- GlobalMetadataState : coordinated allocation with cache-hit semantics
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
||||
from vllm.distributed.kv_transfer.kv_connector.v1.hf3fs.hf3fs_metadata_server import (
|
||||
GlobalMetadataState,
|
||||
KeyMetadata,
|
||||
RankFileMetadata,
|
||||
)
|
||||
|
||||
# ===========================================================================
|
||||
# TestRankFileMetadata
|
||||
# ===========================================================================
|
||||
|
||||
|
||||
class TestRankFileMetadata:
|
||||
"""Unit tests for RankFileMetadata page allocation primitives."""
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"alloc_count, expected_pages",
|
||||
[(3, 3), (5, 0)],
|
||||
ids=["alloc_partial", "alloc_exceeds"],
|
||||
)
|
||||
def test_allocate_pages(self, alloc_count, expected_pages):
|
||||
"""allocate_pages returns correct pages or empty list when insufficient."""
|
||||
rank_meta = RankFileMetadata(rank_id=0, num_pages=3, free_pages=list(range(3)))
|
||||
pages = rank_meta.allocate_pages(alloc_count)
|
||||
assert len(pages) == expected_pages
|
||||
if expected_pages > 0:
|
||||
rank_meta.release_pages(pages)
|
||||
assert rank_meta.get_free_page_count() == 3
|
||||
|
||||
def test_release_pages_restores_count(self):
|
||||
"""Releasing allocated pages returns them to the free pool."""
|
||||
rank_meta = RankFileMetadata(rank_id=0, num_pages=4, free_pages=list(range(4)))
|
||||
pages = rank_meta.allocate_pages(2)
|
||||
assert rank_meta.get_free_page_count() == 2
|
||||
rank_meta.release_pages(pages)
|
||||
assert rank_meta.get_free_page_count() == 4
|
||||
|
||||
def test_release_pages_no_duplicates(self):
|
||||
"""Releasing the same page twice must not create duplicates."""
|
||||
rank_meta = RankFileMetadata(rank_id=0, num_pages=3, free_pages=list(range(3)))
|
||||
rank_meta.allocate_pages(1) # takes page 0
|
||||
rank_meta.release_pages([0])
|
||||
rank_meta.release_pages([0]) # second release of the same page
|
||||
assert rank_meta.get_free_page_count() == 3
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# TestKeyMetadata
|
||||
# ===========================================================================
|
||||
|
||||
|
||||
class TestKeyMetadata:
|
||||
"""Unit tests for KeyMetadata completion tracking."""
|
||||
|
||||
def test_is_complete_false_until_all_ranks(self):
|
||||
"""is_complete() returns True only when all ranks confirmed."""
|
||||
key_meta = KeyMetadata(key="k", rank_to_page={}, tp_world_size=2)
|
||||
assert key_meta.is_complete() is False
|
||||
key_meta.add_rank_page(0, 5)
|
||||
assert key_meta.is_complete() is False
|
||||
key_meta.add_rank_page(1, 10)
|
||||
assert key_meta.is_complete() is True
|
||||
|
||||
def test_get_rank_page_returns_none_for_missing_rank(self):
|
||||
"""get_rank_page() returns None when the rank has no entry."""
|
||||
key_meta = KeyMetadata(key="k", rank_to_page={0: 3}, tp_world_size=2)
|
||||
assert key_meta.get_rank_page(0) == 3
|
||||
assert key_meta.get_rank_page(1) is None
|
||||
|
||||
def test_get_all_pages(self):
|
||||
"""get_all_pages() returns all (rank, page) pairs."""
|
||||
key_meta = KeyMetadata(key="k", rank_to_page={0: 1, 1: 2}, tp_world_size=2)
|
||||
pairs = key_meta.get_all_pages()
|
||||
assert set(pairs) == {(0, 1), (1, 2)}
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# TestGlobalMetadataStateAllocation
|
||||
# ===========================================================================
|
||||
|
||||
|
||||
class TestGlobalMetadataStateAllocation:
|
||||
"""Tests for GlobalMetadataState allocation and cache-hit semantics."""
|
||||
|
||||
def test_uninitialized_rank_raises_on_allocate(self):
|
||||
"""allocate_pages_for_keys raises ValueError for unknown rank."""
|
||||
state = GlobalMetadataState()
|
||||
with pytest.raises((ValueError, Exception)):
|
||||
state.allocate_pages_for_keys(99, [("key", "")])
|
||||
|
||||
def test_uninitialized_rank_raises_on_get_locations(self):
|
||||
"""get_key_locations raises ValueError for unknown rank."""
|
||||
state = GlobalMetadataState()
|
||||
with pytest.raises((ValueError, Exception)):
|
||||
state.get_key_locations(99, ["any_key"])
|
||||
|
||||
def test_basic_allocation_and_confirm(self):
|
||||
"""Allocating a page and confirming it marks the key as complete."""
|
||||
state = GlobalMetadataState()
|
||||
state.initialize_rank(0, 4)
|
||||
|
||||
results = state.allocate_pages_for_keys(0, [("K", "")])
|
||||
assert results["K"] >= 0
|
||||
|
||||
state.confirm_write_for_keys(0, [("K", results["K"])])
|
||||
assert state.batch_key_exists(["K"]) == [True]
|
||||
locations = state.get_key_locations(0, ["K"])
|
||||
assert locations == [results["K"]]
|
||||
|
||||
def test_allocate_pages_cache_hit_does_not_leak_pages(self):
|
||||
"""Cache-hit key must not consume a page from the free pool;
|
||||
the pre-allocated slot must be returned before reusing the existing page.
|
||||
"""
|
||||
state = GlobalMetadataState()
|
||||
state.initialize_rank(0, 5) # 5 free pages: [0,1,2,3,4]
|
||||
|
||||
# Simulate a key that has already been fully written and confirmed.
|
||||
state.key_metadata["K_cached"] = KeyMetadata(
|
||||
key="K_cached", rank_to_page={0: 2}, tp_world_size=1
|
||||
)
|
||||
|
||||
free_before = state.rank_metadata[0].get_free_page_count() # 5
|
||||
|
||||
results = state.allocate_pages_for_keys(0, [("K_cached", ""), ("K_new", "")])
|
||||
|
||||
free_after = state.rank_metadata[0].get_free_page_count()
|
||||
|
||||
# Cache-hit key must reuse its existing page.
|
||||
assert results["K_cached"] == 2, (
|
||||
f"Cache-hit key should reuse page 2, got {results['K_cached']}"
|
||||
)
|
||||
# New key must receive a valid page.
|
||||
assert results["K_new"] >= 0, (
|
||||
f"New key should get a valid page, got {results['K_new']}"
|
||||
)
|
||||
# Exactly one page consumed from the free pool.
|
||||
assert free_before - free_after == 1, (
|
||||
f"Expected 1 page consumed, got delta={free_before - free_after}"
|
||||
)
|
||||
|
||||
def test_allocate_pages_all_cache_hits_frees_all_slots(self):
|
||||
"""When every key in the batch is a cache hit, no pages are consumed."""
|
||||
state = GlobalMetadataState()
|
||||
state.initialize_rank(0, 5)
|
||||
|
||||
for key, page in (("K1", 0), ("K2", 1)):
|
||||
state.key_metadata[key] = KeyMetadata(
|
||||
key=key, rank_to_page={0: page}, tp_world_size=1
|
||||
)
|
||||
|
||||
free_before = state.rank_metadata[0].get_free_page_count()
|
||||
results = state.allocate_pages_for_keys(0, [("K1", ""), ("K2", "")])
|
||||
free_after = state.rank_metadata[0].get_free_page_count()
|
||||
|
||||
assert results["K1"] == 0
|
||||
assert results["K2"] == 1
|
||||
assert free_after == free_before, (
|
||||
f"All-cache-hit batch must not consume free pages; "
|
||||
f"before={free_before}, after={free_after}"
|
||||
)
|
||||
|
||||
def test_allocate_returns_minus_one_when_pool_exhausted(self):
|
||||
"""If the free pool is exhausted, all new keys receive -1."""
|
||||
state = GlobalMetadataState()
|
||||
state.initialize_rank(0, 1) # only 1 free page
|
||||
|
||||
results = state.allocate_pages_for_keys(0, [("K1", ""), ("K2", "")])
|
||||
# allocate_pages uses all-or-nothing: 2 needed but only 1 available → []
|
||||
assert all(v == -1 for v in results.values()), f"Expected all -1, got {results}"
|
||||
|
||||
def test_confirm_write_releases_pages(self):
|
||||
"""confirm_write_for_keys with pages_to_release returns them to pool."""
|
||||
state = GlobalMetadataState()
|
||||
state.initialize_rank(0, 3)
|
||||
|
||||
results = state.allocate_pages_for_keys(0, [("K", "")])
|
||||
page = results["K"]
|
||||
free_after_alloc = state.rank_metadata[0].get_free_page_count()
|
||||
|
||||
state.confirm_write_for_keys(0, [("K", page)], pages_to_release=[page])
|
||||
free_after_release = state.rank_metadata[0].get_free_page_count()
|
||||
|
||||
assert free_after_release == free_after_alloc + 1
|
||||
@@ -15,6 +15,7 @@ from vllm.v1.kv_offload.abstract import (
|
||||
from vllm.v1.kv_offload.cpu.manager import CPUOffloadingManager
|
||||
from vllm.v1.kv_offload.cpu.policies.arc import ARCCachePolicy
|
||||
from vllm.v1.kv_offload.mediums import CPULoadStoreSpec
|
||||
from vllm.v1.kv_offload.reuse_manager import FilterReusedOffloadingManager
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -243,335 +244,300 @@ def test_cpu_manager():
|
||||
)
|
||||
|
||||
|
||||
def test_arc_manager_basic():
|
||||
"""
|
||||
Tests CPUOffloadingManager with arc policy.
|
||||
Verifies that ARC handles store, load, and lookup operations correctly.
|
||||
"""
|
||||
block_size = 256
|
||||
arc_manager = CPUOffloadingManager(
|
||||
block_size=block_size, num_blocks=4, cache_policy="arc", enable_events=True
|
||||
)
|
||||
arc_policy = arc_manager._policy
|
||||
assert isinstance(arc_policy, ARCCachePolicy)
|
||||
class TestARCPolicy:
|
||||
"""Unit tests for CPUOffloadingManager with ARC eviction policy."""
|
||||
|
||||
# prepare store [1, 2]
|
||||
prepare_store_output = arc_manager.prepare_store(to_hashes([1, 2]))
|
||||
verify_store_output(
|
||||
prepare_store_output,
|
||||
ExpectedPrepareStoreOutput(
|
||||
block_hashes_to_store=[1, 2],
|
||||
store_block_ids=[0, 1],
|
||||
block_hashes_evicted=[],
|
||||
),
|
||||
)
|
||||
def _make_manager(
|
||||
self, num_blocks: int = 4, enable_events: bool = True
|
||||
) -> tuple[CPUOffloadingManager, ARCCachePolicy]:
|
||||
manager = CPUOffloadingManager(
|
||||
block_size=256,
|
||||
num_blocks=num_blocks,
|
||||
cache_policy="arc",
|
||||
enable_events=enable_events,
|
||||
)
|
||||
policy = manager._policy
|
||||
assert isinstance(policy, ARCCachePolicy)
|
||||
return manager, policy
|
||||
|
||||
# lookup [1, 2] -> not ready
|
||||
assert arc_manager.lookup(to_hashes([1, 2])) == 0
|
||||
def test_basic(self):
|
||||
"""
|
||||
Tests CPUOffloadingManager with arc policy.
|
||||
Verifies that ARC handles store, load, and lookup operations correctly.
|
||||
"""
|
||||
cpu_manager, arc_policy = self._make_manager()
|
||||
|
||||
# no events so far
|
||||
assert list(arc_manager.take_events()) == []
|
||||
# prepare store [1, 2]
|
||||
prepare_store_output = cpu_manager.prepare_store(to_hashes([1, 2]))
|
||||
verify_store_output(
|
||||
prepare_store_output,
|
||||
ExpectedPrepareStoreOutput(
|
||||
block_hashes_to_store=[1, 2],
|
||||
store_block_ids=[0, 1],
|
||||
block_hashes_evicted=[],
|
||||
),
|
||||
)
|
||||
|
||||
# complete store [1, 2]
|
||||
arc_manager.complete_store(to_hashes([1, 2]))
|
||||
verify_events(
|
||||
arc_manager.take_events(), block_size=block_size, expected_stores=({1, 2},)
|
||||
)
|
||||
# lookup [1, 2] -> not ready
|
||||
assert cpu_manager.lookup(to_hashes([1, 2])) == 0
|
||||
|
||||
# lookup [1, 2]
|
||||
assert arc_manager.lookup(to_hashes([1])) == 1
|
||||
assert arc_manager.lookup(to_hashes([1, 2])) == 2
|
||||
assert arc_manager.lookup(to_hashes([1, 2, 3])) == 2
|
||||
# no events so far
|
||||
assert list(cpu_manager.take_events()) == []
|
||||
|
||||
# blocks should be in T1 (recent)
|
||||
assert len(arc_policy.t1) == 2
|
||||
assert len(arc_policy.t2) == 0
|
||||
# complete store [1, 2]
|
||||
cpu_manager.complete_store(to_hashes([1, 2]))
|
||||
verify_events(
|
||||
cpu_manager.take_events(), block_size=256, expected_stores=({1, 2},)
|
||||
)
|
||||
|
||||
# lookup [1, 2]
|
||||
assert cpu_manager.lookup(to_hashes([1])) == 1
|
||||
assert cpu_manager.lookup(to_hashes([1, 2])) == 2
|
||||
assert cpu_manager.lookup(to_hashes([1, 2, 3])) == 2
|
||||
|
||||
def test_arc_manager_t1_to_t2_promotion():
|
||||
"""
|
||||
Tests that accessing a block in T1 promotes it to T2 (frequent).
|
||||
This is a key feature of ARC's adaptive behavior.
|
||||
"""
|
||||
block_size = 256
|
||||
arc_manager = CPUOffloadingManager(
|
||||
block_size=block_size, num_blocks=4, cache_policy="arc", enable_events=False
|
||||
)
|
||||
arc_policy = arc_manager._policy
|
||||
assert isinstance(arc_policy, ARCCachePolicy)
|
||||
# blocks should be in T1 (recent)
|
||||
assert len(arc_policy.t1) == 2
|
||||
assert len(arc_policy.t2) == 0
|
||||
|
||||
# store and complete block 1
|
||||
arc_manager.prepare_store(to_hashes([1]))
|
||||
arc_manager.complete_store(to_hashes([1]))
|
||||
def test_t1_to_t2_promotion(self):
|
||||
"""
|
||||
Tests that accessing a block in T1 promotes it to T2 (frequent).
|
||||
This is a key feature of ARC's adaptive behavior.
|
||||
"""
|
||||
cpu_manager, arc_policy = self._make_manager(enable_events=False)
|
||||
|
||||
# block 1 starts in T1 (recent)
|
||||
assert to_hashes([1])[0] in arc_policy.t1
|
||||
assert to_hashes([1])[0] not in arc_policy.t2
|
||||
# store and complete block 1
|
||||
cpu_manager.prepare_store(to_hashes([1]))
|
||||
cpu_manager.complete_store(to_hashes([1]))
|
||||
|
||||
# touch block 1 (simulate second access)
|
||||
arc_manager.touch(to_hashes([1]))
|
||||
# block 1 starts in T1 (recent)
|
||||
assert to_hashes([1])[0] in arc_policy.t1
|
||||
assert to_hashes([1])[0] not in arc_policy.t2
|
||||
|
||||
# block 1 should now be in T2 (frequent)
|
||||
assert to_hashes([1])[0] not in arc_policy.t1
|
||||
assert to_hashes([1])[0] in arc_policy.t2
|
||||
# touch block 1 (simulate second access)
|
||||
cpu_manager.touch(to_hashes([1]))
|
||||
|
||||
# block 1 should now be in T2 (frequent)
|
||||
assert to_hashes([1])[0] not in arc_policy.t1
|
||||
assert to_hashes([1])[0] in arc_policy.t2
|
||||
|
||||
def test_arc_manager_eviction_with_load():
|
||||
"""
|
||||
Tests ARC eviction behavior similar to LRU test.
|
||||
Verifies that blocks being loaded (ref_cnt > 0) cannot be evicted.
|
||||
"""
|
||||
block_size = 256
|
||||
arc_manager = CPUOffloadingManager(
|
||||
block_size=block_size, num_blocks=4, cache_policy="arc", enable_events=True
|
||||
)
|
||||
def test_eviction_with_load(self):
|
||||
"""
|
||||
Tests ARC eviction behavior similar to LRU test.
|
||||
Verifies that blocks being loaded (ref_cnt > 0) cannot be evicted.
|
||||
"""
|
||||
cpu_manager, _ = self._make_manager()
|
||||
|
||||
# prepare and complete store [1, 2, 3, 4]
|
||||
prepare_store_output = arc_manager.prepare_store(to_hashes([1, 2, 3, 4]))
|
||||
verify_store_output(
|
||||
prepare_store_output,
|
||||
ExpectedPrepareStoreOutput(
|
||||
block_hashes_to_store=[1, 2, 3, 4],
|
||||
store_block_ids=[0, 1, 2, 3],
|
||||
block_hashes_evicted=[],
|
||||
),
|
||||
)
|
||||
arc_manager.complete_store(to_hashes([1, 2, 3, 4]))
|
||||
# prepare and complete store [1, 2, 3, 4]
|
||||
prepare_store_output = cpu_manager.prepare_store(to_hashes([1, 2, 3, 4]))
|
||||
verify_store_output(
|
||||
prepare_store_output,
|
||||
ExpectedPrepareStoreOutput(
|
||||
block_hashes_to_store=[1, 2, 3, 4],
|
||||
store_block_ids=[0, 1, 2, 3],
|
||||
block_hashes_evicted=[],
|
||||
),
|
||||
)
|
||||
cpu_manager.complete_store(to_hashes([1, 2, 3, 4]))
|
||||
|
||||
# prepare load [2, 3] (increases ref_cnt)
|
||||
prepare_load_output = arc_manager.prepare_load(to_hashes([2, 3]))
|
||||
verify_load_output(prepare_load_output, [1, 2])
|
||||
# prepare load [2, 3] (increases ref_cnt)
|
||||
prepare_load_output = cpu_manager.prepare_load(to_hashes([2, 3]))
|
||||
verify_load_output(prepare_load_output, [1, 2])
|
||||
|
||||
# prepare store [5, 6, 7] with [2, 3] being loaded
|
||||
# should fail because [2, 3] have ref_cnt > 0
|
||||
assert arc_manager.prepare_store(to_hashes([5, 6, 7])) is None
|
||||
# prepare store [5, 6, 7] with [2, 3] being loaded
|
||||
# should fail because [2, 3] have ref_cnt > 0
|
||||
assert cpu_manager.prepare_store(to_hashes([5, 6, 7])) is None
|
||||
|
||||
# complete load [2, 3]
|
||||
arc_manager.complete_load(to_hashes([2, 3]))
|
||||
# complete load [2, 3]
|
||||
cpu_manager.complete_load(to_hashes([2, 3]))
|
||||
|
||||
# now prepare store [5, 6, 7] should succeed
|
||||
# ARC will evict blocks one at a time from T1 as needed
|
||||
prepare_store_output = arc_manager.prepare_store(to_hashes([5, 6, 7]))
|
||||
assert prepare_store_output is not None
|
||||
# Should successfully evict enough blocks to make room (at least 1)
|
||||
assert len(prepare_store_output.block_hashes_evicted) >= 1
|
||||
# now prepare store [5, 6, 7] should succeed
|
||||
# ARC will evict blocks one at a time from T1 as needed
|
||||
prepare_store_output = cpu_manager.prepare_store(to_hashes([5, 6, 7]))
|
||||
assert prepare_store_output is not None
|
||||
# Should successfully evict enough blocks to make room (at least 1)
|
||||
assert len(prepare_store_output.block_hashes_evicted) >= 1
|
||||
|
||||
def test_adaptive_target(self):
|
||||
"""
|
||||
Tests ARC's adaptive target adjustment via ghost lists.
|
||||
When a block in B1 (ghost list) is accessed, target_t1_size increases.
|
||||
When a block in B2 is accessed, target_t1_size decreases.
|
||||
"""
|
||||
cpu_manager, arc_policy = self._make_manager(num_blocks=2, enable_events=False)
|
||||
|
||||
def test_arc_manager_adaptive_target():
|
||||
"""
|
||||
Tests ARC's adaptive target adjustment via ghost lists.
|
||||
When a block in B1 (ghost list) is accessed, target_t1_size increases.
|
||||
When a block in B2 is accessed, target_t1_size decreases.
|
||||
"""
|
||||
block_size = 256
|
||||
arc_manager = CPUOffloadingManager(
|
||||
block_size=block_size, num_blocks=2, cache_policy="arc", enable_events=False
|
||||
)
|
||||
arc_policy = arc_manager._policy
|
||||
assert isinstance(arc_policy, ARCCachePolicy)
|
||||
# store blocks 1, 2 (fills cache)
|
||||
cpu_manager.prepare_store(to_hashes([1, 2]))
|
||||
cpu_manager.complete_store(to_hashes([1, 2]))
|
||||
|
||||
# store blocks 1, 2 (fills cache)
|
||||
arc_manager.prepare_store(to_hashes([1, 2]))
|
||||
arc_manager.complete_store(to_hashes([1, 2]))
|
||||
initial_target = arc_policy.target_t1_size
|
||||
|
||||
initial_target = arc_policy.target_t1_size
|
||||
# store block 3, evicting block 1 (moves to B1 ghost list)
|
||||
cpu_manager.prepare_store(to_hashes([3]))
|
||||
cpu_manager.complete_store(to_hashes([3]))
|
||||
|
||||
# store block 3, evicting block 1 (moves to B1 ghost list)
|
||||
arc_manager.prepare_store(to_hashes([3]))
|
||||
arc_manager.complete_store(to_hashes([3]))
|
||||
# block 1 should be in B1 (ghost list)
|
||||
assert to_hashes([1])[0] in arc_policy.b1
|
||||
|
||||
# block 1 should be in B1 (ghost list)
|
||||
assert to_hashes([1])[0] in arc_policy.b1
|
||||
# touch block 1 (cache miss, but in B1)
|
||||
# this should increase target_t1_size (favor recency)
|
||||
cpu_manager.touch(to_hashes([1]))
|
||||
|
||||
# touch block 1 (cache miss, but in B1)
|
||||
# this should increase target_t1_size (favor recency)
|
||||
arc_manager.touch(to_hashes([1]))
|
||||
# target should have increased
|
||||
assert arc_policy.target_t1_size > initial_target
|
||||
|
||||
# target should have increased
|
||||
assert arc_policy.target_t1_size > initial_target
|
||||
def test_t1_t2_eviction_policy(self):
|
||||
"""
|
||||
Tests that ARC evicts from T1 or T2 based on target_t1_size.
|
||||
If |T1| >= target_t1_size, evict from T1, otherwise from T2.
|
||||
"""
|
||||
cpu_manager, arc_policy = self._make_manager(enable_events=False)
|
||||
|
||||
# store blocks 1, 2, 3, 4
|
||||
cpu_manager.prepare_store(to_hashes([1, 2, 3, 4]))
|
||||
cpu_manager.complete_store(to_hashes([1, 2, 3, 4]))
|
||||
|
||||
def test_arc_manager_t1_t2_eviction_policy():
|
||||
"""
|
||||
Tests that ARC evicts from T1 or T2 based on target_t1_size.
|
||||
If |T1| >= target_t1_size, evict from T1, otherwise from T2.
|
||||
"""
|
||||
block_size = 256
|
||||
arc_manager = CPUOffloadingManager(
|
||||
block_size=block_size, num_blocks=4, cache_policy="arc", enable_events=False
|
||||
)
|
||||
arc_policy = arc_manager._policy
|
||||
assert isinstance(arc_policy, ARCCachePolicy)
|
||||
# promote blocks 3, 4 to T2 by touching them
|
||||
cpu_manager.touch(to_hashes([3, 4]))
|
||||
|
||||
# store blocks 1, 2, 3, 4
|
||||
arc_manager.prepare_store(to_hashes([1, 2, 3, 4]))
|
||||
arc_manager.complete_store(to_hashes([1, 2, 3, 4]))
|
||||
# now: T1 = {1, 2}, T2 = {3, 4}
|
||||
assert len(arc_policy.t1) == 2
|
||||
assert len(arc_policy.t2) == 2
|
||||
|
||||
# promote blocks 3, 4 to T2 by touching them
|
||||
arc_manager.touch(to_hashes([3, 4]))
|
||||
# set target_t1_size to prefer evicting from T1
|
||||
# (when |T1| >= target, evict from T1)
|
||||
arc_policy.target_t1_size = 1
|
||||
|
||||
# now: T1 = {1, 2}, T2 = {3, 4}
|
||||
assert len(arc_policy.t1) == 2
|
||||
assert len(arc_policy.t2) == 2
|
||||
# store block 5, should evict from T1 (block 1, LRU in T1)
|
||||
output = cpu_manager.prepare_store(to_hashes([5]))
|
||||
assert output is not None
|
||||
assert to_hashes([1]) == output.block_hashes_evicted
|
||||
|
||||
# set target_t1_size to prefer evicting from T1
|
||||
# (when |T1| >= target, evict from T1)
|
||||
arc_policy.target_t1_size = 1
|
||||
cpu_manager.complete_store(to_hashes([5]))
|
||||
|
||||
# store block 5, should evict from T1 (block 1, LRU in T1)
|
||||
output = arc_manager.prepare_store(to_hashes([5]))
|
||||
assert output is not None
|
||||
assert to_hashes([1]) == output.block_hashes_evicted
|
||||
# block 1 should be in B1 (ghost list)
|
||||
assert to_hashes([1])[0] in arc_policy.b1
|
||||
# block 5 should be in T1
|
||||
assert to_hashes([5])[0] in arc_policy.t1
|
||||
|
||||
arc_manager.complete_store(to_hashes([5]))
|
||||
def test_ghost_list_bounds(self):
|
||||
"""
|
||||
Tests that ghost lists (B1, B2) don't grow unbounded.
|
||||
They should be capped at cache_capacity.
|
||||
"""
|
||||
cpu_manager, arc_policy = self._make_manager(num_blocks=2, enable_events=False)
|
||||
|
||||
# block 1 should be in B1 (ghost list)
|
||||
assert to_hashes([1])[0] in arc_policy.b1
|
||||
# block 5 should be in T1
|
||||
assert to_hashes([5])[0] in arc_policy.t1
|
||||
# fill cache with blocks 1, 2
|
||||
cpu_manager.prepare_store(to_hashes([1, 2]))
|
||||
cpu_manager.complete_store(to_hashes([1, 2]))
|
||||
|
||||
# store many blocks to fill ghost lists
|
||||
for i in range(3, 20):
|
||||
cpu_manager.prepare_store(to_hashes([i]))
|
||||
cpu_manager.complete_store(to_hashes([i]))
|
||||
|
||||
def test_arc_manager_ghost_list_bounds():
|
||||
"""
|
||||
Tests that ghost lists (B1, B2) don't grow unbounded.
|
||||
They should be capped at cache_capacity.
|
||||
"""
|
||||
block_size = 256
|
||||
arc_manager = CPUOffloadingManager(
|
||||
block_size=block_size, num_blocks=2, cache_policy="arc", enable_events=False
|
||||
)
|
||||
arc_policy = arc_manager._policy
|
||||
assert isinstance(arc_policy, ARCCachePolicy)
|
||||
# ghost lists should not exceed cache_capacity
|
||||
assert len(arc_policy.b1) <= arc_policy.cache_capacity
|
||||
assert len(arc_policy.b2) <= arc_policy.cache_capacity
|
||||
|
||||
# fill cache with blocks 1, 2
|
||||
arc_manager.prepare_store(to_hashes([1, 2]))
|
||||
arc_manager.complete_store(to_hashes([1, 2]))
|
||||
def test_touch_ordering(self):
|
||||
"""
|
||||
Tests that touch() correctly updates access patterns.
|
||||
Similar to LRU test but verifies T1/T2 ordering.
|
||||
"""
|
||||
cpu_manager, arc_policy = self._make_manager()
|
||||
|
||||
# store many blocks to fill ghost lists
|
||||
for i in range(3, 20):
|
||||
arc_manager.prepare_store(to_hashes([i]))
|
||||
arc_manager.complete_store(to_hashes([i]))
|
||||
# store blocks 1, 2, 3, 4
|
||||
cpu_manager.prepare_store(to_hashes([1, 2, 3, 4]))
|
||||
cpu_manager.complete_store(to_hashes([1, 2, 3, 4]))
|
||||
|
||||
# ghost lists should not exceed cache_capacity
|
||||
assert len(arc_policy.b1) <= arc_policy.cache_capacity
|
||||
assert len(arc_policy.b2) <= arc_policy.cache_capacity
|
||||
# promote 3, 4 to T2
|
||||
cpu_manager.touch(to_hashes([3, 4]))
|
||||
|
||||
# T1 = {1, 2}, T2 = {3, 4}
|
||||
# touch [1, 3, 4] - should promote 1 to T2, and move 3,4 to end of T2
|
||||
cpu_manager.touch(to_hashes([1, 3, 4]))
|
||||
|
||||
def test_arc_manager_touch_ordering():
|
||||
"""
|
||||
Tests that touch() correctly updates access patterns.
|
||||
Similar to LRU test but verifies T1/T2 ordering.
|
||||
"""
|
||||
block_size = 256
|
||||
arc_manager = CPUOffloadingManager(
|
||||
block_size=block_size, num_blocks=4, cache_policy="arc", enable_events=True
|
||||
)
|
||||
arc_policy = arc_manager._policy
|
||||
assert isinstance(arc_policy, ARCCachePolicy)
|
||||
# T1 = {2}, T2 = {1, 3, 4} (in that order, with 4 most recent)
|
||||
assert len(arc_policy.t1) == 1
|
||||
assert len(arc_policy.t2) == 3
|
||||
|
||||
# store blocks 1, 2, 3, 4
|
||||
arc_manager.prepare_store(to_hashes([1, 2, 3, 4]))
|
||||
arc_manager.complete_store(to_hashes([1, 2, 3, 4]))
|
||||
# store block 5, should evict from T1 (block 2, only one in T1)
|
||||
prepare_store_output = cpu_manager.prepare_store(to_hashes([5]))
|
||||
verify_store_output(
|
||||
prepare_store_output,
|
||||
ExpectedPrepareStoreOutput(
|
||||
block_hashes_to_store=[5],
|
||||
store_block_ids=[1], # reuses block 2's storage
|
||||
block_hashes_evicted=[2],
|
||||
),
|
||||
)
|
||||
|
||||
# promote 3, 4 to T2
|
||||
arc_manager.touch(to_hashes([3, 4]))
|
||||
def test_failed_store(self):
|
||||
"""
|
||||
Tests that failed store operations clean up correctly.
|
||||
Similar to LRU test but for ARC.
|
||||
"""
|
||||
cpu_manager, arc_policy = self._make_manager()
|
||||
|
||||
# T1 = {1, 2}, T2 = {3, 4}
|
||||
# touch [1, 3, 4] - should promote 1 to T2, and move 3,4 to end of T2
|
||||
arc_manager.touch(to_hashes([1, 3, 4]))
|
||||
# store blocks 1, 2, 3, 4
|
||||
cpu_manager.prepare_store(to_hashes([1, 2, 3, 4]))
|
||||
cpu_manager.complete_store(to_hashes([1, 2, 3, 4]))
|
||||
|
||||
# T1 = {2}, T2 = {1, 3, 4} (in that order, with 4 most recent)
|
||||
assert len(arc_policy.t1) == 1
|
||||
assert len(arc_policy.t2) == 3
|
||||
# prepare store block 5 (will evict block 1)
|
||||
prepare_store_output = cpu_manager.prepare_store(to_hashes([5]))
|
||||
assert prepare_store_output is not None
|
||||
assert len(prepare_store_output.block_hashes_evicted) == 1
|
||||
|
||||
# store block 5, should evict from T1 (block 2, only one in T1)
|
||||
prepare_store_output = arc_manager.prepare_store(to_hashes([5]))
|
||||
verify_store_output(
|
||||
prepare_store_output,
|
||||
ExpectedPrepareStoreOutput(
|
||||
block_hashes_to_store=[5],
|
||||
store_block_ids=[1], # reuses block 2's storage
|
||||
block_hashes_evicted=[2],
|
||||
),
|
||||
)
|
||||
# complete store with failure
|
||||
cpu_manager.complete_store(to_hashes([5]), success=False)
|
||||
|
||||
# block 5 should not be in cache
|
||||
assert cpu_manager.lookup(to_hashes([5])) == 0
|
||||
# block 5 should not be in T1 or T2
|
||||
assert to_hashes([5])[0] not in arc_policy.t1
|
||||
assert to_hashes([5])[0] not in arc_policy.t2
|
||||
|
||||
def test_arc_manager_failed_store():
|
||||
"""
|
||||
Tests that failed store operations clean up correctly.
|
||||
Similar to LRU test but for ARC.
|
||||
"""
|
||||
block_size = 256
|
||||
arc_manager = CPUOffloadingManager(
|
||||
block_size=block_size, num_blocks=4, cache_policy="arc", enable_events=True
|
||||
)
|
||||
arc_policy = arc_manager._policy
|
||||
assert isinstance(arc_policy, ARCCachePolicy)
|
||||
# evicted block should still be gone (in B1 ghost list)
|
||||
evicted_hash = prepare_store_output.block_hashes_evicted[0]
|
||||
assert evicted_hash in arc_policy.b1
|
||||
|
||||
# store blocks 1, 2, 3, 4
|
||||
arc_manager.prepare_store(to_hashes([1, 2, 3, 4]))
|
||||
arc_manager.complete_store(to_hashes([1, 2, 3, 4]))
|
||||
def test_full_scenario(self):
|
||||
"""
|
||||
Comprehensive test covering multiple ARC operations in sequence.
|
||||
Similar to the full LRU test but adapted for ARC behavior.
|
||||
"""
|
||||
cpu_manager, arc_policy = self._make_manager()
|
||||
|
||||
# prepare store block 5 (will evict block 1)
|
||||
prepare_store_output = arc_manager.prepare_store(to_hashes([5]))
|
||||
assert prepare_store_output is not None
|
||||
assert len(prepare_store_output.block_hashes_evicted) == 1
|
||||
# store [1, 2]
|
||||
cpu_manager.prepare_store(to_hashes([1, 2]))
|
||||
cpu_manager.complete_store(to_hashes([1, 2]))
|
||||
|
||||
# complete store with failure
|
||||
arc_manager.complete_store(to_hashes([5]), success=False)
|
||||
# store [3, 4, 5] -> evicts [1]
|
||||
prepare_store_output = cpu_manager.prepare_store(to_hashes([3, 4, 5]))
|
||||
assert prepare_store_output is not None
|
||||
assert len(prepare_store_output.block_hashes_evicted) == 1
|
||||
cpu_manager.complete_store(to_hashes([3, 4, 5]))
|
||||
|
||||
# block 5 should not be in cache
|
||||
assert arc_manager.lookup(to_hashes([5])) == 0
|
||||
# block 5 should not be in T1 or T2
|
||||
assert to_hashes([5])[0] not in arc_policy.t1
|
||||
assert to_hashes([5])[0] not in arc_policy.t2
|
||||
# promote some blocks to T2
|
||||
cpu_manager.touch(to_hashes([2, 3]))
|
||||
|
||||
# evicted block should still be gone (in B1 ghost list)
|
||||
evicted_hash = prepare_store_output.block_hashes_evicted[0]
|
||||
assert evicted_hash in arc_policy.b1
|
||||
# T1 has {4, 5}, T2 has {2, 3}
|
||||
assert len(arc_policy.t1) == 2
|
||||
assert len(arc_policy.t2) == 2
|
||||
|
||||
# store [6] -> should evict from T1 (4 is oldest in T1)
|
||||
prepare_store_output = cpu_manager.prepare_store(to_hashes([6]))
|
||||
assert prepare_store_output is not None
|
||||
cpu_manager.complete_store(to_hashes([6]))
|
||||
|
||||
def test_arc_manager_full_scenario():
|
||||
"""
|
||||
Comprehensive test covering multiple ARC operations in sequence.
|
||||
Similar to the full LRU test but adapted for ARC behavior.
|
||||
"""
|
||||
block_size = 256
|
||||
arc_manager = CPUOffloadingManager(
|
||||
block_size=block_size, num_blocks=4, cache_policy="arc", enable_events=True
|
||||
)
|
||||
arc_policy = arc_manager._policy
|
||||
assert isinstance(arc_policy, ARCCachePolicy)
|
||||
# verify blocks 2, 3 (in T2) are still present
|
||||
assert cpu_manager.lookup(to_hashes([2])) == 1
|
||||
assert cpu_manager.lookup(to_hashes([3])) == 1
|
||||
|
||||
# store [1, 2]
|
||||
arc_manager.prepare_store(to_hashes([1, 2]))
|
||||
arc_manager.complete_store(to_hashes([1, 2]))
|
||||
|
||||
# store [3, 4, 5] -> evicts [1]
|
||||
prepare_store_output = arc_manager.prepare_store(to_hashes([3, 4, 5]))
|
||||
assert prepare_store_output is not None
|
||||
assert len(prepare_store_output.block_hashes_evicted) == 1
|
||||
arc_manager.complete_store(to_hashes([3, 4, 5]))
|
||||
|
||||
# promote some blocks to T2
|
||||
arc_manager.touch(to_hashes([2, 3]))
|
||||
|
||||
# T1 has {4, 5}, T2 has {2, 3}
|
||||
assert len(arc_policy.t1) == 2
|
||||
assert len(arc_policy.t2) == 2
|
||||
|
||||
# store [6] -> should evict from T1 (4 is oldest in T1)
|
||||
prepare_store_output = arc_manager.prepare_store(to_hashes([6]))
|
||||
assert prepare_store_output is not None
|
||||
arc_manager.complete_store(to_hashes([6]))
|
||||
|
||||
# verify blocks 2, 3 (in T2) are still present
|
||||
assert arc_manager.lookup(to_hashes([2])) == 1
|
||||
assert arc_manager.lookup(to_hashes([3])) == 1
|
||||
|
||||
# verify events
|
||||
events = list(arc_manager.take_events())
|
||||
assert len(events) > 0 # should have store and eviction events
|
||||
# verify events
|
||||
events = list(cpu_manager.take_events())
|
||||
assert len(events) > 0 # should have store and eviction events
|
||||
|
||||
|
||||
def test_filter_reused_manager():
|
||||
@@ -583,8 +549,6 @@ def test_filter_reused_manager():
|
||||
block_size=block_size, num_blocks=4, cache_policy="lru", enable_events=True
|
||||
)
|
||||
|
||||
from vllm.v1.kv_offload.reuse_manager import FilterReusedOffloadingManager
|
||||
|
||||
manager = FilterReusedOffloadingManager(
|
||||
backing=lru_manager, store_threshold=2, max_tracker_size=3
|
||||
)
|
||||
|
||||
@@ -12,6 +12,9 @@ from torch._higher_order_ops.auto_functionalize import auto_functionalized
|
||||
from torch._inductor.pattern_matcher import PatternMatcherPass
|
||||
|
||||
import vllm.ir.ops
|
||||
from vllm.compilation.passes.fusion.rms_quant_fusion import (
|
||||
_rms_input_weight_dtype_match,
|
||||
)
|
||||
from vllm.config import VllmConfig
|
||||
from vllm.config.utils import Range
|
||||
from vllm.distributed import get_tp_group, tensor_model_parallel_all_reduce
|
||||
@@ -320,7 +323,12 @@ class AllReduceRMSNormPattern(BasePattern):
|
||||
return allreduce[3], allreduce[1]
|
||||
|
||||
pm.register_replacement(
|
||||
pattern, replacement, self.get_inputs(), pm.fwd_only, pm_pass
|
||||
pattern,
|
||||
replacement,
|
||||
self.get_inputs(),
|
||||
pm.fwd_only,
|
||||
pm_pass,
|
||||
extra_check=_rms_input_weight_dtype_match,
|
||||
)
|
||||
|
||||
|
||||
@@ -459,7 +467,12 @@ class AllReduceFusedRMSNormStaticQuantFP8Pattern(BasePattern):
|
||||
return allreduce[4], allreduce[1]
|
||||
|
||||
pm.register_replacement(
|
||||
pattern, replacement, self.get_inputs(), pm.fwd_only, pm_pass
|
||||
pattern,
|
||||
replacement,
|
||||
self.get_inputs(),
|
||||
pm.fwd_only,
|
||||
pm_pass,
|
||||
extra_check=_rms_input_weight_dtype_match,
|
||||
)
|
||||
|
||||
|
||||
@@ -621,7 +634,12 @@ class AllReduceFusedRMSNormStaticQuantNVFP4Pattern(BasePattern):
|
||||
return allreduce[4], allreduce[1], allreduce[5]
|
||||
|
||||
pm.register_replacement(
|
||||
pattern, replacement, self.get_inputs(), pm.fwd_only, pm_pass
|
||||
pattern,
|
||||
replacement,
|
||||
self.get_inputs(),
|
||||
pm.fwd_only,
|
||||
pm_pass,
|
||||
extra_check=_rms_input_weight_dtype_match,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -38,6 +38,22 @@ FP8_DTYPE = current_platform.fp8_dtype()
|
||||
FP4_DTYPE = torch.uint8
|
||||
|
||||
|
||||
_RMS_NORM_OP = torch.ops.vllm_ir.rms_norm.default
|
||||
|
||||
|
||||
# TODO: extend rmsnorm quant kernels to support mixed input/weight dtypes,
|
||||
# and remove this check.
|
||||
def _rms_input_weight_dtype_match(match: pm.Match) -> bool:
|
||||
"""Prevent fusion when rms_norm input and weight dtypes differ."""
|
||||
for node in match.nodes:
|
||||
if node.target == _RMS_NORM_OP:
|
||||
# rms_norm(x, weight, epsilon, variance_size)
|
||||
x, weight = node.args[0], node.args[1]
|
||||
if isinstance(x, fx.Node) and isinstance(weight, fx.Node):
|
||||
return x.meta["val"].dtype == weight.meta["val"].dtype
|
||||
return True
|
||||
|
||||
|
||||
def empty_bf16(*args: Any, **kwargs: Any) -> torch.Tensor:
|
||||
return torch.empty(*args, **kwargs, dtype=torch.bfloat16, device="cuda")
|
||||
|
||||
@@ -186,7 +202,14 @@ class RMSNormStaticQuantPattern(RMSNormQuantPattern):
|
||||
]
|
||||
pattern(*inputs)
|
||||
|
||||
pm.register_replacement(pattern, replacement, inputs, pm.fwd_only, pm_pass)
|
||||
pm.register_replacement(
|
||||
pattern,
|
||||
replacement,
|
||||
inputs,
|
||||
pm.fwd_only,
|
||||
pm_pass,
|
||||
extra_check=_rms_input_weight_dtype_match,
|
||||
)
|
||||
|
||||
|
||||
class FusedAddRMSNormStaticQuantPattern(RMSNormQuantPattern):
|
||||
@@ -249,6 +272,7 @@ class FusedAddRMSNormStaticQuantPattern(RMSNormQuantPattern):
|
||||
inputs,
|
||||
pm.fwd_only,
|
||||
pm_pass,
|
||||
extra_check=_rms_input_weight_dtype_match,
|
||||
)
|
||||
|
||||
|
||||
@@ -350,6 +374,7 @@ class FusedAddRMSNormGroupQuantPattern(RMSNormQuantPattern):
|
||||
self.rmsnorm_matcher.inputs() + [scale],
|
||||
pm.fwd_only,
|
||||
pm_pass,
|
||||
extra_check=_rms_input_weight_dtype_match,
|
||||
)
|
||||
|
||||
|
||||
@@ -445,6 +470,7 @@ class RMSNormGroupQuantPattern(RMSNormQuantPattern):
|
||||
],
|
||||
pm.fwd_only,
|
||||
pm_pass,
|
||||
extra_check=_rms_input_weight_dtype_match,
|
||||
)
|
||||
|
||||
|
||||
@@ -503,6 +529,7 @@ class RMSNormDynamicQuantPattern(RMSNormQuantPattern):
|
||||
],
|
||||
pm.fwd_only,
|
||||
pm_pass,
|
||||
extra_check=_rms_input_weight_dtype_match,
|
||||
)
|
||||
|
||||
|
||||
@@ -559,6 +586,7 @@ class FusedAddRMSNormDynamicQuantPattern(RMSNormQuantPattern):
|
||||
self.rmsnorm_matcher.inputs(),
|
||||
pm.fwd_only,
|
||||
pm_pass,
|
||||
extra_check=_rms_input_weight_dtype_match,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -30,7 +30,7 @@ class AttentionConfig:
|
||||
use_cudnn_prefill: bool = False
|
||||
"""Whether to use cudnn prefill."""
|
||||
|
||||
use_trtllm_ragged_deepseek_prefill: bool = True
|
||||
use_trtllm_ragged_deepseek_prefill: bool = False
|
||||
"""Whether to use TRTLLM ragged deepseek prefill."""
|
||||
|
||||
use_trtllm_attention: bool | None = None
|
||||
|
||||
@@ -107,12 +107,11 @@ class KVConnectorFactory:
|
||||
if connector_name is None:
|
||||
raise ValueError("Connector name is not set in KVTransferConfig")
|
||||
compat_sig = False
|
||||
if connector_name in cls._registry:
|
||||
connector_cls = cls._registry[connector_name]()
|
||||
else:
|
||||
connector_module_path = kv_transfer_config.kv_connector_module_path
|
||||
if connector_module_path is None:
|
||||
raise ValueError(f"Unsupported connector type: {connector_name}")
|
||||
connector_module_path = kv_transfer_config.kv_connector_module_path
|
||||
if connector_module_path is not None and not connector_module_path:
|
||||
raise ValueError("kv_connector_module_path cannot be an empty string.")
|
||||
if connector_module_path:
|
||||
# External module path takes priority over internal registry.
|
||||
connector_module = importlib.import_module(connector_module_path)
|
||||
try:
|
||||
connector_cls = getattr(connector_module, connector_name)
|
||||
@@ -128,6 +127,10 @@ class KVConnectorFactory:
|
||||
"Please update to include kv_cache_config as the second argument.",
|
||||
connector_cls.__name__,
|
||||
)
|
||||
elif connector_name in cls._registry:
|
||||
connector_cls = cls._registry[connector_name]()
|
||||
else:
|
||||
raise ValueError(f"Unsupported connector type: {connector_name}")
|
||||
return connector_cls, compat_sig
|
||||
|
||||
@classmethod
|
||||
@@ -208,15 +211,18 @@ KVConnectorFactory.register_connector(
|
||||
"vllm.distributed.kv_transfer.kv_connector.v1.mooncake.mooncake_connector",
|
||||
"MooncakeConnector",
|
||||
)
|
||||
|
||||
KVConnectorFactory.register_connector(
|
||||
"FlexKVConnectorV1",
|
||||
"vllm.distributed.kv_transfer.kv_connector.v1.flexkv_connector",
|
||||
"FlexKVConnectorV1",
|
||||
)
|
||||
|
||||
KVConnectorFactory.register_connector(
|
||||
"SimpleCPUOffloadConnector",
|
||||
"vllm.distributed.kv_transfer.kv_connector.v1.simple_cpu_offload_connector",
|
||||
"SimpleCPUOffloadConnector",
|
||||
)
|
||||
KVConnectorFactory.register_connector(
|
||||
"HF3FSKVConnector",
|
||||
"vllm.distributed.kv_transfer.kv_connector.v1.hf3fs.hf3fs_connector",
|
||||
"HF3FSKVConnector",
|
||||
)
|
||||
|
||||
@@ -0,0 +1,298 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
import logging
|
||||
import multiprocessing
|
||||
import os
|
||||
import threading
|
||||
from functools import wraps
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
import torch.utils.cpp_extension
|
||||
from torch.utils.cpp_extension import load
|
||||
|
||||
root = Path(__file__).parent.resolve()
|
||||
cuda_include_path = os.path.join(torch.utils.cpp_extension.CUDA_HOME, "include")
|
||||
hf3fs_utils = load(
|
||||
name="hf3fs_utils",
|
||||
sources=[f"{root}/utils/hf3fs_utils.cpp"],
|
||||
extra_include_paths=[cuda_include_path],
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
HF3FS_AVAILABLE = True
|
||||
try:
|
||||
from hf3fs_fuse.io import (
|
||||
deregister_fd,
|
||||
extract_mount_point,
|
||||
make_ioring,
|
||||
make_iovec,
|
||||
register_fd,
|
||||
)
|
||||
except ImportError:
|
||||
HF3FS_AVAILABLE = False
|
||||
|
||||
|
||||
def rsynchronized():
|
||||
def _decorator(func):
|
||||
@wraps(func)
|
||||
def wrapper(self, *args, **kwargs):
|
||||
with self.rlock:
|
||||
return func(self, *args, **kwargs)
|
||||
|
||||
return wrapper
|
||||
|
||||
return _decorator
|
||||
|
||||
|
||||
def wsynchronized():
|
||||
def _decorator(func):
|
||||
@wraps(func)
|
||||
def wrapper(self, *args, **kwargs):
|
||||
with self.wlock:
|
||||
return func(self, *args, **kwargs)
|
||||
|
||||
return wrapper
|
||||
|
||||
return _decorator
|
||||
|
||||
|
||||
class Hf3fsClient:
|
||||
def __init__(self, path: str, size: int, bytes_per_page: int, entries: int):
|
||||
"""Initialize the HF3FS client with hf3fs_fuse.
|
||||
|
||||
Args:
|
||||
path: Path to the file used for storage
|
||||
size: Total size of the storage file in bytes
|
||||
bytes_per_page: Size of each page in bytes
|
||||
entries: Maximum number of concurrent operations
|
||||
"""
|
||||
if not HF3FS_AVAILABLE:
|
||||
raise ImportError(
|
||||
"hf3fs_fuse.io is not available. Please install the hf3fs_fuse package."
|
||||
)
|
||||
|
||||
self.path = path
|
||||
self.size = size
|
||||
self.bytes_per_page = bytes_per_page
|
||||
self.entries = entries
|
||||
|
||||
self._closed = False
|
||||
|
||||
self.file = None
|
||||
self.shm_r = None
|
||||
self.shm_w = None
|
||||
self.ior_r = None
|
||||
self.ior_w = None
|
||||
self.iov_r = None
|
||||
self.iov_w = None
|
||||
try:
|
||||
# Create the file if it doesn't exist and set its size
|
||||
self.file = os.open(self.path, os.O_RDWR | os.O_CREAT)
|
||||
os.ftruncate(self.file, size)
|
||||
register_fd(self.file)
|
||||
|
||||
self.hf3fs_mount_point = extract_mount_point(path)
|
||||
self.bs = self.bytes_per_page
|
||||
self.shm_r = multiprocessing.shared_memory.SharedMemory(
|
||||
size=self.bs * self.entries, create=True
|
||||
)
|
||||
self.shm_w = multiprocessing.shared_memory.SharedMemory(
|
||||
size=self.bs * self.entries, create=True
|
||||
)
|
||||
|
||||
self.shm_r_tensor = torch.frombuffer(self.shm_r.buf, dtype=torch.uint8)
|
||||
self.shm_w_tensor = torch.frombuffer(self.shm_w.buf, dtype=torch.uint8)
|
||||
|
||||
numel = self.bs * self.entries
|
||||
self.r_pinned = torch.empty(
|
||||
numel,
|
||||
dtype=torch.uint8,
|
||||
device="cpu",
|
||||
pin_memory=True,
|
||||
)
|
||||
self.w_pinned = torch.empty(
|
||||
numel,
|
||||
dtype=torch.uint8,
|
||||
device="cpu",
|
||||
pin_memory=True,
|
||||
)
|
||||
|
||||
self.numa = -1
|
||||
self.ior_r = make_ioring(
|
||||
self.hf3fs_mount_point,
|
||||
self.entries,
|
||||
for_read=True,
|
||||
timeout=1,
|
||||
numa=self.numa,
|
||||
)
|
||||
self.ior_w = make_ioring(
|
||||
self.hf3fs_mount_point,
|
||||
self.entries,
|
||||
for_read=False,
|
||||
timeout=1,
|
||||
numa=self.numa,
|
||||
)
|
||||
self.iov_r = make_iovec(self.shm_r, self.hf3fs_mount_point)
|
||||
self.iov_w = make_iovec(self.shm_w, self.hf3fs_mount_point)
|
||||
self.shm_r.unlink()
|
||||
self.shm_w.unlink()
|
||||
|
||||
self.rlock = threading.RLock()
|
||||
self.wlock = threading.RLock()
|
||||
|
||||
self.stream = torch.cuda.Stream()
|
||||
self.stream_ptr_int = self.stream.cuda_stream
|
||||
|
||||
except Exception:
|
||||
self._release_resources()
|
||||
raise
|
||||
|
||||
logger.debug(
|
||||
"Initialized HF3FS client with file: %s, size: %s bytes", path, size
|
||||
)
|
||||
|
||||
def _release_resources(self) -> None:
|
||||
"""Release all acquired resources safely"""
|
||||
# iov must be released before ioring and shm
|
||||
for attr in ("iov_r", "iov_w", "ior_r", "ior_w"):
|
||||
obj = getattr(self, attr, None)
|
||||
if obj is not None:
|
||||
del obj
|
||||
setattr(self, attr, None)
|
||||
|
||||
for attr in ("shm_r", "shm_w"):
|
||||
shm = getattr(self, attr, None)
|
||||
if shm is not None:
|
||||
try:
|
||||
shm.close()
|
||||
except Exception as e:
|
||||
logger.warning("Failed to close %s: %s", attr, e)
|
||||
setattr(self, attr, None)
|
||||
|
||||
if self.file is not None:
|
||||
try:
|
||||
deregister_fd(self.file)
|
||||
except Exception as e:
|
||||
logger.warning("deregister_fd failed: %s", e)
|
||||
try:
|
||||
os.close(self.file)
|
||||
except OSError as e:
|
||||
logger.warning("os.close failed: %s", e)
|
||||
self.file = None
|
||||
|
||||
@rsynchronized()
|
||||
def batch_read(self, offsets: list[int], tensors: list[torch.Tensor]) -> list[int]:
|
||||
"""Read data from the file at specified offsets into tensors.
|
||||
|
||||
Args:
|
||||
offsets: List of byte offsets to read from
|
||||
tensors: List of tensors to read data into
|
||||
|
||||
Returns:
|
||||
List of operation results (0 for success, non-zero for error)
|
||||
"""
|
||||
self.check(offsets, tensors)
|
||||
assert self.ior_r is not None
|
||||
assert self.iov_r is not None
|
||||
|
||||
# prepare
|
||||
current = 0
|
||||
for offset, tensor in zip(offsets, tensors):
|
||||
size = tensor.numel() * tensor.itemsize
|
||||
self.ior_r.prepare(
|
||||
self.iov_r[current : current + size], True, self.file, offset
|
||||
)
|
||||
current += size
|
||||
|
||||
# submit
|
||||
ionum = len(offsets)
|
||||
resv = self.ior_r.submit().wait(min_results=ionum)
|
||||
|
||||
# results
|
||||
with torch.cuda.stream(self.stream):
|
||||
hf3fs_utils.read_shm(
|
||||
self.shm_r_tensor, self.r_pinned, tensors, self.stream_ptr_int
|
||||
)
|
||||
results = [res.result for res in resv]
|
||||
|
||||
return results
|
||||
|
||||
@wsynchronized()
|
||||
def batch_write(
|
||||
self, offsets: list[int], tensors: list[torch.Tensor], event: torch.cuda.Event
|
||||
) -> list[int]:
|
||||
"""Write data from tensors to the file at specified offsets.
|
||||
|
||||
Args:
|
||||
offsets: List of byte offsets to write to
|
||||
tensors: List of tensors containing data to write
|
||||
|
||||
Returns:
|
||||
List of operation results (0 for success, non-zero for error)
|
||||
"""
|
||||
|
||||
self.check(offsets, tensors)
|
||||
assert self.ior_w is not None
|
||||
assert self.iov_w is not None
|
||||
|
||||
# prepare
|
||||
with torch.cuda.stream(self.stream):
|
||||
self.stream.wait_event(event)
|
||||
hf3fs_utils.write_shm(
|
||||
tensors, self.shm_w_tensor, self.w_pinned, self.stream_ptr_int
|
||||
)
|
||||
|
||||
current = 0
|
||||
for offset, tensor in zip(offsets, tensors):
|
||||
size = tensor.numel() * tensor.itemsize
|
||||
self.ior_w.prepare(
|
||||
self.iov_w[current : current + size], False, self.file, offset
|
||||
)
|
||||
current += size
|
||||
|
||||
# submit
|
||||
ionum = len(offsets)
|
||||
resv = self.ior_w.submit().wait(min_results=ionum)
|
||||
|
||||
# results
|
||||
results = [res.result for res in resv]
|
||||
|
||||
return results
|
||||
|
||||
def check(self, offsets: list[int], tensors: list[torch.Tensor]) -> None:
|
||||
sizes = [t.numel() * t.itemsize for t in tensors]
|
||||
if any(
|
||||
[
|
||||
len(offsets) > self.entries,
|
||||
len(offsets) != len(sizes),
|
||||
any(
|
||||
offset < 0 or offset + size > self.size
|
||||
for offset, size in zip(offsets, sizes)
|
||||
),
|
||||
any(size > self.bytes_per_page for size in sizes),
|
||||
]
|
||||
):
|
||||
self.close()
|
||||
raise ValueError("Hf3fsClient.check Failed")
|
||||
|
||||
def get_size(self) -> int:
|
||||
"""Get the total size of the storage file.
|
||||
|
||||
Returns:
|
||||
Size of the file in bytes
|
||||
"""
|
||||
return self.size
|
||||
|
||||
def close(self) -> None:
|
||||
"""Close the client and clean up resources."""
|
||||
if self._closed:
|
||||
return
|
||||
self._closed = True
|
||||
self._release_resources()
|
||||
|
||||
def flush(self) -> None:
|
||||
"""Flush any pending writes to disk."""
|
||||
if not self._closed and self.file is not None:
|
||||
os.fsync(self.file)
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,530 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""
|
||||
HF3FS Metadata Server with key-based organization.
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import logging
|
||||
import threading
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass
|
||||
|
||||
try:
|
||||
import orjson
|
||||
|
||||
HAS_ORJSON = True
|
||||
except ImportError:
|
||||
import json as orjson # type: ignore
|
||||
|
||||
HAS_ORJSON = False
|
||||
|
||||
import requests
|
||||
from fastapi import FastAPI, HTTPException, Request, Response
|
||||
from fastapi.responses import ORJSONResponse
|
||||
from requests.adapters import HTTPAdapter
|
||||
from urllib3.util.retry import Retry
|
||||
|
||||
logging.basicConfig(
|
||||
level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s"
|
||||
)
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class RankFileMetadata:
|
||||
"""Manages file page allocation for a single rank."""
|
||||
|
||||
rank_id: int
|
||||
num_pages: int
|
||||
free_pages: list[int]
|
||||
|
||||
def allocate_pages(self, num_pages: int) -> list[int]:
|
||||
"""Allocate specified number of free pages."""
|
||||
if len(self.free_pages) < num_pages:
|
||||
return []
|
||||
|
||||
allocated = self.free_pages[:num_pages]
|
||||
self.free_pages = self.free_pages[num_pages:]
|
||||
return allocated
|
||||
|
||||
def release_pages(self, page_indices: list[int]) -> None:
|
||||
"""Release pages back to free pool."""
|
||||
for page_idx in page_indices:
|
||||
if page_idx not in self.free_pages:
|
||||
self.free_pages.append(page_idx)
|
||||
|
||||
def get_free_page_count(self) -> int:
|
||||
"""Get current number of free pages."""
|
||||
return len(self.free_pages)
|
||||
|
||||
|
||||
@dataclass
|
||||
class KeyMetadata:
|
||||
"""Manages metadata for a single key across multiple ranks."""
|
||||
|
||||
key: str
|
||||
rank_to_page: dict[int, int] # rank -> allocated page index
|
||||
tp_world_size: int
|
||||
|
||||
def add_rank_page(self, rank: int, page_index: int) -> None:
|
||||
"""Add page allocation for a specific rank."""
|
||||
self.rank_to_page[rank] = page_index
|
||||
|
||||
def get_all_pages(self) -> list[tuple[int, int]]:
|
||||
"""Get all (rank, page) pairs for this key."""
|
||||
return [(rank, page) for rank, page in self.rank_to_page.items()]
|
||||
|
||||
def get_rank_page(self, rank: int) -> int | None:
|
||||
"""Get page index for a specific rank."""
|
||||
return self.rank_to_page.get(rank)
|
||||
|
||||
def is_complete(self) -> bool:
|
||||
"""Check if all ranks in the TP world have allocated pages."""
|
||||
return len(self.rank_to_page) == self.tp_world_size
|
||||
|
||||
|
||||
class GlobalMetadataState:
|
||||
"""Manages global metadata state across all ranks and keys."""
|
||||
|
||||
def __init__(self):
|
||||
self.global_lock = threading.RLock()
|
||||
self.rank_metadata: dict[int, RankFileMetadata] = {}
|
||||
self.key_metadata: dict[str, KeyMetadata] = {}
|
||||
|
||||
def clear(self) -> None:
|
||||
"""Clear all metadata state."""
|
||||
with self.global_lock:
|
||||
self.rank_metadata.clear()
|
||||
self.key_metadata.clear()
|
||||
logger.info("Cleared all metadata state")
|
||||
|
||||
def initialize_rank(self, rank: int, num_pages: int) -> None:
|
||||
"""Initialize a new rank with specified number of pages."""
|
||||
with self.global_lock:
|
||||
if rank not in self.rank_metadata:
|
||||
self.rank_metadata[rank] = RankFileMetadata(
|
||||
rank, num_pages, list(range(num_pages))
|
||||
)
|
||||
logger.info("Initialized rank %s with %s pages", rank, num_pages)
|
||||
|
||||
def allocate_pages_for_keys(
|
||||
self, rank: int, keys: list[tuple[str, str]]
|
||||
) -> dict[str, int]:
|
||||
"""Allocate one page for each key on the specified rank.
|
||||
|
||||
Args:
|
||||
rank: Rank ID to allocate pages on
|
||||
keys: List of keys to allocate pages for
|
||||
|
||||
Returns:
|
||||
Dictionary mapping key -> allocated page index
|
||||
"""
|
||||
with self.global_lock:
|
||||
if rank not in self.rank_metadata:
|
||||
raise ValueError(f"Rank {rank} not initialized")
|
||||
|
||||
# Batch allocate pages for all keys
|
||||
num_pages_needed = len(keys)
|
||||
allocated_pages = self.rank_metadata[rank].allocate_pages(num_pages_needed)
|
||||
|
||||
if len(allocated_pages) < num_pages_needed:
|
||||
logger.warning(
|
||||
"Rank %s only allocated %s pages for %s keys",
|
||||
rank,
|
||||
len(allocated_pages),
|
||||
num_pages_needed,
|
||||
)
|
||||
|
||||
allocation_results = {}
|
||||
for i, (key, prefix_key) in enumerate(keys):
|
||||
if key in self.key_metadata:
|
||||
key_meta = self.key_metadata[key]
|
||||
if key_meta.is_complete() and rank in key_meta.rank_to_page:
|
||||
# key is already fully written, reuse the existing page
|
||||
# and release the allocated pages back to the free pool.
|
||||
if i < len(allocated_pages):
|
||||
self.rank_metadata[rank].release_pages([allocated_pages[i]])
|
||||
allocation_results[key] = key_meta.rank_to_page[rank]
|
||||
continue
|
||||
|
||||
if i < len(allocated_pages):
|
||||
allocation_results[key] = allocated_pages[i]
|
||||
else:
|
||||
allocation_results[key] = -1 # No pages available
|
||||
|
||||
return allocation_results
|
||||
|
||||
def confirm_write_for_keys(
|
||||
self,
|
||||
rank: int,
|
||||
key_confirmations: list[tuple[str, int]],
|
||||
pages_to_release: list[int] | None = None,
|
||||
) -> None:
|
||||
"""Confirm write operations for keys and update metadata.
|
||||
|
||||
Args:
|
||||
rank: Rank ID that confirmed the writes
|
||||
key_confirmations: List of (key, page_index) tuples
|
||||
pages_to_release: List of page indices to release back to free pool
|
||||
"""
|
||||
with self.global_lock:
|
||||
# Confirm successful writes
|
||||
for key, page_index in key_confirmations:
|
||||
if key not in self.key_metadata:
|
||||
# Need to determine tp_world_size from rank_metadata
|
||||
tp_world_size = len(self.rank_metadata)
|
||||
self.key_metadata[key] = KeyMetadata(key, {}, tp_world_size)
|
||||
|
||||
# Add confirmed page to key metadata
|
||||
self.key_metadata[key].add_rank_page(rank, page_index)
|
||||
|
||||
# Release specified pages back to free pool
|
||||
if pages_to_release:
|
||||
self.rank_metadata[rank].release_pages(pages_to_release)
|
||||
logger.debug(
|
||||
"Released %s pages on rank %s: %s",
|
||||
len(pages_to_release),
|
||||
rank,
|
||||
pages_to_release,
|
||||
)
|
||||
|
||||
def batch_key_exists(self, keys: list[str]) -> list[bool]:
|
||||
"""Check if keys exist in metadata and all ranks have confirmed writes.
|
||||
|
||||
Args:
|
||||
keys: List of keys to check
|
||||
|
||||
Returns:
|
||||
List of boolean values indicating key existence and completion
|
||||
"""
|
||||
with self.global_lock:
|
||||
results = []
|
||||
for key in keys:
|
||||
if key not in self.key_metadata:
|
||||
results.append(False)
|
||||
else:
|
||||
# Check if all ranks in the TP world have confirmed writes
|
||||
key_meta = self.key_metadata[key]
|
||||
results.append(key_meta.is_complete())
|
||||
return results
|
||||
|
||||
def get_key_locations(self, rank: int, keys: list[str]) -> list[int | None]:
|
||||
"""Get page indices for keys on a specific rank.
|
||||
|
||||
Args:
|
||||
rank: Rank ID to query
|
||||
keys: List of keys to look up
|
||||
|
||||
Returns:
|
||||
List of page indices in the same order as input keys (None if key not found)
|
||||
"""
|
||||
with self.global_lock:
|
||||
if rank not in self.rank_metadata:
|
||||
raise ValueError(f"Rank {rank} not initialized")
|
||||
|
||||
results = []
|
||||
for key in keys:
|
||||
if key in self.key_metadata:
|
||||
key_meta = self.key_metadata[key]
|
||||
if key_meta.is_complete():
|
||||
page_index = key_meta.get_rank_page(rank)
|
||||
else:
|
||||
page_index = None
|
||||
|
||||
results.append(page_index)
|
||||
else:
|
||||
results.append(None)
|
||||
|
||||
return results
|
||||
|
||||
|
||||
class Hf3fsMetadataServer:
|
||||
"""HF3FS Metadata Server with improved key-based organization."""
|
||||
|
||||
def __init__(self, persistence_path: str | None = None, save_interval: int = 60):
|
||||
self.state = GlobalMetadataState()
|
||||
if HAS_ORJSON:
|
||||
self.app = FastAPI(default_response_class=ORJSONResponse)
|
||||
else:
|
||||
self.app = FastAPI()
|
||||
self._setup_routes()
|
||||
|
||||
async def _read_json(self, request: Request) -> dict:
|
||||
"""Parse request JSON using orjson if available."""
|
||||
body = await request.body()
|
||||
return orjson.loads(body)
|
||||
|
||||
def _json_response(self, content: dict):
|
||||
"""Return ORJSONResponse when available to bypass jsonable_encoder."""
|
||||
if HAS_ORJSON:
|
||||
return ORJSONResponse(content)
|
||||
else:
|
||||
return content
|
||||
|
||||
def _setup_routes(self):
|
||||
"""Setup FastAPI routes for new API design."""
|
||||
self.app.post("/rank/{rank}/initialize")(self.initialize_rank)
|
||||
self.app.post("/keys/batch_allocate")(self.batch_allocate_pages_for_keys)
|
||||
self.app.post("/keys/confirm_write")(self.confirm_write_for_keys)
|
||||
self.app.post("/keys/batch_exists")(self.batch_key_exists)
|
||||
self.app.post("/keys/get_locations")(self.get_key_locations)
|
||||
self.app.post("/clear")(self.clear)
|
||||
|
||||
async def initialize_rank(self, rank: int, request: Request):
|
||||
"""Initialize a rank with specified number of pages."""
|
||||
data = await self._read_json(request)
|
||||
role = data.get("role", "worker")
|
||||
num_pages = data.get("num_pages", 0)
|
||||
|
||||
if role == "scheduler":
|
||||
return self._json_response(
|
||||
{"message": "Scheduler role does not require initialization"}
|
||||
)
|
||||
|
||||
if role == "worker" and num_pages > 0:
|
||||
self.state.initialize_rank(rank, num_pages)
|
||||
return self._json_response(
|
||||
{"message": f"Rank {rank} initialized with {num_pages} pages"}
|
||||
)
|
||||
else:
|
||||
raise HTTPException(
|
||||
status_code=400, detail="Invalid initialization parameters"
|
||||
)
|
||||
|
||||
async def batch_allocate_pages_for_keys(self, request: Request):
|
||||
"""Allocate one page for each key on a specific rank."""
|
||||
data = await self._read_json(request)
|
||||
rank = data.get("rank")
|
||||
keys = data.get("keys", [])
|
||||
|
||||
# Validate input format
|
||||
if rank is None or not isinstance(keys, list):
|
||||
raise HTTPException(
|
||||
status_code=400, detail="Invalid request format: need 'rank' and 'keys'"
|
||||
)
|
||||
|
||||
try:
|
||||
# Perform allocation
|
||||
results = self.state.allocate_pages_for_keys(rank, keys)
|
||||
|
||||
# Convert results to response format
|
||||
response = {"rank": rank, "results": list(results.items())}
|
||||
return self._json_response(response)
|
||||
except Exception as e:
|
||||
raise HTTPException(
|
||||
status_code=500, detail=f"Allocation failed: {str(e)}"
|
||||
) from e
|
||||
|
||||
async def confirm_write_for_keys(self, request: Request):
|
||||
"""Confirm write operations for keys."""
|
||||
data = await self._read_json(request)
|
||||
rank = data.get("rank")
|
||||
confirmations = data.get("confirmations", [])
|
||||
pages_to_release = data.get("pages_to_release", [])
|
||||
|
||||
# Validate input format
|
||||
if rank is None or not isinstance(confirmations, list):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="Invalid request format: need 'rank' and 'confirmations'",
|
||||
)
|
||||
|
||||
try:
|
||||
self.state.confirm_write_for_keys(rank, confirmations, pages_to_release)
|
||||
|
||||
return Response(status_code=204)
|
||||
|
||||
except Exception as e:
|
||||
logger.error("Confirm write for keys failed: %s", e)
|
||||
raise HTTPException(
|
||||
status_code=500, detail=f"Confirmation failed: {str(e)}"
|
||||
) from e
|
||||
|
||||
async def batch_key_exists(self, request: Request):
|
||||
"""Check if multiple keys exist in metadata."""
|
||||
data = await self._read_json(request)
|
||||
keys = data.get("keys", [])
|
||||
|
||||
if not isinstance(keys, list):
|
||||
raise HTTPException(status_code=400, detail="Invalid keys format")
|
||||
|
||||
try:
|
||||
exists_results = self.state.batch_key_exists(keys)
|
||||
return self._json_response({"exists": exists_results})
|
||||
except Exception as e:
|
||||
raise HTTPException(
|
||||
status_code=500, detail=f"Key existence check failed: {str(e)}"
|
||||
) from e
|
||||
|
||||
async def get_key_locations(self, request: Request):
|
||||
"""Get page indices for keys on a specific rank."""
|
||||
data = await self._read_json(request)
|
||||
rank = data.get("rank")
|
||||
keys = data.get("keys", [])
|
||||
|
||||
# Validate input format
|
||||
if rank is None or not isinstance(keys, list):
|
||||
raise HTTPException(
|
||||
status_code=400, detail="Invalid request format: need 'rank' and 'keys'"
|
||||
)
|
||||
|
||||
try:
|
||||
# Get key locations
|
||||
locations = self.state.get_key_locations(rank, keys)
|
||||
return self._json_response({"locations": locations})
|
||||
except Exception as e:
|
||||
raise HTTPException(
|
||||
status_code=500, detail=f"Failed to get key locations: {str(e)}"
|
||||
) from e
|
||||
|
||||
async def clear(self, request: Request):
|
||||
"""Clear the metadata server."""
|
||||
self.state.clear()
|
||||
return Response(status_code=204)
|
||||
|
||||
def run(self, host: str = "0.0.0.0", port: int = 18000):
|
||||
"""Run the metadata server."""
|
||||
import uvicorn
|
||||
|
||||
logger.info("Starting improved metadata server on http://%s:%s", host, port)
|
||||
uvicorn.run(self.app, host=host, port=port)
|
||||
|
||||
|
||||
# --- Client implementation ---
|
||||
class Hf3fsMetadataInterface(ABC):
|
||||
"""Interface for HF3FS metadata operations."""
|
||||
|
||||
@abstractmethod
|
||||
def initialize(self, rank: int, num_pages: int = 0, role: str = "worker") -> None:
|
||||
"""Initialize the metadata service with specified number of pages."""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def allocate_pages_for_keys(
|
||||
self, rank: int, keys: list[tuple[str, str]]
|
||||
) -> list[tuple[str, int]]:
|
||||
"""Allocate one page for each key on the specified rank."""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def confirm_write_for_keys(
|
||||
self,
|
||||
rank: int,
|
||||
key_confirmations: list[tuple[str, int]],
|
||||
pages_to_release: list[int] | None = None,
|
||||
) -> None:
|
||||
"""Confirm write operations for keys and optionally release pages."""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def batch_key_exists(self, keys: list[str]) -> list[bool]:
|
||||
"""Check if keys exist and are complete across all ranks."""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def get_key_locations(self, rank: int, keys: list[str]) -> list[int]:
|
||||
"""Get page indices for keys on a specific rank."""
|
||||
pass
|
||||
|
||||
|
||||
class Hf3fsGlobalMetadataClient(Hf3fsMetadataInterface):
|
||||
"""Global HTTP metadata client for HF3FS."""
|
||||
|
||||
def __init__(self, base_url: str = "http://localhost:18000", max_retries: int = 3):
|
||||
self.base_url = base_url.rstrip("/")
|
||||
self._session = requests.Session()
|
||||
|
||||
retry_strategy = Retry(
|
||||
total=max_retries,
|
||||
backoff_factor=0.3,
|
||||
status_forcelist=[500, 502, 503, 504],
|
||||
allowed_methods=["GET", "POST"],
|
||||
)
|
||||
adapter = HTTPAdapter(max_retries=retry_strategy)
|
||||
self._session.mount("http://", adapter)
|
||||
|
||||
def _post(self, endpoint: str, json_data: dict) -> dict:
|
||||
"""Make POST request to metadata server."""
|
||||
try:
|
||||
url = f"{self.base_url}/{endpoint}"
|
||||
headers = {"Content-Type": "application/json"}
|
||||
if HAS_ORJSON:
|
||||
payload = orjson.dumps(json_data)
|
||||
else:
|
||||
import json
|
||||
|
||||
payload = json.dumps(json_data).encode("utf-8")
|
||||
response = self._session.post(url, data=payload, headers=headers)
|
||||
response.raise_for_status()
|
||||
|
||||
if response.status_code == 204 or not response.content:
|
||||
return {}
|
||||
if HAS_ORJSON:
|
||||
return orjson.loads(response.content)
|
||||
else:
|
||||
return response.json()
|
||||
except requests.exceptions.RequestException as e:
|
||||
logger.error("Failed to POST to %s after retries: %s", endpoint, e)
|
||||
raise RuntimeError(f"Failed to connect to metadata server: {e}") from e
|
||||
|
||||
def initialize(self, rank: int, num_pages: int = 0, role: str = "worker") -> None:
|
||||
"""Initialize a rank with specified number of pages."""
|
||||
self._post(f"rank/{rank}/initialize", {"num_pages": num_pages, "role": role})
|
||||
|
||||
def allocate_pages_for_keys(
|
||||
self, rank: int, keys: list[tuple[str, str]]
|
||||
) -> list[tuple[str, int]]:
|
||||
"""Allocate pages for keys on the specified rank."""
|
||||
response = self._post("keys/batch_allocate", {"rank": rank, "keys": keys})
|
||||
|
||||
# Convert response to expected format
|
||||
return response.get("results", {})
|
||||
|
||||
def confirm_write_for_keys(
|
||||
self,
|
||||
rank: int,
|
||||
key_confirmations: list[tuple[str, int]],
|
||||
pages_to_release: list[int] | None = None,
|
||||
) -> None:
|
||||
"""Confirm write operations for keys and optionally release pages."""
|
||||
payload = {
|
||||
"rank": rank,
|
||||
"confirmations": key_confirmations,
|
||||
"pages_to_release": pages_to_release or [],
|
||||
}
|
||||
|
||||
self._post("keys/confirm_write", payload)
|
||||
|
||||
def batch_key_exists(self, keys: list[str]) -> list[bool]:
|
||||
"""Check if keys exist and are complete across all ranks."""
|
||||
response = self._post("keys/batch_exists", {"keys": keys})
|
||||
return response.get("exists", [])
|
||||
|
||||
def get_key_locations(self, rank: int, keys: list[str]) -> list[int]:
|
||||
"""Get page indices for keys on a specific rank."""
|
||||
response = self._post("keys/get_locations", {"rank": rank, "keys": keys})
|
||||
return response.get("locations", [])
|
||||
|
||||
|
||||
def run_metadata_server(
|
||||
host: str = "0.0.0.0",
|
||||
port: int = 18000,
|
||||
):
|
||||
"""Run the improved HF3FS metadata server."""
|
||||
server = Hf3fsMetadataServer()
|
||||
server.run(host=host, port=port)
|
||||
|
||||
|
||||
# --- Main Execution ---
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser(description="Improved HF3FS Metadata Server")
|
||||
parser.add_argument(
|
||||
"--host", type=str, default="0.0.0.0", help="Host to bind the server to."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--port", type=int, default=18000, help="Port to run the server on."
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
run_metadata_server(args.host, args.port)
|
||||
@@ -0,0 +1,139 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
import threading
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Optional
|
||||
|
||||
from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorMetadata
|
||||
from vllm.v1.request import Request
|
||||
|
||||
|
||||
class AtomicCounter:
|
||||
"""Thread-safe atomic counter for round-robin operations."""
|
||||
|
||||
def __init__(self, n: int):
|
||||
assert n > 0, "Counter size must be positive"
|
||||
self._n = n
|
||||
self._value = 0
|
||||
self._lock = threading.Lock()
|
||||
|
||||
def next(self) -> int:
|
||||
"""Get next value in round-robin fashion."""
|
||||
with self._lock:
|
||||
current = self._value
|
||||
self._value = (current + 1) % self._n
|
||||
return current
|
||||
|
||||
|
||||
@dataclass
|
||||
class LoadBlockInfo:
|
||||
"""Operation for loading blocks from external storage."""
|
||||
|
||||
num_computed_blocks: int
|
||||
num_blocks_to_load: int
|
||||
need_fetch_block_ids: list[int]
|
||||
|
||||
|
||||
@dataclass
|
||||
class SaveBlockInfo:
|
||||
"""Operation for saving blocks to external storage."""
|
||||
|
||||
skip_leading_blocks: int
|
||||
|
||||
|
||||
@dataclass
|
||||
class RequestSchedulingState:
|
||||
"""Unified request scheduling state management."""
|
||||
|
||||
request_id: str
|
||||
request: Request | None = None
|
||||
|
||||
# Token and block tracking
|
||||
token_ids: list[int] = field(default_factory=list)
|
||||
allocated_block_ids: list[int] = field(default_factory=list)
|
||||
num_saved_blocks: int = 0
|
||||
|
||||
# Load operation info
|
||||
load_op: LoadBlockInfo | None = None
|
||||
|
||||
# Scheduling phase
|
||||
phase: str = "NEW" # NEW -> WAITING_TO_LOAD -> ACTIVE -> FINISHED
|
||||
|
||||
def needs_loading(self) -> bool:
|
||||
"""Check if request needs loading."""
|
||||
return self.load_op is not None and self.load_op.num_blocks_to_load > 0
|
||||
|
||||
def is_ready_to_load(self) -> bool:
|
||||
"""Check if request is ready for loading."""
|
||||
return self.phase == "WAITING_TO_LOAD" and self.needs_loading()
|
||||
|
||||
def update_tokens_and_blocks(self, new_token_ids: list[int], new_block_ids) -> None:
|
||||
"""Update with new tokens and blocks."""
|
||||
if new_token_ids:
|
||||
self.token_ids.extend(new_token_ids)
|
||||
|
||||
if new_block_ids is not None:
|
||||
normalized_block_ids = self._normalize_block_ids(new_block_ids)
|
||||
self.allocated_block_ids.extend(normalized_block_ids)
|
||||
|
||||
def _normalize_block_ids(self, block_ids) -> list[int]:
|
||||
"""Normalize block_ids to list format."""
|
||||
if not block_ids:
|
||||
return []
|
||||
if isinstance(block_ids, tuple):
|
||||
return block_ids[0] if block_ids else []
|
||||
if isinstance(block_ids, list):
|
||||
return block_ids
|
||||
return []
|
||||
|
||||
|
||||
@dataclass
|
||||
class HF3FSRequestMetadata:
|
||||
"""Metadata for a single request in HF3FS connector."""
|
||||
|
||||
request_id: str
|
||||
token_ids: list[int]
|
||||
block_ids: list[int]
|
||||
load_block_op: LoadBlockInfo | None = None
|
||||
save_block_op: SaveBlockInfo | None = None
|
||||
|
||||
@staticmethod
|
||||
def from_scheduling_state(
|
||||
state: "RequestSchedulingState",
|
||||
block_size: int,
|
||||
load_op: LoadBlockInfo | None = None,
|
||||
skip_leading_blocks: int | None = None,
|
||||
) -> Optional["HF3FSRequestMetadata"]:
|
||||
"""Create request metadata from scheduling state."""
|
||||
token_count = len(state.token_ids)
|
||||
total_blocks = token_count // block_size
|
||||
|
||||
skip_blocks = (
|
||||
state.num_saved_blocks
|
||||
if skip_leading_blocks is None
|
||||
else skip_leading_blocks
|
||||
)
|
||||
|
||||
new_blocks_to_save = total_blocks - state.num_saved_blocks
|
||||
if new_blocks_to_save <= 0 and load_op is None:
|
||||
return None
|
||||
|
||||
state.num_saved_blocks = total_blocks
|
||||
return HF3FSRequestMetadata(
|
||||
request_id=state.request_id,
|
||||
token_ids=state.token_ids,
|
||||
block_ids=state.allocated_block_ids,
|
||||
load_block_op=load_op,
|
||||
save_block_op=SaveBlockInfo(skip_leading_blocks=skip_blocks),
|
||||
)
|
||||
|
||||
|
||||
class HF3FSConnectorMetadata(KVConnectorMetadata):
|
||||
"""Container for HF3FS connector metadata."""
|
||||
|
||||
def __init__(self):
|
||||
self.requests: list[HF3FSRequestMetadata] = []
|
||||
|
||||
def add_request(self, request_metadata: HF3FSRequestMetadata) -> None:
|
||||
"""Add request to metadata."""
|
||||
self.requests.append(request_metadata)
|
||||
@@ -0,0 +1,288 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
import torch
|
||||
|
||||
from vllm.logger import init_logger
|
||||
from vllm.triton_utils import tl, triton
|
||||
|
||||
|
||||
@triton.jit
|
||||
def kv_cache_scatter_kernel(
|
||||
kv_cache_ptrs_ptr,
|
||||
source_ptr,
|
||||
token_indices_ptr,
|
||||
num_tokens_in_block,
|
||||
hidden_size,
|
||||
total_token_in_kvcache,
|
||||
num_layers,
|
||||
is_mla,
|
||||
BLOCK_SIZE: tl.constexpr,
|
||||
):
|
||||
layer_idx = tl.program_id(0)
|
||||
token_pos = tl.program_id(1)
|
||||
|
||||
if layer_idx >= num_layers or token_pos >= num_tokens_in_block:
|
||||
return
|
||||
|
||||
token_idx = tl.load(token_indices_ptr + token_pos)
|
||||
kv_cache_ptr = tl.cast(tl.load(kv_cache_ptrs_ptr + layer_idx), source_ptr.dtype)
|
||||
|
||||
if token_idx >= total_token_in_kvcache:
|
||||
return
|
||||
|
||||
if is_mla:
|
||||
# MLA format: source [num_layers, num_tokens_in_block, hidden_size]
|
||||
# MLA format: target [total_token_in_kvcache, hidden_size] (per layer)
|
||||
source_offset = (layer_idx * num_tokens_in_block + token_pos) * hidden_size
|
||||
target_offset = token_idx * hidden_size
|
||||
|
||||
for i in range(0, hidden_size, BLOCK_SIZE):
|
||||
offset = i + tl.arange(0, BLOCK_SIZE)
|
||||
mask = offset < hidden_size
|
||||
val = tl.load(source_ptr + source_offset + offset, mask=mask)
|
||||
tl.store(kv_cache_ptr + target_offset + offset, val, mask=mask)
|
||||
else:
|
||||
# MHA format: source [num_layers, 2, num_tokens_in_block, hidden_size]
|
||||
# MHA format: target [2, total_token_in_kvcache, hidden_size]
|
||||
source_offset_k = (
|
||||
layer_idx * num_tokens_in_block * 2 + token_pos
|
||||
) * hidden_size
|
||||
source_offset_v = (
|
||||
layer_idx * num_tokens_in_block * 2 + num_tokens_in_block + token_pos
|
||||
) * hidden_size
|
||||
|
||||
target_offset_k = token_idx * hidden_size
|
||||
target_offset_v = (total_token_in_kvcache + token_idx) * hidden_size
|
||||
|
||||
for i in range(0, hidden_size, BLOCK_SIZE):
|
||||
offset = i + tl.arange(0, BLOCK_SIZE)
|
||||
mask = offset < hidden_size
|
||||
|
||||
val_k = tl.load(source_ptr + source_offset_k + offset, mask=mask)
|
||||
val_v = tl.load(source_ptr + source_offset_v + offset, mask=mask)
|
||||
|
||||
tl.store(kv_cache_ptr + target_offset_k + offset, val_k, mask=mask)
|
||||
tl.store(kv_cache_ptr + target_offset_v + offset, val_v, mask=mask)
|
||||
|
||||
|
||||
@triton.jit
|
||||
def kv_cache_gather_kernel(
|
||||
kv_cache_ptrs_ptr,
|
||||
dst_ptr,
|
||||
token_indices_ptr,
|
||||
num_tokens_in_block,
|
||||
hidden_size,
|
||||
total_token_in_kvcache,
|
||||
num_layers,
|
||||
is_mla,
|
||||
BLOCK_SIZE: tl.constexpr,
|
||||
):
|
||||
layer_idx = tl.program_id(0)
|
||||
token_pos = tl.program_id(1)
|
||||
|
||||
if layer_idx >= num_layers or token_pos >= num_tokens_in_block:
|
||||
return
|
||||
|
||||
token_idx = tl.load(token_indices_ptr + token_pos)
|
||||
kv_cache_ptr = tl.cast(tl.load(kv_cache_ptrs_ptr + layer_idx), dst_ptr.dtype)
|
||||
|
||||
if token_idx >= total_token_in_kvcache:
|
||||
return
|
||||
|
||||
if is_mla:
|
||||
# MLA format: source [total_token_in_kvcache, hidden_size] (per layer)
|
||||
# MLA format: dst [num_layers, num_tokens_in_block, hidden_size]
|
||||
kvcache_offset = token_idx * hidden_size
|
||||
dst_offset = (layer_idx * num_tokens_in_block + token_pos) * hidden_size
|
||||
|
||||
for i in range(0, hidden_size, BLOCK_SIZE):
|
||||
offset = i + tl.arange(0, BLOCK_SIZE)
|
||||
mask = offset < hidden_size
|
||||
val = tl.load(kv_cache_ptr + kvcache_offset + offset, mask=mask)
|
||||
tl.store(dst_ptr + dst_offset + offset, val, mask=mask)
|
||||
else:
|
||||
# MHA format: source [2, total_token_in_kvcache, hidden_size]
|
||||
# MHA format: dst [num_layers, 2, num_tokens_in_block, hidden_size]
|
||||
dst_offset_k = (layer_idx * num_tokens_in_block * 2 + token_pos) * hidden_size
|
||||
dst_offset_v = (
|
||||
layer_idx * num_tokens_in_block * 2 + num_tokens_in_block + token_pos
|
||||
) * hidden_size
|
||||
|
||||
kvcache_offset_k = token_idx * hidden_size
|
||||
kvcache_offset_v = (total_token_in_kvcache + token_idx) * hidden_size
|
||||
|
||||
for i in range(0, hidden_size, BLOCK_SIZE):
|
||||
offset = i + tl.arange(0, BLOCK_SIZE)
|
||||
mask = offset < hidden_size
|
||||
|
||||
val_k = tl.load(kv_cache_ptr + kvcache_offset_k + offset, mask=mask)
|
||||
val_v = tl.load(kv_cache_ptr + kvcache_offset_v + offset, mask=mask)
|
||||
|
||||
tl.store(dst_ptr + dst_offset_k + offset, val_k, mask=mask)
|
||||
tl.store(dst_ptr + dst_offset_v + offset, val_v, mask=mask)
|
||||
|
||||
|
||||
def scatter_kv_caches(
|
||||
kv_caches_ptrs: torch.Tensor,
|
||||
total_token_in_kvcache: int,
|
||||
src_tensor: torch.Tensor,
|
||||
token_indices: list[int],
|
||||
is_mla: bool = False,
|
||||
) -> None:
|
||||
"""Scatter KV cache data from source tensor to KV cache storage.
|
||||
|
||||
Args:
|
||||
kv_caches_ptrs: Tensor of KV cache pointers (one per layer)
|
||||
total_token_in_kvcache: Total number of tokens in KV cache
|
||||
src_tensor: Source tensor containing data to scatter
|
||||
- MHA format: [num_layers, 2, num_tokens_in_block, hidden_size]
|
||||
- MLA format: [num_layers, num_tokens_in_block, hidden_size]
|
||||
token_indices: List of token positions to update
|
||||
is_mla: Whether using MLA model format
|
||||
"""
|
||||
num_layers = len(kv_caches_ptrs)
|
||||
num_tokens_in_block = len(token_indices)
|
||||
|
||||
if is_mla:
|
||||
# MLA: src_tensor is [num_layers, num_tokens_in_block, hidden_size]
|
||||
assert len(src_tensor.shape) == 3, (
|
||||
f"MLA src_tensor should be 3D, got {src_tensor.shape}"
|
||||
)
|
||||
hidden_size = src_tensor.shape[2]
|
||||
else:
|
||||
# MHA: src_tensor is [num_layers, 2, num_tokens_in_block, hidden_size]
|
||||
assert len(src_tensor.shape) == 4, (
|
||||
f"MHA src_tensor should be 4D, got {src_tensor.shape}"
|
||||
)
|
||||
hidden_size = src_tensor.shape[3]
|
||||
|
||||
device = src_tensor.device
|
||||
token_indices_tensor = torch.tensor(
|
||||
token_indices, dtype=torch.int32, device="cpu"
|
||||
).to(device, non_blocking=True)
|
||||
|
||||
grid = (num_layers, num_tokens_in_block)
|
||||
BLOCK_SIZE = 128
|
||||
|
||||
kv_cache_scatter_kernel[grid](
|
||||
kv_caches_ptrs,
|
||||
src_tensor,
|
||||
token_indices_tensor,
|
||||
num_tokens_in_block,
|
||||
hidden_size,
|
||||
total_token_in_kvcache,
|
||||
num_layers,
|
||||
is_mla,
|
||||
BLOCK_SIZE=BLOCK_SIZE,
|
||||
)
|
||||
|
||||
|
||||
def gather_kv_caches(
|
||||
kv_caches_ptrs: torch.Tensor,
|
||||
total_token_in_kvcache: int,
|
||||
dst_tensor: torch.Tensor,
|
||||
token_indices: list[int],
|
||||
is_mla: bool = False,
|
||||
) -> None:
|
||||
"""Gather KV cache data from KV cache storage to destination tensor.
|
||||
|
||||
Args:
|
||||
kv_caches_ptrs: Tensor of KV cache pointers (one per layer)
|
||||
total_token_in_kvcache: Total number of tokens in KV cache
|
||||
dst_tensor: Destination tensor to store gathered data
|
||||
- MHA format: [num_layers, 2, num_tokens_in_block, hidden_size]
|
||||
- MLA format: [num_layers, num_tokens_in_block, hidden_size]
|
||||
token_indices: List of token positions to gather
|
||||
is_mla: Whether using MLA model format
|
||||
"""
|
||||
num_layers = kv_caches_ptrs.shape[0]
|
||||
num_tokens_in_block = len(token_indices)
|
||||
|
||||
if is_mla:
|
||||
# MLA: dst_tensor is [num_layers, num_tokens_in_block, hidden_size]
|
||||
assert len(dst_tensor.shape) == 3, (
|
||||
f"MLA dst_tensor should be 3D, got {dst_tensor.shape}"
|
||||
)
|
||||
assert dst_tensor.shape[0] == num_layers, (
|
||||
f"Layer count mismatch: {dst_tensor.shape[0]} vs {num_layers}"
|
||||
)
|
||||
assert dst_tensor.shape[1] == num_tokens_in_block, (
|
||||
f"Token count mismatch: {dst_tensor.shape[1]} vs {num_tokens_in_block}"
|
||||
)
|
||||
hidden_size = dst_tensor.shape[2]
|
||||
else:
|
||||
# MHA: dst_tensor is [num_layers, 2, num_tokens_in_block, hidden_size]
|
||||
assert len(dst_tensor.shape) == 4, (
|
||||
f"MHA dst_tensor should be 4D, got {dst_tensor.shape}"
|
||||
)
|
||||
assert dst_tensor.shape[0] == num_layers, (
|
||||
f"Layer count mismatch: {dst_tensor.shape[0]} vs {num_layers}"
|
||||
)
|
||||
assert dst_tensor.shape[1] == 2, (
|
||||
f"MHA should have 2 (K,V) components, got {dst_tensor.shape[1]}"
|
||||
)
|
||||
assert dst_tensor.shape[2] == num_tokens_in_block, (
|
||||
f"Token count mismatch: {dst_tensor.shape[2]} vs {num_tokens_in_block}"
|
||||
)
|
||||
hidden_size = dst_tensor.shape[3]
|
||||
|
||||
device = dst_tensor.device
|
||||
token_indices_tensor = torch.tensor(
|
||||
token_indices, dtype=torch.int32, device="cpu"
|
||||
).to(device, non_blocking=True)
|
||||
|
||||
grid = (num_layers, num_tokens_in_block)
|
||||
BLOCK_SIZE = 128
|
||||
|
||||
kv_cache_gather_kernel[grid](
|
||||
kv_caches_ptrs,
|
||||
dst_tensor,
|
||||
token_indices_tensor,
|
||||
num_tokens_in_block,
|
||||
hidden_size,
|
||||
total_token_in_kvcache,
|
||||
num_layers,
|
||||
is_mla,
|
||||
BLOCK_SIZE=BLOCK_SIZE,
|
||||
)
|
||||
|
||||
|
||||
class CopyBufferAllocator:
|
||||
"""Memory pool for tensor buffers to avoid frequent allocation/deallocation."""
|
||||
|
||||
def __init__(
|
||||
self, device: torch.device, dtype: torch.dtype, shape: list, max_count: int
|
||||
):
|
||||
self._shape = shape
|
||||
self._max_count = max_count
|
||||
self._device = device
|
||||
self._free_buffers = [
|
||||
torch.empty(shape, dtype=dtype, device=device) for _ in range(max_count)
|
||||
]
|
||||
self._inuse_count = 0
|
||||
|
||||
def alloc_buffer(self, count: int) -> list[torch.Tensor] | None:
|
||||
"""Allocate buffers from the pool."""
|
||||
if count == 0:
|
||||
return []
|
||||
|
||||
if self._inuse_count + count <= self._max_count:
|
||||
self._inuse_count += count
|
||||
result = self._free_buffers[-count:]
|
||||
del self._free_buffers[-count:]
|
||||
return result
|
||||
return None
|
||||
|
||||
def free_buffer(self, buffers: list[torch.Tensor]) -> None:
|
||||
"""Return buffers to the pool."""
|
||||
if not buffers:
|
||||
return
|
||||
|
||||
if self._inuse_count >= len(buffers):
|
||||
self._inuse_count -= len(buffers)
|
||||
self._free_buffers.extend(buffers)
|
||||
else:
|
||||
raise RuntimeError("Attempted to free more buffers than allocated")
|
||||
|
||||
|
||||
logger = init_logger(__name__)
|
||||
@@ -0,0 +1,133 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
import logging
|
||||
import os
|
||||
|
||||
import torch
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
HF3FS_AVAILABLE = True
|
||||
|
||||
|
||||
class Hf3fsClient:
|
||||
"""Mock HF3FS client using file backend for debugging and testing."""
|
||||
|
||||
def __init__(self, path: str, size: int, bytes_per_page: int, entries: int):
|
||||
self._size = size
|
||||
self._bytes_per_page = bytes_per_page
|
||||
self._entries = entries
|
||||
self._file_path = path
|
||||
|
||||
self._ensure_file_exists()
|
||||
logger.debug("Initialized mock HF3FS client: %s (%d bytes)", path, size)
|
||||
|
||||
def _ensure_file_exists(self) -> None:
|
||||
"""Create file if it doesn't exist."""
|
||||
if not os.path.exists(self._file_path):
|
||||
with open(self._file_path, "w+b") as f:
|
||||
f.truncate(self._size)
|
||||
|
||||
def batch_read(self, offsets: list[int], tensors: list[torch.Tensor]) -> list[int]:
|
||||
"""Read data from file at specified offsets into tensors."""
|
||||
results = []
|
||||
|
||||
try:
|
||||
with open(self._file_path, "rb") as f:
|
||||
for offset, tensor in zip(offsets, tensors):
|
||||
num_bytes = tensor.numel() * tensor.element_size()
|
||||
|
||||
if offset < 0 or offset + num_bytes > self._size:
|
||||
results.append(-1)
|
||||
continue
|
||||
|
||||
f.seek(offset)
|
||||
buffer_data = f.read(num_bytes)
|
||||
|
||||
if len(buffer_data) == num_bytes == self._bytes_per_page:
|
||||
tensor_data = self._convert_buffer_to_tensor(
|
||||
buffer_data, tensor.dtype
|
||||
)
|
||||
tensor.copy_(
|
||||
tensor_data.reshape(tensor.shape).to(tensor.device)
|
||||
)
|
||||
results.append(self._bytes_per_page)
|
||||
else:
|
||||
logger.error(
|
||||
"Read size mismatch: got %d, expected %d",
|
||||
len(buffer_data),
|
||||
num_bytes,
|
||||
)
|
||||
results.append(-1)
|
||||
except Exception as e:
|
||||
logger.error("Batch read error: %s", e)
|
||||
results.extend([-1] * (len(offsets) - len(results)))
|
||||
|
||||
return results
|
||||
|
||||
def _convert_buffer_to_tensor(
|
||||
self, buffer_data: bytes, dtype: torch.dtype
|
||||
) -> torch.Tensor:
|
||||
"""Convert buffer data to tensor with proper dtype handling."""
|
||||
if dtype == torch.bfloat16:
|
||||
tensor_data = torch.frombuffer(buffer_data, dtype=torch.uint16)
|
||||
return tensor_data.view(dtype=torch.bfloat16)
|
||||
else:
|
||||
return torch.frombuffer(buffer_data, dtype=dtype)
|
||||
|
||||
def batch_write(
|
||||
self, offsets: list[int], tensors: list[torch.Tensor], event: torch.cuda.Event
|
||||
) -> list[int]:
|
||||
"""Write data from tensors to file at specified offsets."""
|
||||
results = []
|
||||
|
||||
try:
|
||||
torch.cuda.current_stream().wait_event(event)
|
||||
|
||||
# Convert tensors to bytes
|
||||
data_bytes_list = [self._tensor_to_bytes(tensor) for tensor in tensors]
|
||||
|
||||
# Write to file
|
||||
with open(self._file_path, "r+b") as f:
|
||||
for offset, data_bytes in zip(offsets, data_bytes_list):
|
||||
if offset < 0 or offset + len(data_bytes) > self._size:
|
||||
results.append(-1)
|
||||
continue
|
||||
|
||||
f.seek(offset)
|
||||
bytes_written = f.write(data_bytes)
|
||||
|
||||
if bytes_written == len(data_bytes) == self._bytes_per_page:
|
||||
results.append(self._bytes_per_page)
|
||||
else:
|
||||
logger.error(
|
||||
"Write size mismatch: wrote %d, expected %d",
|
||||
bytes_written,
|
||||
self._bytes_per_page,
|
||||
)
|
||||
results.append(-1)
|
||||
|
||||
except Exception as e:
|
||||
logger.error("Batch write error: %s", e)
|
||||
results.extend([-1] * (len(offsets) - len(results)))
|
||||
|
||||
return results
|
||||
|
||||
def _tensor_to_bytes(self, tensor: torch.Tensor) -> bytes:
|
||||
"""Convert tensor to bytes with proper dtype handling."""
|
||||
cpu_tensor = tensor.cpu()
|
||||
if cpu_tensor.dtype == torch.bfloat16:
|
||||
return cpu_tensor.view(dtype=torch.uint16).numpy().tobytes()
|
||||
else:
|
||||
return cpu_tensor.numpy().tobytes()
|
||||
|
||||
def get_size(self) -> int:
|
||||
"""Get the total size of the storage file."""
|
||||
return self._size
|
||||
|
||||
def close(self) -> None:
|
||||
"""Close the client (no-op for file backend)."""
|
||||
pass
|
||||
|
||||
def flush(self) -> None:
|
||||
"""Flush any pending writes (no-op for file backend)."""
|
||||
pass
|
||||
@@ -0,0 +1,57 @@
|
||||
#include <cuda_runtime.h>
|
||||
#include <torch/extension.h>
|
||||
#include <cstring>
|
||||
#include <vector>
|
||||
|
||||
void read_shm(const torch::Tensor& shm, const torch::Tensor& pin,
|
||||
std::vector<torch::Tensor> dst, uint64_t stream_ptr) {
|
||||
py::gil_scoped_release release;
|
||||
|
||||
cudaStream_t stream = reinterpret_cast<cudaStream_t>(stream_ptr);
|
||||
|
||||
// Copy from shared memory to pinned memory
|
||||
char* shm_ptr = static_cast<char*>(shm.data_ptr());
|
||||
char* src_ptr = static_cast<char*>(pin.data_ptr());
|
||||
std::memcpy(src_ptr, shm_ptr, shm.numel() * shm.element_size());
|
||||
|
||||
// Copy from pinned memory to GPU tensors
|
||||
size_t current = 0;
|
||||
for (size_t i = 0; i < dst.size(); ++i) {
|
||||
auto& t = dst[i];
|
||||
size_t t_bytes = t.numel() * t.element_size();
|
||||
char* dst_ptr = static_cast<char*>(t.data_ptr());
|
||||
cudaMemcpyAsync(dst_ptr, src_ptr + current, t_bytes, cudaMemcpyHostToDevice,
|
||||
stream);
|
||||
current += t_bytes;
|
||||
}
|
||||
cudaStreamSynchronize(stream);
|
||||
}
|
||||
|
||||
void write_shm(const std::vector<torch::Tensor> src, torch::Tensor& shm,
|
||||
const torch::Tensor& pin, uint64_t stream_ptr) {
|
||||
py::gil_scoped_release release;
|
||||
|
||||
cudaStream_t stream = reinterpret_cast<cudaStream_t>(stream_ptr);
|
||||
|
||||
// Copy from GPU tensors to pinned memory
|
||||
char* dst_ptr = static_cast<char*>(pin.data_ptr());
|
||||
size_t current = 0;
|
||||
for (size_t i = 0; i < src.size(); ++i) {
|
||||
auto& t = src[i];
|
||||
size_t t_bytes = t.numel() * t.element_size();
|
||||
char* src_ptr = static_cast<char*>(t.data_ptr());
|
||||
cudaMemcpyAsync(dst_ptr + current, src_ptr, t_bytes, cudaMemcpyDeviceToHost,
|
||||
stream);
|
||||
current += t_bytes;
|
||||
}
|
||||
cudaStreamSynchronize(stream);
|
||||
|
||||
// Copy from pinned memory to shared memory
|
||||
char* shm_ptr = static_cast<char*>(shm.data_ptr());
|
||||
std::memcpy(shm_ptr, dst_ptr, shm.numel() * shm.element_size());
|
||||
}
|
||||
|
||||
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||
m.def("read_shm", &read_shm, "Read tensors from shared memory");
|
||||
m.def("write_shm", &write_shm, "Write tensors to shared memory");
|
||||
}
|
||||
@@ -1001,7 +1001,7 @@ class OpenAIServingResponses(OpenAIServing):
|
||||
|
||||
# Use parser to extract and create response output items
|
||||
if self.parser:
|
||||
parser = self.parser(tokenizer)
|
||||
parser = self.parser(tokenizer, request.tools)
|
||||
return parser.extract_response_outputs(
|
||||
model_output=final_output.text,
|
||||
model_output_token_ids=final_output.token_ids,
|
||||
|
||||
@@ -500,7 +500,7 @@ if is_torch_equal("2.9.0"):
|
||||
# This mirrors the fix in https://github.com/pytorch/pytorch/pull/177558
|
||||
# and can be removed once torch >=2.12 is the minimum supported version.
|
||||
|
||||
if not is_torch_equal_or_newer("2.12.0"):
|
||||
if is_torch_equal_or_newer("2.10.0") and not is_torch_equal_or_newer("2.12.0"):
|
||||
import builtins as _builtins
|
||||
import pickle
|
||||
|
||||
|
||||
@@ -1464,6 +1464,10 @@ environment_variables: dict[str, Callable[[], Any]] = {
|
||||
# - "flashinfer-trtllm": use flashinfer trtllm GEMM backend
|
||||
# - "flashinfer-cutlass": use flashinfer cutlass GEMM backend
|
||||
# - "marlin": use marlin GEMM backend (for GPUs without native FP4 support)
|
||||
# - "emulation":
|
||||
# use BF16/FP16 GEMM, dequantizing weights and running QDQ on activations.
|
||||
# This is only meant for research purposes to run on devices where NVFP4
|
||||
# GEMM kernels are not available.
|
||||
# - <none>: automatically pick an available backend
|
||||
"VLLM_NVFP4_GEMM_BACKEND": env_with_choices(
|
||||
"VLLM_NVFP4_GEMM_BACKEND",
|
||||
@@ -1474,6 +1478,7 @@ environment_variables: dict[str, Callable[[], Any]] = {
|
||||
"flashinfer-cutlass",
|
||||
"cutlass",
|
||||
"marlin",
|
||||
"emulation",
|
||||
],
|
||||
),
|
||||
# Controls garbage collection during CUDA graph capture.
|
||||
|
||||
@@ -16,7 +16,6 @@ def rms_norm(
|
||||
x_var = x if variance_size is None else x[..., :variance_size]
|
||||
variance = x_var.pow(2).mean(dim=-1, keepdim=True)
|
||||
x = x * torch.rsqrt(variance + epsilon)
|
||||
x = x.to(orig_dtype)
|
||||
if weight is not None:
|
||||
x = x * weight
|
||||
return x
|
||||
x = x.to(weight.dtype) * weight
|
||||
return x.to(orig_dtype)
|
||||
|
||||
@@ -36,13 +36,11 @@ AITER_SUPPORTED = is_aiter_found()
|
||||
|
||||
rms_no_var_16bit_only = (
|
||||
lambda x, weight, epsilon, variance_size=None: variance_size is None
|
||||
and x.dtype
|
||||
in (
|
||||
torch.float16,
|
||||
torch.bfloat16,
|
||||
)
|
||||
and x.dtype in (torch.float16, torch.bfloat16)
|
||||
and (weight is None or weight.dtype == x.dtype)
|
||||
)
|
||||
"""AITER rms_norm only supports float16 and bfloat16 acts and no var_size override."""
|
||||
"""AITER rms_norm only supports float16 and bfloat16 acts, no var_size override,
|
||||
and requires weight dtype to match x dtype."""
|
||||
|
||||
|
||||
@ir.ops.rms_norm.register_impl(
|
||||
|
||||
@@ -11,8 +11,11 @@ current_platform.import_kernels()
|
||||
CUDA_ALIKE = current_platform.is_cuda_alike()
|
||||
"""Most kernels in this file are supported on all CUDA-alike platforms."""
|
||||
|
||||
rms_no_var_size = lambda x, weight, epsilon, variance_size=None: variance_size is None
|
||||
"""vLLM kernel does not support variance_size parameter."""
|
||||
rms_no_var_size = (
|
||||
lambda x, weight, epsilon, variance_size=None: variance_size is None
|
||||
and (weight is None or weight.dtype == x.dtype)
|
||||
)
|
||||
"""vLLM kernel requires no variance_size override and matching input/weight dtype."""
|
||||
|
||||
|
||||
@ir.ops.rms_norm.register_impl(
|
||||
|
||||
@@ -18,7 +18,9 @@ def is_xpu_kernels_found() -> bool:
|
||||
XPU_KERNELS_SUPPORTED = is_xpu_kernels_found()
|
||||
"""Kernels in this file are supported if vLLM XPU kernels are installed."""
|
||||
|
||||
rms_no_var = lambda x, weight, epsilon, variance_size=None: variance_size is None
|
||||
rms_no_var = lambda x, weight, epsilon, variance_size=None: variance_size is None and (
|
||||
weight is None or weight.dtype == x.dtype
|
||||
)
|
||||
|
||||
|
||||
@ir.ops.rms_norm.register_impl(
|
||||
|
||||
@@ -47,7 +47,6 @@ if has_triton_kernels():
|
||||
BIT,
|
||||
Bitmatrix,
|
||||
)
|
||||
from triton_kernels.topk import topk
|
||||
|
||||
try:
|
||||
from triton_kernels.tensor import (
|
||||
@@ -89,6 +88,7 @@ def pack_bitmatrix(
|
||||
offsets = offsets_m[:, None] * n_expts_act + offsets_k[None, :]
|
||||
mask = (offsets_m < n_rows)[:, None] & (offsets_k < n_expts_act)[None, :]
|
||||
indices = tl.load(topk_ids + offsets, mask=mask, other=-1)
|
||||
valid = indices >= 0
|
||||
div = indices // 32
|
||||
rem = indices % 32
|
||||
one = tl.cast(1, tl.uint32)
|
||||
@@ -99,8 +99,13 @@ def pack_bitmatrix(
|
||||
offs = tl.arange(0, BLOCK_SIZE_K // 32) + i * (BLOCK_SIZE_K // 32)
|
||||
# All topks that need to go into this column has the correct bit set.
|
||||
# Other bits are 0. x is a 2D tensor.
|
||||
# Guard with `valid` to prevent negative indices from producing
|
||||
# spurious bits (on HIP, -1 // 32 == 0 and 1 << (-1 % 32) sets
|
||||
# bit 31).
|
||||
x = tl.where(
|
||||
div[:, :, None] == offs[None, None, :], (one << rem)[:, :, None], 0
|
||||
valid[:, :, None] & (div[:, :, None] == offs[None, None, :]),
|
||||
(one << rem)[:, :, None],
|
||||
0,
|
||||
)
|
||||
# Reduce x to get a single int32_t bitpack.
|
||||
y = tl.reduce_or(x, axis=1)
|
||||
@@ -108,93 +113,6 @@ def pack_bitmatrix(
|
||||
tl.store(bitmatrix_ptrs, y, mask=offsets_m[:, None] < n_rows)
|
||||
|
||||
|
||||
def legacy_routing_from_bitmatrix(
|
||||
bitmatrix: "Bitmatrix",
|
||||
expt_scal: torch.Tensor,
|
||||
expt_indx: torch.Tensor,
|
||||
n_expts_tot: int,
|
||||
n_expts_act: int,
|
||||
) -> tuple["RoutingData", "GatherIndx", "ScatterIndx"]:
|
||||
"""
|
||||
Replacement for the removed triton_kernels.routing.routing_from_bitmatrix.
|
||||
Creates routing data from a bitmatrix representation.
|
||||
"""
|
||||
if use_legacy_triton_kernels:
|
||||
from triton_kernels.routing import routing_from_bitmatrix
|
||||
|
||||
return routing_from_bitmatrix(
|
||||
bitmatrix, expt_scal, expt_indx, n_expts_tot, n_expts_act
|
||||
)
|
||||
sparse_logits = SparseMatrix(indx=expt_indx, vals=expt_scal, mask=bitmatrix)
|
||||
dispatch_indx = sparse_logits.mask_metadata.row_sorted_indx
|
||||
combine_indx = sparse_logits.mask_metadata.col_sorted_indx
|
||||
ragged_batch_metadata = make_ragged_tensor_metadata(
|
||||
sparse_logits.mask_metadata.col_sum,
|
||||
dispatch_indx.shape[0],
|
||||
)
|
||||
gate_scal = sparse_logits.vals.flatten()[combine_indx]
|
||||
routing_data = RoutingData(
|
||||
gate_scal,
|
||||
ragged_batch_metadata.block_sizes,
|
||||
n_expts_tot,
|
||||
n_expts_act,
|
||||
ragged_batch_metadata,
|
||||
)
|
||||
gather_idx = GatherIndx(combine_indx, dispatch_indx)
|
||||
scatter_idx = ScatterIndx(dispatch_indx, combine_indx)
|
||||
return routing_data, gather_idx, scatter_idx
|
||||
|
||||
|
||||
def legacy_routing_from_sparsematrix(
|
||||
sparse_logits: "SparseMatrix",
|
||||
n_expts_tot: int,
|
||||
n_expts_act: int,
|
||||
) -> tuple["RoutingData", "GatherIndx", "ScatterIndx"]:
|
||||
"""
|
||||
Creates routing data from a SparseMatrix representation.
|
||||
"""
|
||||
dispatch_indx = sparse_logits.mask_metadata.row_sorted_indx
|
||||
combine_indx = sparse_logits.mask_metadata.col_sorted_indx
|
||||
ragged_batch_metadata = make_ragged_tensor_metadata(
|
||||
sparse_logits.mask_metadata.col_sum,
|
||||
dispatch_indx.shape[0],
|
||||
)
|
||||
gate_scal = sparse_logits.vals.flatten()[combine_indx]
|
||||
routing_data = RoutingData(
|
||||
gate_scal,
|
||||
ragged_batch_metadata.block_sizes,
|
||||
n_expts_tot,
|
||||
n_expts_act,
|
||||
ragged_batch_metadata,
|
||||
)
|
||||
gather_idx = GatherIndx(combine_indx, dispatch_indx)
|
||||
scatter_idx = ScatterIndx(dispatch_indx, combine_indx)
|
||||
return routing_data, gather_idx, scatter_idx
|
||||
|
||||
|
||||
def legacy_routing(
|
||||
logits: torch.Tensor,
|
||||
n_expts_act: int,
|
||||
sm_first: bool = False,
|
||||
) -> tuple["RoutingData", "GatherIndx", "ScatterIndx"]:
|
||||
"""
|
||||
Replacement for the removed triton_kernels.routing.routing function.
|
||||
Computes routing data from gating logits.
|
||||
"""
|
||||
if use_legacy_triton_kernels:
|
||||
from triton_kernels.routing import routing
|
||||
|
||||
return routing(logits, n_expts_act, sm_first=sm_first)
|
||||
if sm_first:
|
||||
logits = torch.softmax(logits, dim=-1)
|
||||
sparse_logits = topk(logits, n_expts_act, apply_softmax=not sm_first)
|
||||
return legacy_routing_from_sparsematrix(
|
||||
sparse_logits,
|
||||
logits.shape[-1],
|
||||
n_expts_act,
|
||||
)
|
||||
|
||||
|
||||
def triton_kernel_moe_forward(
|
||||
hidden_states: torch.Tensor,
|
||||
w1, # Tensor or triton_kernels.Tensor
|
||||
@@ -241,26 +159,22 @@ def triton_kernel_moe_forward(
|
||||
unpadded_K_w2=unpadded_K_w2,
|
||||
)
|
||||
|
||||
if expert_map is not None:
|
||||
# With expert parallelism, legacy_routing produces routing data
|
||||
# using global expert IDs which don't correspond to local weight
|
||||
# indices. Split the routing into topk selection + expert_map
|
||||
# remapping + local routing data construction (matching the
|
||||
# approach used by OAITritonExperts.apply).
|
||||
from triton_kernels.topk import topk as topk_fn
|
||||
from triton_kernels.topk import topk as topk_fn
|
||||
|
||||
sm_first = not renormalize
|
||||
logits = gating_output
|
||||
if sm_first:
|
||||
logits = torch.softmax(logits, dim=-1)
|
||||
topk_result = topk_fn(logits, topk, apply_softmax=not sm_first)
|
||||
# topk may return a tuple (vals, indx, bitmatrix) or a
|
||||
# SparseMatrix depending on the triton_kernels version.
|
||||
if isinstance(topk_result, tuple):
|
||||
topk_weights, topk_ids_raw, _ = topk_result
|
||||
else:
|
||||
topk_weights = topk_result.vals
|
||||
topk_ids_raw = topk_result.indx
|
||||
sm_first = not renormalize
|
||||
logits = gating_output
|
||||
if sm_first:
|
||||
logits = torch.softmax(logits, dim=-1)
|
||||
topk_result = topk_fn(logits, topk, apply_softmax=not sm_first)
|
||||
# topk may return a tuple (vals, indx, bitmatrix) or a
|
||||
# SparseMatrix depending on the triton_kernels version.
|
||||
if isinstance(topk_result, tuple):
|
||||
topk_weights, topk_ids_raw, _ = topk_result
|
||||
else:
|
||||
topk_weights = topk_result.vals
|
||||
topk_ids_raw = topk_result.indx
|
||||
|
||||
if expert_map is not None:
|
||||
# topk_ids_raw contains global expert IDs - remap to local.
|
||||
topk_ids = expert_map[topk_ids_raw.to(torch.long)]
|
||||
local_num_experts = w1.shape[0]
|
||||
@@ -271,8 +185,9 @@ def triton_kernel_moe_forward(
|
||||
effective_expert_map = None
|
||||
effective_global_num_experts = local_num_experts
|
||||
else:
|
||||
routing_data, gather_idx, scatter_idx = legacy_routing(
|
||||
gating_output, topk, sm_first=not renormalize
|
||||
topk_ids = topk_ids_raw.to(torch.long)
|
||||
routing_data, gather_idx, scatter_idx = make_routing_data(
|
||||
topk_ids, topk_weights, gating_output.shape[-1]
|
||||
)
|
||||
effective_expert_map = expert_map
|
||||
effective_global_num_experts = global_num_experts
|
||||
@@ -539,10 +454,31 @@ def make_routing_data(
|
||||
|
||||
# matmul_ogs expects invalid topk_weights to be -1s
|
||||
topk_weights = torch.where(topk_ids == -1, -1.0, topk_weights)
|
||||
routing_data, gather_indx, scatter_indx = legacy_routing_from_bitmatrix(
|
||||
bitmatrix, topk_weights, topk_ids, num_local_experts, num_topk
|
||||
)
|
||||
|
||||
if use_legacy_triton_kernels:
|
||||
from triton_kernels.routing import routing_from_bitmatrix
|
||||
|
||||
return routing_from_bitmatrix(
|
||||
bitmatrix, topk_weights, topk_ids, num_local_experts, num_topk
|
||||
)
|
||||
|
||||
sparse_logits = SparseMatrix(indx=topk_ids, vals=topk_weights, mask=bitmatrix)
|
||||
dispatch_indx = sparse_logits.mask_metadata.row_sorted_indx
|
||||
combine_indx = sparse_logits.mask_metadata.col_sorted_indx
|
||||
ragged_batch_metadata = make_ragged_tensor_metadata(
|
||||
sparse_logits.mask_metadata.col_sum,
|
||||
dispatch_indx.shape[0],
|
||||
)
|
||||
gate_scal = sparse_logits.vals.flatten()[combine_indx]
|
||||
routing_data = RoutingData(
|
||||
gate_scal,
|
||||
ragged_batch_metadata.block_sizes,
|
||||
num_local_experts,
|
||||
num_topk,
|
||||
ragged_batch_metadata,
|
||||
)
|
||||
gather_indx = GatherIndx(combine_indx, dispatch_indx)
|
||||
scatter_indx = ScatterIndx(dispatch_indx, combine_indx)
|
||||
return routing_data, gather_indx, scatter_indx
|
||||
|
||||
|
||||
|
||||
@@ -222,7 +222,7 @@ class UnquantizedFusedMoEMethod(FusedMoEMethodBase, CustomOp):
|
||||
self.cpu_fused_moe = cpu_fused_moe.CPUFusedMOE(layer)
|
||||
else:
|
||||
self.cpu_fused_moe = cpu_fused_moe.CPUFusedMOE(layer)
|
||||
elif current_platform.is_xpu():
|
||||
elif self.unquantized_backend == UnquantizedMoeBackend.XPU:
|
||||
w13 = layer.w13_weight
|
||||
w2 = layer.w2_weight
|
||||
|
||||
|
||||
@@ -376,77 +376,32 @@ class GemmaRMSNorm(CustomOp):
|
||||
self.weight = nn.Parameter(torch.zeros(hidden_size))
|
||||
self.variance_epsilon = eps
|
||||
|
||||
@staticmethod
|
||||
def _forward_static_no_residual(
|
||||
weight: torch.Tensor,
|
||||
variance_epsilon: float,
|
||||
x: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""PyTorch-native implementation equivalent to forward() without residual."""
|
||||
orig_dtype = x.dtype
|
||||
x = x.float()
|
||||
variance = x.pow(2).mean(dim=-1, keepdim=True)
|
||||
x = x * torch.rsqrt(variance + variance_epsilon)
|
||||
x = x * (1.0 + weight.float())
|
||||
x = x.to(orig_dtype)
|
||||
return x
|
||||
|
||||
@staticmethod
|
||||
def _forward_static_with_residual(
|
||||
weight: torch.Tensor,
|
||||
variance_epsilon: float,
|
||||
x: torch.Tensor,
|
||||
residual: torch.Tensor,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""PyTorch-native implementation equivalent to forward() with residual."""
|
||||
orig_dtype = x.dtype
|
||||
x = (
|
||||
x.float() + residual.float()
|
||||
if orig_dtype == torch.float16
|
||||
else x + residual
|
||||
)
|
||||
residual = x
|
||||
|
||||
x = x.float()
|
||||
variance = x.pow(2).mean(dim=-1, keepdim=True)
|
||||
x = x * torch.rsqrt(variance + variance_epsilon)
|
||||
# Llama does x.to(float16) * w whilst Gemma is (x * w).to(float16)
|
||||
# See https://github.com/huggingface/transformers/pull/29402
|
||||
x = x * (1.0 + weight.float())
|
||||
x = x.to(orig_dtype)
|
||||
return x, residual
|
||||
|
||||
def forward_native(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
residual: torch.Tensor | None = None,
|
||||
) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]:
|
||||
"""PyTorch-native implementation equivalent to forward()."""
|
||||
if residual is None:
|
||||
return self._forward_static_no_residual(
|
||||
self.weight.data, self.variance_epsilon, x
|
||||
)
|
||||
else:
|
||||
return self._forward_static_with_residual(
|
||||
self.weight.data, self.variance_epsilon, x, residual
|
||||
orig_dtype = x.dtype
|
||||
weight = self.weight.data.float() + 1.0
|
||||
if residual is not None:
|
||||
x = (
|
||||
x.float() + residual.float()
|
||||
if orig_dtype == torch.float16
|
||||
else x + residual
|
||||
)
|
||||
residual = x
|
||||
# ir.ops.rms_norm handles fp32 upcast internally
|
||||
out = ir.ops.rms_norm(x, weight, self.variance_epsilon)
|
||||
return (
|
||||
out.to(orig_dtype) if residual is None else (out.to(orig_dtype), residual)
|
||||
)
|
||||
|
||||
def forward_cuda(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
residual: torch.Tensor | None = None,
|
||||
) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]:
|
||||
if torch.compiler.is_compiling():
|
||||
return self.forward_native(x, residual)
|
||||
|
||||
if not getattr(self, "_is_compiled", False):
|
||||
self._forward_static_no_residual = torch.compile( # type: ignore
|
||||
self._forward_static_no_residual
|
||||
)
|
||||
self._forward_static_with_residual = torch.compile( # type: ignore
|
||||
self._forward_static_with_residual
|
||||
)
|
||||
self._is_compiled = True
|
||||
return self.forward_native(x, residual)
|
||||
|
||||
|
||||
|
||||
@@ -910,7 +910,15 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
|
||||
self.validate_shard_id(loaded_shard_id)
|
||||
if loaded_shard_id is None or isinstance(loaded_shard_id, tuple):
|
||||
if isinstance(param, PerTensorScaleParameter):
|
||||
param.load_merged_column_weight(loaded_weight=loaded_weight, shard_id=0)
|
||||
if isinstance(loaded_shard_id, tuple):
|
||||
for idx in loaded_shard_id:
|
||||
param.load_merged_column_weight(
|
||||
loaded_weight=loaded_weight, shard_id=idx
|
||||
)
|
||||
else:
|
||||
param.load_merged_column_weight(
|
||||
loaded_weight=loaded_weight, shard_id=0
|
||||
)
|
||||
return
|
||||
elif type(param) in (RowvLLMParameter, BasevLLMParameter):
|
||||
param.load_merged_column_weight(loaded_weight=loaded_weight)
|
||||
|
||||
-2541
File diff suppressed because it is too large
Load Diff
+10
@@ -0,0 +1,10 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
from vllm.model_executor.layers.quantization.compressed_tensors.compressed_tensors_moe.compressed_tensors_moe import ( # noqa: E501
|
||||
CompressedTensorsMoEMethod,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"CompressedTensorsMoEMethod",
|
||||
]
|
||||
+175
@@ -0,0 +1,175 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
|
||||
import torch
|
||||
from compressed_tensors import CompressionFormat
|
||||
from compressed_tensors.quantization import (
|
||||
ActivationOrdering,
|
||||
QuantizationStrategy,
|
||||
)
|
||||
|
||||
from vllm.logger import init_logger
|
||||
from vllm.model_executor.layers.fused_moe import (
|
||||
FusedMoEMethodBase,
|
||||
UnquantizedFusedMoEMethod,
|
||||
)
|
||||
from vllm.model_executor.layers.quantization.compressed_tensors.schemes.compressed_tensors_wNa16 import ( # noqa
|
||||
WNA16_SUPPORTED_BITS,
|
||||
)
|
||||
from vllm.model_executor.layers.quantization.utils.marlin_utils import (
|
||||
check_moe_marlin_supports_layer,
|
||||
)
|
||||
from vllm.platforms import current_platform
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class CompressedTensorsMoEMethod(FusedMoEMethodBase):
|
||||
@staticmethod
|
||||
def get_moe_method(
|
||||
quant_config: "CompressedTensorsConfig", # type: ignore # noqa E501
|
||||
layer: torch.nn.Module,
|
||||
layer_name: str,
|
||||
) -> FusedMoEMethodBase:
|
||||
# FusedMoE was made by combining multiple Linears so need to
|
||||
# make sure quantization config for Linear can target it
|
||||
quant_config._add_fused_moe_to_target_scheme_map()
|
||||
unfused_names = [
|
||||
layer_name + proj_name
|
||||
for proj_name in [".0.gate_proj", ".0.up_proj", ".0.down_proj"]
|
||||
]
|
||||
# TODO: refactor this to use expert_mapping and check all layer numbers
|
||||
all_scheme_dicts = [
|
||||
quant_config.get_scheme_dict(layer, name) for name in unfused_names
|
||||
]
|
||||
scheme_dict = all_scheme_dicts.pop()
|
||||
|
||||
# multiple schemes found
|
||||
if not all([cur_dict == scheme_dict for cur_dict in all_scheme_dicts]):
|
||||
raise ValueError(
|
||||
"All MoE projections need to have same "
|
||||
"quantization scheme but found multiple"
|
||||
)
|
||||
|
||||
if scheme_dict is None: # ignored layer
|
||||
return UnquantizedFusedMoEMethod(layer.moe_config)
|
||||
|
||||
# TODO: @dsikka: refactor this to use schemes as other kernels
|
||||
# are supported + check if the layer is being ignored.
|
||||
weight_quant = scheme_dict.get("weights")
|
||||
input_quant = scheme_dict.get("input_activations")
|
||||
format = scheme_dict.get("format")
|
||||
|
||||
if quant_config._is_mxfp4(weight_quant):
|
||||
from .compressed_tensors_moe_w4a4_mxfp4 import (
|
||||
CompressedTensorsW4A4Mxfp4MoEMethod,
|
||||
)
|
||||
|
||||
return CompressedTensorsW4A4Mxfp4MoEMethod(layer.moe_config)
|
||||
|
||||
if quant_config._is_wNa16_group_channel(weight_quant, input_quant):
|
||||
# group_size=None means channelwise
|
||||
group_size = weight_quant.group_size or -1
|
||||
|
||||
valid_format_and_bits = (
|
||||
weight_quant.num_bits in WNA16_SUPPORTED_BITS
|
||||
and format == CompressionFormat.pack_quantized.value
|
||||
)
|
||||
|
||||
if not valid_format_and_bits:
|
||||
raise ValueError(
|
||||
"For Fused MoE layers, only format: ",
|
||||
f"{CompressionFormat.pack_quantized.value} ",
|
||||
f" and bits: {WNA16_SUPPORTED_BITS} is supported ",
|
||||
f"but got format: {CompressionFormat.pack_quantized.value} "
|
||||
f" and bits: {weight_quant.num_bits}",
|
||||
)
|
||||
|
||||
# Prefer to use the MarlinMoE kernel when it is supported.
|
||||
if (
|
||||
not check_moe_marlin_supports_layer(layer, group_size)
|
||||
or current_platform.is_rocm()
|
||||
):
|
||||
from .compressed_tensors_moe_wna16 import (
|
||||
CompressedTensorsWNA16MoEMethod,
|
||||
)
|
||||
|
||||
if (
|
||||
weight_quant.strategy == QuantizationStrategy.GROUP
|
||||
and weight_quant.actorder
|
||||
in (ActivationOrdering.GROUP, ActivationOrdering.DYNAMIC)
|
||||
):
|
||||
raise ValueError(
|
||||
"WNA16MoE is not supported with actorder=group/dynamic."
|
||||
)
|
||||
logger.info_once("Using CompressedTensorsWNA16MoEMethod")
|
||||
return CompressedTensorsWNA16MoEMethod(
|
||||
weight_quant, input_quant, layer.moe_config
|
||||
)
|
||||
else:
|
||||
from .compressed_tensors_moe_wna16_marlin import (
|
||||
CompressedTensorsWNA16MarlinMoEMethod,
|
||||
)
|
||||
|
||||
logger.info_once("Using CompressedTensorsWNA16MarlinMoEMethod")
|
||||
return CompressedTensorsWNA16MarlinMoEMethod(
|
||||
weight_quant, input_quant, layer.moe_config
|
||||
)
|
||||
elif quant_config._is_nvfp4_format(weight_quant):
|
||||
from .compressed_tensors_moe_w4a4_nvfp4 import (
|
||||
CompressedTensorsW4A4Nvfp4MoEMethod,
|
||||
)
|
||||
|
||||
_is_valid_nvfp4_activations = (
|
||||
quant_config._is_nvfp4_format(input_quant) or input_quant is None
|
||||
)
|
||||
if not _is_valid_nvfp4_activations:
|
||||
raise ValueError(
|
||||
"For NVFP4 weights, input quantization must also be NVFP4 format ",
|
||||
f"or None for NVFP4A16, found {input_quant}",
|
||||
)
|
||||
return CompressedTensorsW4A4Nvfp4MoEMethod(
|
||||
layer.moe_config, layer_name, use_a16=(input_quant is None)
|
||||
)
|
||||
elif (
|
||||
quant_config._is_fp8_w8a8_sm90(weight_quant, input_quant)
|
||||
or quant_config._is_fp8_w8a8_sm100(weight_quant, input_quant)
|
||||
or quant_config._is_fp8_w8a8(weight_quant, input_quant)
|
||||
):
|
||||
from .compressed_tensors_moe_w8a8_fp8 import (
|
||||
CompressedTensorsW8A8Fp8MoEMethod,
|
||||
)
|
||||
|
||||
return CompressedTensorsW8A8Fp8MoEMethod(
|
||||
weight_quant, input_quant, layer.moe_config
|
||||
)
|
||||
elif quant_config._is_dynamic_token_w8a8(weight_quant, input_quant):
|
||||
from .compressed_tensors_moe_w8a8_int8 import (
|
||||
CompressedTensorsW8A8Int8MoEMethod,
|
||||
)
|
||||
|
||||
return CompressedTensorsW8A8Int8MoEMethod(
|
||||
weight_quant, input_quant, layer.moe_config
|
||||
)
|
||||
elif quant_config._is_fp8_w4a8_sm90(weight_quant, input_quant):
|
||||
from .compressed_tensors_moe_w4a8_fp8 import (
|
||||
CompressedTensorsW4A8Fp8MoEMethod,
|
||||
)
|
||||
|
||||
logger.info_once("Using CompressedTensorsW4A8Fp8MoEMethod")
|
||||
return CompressedTensorsW4A8Fp8MoEMethod(
|
||||
weight_quant, input_quant, layer.moe_config
|
||||
)
|
||||
elif quant_config._is_dynamic_token_w4a8_int(weight_quant, input_quant):
|
||||
from .compressed_tensors_moe_w4a8_int8 import (
|
||||
CompressedTensorsW4A8Int8MoEMethod,
|
||||
)
|
||||
|
||||
return CompressedTensorsW4A8Int8MoEMethod(
|
||||
weight_quant, input_quant, layer.moe_config
|
||||
)
|
||||
else:
|
||||
raise RuntimeError(
|
||||
f"Unsupported FusedMoe scheme: {weight_quant}, {input_quant}"
|
||||
)
|
||||
+168
@@ -0,0 +1,168 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
|
||||
import torch
|
||||
|
||||
from vllm.logger import init_logger
|
||||
from vllm.model_executor.layers.fused_moe import (
|
||||
FusedMoE,
|
||||
FusedMoeWeightScaleSupported,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.config import (
|
||||
FusedMoEQuantConfig,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.fused_marlin_moe import (
|
||||
MarlinExperts,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.oracle.mxfp4 import (
|
||||
Mxfp4MoeBackend,
|
||||
make_mxfp4_moe_kernel,
|
||||
make_mxfp4_moe_quant_config,
|
||||
)
|
||||
from vllm.model_executor.layers.quantization.compressed_tensors.compressed_tensors_moe import ( # noqa E501
|
||||
CompressedTensorsMoEMethod,
|
||||
)
|
||||
from vllm.model_executor.layers.quantization.utils.marlin_utils_fp4 import (
|
||||
prepare_moe_fp4_layer_for_marlin,
|
||||
)
|
||||
from vllm.model_executor.utils import set_weight_attrs
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class CompressedTensorsW4A4Mxfp4MoEMethod(CompressedTensorsMoEMethod):
|
||||
def __init__(self, moe):
|
||||
super().__init__(moe)
|
||||
self.group_size = 32
|
||||
self.mxfp4_backend = Mxfp4MoeBackend.MARLIN
|
||||
self.experts_cls = MarlinExperts
|
||||
|
||||
def create_weights(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
num_experts: int,
|
||||
hidden_size: int,
|
||||
intermediate_size_per_partition: int,
|
||||
params_dtype: torch.dtype,
|
||||
**extra_weight_attrs,
|
||||
):
|
||||
layer.num_experts = num_experts
|
||||
layer.params_dtype = params_dtype
|
||||
|
||||
w13_weight = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts,
|
||||
2 * intermediate_size_per_partition,
|
||||
# 2 fp4 items are packed in the input dimension
|
||||
hidden_size // 2,
|
||||
requires_grad=False,
|
||||
dtype=torch.uint8,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w13_weight_packed", w13_weight)
|
||||
set_weight_attrs(w13_weight, extra_weight_attrs)
|
||||
|
||||
w2_weight = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts,
|
||||
hidden_size,
|
||||
# 2 fp4 items are packed in the input dimension
|
||||
intermediate_size_per_partition // 2,
|
||||
dtype=torch.uint8,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w2_weight_packed", w2_weight)
|
||||
set_weight_attrs(w2_weight, extra_weight_attrs)
|
||||
|
||||
w13_weight_scale = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts,
|
||||
2 * intermediate_size_per_partition,
|
||||
# 2 fp4 items are packed in the input dimension
|
||||
hidden_size // self.group_size,
|
||||
dtype=torch.uint8,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w13_weight_scale", w13_weight_scale)
|
||||
extra_weight_attrs.update(
|
||||
{"quant_method": FusedMoeWeightScaleSupported.GROUP.value}
|
||||
)
|
||||
set_weight_attrs(w13_weight_scale, extra_weight_attrs)
|
||||
|
||||
w2_weight_scale = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts,
|
||||
hidden_size,
|
||||
# 2 fp4 items are packed in the input dimension
|
||||
intermediate_size_per_partition // self.group_size,
|
||||
dtype=torch.uint8,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w2_weight_scale", w2_weight_scale)
|
||||
set_weight_attrs(w2_weight_scale, extra_weight_attrs)
|
||||
|
||||
def get_fused_moe_quant_config(
|
||||
self, layer: torch.nn.Module
|
||||
) -> FusedMoEQuantConfig | None:
|
||||
return make_mxfp4_moe_quant_config(
|
||||
mxfp4_backend=self.mxfp4_backend,
|
||||
w1_scale=layer.w13_weight_scale,
|
||||
w2_scale=layer.w2_weight_scale,
|
||||
)
|
||||
|
||||
def process_weights_after_loading(self, layer: FusedMoE) -> None:
|
||||
layer.w13_weight = torch.nn.Parameter(
|
||||
layer.w13_weight_packed.data, requires_grad=False
|
||||
)
|
||||
delattr(layer, "w13_weight_packed")
|
||||
|
||||
layer.w2_weight = torch.nn.Parameter(
|
||||
layer.w2_weight_packed.data, requires_grad=False
|
||||
)
|
||||
delattr(layer, "w2_weight_packed")
|
||||
|
||||
logger.warning_once(
|
||||
"Your GPU does not have native support for FP4 computation but "
|
||||
"FP4 quantization is being used. Weight-only FP4 compression "
|
||||
"will be used leveraging the Marlin kernel. This may degrade "
|
||||
"performance for compute-heavy workloads."
|
||||
)
|
||||
prepare_moe_fp4_layer_for_marlin(layer)
|
||||
|
||||
self.moe_quant_config = self.get_fused_moe_quant_config(layer)
|
||||
if self.moe_quant_config is not None:
|
||||
self.moe_kernel = make_mxfp4_moe_kernel(
|
||||
moe_quant_config=self.moe_quant_config,
|
||||
moe_config=self.moe,
|
||||
experts_cls=self.experts_cls,
|
||||
mxfp4_backend=self.mxfp4_backend,
|
||||
shared_experts=layer.shared_experts,
|
||||
routing_tables=layer._maybe_init_expert_routing_tables(),
|
||||
)
|
||||
|
||||
def apply(
|
||||
self,
|
||||
layer: FusedMoE,
|
||||
x: torch.Tensor,
|
||||
topk_weights: torch.Tensor,
|
||||
topk_ids: torch.Tensor,
|
||||
shared_experts_input: torch.Tensor | None,
|
||||
) -> torch.Tensor:
|
||||
assert self.moe_kernel is not None
|
||||
return self.moe_kernel.apply(
|
||||
x,
|
||||
layer.w13_weight,
|
||||
layer.w2_weight,
|
||||
topk_weights,
|
||||
topk_ids,
|
||||
activation=layer.activation,
|
||||
global_num_experts=layer.global_num_experts,
|
||||
expert_map=layer.expert_map,
|
||||
apply_router_weight_on_input=layer.apply_router_weight_on_input,
|
||||
shared_experts_input=shared_experts_input,
|
||||
)
|
||||
+306
@@ -0,0 +1,306 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
|
||||
import torch
|
||||
|
||||
import vllm.model_executor.layers.fused_moe.modular_kernel as mk
|
||||
from vllm.logger import init_logger
|
||||
from vllm.model_executor.layers.fused_moe import (
|
||||
FusedMoE,
|
||||
FusedMoeWeightScaleSupported,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.config import (
|
||||
FusedMoEConfig,
|
||||
FusedMoEQuantConfig,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.oracle.nvfp4 import (
|
||||
convert_to_nvfp4_moe_kernel_format,
|
||||
is_global_sf_supported_for_nvfp4_backend,
|
||||
make_nvfp4_moe_kernel,
|
||||
make_nvfp4_moe_quant_config,
|
||||
select_nvfp4_moe_backend,
|
||||
)
|
||||
from vllm.model_executor.layers.quantization.compressed_tensors.compressed_tensors_moe import ( # noqa E501
|
||||
CompressedTensorsMoEMethod,
|
||||
)
|
||||
from vllm.model_executor.layers.quantization.utils.quant_utils import (
|
||||
kNvfp4Dynamic,
|
||||
kNvfp4Static,
|
||||
)
|
||||
from vllm.model_executor.utils import replace_parameter, set_weight_attrs
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class CompressedTensorsW4A4Nvfp4MoEMethod(CompressedTensorsMoEMethod):
|
||||
def __init__(
|
||||
self,
|
||||
moe: FusedMoEConfig,
|
||||
layer_name: str | None = None,
|
||||
use_a16: bool = False,
|
||||
):
|
||||
super().__init__(moe)
|
||||
self.group_size = 16
|
||||
|
||||
# Select experts implementation.
|
||||
self.nvfp4_backend, self.experts_cls = select_nvfp4_moe_backend(
|
||||
config=self.moe,
|
||||
weight_key=kNvfp4Static,
|
||||
activation_key=None if use_a16 else kNvfp4Dynamic,
|
||||
)
|
||||
|
||||
self.use_global_sf = is_global_sf_supported_for_nvfp4_backend(
|
||||
self.nvfp4_backend
|
||||
)
|
||||
|
||||
def create_weights(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
num_experts: int,
|
||||
hidden_size: int,
|
||||
intermediate_size_per_partition: int,
|
||||
params_dtype: torch.dtype,
|
||||
**extra_weight_attrs,
|
||||
):
|
||||
layer.num_experts = num_experts
|
||||
layer.params_dtype = params_dtype
|
||||
w13_num_shards = 2 if self.moe.is_act_and_mul else 1
|
||||
|
||||
w13_weight = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts,
|
||||
w13_num_shards * intermediate_size_per_partition,
|
||||
# 2 fp4 items are packed in the input dimension
|
||||
hidden_size // 2,
|
||||
requires_grad=False,
|
||||
dtype=torch.uint8,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w13_weight_packed", w13_weight)
|
||||
set_weight_attrs(w13_weight, extra_weight_attrs)
|
||||
|
||||
w2_weight = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts,
|
||||
hidden_size,
|
||||
# 2 fp4 items are packed in the input dimension
|
||||
intermediate_size_per_partition // 2,
|
||||
dtype=torch.uint8,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w2_weight_packed", w2_weight)
|
||||
set_weight_attrs(w2_weight, extra_weight_attrs)
|
||||
|
||||
# Weight Scales
|
||||
w13_weight_scale = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts,
|
||||
w13_num_shards * intermediate_size_per_partition,
|
||||
# 2 fp4 items are packed in the input dimension
|
||||
hidden_size // self.group_size,
|
||||
dtype=torch.float8_e4m3fn,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w13_weight_scale", w13_weight_scale)
|
||||
extra_weight_attrs.update(
|
||||
{"quant_method": FusedMoeWeightScaleSupported.GROUP.value}
|
||||
)
|
||||
set_weight_attrs(w13_weight_scale, extra_weight_attrs)
|
||||
|
||||
w2_weight_scale = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts,
|
||||
hidden_size,
|
||||
# 2 fp4 items are packed in the input dimension
|
||||
intermediate_size_per_partition // self.group_size,
|
||||
dtype=torch.float8_e4m3fn,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w2_weight_scale", w2_weight_scale)
|
||||
extra_weight_attrs.update(
|
||||
{"quant_method": FusedMoeWeightScaleSupported.GROUP.value}
|
||||
)
|
||||
set_weight_attrs(w2_weight_scale, extra_weight_attrs)
|
||||
|
||||
# Weight Global Scales
|
||||
w13_weight_scale_2 = torch.nn.Parameter(
|
||||
torch.empty(num_experts, w13_num_shards, dtype=torch.float32),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w13_weight_global_scale", w13_weight_scale_2)
|
||||
extra_weight_attrs.update(
|
||||
{"quant_method": FusedMoeWeightScaleSupported.TENSOR.value}
|
||||
)
|
||||
set_weight_attrs(w13_weight_scale_2, extra_weight_attrs)
|
||||
|
||||
w2_weight_scale_2 = torch.nn.Parameter(
|
||||
torch.empty(num_experts, dtype=torch.float32), requires_grad=False
|
||||
)
|
||||
layer.register_parameter("w2_weight_global_scale", w2_weight_scale_2)
|
||||
extra_weight_attrs.update(
|
||||
{"quant_method": FusedMoeWeightScaleSupported.TENSOR.value}
|
||||
)
|
||||
set_weight_attrs(w2_weight_scale_2, extra_weight_attrs)
|
||||
|
||||
# Input Global Scales
|
||||
w13_input_scale = torch.nn.Parameter(
|
||||
torch.empty(num_experts, w13_num_shards, dtype=torch.float32),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w13_input_global_scale", w13_input_scale)
|
||||
extra_weight_attrs.update(
|
||||
{"quant_method": FusedMoeWeightScaleSupported.TENSOR.value}
|
||||
)
|
||||
set_weight_attrs(w13_input_scale, extra_weight_attrs)
|
||||
|
||||
w2_input_scale = torch.nn.Parameter(
|
||||
torch.empty(num_experts, dtype=torch.float32), requires_grad=False
|
||||
)
|
||||
layer.register_parameter("w2_input_global_scale", w2_input_scale)
|
||||
extra_weight_attrs.update(
|
||||
{"quant_method": FusedMoeWeightScaleSupported.TENSOR.value}
|
||||
)
|
||||
set_weight_attrs(w2_input_scale, extra_weight_attrs)
|
||||
|
||||
def process_weights_after_loading(self, layer: FusedMoE) -> None:
|
||||
"""
|
||||
Convert NVFP4 MoE weights into kernel format and setup the kernel.
|
||||
"""
|
||||
# NOTE(rob): wN_weight_packed -> wN_weight is because ModularKernelMethod
|
||||
# requires this naming convention. However, the name change breaks
|
||||
# reloading because the state dict no longer matches disk. Once we
|
||||
# remove MKM, we should revert this change to ensure compatibility.
|
||||
layer.w13_weight = torch.nn.Parameter(
|
||||
layer.w13_weight_packed.data, requires_grad=False
|
||||
)
|
||||
delattr(layer, "w13_weight_packed")
|
||||
|
||||
layer.w2_weight = torch.nn.Parameter(
|
||||
layer.w2_weight_packed.data, requires_grad=False
|
||||
)
|
||||
delattr(layer, "w2_weight_packed")
|
||||
|
||||
# Use a single gscale for w13.
|
||||
if self.moe.is_act_and_mul and not torch.allclose(
|
||||
layer.w13_weight_global_scale[:, 0], layer.w13_weight_global_scale[:, 1]
|
||||
):
|
||||
logger.warning_once(
|
||||
"w1_weight_global_scale must match w3_weight_global_scale. "
|
||||
"Accuracy may be affected.",
|
||||
)
|
||||
w13_weight_global_scale = layer.w13_weight_global_scale[:, 0].contiguous()
|
||||
|
||||
# Shuffle weights into the NvFp4 kernel format.
|
||||
(
|
||||
w13,
|
||||
w13_scale,
|
||||
w13_scale_2,
|
||||
a13_scale,
|
||||
w2,
|
||||
w2_scale,
|
||||
w2_scale_2,
|
||||
a2_scale,
|
||||
) = convert_to_nvfp4_moe_kernel_format(
|
||||
nvfp4_backend=self.nvfp4_backend,
|
||||
layer=layer,
|
||||
w13=layer.w13_weight,
|
||||
w13_scale=layer.w13_weight_scale,
|
||||
w13_scale_2=(1.0 / w13_weight_global_scale),
|
||||
a13_scale=(1.0 / layer.w13_input_global_scale),
|
||||
w2=layer.w2_weight,
|
||||
w2_scale=layer.w2_weight_scale,
|
||||
w2_scale_2=(1.0 / layer.w2_weight_global_scale),
|
||||
a2_scale=(1.0 / layer.w2_input_global_scale),
|
||||
is_act_and_mul=self.moe.is_act_and_mul,
|
||||
)
|
||||
|
||||
replace_parameter(layer, "w13_weight", w13)
|
||||
replace_parameter(layer, "w13_weight_scale", w13_scale)
|
||||
replace_parameter(layer, "w2_weight", w2)
|
||||
replace_parameter(layer, "w2_weight_scale", w2_scale)
|
||||
layer.w13_weight_scale_2 = w13_scale_2
|
||||
layer.w2_weight_scale_2 = w2_scale_2
|
||||
layer.w13_input_scale = a13_scale
|
||||
layer.w2_input_scale = a2_scale
|
||||
|
||||
# Setup modular kernel.
|
||||
self.moe_quant_config = self.get_fused_moe_quant_config(layer)
|
||||
assert self.experts_cls is not None
|
||||
self.moe_kernel = make_nvfp4_moe_kernel(
|
||||
moe_quant_config=self.moe_quant_config,
|
||||
moe_config=self.moe,
|
||||
experts_cls=self.experts_cls,
|
||||
shared_experts=layer.shared_experts,
|
||||
routing_tables=layer._maybe_init_expert_routing_tables(),
|
||||
)
|
||||
self.moe_kernel.fused_experts.process_weights_after_loading(layer)
|
||||
|
||||
def maybe_make_prepare_finalize(
|
||||
self,
|
||||
routing_tables: tuple[torch.Tensor, torch.Tensor, torch.Tensor] | None = None,
|
||||
) -> mk.FusedMoEPrepareAndFinalizeModular | None:
|
||||
raise ValueError(
|
||||
f"{self.__class__.__name__} uses the new modular kernel initialization "
|
||||
"logic. This function should not be called."
|
||||
)
|
||||
|
||||
def get_fused_moe_quant_config(self, layer: torch.nn.Module) -> FusedMoEQuantConfig:
|
||||
return make_nvfp4_moe_quant_config(
|
||||
backend=self.nvfp4_backend,
|
||||
w13_scale=layer.w13_weight_scale,
|
||||
w2_scale=layer.w2_weight_scale,
|
||||
w13_scale_2=layer.w13_weight_scale_2,
|
||||
w2_scale_2=layer.w2_weight_scale_2,
|
||||
a13_scale=layer.w13_input_scale,
|
||||
a2_scale=layer.w2_input_scale,
|
||||
)
|
||||
|
||||
def apply_monolithic(
|
||||
self,
|
||||
layer: FusedMoE,
|
||||
x: torch.Tensor,
|
||||
router_logits: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
assert self.is_monolithic
|
||||
assert self.moe_kernel is not None
|
||||
return self.moe_kernel.apply_monolithic(
|
||||
x,
|
||||
layer.w13_weight,
|
||||
layer.w2_weight,
|
||||
router_logits,
|
||||
activation=layer.activation,
|
||||
global_num_experts=layer.global_num_experts,
|
||||
expert_map=layer.expert_map,
|
||||
apply_router_weight_on_input=layer.apply_router_weight_on_input,
|
||||
num_expert_group=layer.num_expert_group,
|
||||
topk_group=layer.topk_group,
|
||||
e_score_correction_bias=layer.e_score_correction_bias,
|
||||
routed_scaling_factor=layer.routed_scaling_factor,
|
||||
)
|
||||
|
||||
def apply(
|
||||
self,
|
||||
layer: FusedMoE,
|
||||
x: torch.Tensor,
|
||||
topk_weights: torch.Tensor,
|
||||
topk_ids: torch.Tensor,
|
||||
shared_experts_input: torch.Tensor | None,
|
||||
) -> torch.Tensor:
|
||||
assert self.moe_kernel is not None
|
||||
return self.moe_kernel.apply(
|
||||
x,
|
||||
layer.w13_weight,
|
||||
layer.w2_weight,
|
||||
topk_weights,
|
||||
topk_ids,
|
||||
activation=layer.activation,
|
||||
global_num_experts=layer.global_num_experts,
|
||||
expert_map=layer.expert_map,
|
||||
apply_router_weight_on_input=layer.apply_router_weight_on_input,
|
||||
shared_experts_input=shared_experts_input,
|
||||
)
|
||||
+343
@@ -0,0 +1,343 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
|
||||
import torch
|
||||
from compressed_tensors.quantization import (
|
||||
QuantizationArgs,
|
||||
)
|
||||
|
||||
import vllm.model_executor.layers.fused_moe.modular_kernel as mk
|
||||
from vllm import _custom_ops as ops
|
||||
from vllm.logger import init_logger
|
||||
from vllm.model_executor.layers.fused_moe import (
|
||||
FusedMoE,
|
||||
FusedMoEActivationFormat,
|
||||
FusedMoEExpertsModular,
|
||||
FusedMoeWeightScaleSupported,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.config import (
|
||||
FusedMoEConfig,
|
||||
FusedMoEQuantConfig,
|
||||
int4_w4afp8_moe_quant_config,
|
||||
)
|
||||
from vllm.model_executor.layers.quantization.compressed_tensors.compressed_tensors_moe import ( # noqa E501
|
||||
CompressedTensorsMoEMethod,
|
||||
)
|
||||
from vllm.model_executor.layers.quantization.utils.quant_utils import (
|
||||
convert_bf16_scales_to_fp8,
|
||||
convert_packed_uint4b8_to_signed_int4_inplace,
|
||||
)
|
||||
from vllm.model_executor.utils import replace_parameter, set_weight_attrs
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class CompressedTensorsW4A8Fp8MoEMethod(CompressedTensorsMoEMethod):
|
||||
def __init__(
|
||||
self,
|
||||
weight_quant: QuantizationArgs,
|
||||
input_quant: QuantizationArgs,
|
||||
moe: FusedMoEConfig,
|
||||
layer_name: str | None = None,
|
||||
):
|
||||
super().__init__(moe)
|
||||
self.weight_quant = weight_quant
|
||||
self.input_quant = input_quant
|
||||
|
||||
self.group_size = self.weight_quant.group_size
|
||||
self.num_bits = self.weight_quant.num_bits
|
||||
self.packed_factor = 32 // self.num_bits
|
||||
|
||||
assert self.weight_quant.symmetric, (
|
||||
"Only symmetric quantization is supported for W4A8 MoE"
|
||||
)
|
||||
assert self.weight_quant.actorder != "group"
|
||||
assert self.group_size == 128, "Only group size 128 supported for W4A8 MoE"
|
||||
|
||||
self.disable_expert_map = False
|
||||
self.layer_name = layer_name
|
||||
|
||||
from vllm.model_executor.layers.quantization.input_quant_fp8 import QuantFP8
|
||||
from vllm.model_executor.layers.quantization.utils.quant_utils import (
|
||||
GroupShape,
|
||||
)
|
||||
|
||||
self.quant_fp8 = QuantFP8(static=False, group_shape=GroupShape.PER_TOKEN)
|
||||
|
||||
def create_weights(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
num_experts: int,
|
||||
hidden_size: int,
|
||||
intermediate_size_per_partition: int,
|
||||
params_dtype: torch.dtype,
|
||||
**extra_weight_attrs,
|
||||
):
|
||||
layer.num_experts = num_experts
|
||||
layer.orig_dtype = params_dtype
|
||||
layer.weight_block_size = None
|
||||
|
||||
# requirement for CUTLASS reorder_tensor
|
||||
assert hidden_size % 256 == 0, f"{hidden_size=} must be divisible by 256"
|
||||
assert intermediate_size_per_partition % 256 == 0, (
|
||||
f"{intermediate_size_per_partition=} must be divisible by 256"
|
||||
)
|
||||
# storage type, pack 8xint4 into int32
|
||||
params_dtype = torch.int32
|
||||
|
||||
# WEIGHTS
|
||||
w13_weight_packed = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts,
|
||||
2 * intermediate_size_per_partition,
|
||||
hidden_size // self.packed_factor,
|
||||
dtype=params_dtype,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w13_weight_packed", w13_weight_packed)
|
||||
set_weight_attrs(w13_weight_packed, extra_weight_attrs)
|
||||
|
||||
w2_weight_packed = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts,
|
||||
hidden_size,
|
||||
intermediate_size_per_partition // self.packed_factor,
|
||||
dtype=params_dtype,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w2_weight_packed", w2_weight_packed)
|
||||
set_weight_attrs(w2_weight_packed, extra_weight_attrs)
|
||||
|
||||
# SCALES
|
||||
# weight_scale refers to the group-wise scales
|
||||
# they are initially loaded as bf16, we will convert to fp8
|
||||
# after loading
|
||||
w13_weight_scale = torch.nn.Parameter(
|
||||
torch.ones(
|
||||
num_experts,
|
||||
2 * intermediate_size_per_partition,
|
||||
hidden_size // self.group_size,
|
||||
dtype=layer.orig_dtype,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w13_weight_scale", w13_weight_scale)
|
||||
|
||||
w2_weight_scale = torch.nn.Parameter(
|
||||
torch.ones(
|
||||
num_experts,
|
||||
hidden_size,
|
||||
intermediate_size_per_partition // self.group_size,
|
||||
dtype=layer.orig_dtype,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w2_weight_scale", w2_weight_scale)
|
||||
# Add PER-GROUP quantization for FusedMoE.weight_loader.
|
||||
extra_weight_attrs.update(
|
||||
{"quant_method": FusedMoeWeightScaleSupported.GROUP.value}
|
||||
)
|
||||
set_weight_attrs(w13_weight_scale, extra_weight_attrs)
|
||||
set_weight_attrs(w2_weight_scale, extra_weight_attrs)
|
||||
|
||||
# weight shapes
|
||||
w2_weight_shape = torch.nn.Parameter(
|
||||
torch.empty(num_experts, 2), requires_grad=False
|
||||
)
|
||||
layer.register_parameter("w2_weight_shape", w2_weight_shape)
|
||||
set_weight_attrs(w2_weight_shape, extra_weight_attrs)
|
||||
w13_weight_shape = torch.nn.Parameter(
|
||||
torch.empty(num_experts, 2), requires_grad=False
|
||||
)
|
||||
layer.register_parameter("w13_weight_shape", w13_weight_shape)
|
||||
set_weight_attrs(w13_weight_shape, extra_weight_attrs)
|
||||
|
||||
# don't use input scales
|
||||
layer.w13_input_scale = None
|
||||
layer.w2_input_scale = None
|
||||
|
||||
def process_weights_after_loading(self, layer):
|
||||
device = layer.w13_weight_packed.device
|
||||
|
||||
# STRIDES
|
||||
# A, C
|
||||
self.a_strides1_c_strides2 = torch.full(
|
||||
(layer.local_num_experts,),
|
||||
layer.hidden_size,
|
||||
device=device,
|
||||
dtype=torch.int64,
|
||||
)
|
||||
self.a_strides2 = torch.full(
|
||||
(layer.local_num_experts,),
|
||||
layer.intermediate_size_per_partition,
|
||||
device=device,
|
||||
dtype=torch.int64,
|
||||
)
|
||||
self.c_strides1 = torch.full(
|
||||
(layer.local_num_experts,),
|
||||
2 * layer.intermediate_size_per_partition,
|
||||
device=device,
|
||||
dtype=torch.int64,
|
||||
)
|
||||
|
||||
# S (group-wise scales)
|
||||
# sizeof(StrideS) = 16 bytes, so we need to use 2xint64 to encode it
|
||||
self.s_strides1 = torch.zeros(
|
||||
(layer.local_num_experts, 2), device=device, dtype=torch.int64
|
||||
)
|
||||
self.s_strides1[:, 0] = 2 * layer.intermediate_size_per_partition
|
||||
|
||||
self.s_strides2 = torch.zeros(
|
||||
(layer.local_num_experts, 2), device=device, dtype=torch.int64
|
||||
)
|
||||
self.s_strides2[:, 0] = layer.hidden_size
|
||||
|
||||
# encode and reorder weight tensors, and get the layout to pass to
|
||||
# the grouped gemm kernel. `b_strides1/2` specifies the entire layout
|
||||
convert_packed_uint4b8_to_signed_int4_inplace(layer.w13_weight_packed)
|
||||
w13_weight_shuffled, self.b_strides1 = (
|
||||
ops.cutlass_encode_and_reorder_int4b_grouped(layer.w13_weight_packed)
|
||||
)
|
||||
replace_parameter(layer, "w13_weight_packed", w13_weight_shuffled)
|
||||
convert_packed_uint4b8_to_signed_int4_inplace(layer.w2_weight_packed)
|
||||
w2_weight_shuffled, self.b_strides2 = (
|
||||
ops.cutlass_encode_and_reorder_int4b_grouped(layer.w2_weight_packed)
|
||||
)
|
||||
replace_parameter(layer, "w2_weight_packed", w2_weight_shuffled)
|
||||
|
||||
# convert bf16 scales to (fp8_scales, channel_scales)
|
||||
w13_weight_scale, w13_weight_chan_scale = convert_bf16_scales_to_fp8(
|
||||
self.quant_fp8, layer.w13_weight_scale
|
||||
)
|
||||
w2_weight_scale, w2_weight_chan_scale = convert_bf16_scales_to_fp8(
|
||||
self.quant_fp8, layer.w2_weight_scale
|
||||
)
|
||||
|
||||
# register channel scales
|
||||
layer.register_parameter(
|
||||
"w13_weight_chan_scale",
|
||||
torch.nn.Parameter(w13_weight_chan_scale, requires_grad=False),
|
||||
)
|
||||
layer.register_parameter(
|
||||
"w2_weight_chan_scale",
|
||||
torch.nn.Parameter(w2_weight_chan_scale, requires_grad=False),
|
||||
)
|
||||
|
||||
# The scales are stored as (E, N, K // 128) but the kernel expects
|
||||
# (E, K // 128, N) in row-major format, so we need to permute the last 2 dims
|
||||
# and make it contiguous
|
||||
w13_weight_scale_packed = ops.cutlass_pack_scale_fp8(
|
||||
w13_weight_scale.permute(0, 2, 1).contiguous()
|
||||
)
|
||||
replace_parameter(layer, "w13_weight_scale", w13_weight_scale_packed)
|
||||
w2_weight_scale_packed = ops.cutlass_pack_scale_fp8(
|
||||
w2_weight_scale.permute(0, 2, 1).contiguous()
|
||||
)
|
||||
replace_parameter(layer, "w2_weight_scale", w2_weight_scale_packed)
|
||||
|
||||
def maybe_make_prepare_finalize(
|
||||
self,
|
||||
routing_tables: tuple[torch.Tensor, torch.Tensor, torch.Tensor] | None = None,
|
||||
) -> mk.FusedMoEPrepareAndFinalizeModular | None:
|
||||
return super().maybe_make_prepare_finalize(routing_tables)
|
||||
|
||||
def get_fused_moe_quant_config(
|
||||
self, layer: torch.nn.Module
|
||||
) -> FusedMoEQuantConfig | None:
|
||||
# Store quantization scales; both per-group and per-channel
|
||||
# Note we haven't specified the group size here because
|
||||
# the quant config logic assumes group-wise scaling
|
||||
# and channel-wise scaling are exclusive.
|
||||
return int4_w4afp8_moe_quant_config(
|
||||
w1_scale=layer.w13_weight_scale, # group scale
|
||||
w2_scale=layer.w2_weight_scale, # group scale
|
||||
g1_alphas=layer.w13_weight_chan_scale,
|
||||
g2_alphas=layer.w2_weight_chan_scale,
|
||||
per_act_token_quant=True, # always use dynamic per-token
|
||||
per_out_ch_quant=True, # always use per-channel
|
||||
)
|
||||
|
||||
def select_gemm_impl(
|
||||
self,
|
||||
prepare_finalize: mk.FusedMoEPrepareAndFinalizeModular,
|
||||
layer: torch.nn.Module,
|
||||
) -> mk.FusedMoEExpertsModular:
|
||||
assert self.moe_quant_config is not None
|
||||
assert (
|
||||
prepare_finalize.activation_format == FusedMoEActivationFormat.Standard
|
||||
), "BatchedExperts not supported"
|
||||
|
||||
from vllm.model_executor.layers.fused_moe import CutlassExpertsW4A8Fp8
|
||||
|
||||
experts: FusedMoEExpertsModular
|
||||
|
||||
logger.debug("CutlassExpertsW4A8Fp8(%s)", self.__class__.__name__)
|
||||
experts = CutlassExpertsW4A8Fp8(
|
||||
out_dtype=self.moe.in_dtype,
|
||||
a_strides1=self.a_strides1_c_strides2,
|
||||
a_strides2=self.a_strides2,
|
||||
b_strides1=self.b_strides1,
|
||||
b_strides2=self.b_strides2,
|
||||
c_strides1=self.c_strides1,
|
||||
c_strides2=self.a_strides1_c_strides2,
|
||||
s_strides1=self.s_strides1,
|
||||
s_strides2=self.s_strides2,
|
||||
moe_config=self.moe,
|
||||
quant_config=self.moe_quant_config,
|
||||
group_size=self.group_size,
|
||||
)
|
||||
|
||||
num_dispatchers = prepare_finalize.num_dispatchers()
|
||||
self.disable_expert_map = (
|
||||
num_dispatchers > 1 or not experts.supports_expert_map()
|
||||
)
|
||||
|
||||
return experts
|
||||
|
||||
def apply(
|
||||
self,
|
||||
layer: FusedMoE,
|
||||
x: torch.Tensor,
|
||||
topk_weights: torch.Tensor,
|
||||
topk_ids: torch.Tensor,
|
||||
shared_experts_input: torch.Tensor | None,
|
||||
) -> torch.Tensor:
|
||||
if layer.enable_eplb:
|
||||
raise NotImplementedError(
|
||||
"EPLB not supported for `CompressedTensorsW4A8Fp8MoEMethod` yet."
|
||||
)
|
||||
assert self.moe_quant_config is not None
|
||||
|
||||
from vllm.model_executor.layers.fused_moe.cutlass_moe import (
|
||||
cutlass_moe_w4a8_fp8,
|
||||
)
|
||||
|
||||
return cutlass_moe_w4a8_fp8(
|
||||
x,
|
||||
layer.w13_weight_packed,
|
||||
layer.w2_weight_packed,
|
||||
topk_weights,
|
||||
topk_ids,
|
||||
moe_config=self.moe,
|
||||
quant_config=self.moe_quant_config,
|
||||
activation=layer.activation,
|
||||
global_num_experts=layer.global_num_experts,
|
||||
expert_map=None if self.disable_expert_map else layer.expert_map,
|
||||
a_strides1=self.a_strides1_c_strides2,
|
||||
a_strides2=self.a_strides2,
|
||||
b_strides1=self.b_strides1,
|
||||
b_strides2=self.b_strides2,
|
||||
c_strides1=self.c_strides1,
|
||||
c_strides2=self.a_strides1_c_strides2,
|
||||
s_strides1=self.s_strides1,
|
||||
s_strides2=self.s_strides2,
|
||||
group_size=self.group_size,
|
||||
apply_router_weight_on_input=layer.apply_router_weight_on_input,
|
||||
)
|
||||
|
||||
@property
|
||||
def supports_eplb(self) -> bool:
|
||||
return False
|
||||
+349
@@ -0,0 +1,349 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
|
||||
import torch
|
||||
from compressed_tensors.quantization import (
|
||||
QuantizationArgs,
|
||||
QuantizationStrategy,
|
||||
)
|
||||
|
||||
from vllm.logger import init_logger
|
||||
from vllm.model_executor.layers.fused_moe import (
|
||||
FusedMoE,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.activation import MoEActivation
|
||||
from vllm.model_executor.layers.fused_moe.config import (
|
||||
FusedMoEConfig,
|
||||
FusedMoEQuantConfig,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.cpu_fused_moe import select_experts
|
||||
from vllm.model_executor.layers.quantization.compressed_tensors.compressed_tensors_moe import ( # noqa E501
|
||||
CompressedTensorsMoEMethod,
|
||||
)
|
||||
from vllm.model_executor.utils import replace_parameter, set_weight_attrs
|
||||
from vllm.platforms import CpuArchEnum, current_platform
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class CompressedTensorsW4A8Int8MoEMethod(CompressedTensorsMoEMethod):
|
||||
"""
|
||||
CPU-only MoE method using dynamic 4-bit matmul kernels on Arm Platform
|
||||
- Weights: int4 (stored as int8 values in [-8,7], packed to uint8 nibbles)
|
||||
- Scales: Fp32 for Channelwise , bf16 for groupwise quantization
|
||||
- Bias: Same data type as original weights
|
||||
- Activations: FP32/Bf16 dynamic per-token (A8 Int),
|
||||
quantized inside the kernel
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
weight_quant: QuantizationArgs,
|
||||
input_quant: QuantizationArgs,
|
||||
moe: FusedMoEConfig,
|
||||
layer_name: str | None = None,
|
||||
):
|
||||
super().__init__(moe)
|
||||
self.has_bias = self.moe.has_bias
|
||||
self.weight_quant = weight_quant
|
||||
self.input_quant = input_quant
|
||||
|
||||
# Validate scheme: weights=W4 (channel or group),
|
||||
# activations=dynamic TOKEN (A8)
|
||||
|
||||
# Must be dynamic per-token activations
|
||||
if (
|
||||
input_quant.strategy != QuantizationStrategy.TOKEN
|
||||
or not input_quant.dynamic
|
||||
):
|
||||
raise ValueError(
|
||||
"W4A8-int MoE needs dynamic per-token activation quantization."
|
||||
)
|
||||
|
||||
# Weight can be channel-wise (group_size=None) or group-wise
|
||||
self.group_size = (
|
||||
weight_quant.group_size if (weight_quant.group_size is not None) else -1
|
||||
)
|
||||
if weight_quant.num_bits != 4:
|
||||
raise ValueError("This method only supports 4-bit weights (num_bits=4).")
|
||||
|
||||
# CPU only
|
||||
if not current_platform.is_cpu():
|
||||
raise ValueError("CompressedTensorsW4A8Int8MoEMethod is CPU-only.")
|
||||
|
||||
# Arm: check _dyn ops availability
|
||||
if current_platform.get_cpu_architecture() == CpuArchEnum.ARM:
|
||||
try:
|
||||
_ = torch.ops.aten._dyn_quant_matmul_4bit
|
||||
_ = torch.ops.aten._dyn_quant_pack_4bit_weight
|
||||
except AttributeError as err:
|
||||
raise RuntimeError(
|
||||
f"""PyTorch {torch.__version__} lacks _dyn_quant_* 4bit ops;
|
||||
install a newer build."""
|
||||
) from err
|
||||
self.static_input_scales = False # always dynamic per token
|
||||
|
||||
# ---- parameter creation ----
|
||||
def create_weights(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
num_experts: int,
|
||||
hidden_size: int,
|
||||
intermediate_size_per_partition: int,
|
||||
params_dtype: torch.dtype,
|
||||
**extra_weight_attrs,
|
||||
):
|
||||
# Shapes per local rank (TP/EP):
|
||||
# w13: [E, 2*I_local, H] int8 (int4 values in [-8,7])
|
||||
# w2 : [E, H, I_local] int8
|
||||
# Scales:
|
||||
# channel-wise: group_size=-1 -> per-output-row, single scale per row
|
||||
# group-wise : group_size=g ->
|
||||
# per-output-row, (in_features/g) scales
|
||||
|
||||
E = num_experts
|
||||
H = hidden_size
|
||||
IN = intermediate_size_per_partition
|
||||
g = self.group_size
|
||||
|
||||
# Per-row scale columns
|
||||
def _n_scale_cols(in_features: int) -> int:
|
||||
return 1 if g == -1 else (in_features // g)
|
||||
|
||||
# Register unpacked int4-as-int8 weights the loader will fill.
|
||||
w13 = torch.nn.Parameter(
|
||||
torch.empty(E, 2 * IN, H, dtype=torch.int8), requires_grad=False
|
||||
)
|
||||
set_weight_attrs(w13, extra_weight_attrs)
|
||||
layer.register_parameter("w13_weight", w13)
|
||||
|
||||
w2 = torch.nn.Parameter(
|
||||
torch.empty(E, H, IN, dtype=torch.int8), requires_grad=False
|
||||
)
|
||||
set_weight_attrs(w2, extra_weight_attrs)
|
||||
layer.register_parameter("w2_weight", w2)
|
||||
|
||||
# Register scales
|
||||
# KleidiAI groupwise kernels accepts float32 scales
|
||||
# KleidiAI groupwise kernels accepts bfloat16 scales
|
||||
scale_dtype = torch.float32 if g == -1 else torch.bfloat16
|
||||
|
||||
w13_s = torch.nn.Parameter(
|
||||
torch.ones(E, 2 * IN, _n_scale_cols(H), dtype=scale_dtype),
|
||||
requires_grad=False,
|
||||
)
|
||||
set_weight_attrs(
|
||||
w13_s,
|
||||
{"quant_method": "channel" if g == -1 else "group", **extra_weight_attrs},
|
||||
)
|
||||
layer.register_parameter("w13_weight_scale", w13_s)
|
||||
|
||||
w2_s = torch.nn.Parameter(
|
||||
torch.ones(E, H, _n_scale_cols(IN), dtype=scale_dtype), requires_grad=False
|
||||
)
|
||||
set_weight_attrs(
|
||||
w2_s,
|
||||
{"quant_method": "channel" if g == -1 else "group", **extra_weight_attrs},
|
||||
)
|
||||
layer.register_parameter("w2_weight_scale", w2_s)
|
||||
|
||||
if self.has_bias:
|
||||
w13_bias = torch.nn.Parameter(
|
||||
torch.zeros(E, 2 * IN, dtype=params_dtype), requires_grad=False
|
||||
)
|
||||
layer.register_parameter("w13_bias", w13_bias)
|
||||
set_weight_attrs(w13_bias, extra_weight_attrs)
|
||||
|
||||
w2_bias = torch.nn.Parameter(
|
||||
torch.zeros(num_experts, hidden_size, dtype=params_dtype),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w2_bias", w2_bias)
|
||||
set_weight_attrs(w2_bias, extra_weight_attrs)
|
||||
|
||||
# Placeholders for packed weights (will be replaced after packing)
|
||||
layer.register_parameter(
|
||||
"w13_weight_packed", torch.nn.Parameter(torch.empty(0), requires_grad=False)
|
||||
)
|
||||
set_weight_attrs(layer.w13_weight_packed, extra_weight_attrs)
|
||||
|
||||
layer.register_parameter(
|
||||
"w2_weight_packed", torch.nn.Parameter(torch.empty(0), requires_grad=False)
|
||||
)
|
||||
set_weight_attrs(layer.w2_weight_packed, extra_weight_attrs)
|
||||
|
||||
# dims for 4 bit fused matmuls
|
||||
layer.w13_in_features = H
|
||||
layer.w13_out_features = 2 * IN
|
||||
layer.w2_in_features = IN
|
||||
layer.w2_out_features = H
|
||||
layer.group_size = g
|
||||
|
||||
# post-load packing to dyn-4bit KleidiAI kernel's format
|
||||
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
|
||||
E = layer.w13_weight.shape[0]
|
||||
H = layer.w13_in_features
|
||||
I2 = layer.w13_out_features
|
||||
IN = layer.w2_in_features
|
||||
g = layer.group_size
|
||||
|
||||
def _pack_matrix(
|
||||
int4_as_int8_2d: torch.Tensor,
|
||||
scales_2d: torch.Tensor,
|
||||
bias_1d: torch.Tensor | None,
|
||||
in_features: int,
|
||||
out_features: int,
|
||||
) -> torch.Tensor:
|
||||
# int4 values are stored as int8 in [-8,7].
|
||||
# Shift to unsigned nibble and pack pairs along input-dim.
|
||||
tmp = int4_as_int8_2d.add(8) # [out, in]
|
||||
uint8_nibbles = ((tmp[:, 1::2] << 4) | tmp[:, ::2]).to(
|
||||
torch.uint8
|
||||
) # [out, in//2]
|
||||
|
||||
# KleidiAI groupwise kernels accepts float32 scales
|
||||
# KleidiAI groupwise kernels accepts bfloat16 scales
|
||||
scale_dtype = torch.float32 if g == -1 else torch.bfloat16
|
||||
scales = scales_2d.to(scale_dtype)
|
||||
bias = None if bias_1d is None else bias_1d.to(torch.float32)
|
||||
return torch.ops.aten._dyn_quant_pack_4bit_weight(
|
||||
uint8_nibbles,
|
||||
scales,
|
||||
bias,
|
||||
g if g != -1 else in_features,
|
||||
in_features,
|
||||
out_features,
|
||||
)
|
||||
|
||||
# Pack per expert
|
||||
w13_packed_list = []
|
||||
w2_packed_list = []
|
||||
|
||||
has_w13_bias = hasattr(layer, "w13_bias") and layer.w13_bias is not None
|
||||
has_w2_bias = hasattr(layer, "w2_bias") and layer.w2_bias is not None
|
||||
|
||||
for e in range(E):
|
||||
w13_packed_list.append(
|
||||
_pack_matrix(
|
||||
layer.w13_weight[e], # [2I, H]
|
||||
layer.w13_weight_scale[e], # [2I, H/g or 1]
|
||||
layer.w13_bias[e] if has_w13_bias else None, # [2I]
|
||||
H,
|
||||
I2,
|
||||
)
|
||||
)
|
||||
w2_packed_list.append(
|
||||
_pack_matrix(
|
||||
# w2 shape is [H, IN]; we need [out, in] == [H, IN].
|
||||
layer.w2_weight[e], # [H, IN]
|
||||
layer.w2_weight_scale[e], # [H, IN/g or 1]
|
||||
layer.w2_bias[e] if has_w2_bias else None, # [H]
|
||||
IN,
|
||||
layer.w2_out_features, # in_features=IN, out_features=H
|
||||
)
|
||||
)
|
||||
|
||||
# each packed tensor has identical shape per expert; stack on dim 0
|
||||
w13_packed = torch.stack(w13_packed_list, dim=0)
|
||||
w2_packed = torch.stack(w2_packed_list, dim=0)
|
||||
|
||||
replace_parameter(
|
||||
layer,
|
||||
"w13_weight_packed",
|
||||
torch.nn.Parameter(w13_packed, requires_grad=False),
|
||||
)
|
||||
replace_parameter(
|
||||
layer,
|
||||
"w2_weight_packed",
|
||||
torch.nn.Parameter(w2_packed, requires_grad=False),
|
||||
)
|
||||
|
||||
# free raw tensors/scales/bias now that they're packed into the payload.
|
||||
replace_parameter(
|
||||
layer, "w13_weight", torch.nn.Parameter(torch.empty(0), requires_grad=False)
|
||||
)
|
||||
replace_parameter(
|
||||
layer, "w2_weight", torch.nn.Parameter(torch.empty(0), requires_grad=False)
|
||||
)
|
||||
replace_parameter(
|
||||
layer,
|
||||
"w13_weight_scale",
|
||||
torch.nn.Parameter(torch.empty(0), requires_grad=False),
|
||||
)
|
||||
replace_parameter(
|
||||
layer,
|
||||
"w2_weight_scale",
|
||||
torch.nn.Parameter(torch.empty(0), requires_grad=False),
|
||||
)
|
||||
if has_w13_bias:
|
||||
replace_parameter(
|
||||
layer,
|
||||
"w13_bias",
|
||||
torch.nn.Parameter(torch.empty(0), requires_grad=False),
|
||||
)
|
||||
if has_w2_bias:
|
||||
replace_parameter(
|
||||
layer,
|
||||
"w2_bias",
|
||||
torch.nn.Parameter(torch.empty(0), requires_grad=False),
|
||||
)
|
||||
|
||||
def get_fused_moe_quant_config(
|
||||
self, layer: torch.nn.Module
|
||||
) -> FusedMoEQuantConfig | None:
|
||||
# CPU dynamic 4-bit MoE path does not use modular kernels or
|
||||
# fused_experts; quant config is not needed.
|
||||
return None
|
||||
|
||||
@property
|
||||
def is_monolithic(self) -> bool:
|
||||
return True
|
||||
|
||||
def apply_monolithic(
|
||||
self,
|
||||
layer: FusedMoE,
|
||||
x: torch.Tensor,
|
||||
router_logits: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
assert not layer.enable_eplb, "EPLB not supported for W4A8-int MoE yet."
|
||||
assert layer.activation in (
|
||||
MoEActivation.SILU,
|
||||
MoEActivation.SWIGLUOAI,
|
||||
MoEActivation.SWIGLUSTEP,
|
||||
), "Only SiLU/SwiGLUGU/SwiGLUUG are supported."
|
||||
assert layer.expert_map is None, """expert_map/EP not implemented
|
||||
for CPU dyn-4bit MoE."""
|
||||
|
||||
def _act_kind(s: MoEActivation) -> int:
|
||||
# 0 = SwiGLU_Gu (SiLU(g)*u), 1 = SwiGLU_Ug (SiLU(u)*g), 2 = SiLU
|
||||
if s == MoEActivation.SWIGLUSTEP:
|
||||
return 0
|
||||
if s == MoEActivation.SWIGLUOAI:
|
||||
return 1
|
||||
if s == MoEActivation.SILU:
|
||||
return 2
|
||||
raise ValueError(f"Unknown activation '{s}'")
|
||||
|
||||
# Apply topk softmax on router output
|
||||
topk_weights, topk_ids = select_experts(
|
||||
hidden_states=x,
|
||||
router_logits=router_logits,
|
||||
top_k=layer.top_k,
|
||||
use_grouped_topk=layer.use_grouped_topk,
|
||||
renormalize=layer.renormalize,
|
||||
)
|
||||
|
||||
return torch.ops._C.dynamic_4bit_int_moe(
|
||||
x,
|
||||
topk_ids.to(torch.long),
|
||||
topk_weights,
|
||||
layer.w13_weight_packed,
|
||||
layer.w2_weight_packed,
|
||||
layer.w2_out_features,
|
||||
layer.w2_in_features,
|
||||
layer.w13_out_features,
|
||||
layer.group_size,
|
||||
layer.apply_router_weight_on_input,
|
||||
int(_act_kind(layer.activation)),
|
||||
)
|
||||
+414
@@ -0,0 +1,414 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
|
||||
import torch
|
||||
from compressed_tensors.quantization import (
|
||||
QuantizationArgs,
|
||||
QuantizationStrategy,
|
||||
)
|
||||
|
||||
import vllm.model_executor.layers.fused_moe.modular_kernel as mk
|
||||
from vllm.distributed import get_tensor_model_parallel_world_size
|
||||
from vllm.logger import init_logger
|
||||
from vllm.model_executor.layers.fused_moe import (
|
||||
FusedMoE,
|
||||
FusedMoeWeightScaleSupported,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.config import (
|
||||
FusedMoEConfig,
|
||||
FusedMoEQuantConfig,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.oracle.fp8 import (
|
||||
convert_to_fp8_moe_kernel_format,
|
||||
make_fp8_moe_kernel,
|
||||
make_fp8_moe_quant_config,
|
||||
select_fp8_moe_backend,
|
||||
)
|
||||
from vllm.model_executor.layers.quantization.compressed_tensors.compressed_tensors_moe import ( # noqa E501
|
||||
CompressedTensorsMoEMethod,
|
||||
)
|
||||
from vllm.model_executor.layers.quantization.utils.fp8_utils import (
|
||||
process_fp8_input_tensor_strategy_moe,
|
||||
process_fp8_weight_tensor_strategy_moe,
|
||||
)
|
||||
from vllm.model_executor.layers.quantization.utils.quant_utils import (
|
||||
kFp8Dynamic128Sym,
|
||||
kFp8DynamicTokenSym,
|
||||
kFp8Static128BlockSym,
|
||||
kFp8StaticChannelSym,
|
||||
kFp8StaticTensorSym,
|
||||
)
|
||||
from vllm.model_executor.layers.quantization.utils.w8a8_utils import (
|
||||
normalize_e4m3fn_to_e4m3fnuz,
|
||||
)
|
||||
from vllm.model_executor.utils import replace_parameter, set_weight_attrs
|
||||
from vllm.platforms import current_platform
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class CompressedTensorsW8A8Fp8MoEMethod(CompressedTensorsMoEMethod):
|
||||
"""W8A8 FP8 MoE quantization using compressed tensors."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
weight_quant: QuantizationArgs,
|
||||
input_quant: QuantizationArgs,
|
||||
moe: FusedMoEConfig,
|
||||
layer_name: str | None = None,
|
||||
):
|
||||
super().__init__(moe)
|
||||
self.weight_quant = weight_quant
|
||||
self.input_quant = input_quant
|
||||
|
||||
per_tensor = (
|
||||
self.weight_quant.strategy == QuantizationStrategy.TENSOR
|
||||
and self.input_quant.strategy == QuantizationStrategy.TENSOR
|
||||
)
|
||||
per_channel = (
|
||||
self.weight_quant.strategy == QuantizationStrategy.CHANNEL
|
||||
and self.input_quant.strategy == QuantizationStrategy.TOKEN
|
||||
)
|
||||
if not (per_tensor or per_channel):
|
||||
assert self.weight_quant.strategy == QuantizationStrategy.BLOCK
|
||||
self.weight_block_size = self.weight_quant.block_structure
|
||||
assert self.weight_quant.dynamic is not None
|
||||
else:
|
||||
self.weight_block_size = None
|
||||
self.block_quant = self.weight_block_size is not None
|
||||
|
||||
self.static_input_scales = not self.input_quant.dynamic
|
||||
if self.static_input_scales and per_channel:
|
||||
raise ValueError(
|
||||
"For FP8 Fused MoE layer, we require either per tensor or "
|
||||
"channelwise, dynamic per token quantization."
|
||||
)
|
||||
|
||||
ct2vllm_weight = {
|
||||
QuantizationStrategy.CHANNEL: kFp8StaticChannelSym,
|
||||
QuantizationStrategy.TENSOR: kFp8StaticTensorSym,
|
||||
QuantizationStrategy.BLOCK: kFp8Static128BlockSym,
|
||||
}
|
||||
ct2vllm_act = {
|
||||
QuantizationStrategy.TOKEN: kFp8DynamicTokenSym,
|
||||
QuantizationStrategy.TENSOR: (
|
||||
kFp8StaticTensorSym if self.static_input_scales else kFp8Dynamic128Sym
|
||||
),
|
||||
}
|
||||
weight_key = ct2vllm_weight[self.weight_quant.strategy]
|
||||
if weight_key == kFp8Static128BlockSym:
|
||||
activation_key = kFp8Dynamic128Sym
|
||||
else:
|
||||
activation_key = ct2vllm_act[self.input_quant.strategy]
|
||||
|
||||
# Select Fp8 MoE backend
|
||||
self.fp8_backend, self.experts_cls = select_fp8_moe_backend(
|
||||
config=self.moe,
|
||||
weight_key=weight_key,
|
||||
activation_key=activation_key,
|
||||
allow_vllm_cutlass=True,
|
||||
)
|
||||
|
||||
def create_weights(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
num_experts: int,
|
||||
hidden_size: int,
|
||||
intermediate_size_per_partition: int,
|
||||
params_dtype: torch.dtype,
|
||||
**extra_weight_attrs,
|
||||
):
|
||||
layer.num_experts = num_experts
|
||||
layer.orig_dtype = params_dtype
|
||||
layer.weight_block_size = None
|
||||
|
||||
params_dtype = torch.float8_e4m3fn
|
||||
w13_num_shards = 2 if self.moe.is_act_and_mul else 1
|
||||
|
||||
if self.block_quant:
|
||||
assert self.weight_block_size is not None
|
||||
layer.weight_block_size = self.weight_block_size
|
||||
tp_size = get_tensor_model_parallel_world_size()
|
||||
block_n, block_k = (
|
||||
self.weight_block_size[0],
|
||||
self.weight_block_size[1],
|
||||
)
|
||||
# NOTE: To ensure proper alignment of the block-wise quantization
|
||||
# scales, the output_size of the weights for both the gate and up
|
||||
# layers must be divisible by block_n.
|
||||
# Required by column parallel or enabling merged weights
|
||||
if intermediate_size_per_partition % block_n != 0:
|
||||
raise ValueError(
|
||||
f"The output_size of gate's and up's weight = "
|
||||
f"{intermediate_size_per_partition} is not divisible by "
|
||||
f"weight quantization block_n = {block_n}."
|
||||
)
|
||||
if tp_size > 1 and intermediate_size_per_partition % block_k != 0:
|
||||
# Required by row parallel
|
||||
raise ValueError(
|
||||
f"The input_size of down's weight = "
|
||||
f"{intermediate_size_per_partition} is not divisible by "
|
||||
f"weight quantization block_k = {block_k}."
|
||||
)
|
||||
|
||||
# WEIGHTS
|
||||
w13_weight = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts,
|
||||
w13_num_shards * intermediate_size_per_partition,
|
||||
hidden_size,
|
||||
dtype=params_dtype,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w13_weight", w13_weight)
|
||||
set_weight_attrs(w13_weight, extra_weight_attrs)
|
||||
|
||||
w2_weight = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts,
|
||||
hidden_size,
|
||||
intermediate_size_per_partition,
|
||||
dtype=params_dtype,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w2_weight", w2_weight)
|
||||
set_weight_attrs(w2_weight, extra_weight_attrs)
|
||||
|
||||
# WEIGHT_SCALES
|
||||
if self.weight_quant.strategy == QuantizationStrategy.TENSOR:
|
||||
# For gated MoE, allocate 2 scales for w1 and w3 respectively.
|
||||
# They will be combined to a single scale after weight loading.
|
||||
# For non-gated MoE, allocate 1 scale for w13.
|
||||
w13_weight_scale = torch.nn.Parameter(
|
||||
torch.ones(num_experts, w13_num_shards, dtype=torch.float32),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w13_weight_scale", w13_weight_scale)
|
||||
w2_weight_scale = torch.nn.Parameter(
|
||||
torch.ones(num_experts, dtype=torch.float32), requires_grad=False
|
||||
)
|
||||
layer.register_parameter("w2_weight_scale", w2_weight_scale)
|
||||
# Add PER-TENSOR quantization for FusedMoE.weight_loader.
|
||||
extra_weight_attrs.update(
|
||||
{"quant_method": FusedMoeWeightScaleSupported.TENSOR.value}
|
||||
)
|
||||
set_weight_attrs(w13_weight_scale, extra_weight_attrs)
|
||||
set_weight_attrs(w2_weight_scale, extra_weight_attrs)
|
||||
|
||||
elif self.weight_quant.strategy == QuantizationStrategy.CHANNEL:
|
||||
w13_weight_scale = torch.nn.Parameter(
|
||||
torch.ones(
|
||||
num_experts,
|
||||
w13_num_shards * intermediate_size_per_partition,
|
||||
1,
|
||||
dtype=torch.float32,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w13_weight_scale", w13_weight_scale)
|
||||
w2_weight_scale = torch.nn.Parameter(
|
||||
torch.ones(num_experts, hidden_size, 1, dtype=torch.float32),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w2_weight_scale", w2_weight_scale)
|
||||
# Add PER-CHANNEL quantization for FusedMoE.weight_loader.
|
||||
extra_weight_attrs.update(
|
||||
{"quant_method": FusedMoeWeightScaleSupported.CHANNEL.value}
|
||||
)
|
||||
set_weight_attrs(w13_weight_scale, extra_weight_attrs)
|
||||
set_weight_attrs(w2_weight_scale, extra_weight_attrs)
|
||||
|
||||
elif self.weight_quant.strategy == QuantizationStrategy.BLOCK:
|
||||
w13_weight_scale = torch.nn.Parameter(
|
||||
torch.ones(
|
||||
num_experts,
|
||||
w13_num_shards
|
||||
* ((intermediate_size_per_partition + block_n - 1) // block_n),
|
||||
(hidden_size + block_k - 1) // block_k,
|
||||
dtype=torch.float32,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w13_weight_scale", w13_weight_scale)
|
||||
w2_weight_scale = torch.nn.Parameter(
|
||||
torch.ones(
|
||||
num_experts,
|
||||
(hidden_size + block_n - 1) // block_n,
|
||||
(intermediate_size_per_partition + block_k - 1) // block_k,
|
||||
dtype=torch.float32,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w2_weight_scale", w2_weight_scale)
|
||||
# Add PER-CHANNEL quantization for FusedMoE.weight_loader.
|
||||
extra_weight_attrs.update(
|
||||
{"quant_method": FusedMoeWeightScaleSupported.BLOCK.value}
|
||||
)
|
||||
set_weight_attrs(w13_weight_scale, extra_weight_attrs)
|
||||
set_weight_attrs(w2_weight_scale, extra_weight_attrs)
|
||||
|
||||
# INPUT_SCALES
|
||||
if self.static_input_scales:
|
||||
w13_input_scale = torch.nn.Parameter(
|
||||
torch.ones(num_experts, dtype=torch.float32), requires_grad=False
|
||||
)
|
||||
layer.register_parameter("w13_input_scale", w13_input_scale)
|
||||
set_weight_attrs(w13_input_scale, extra_weight_attrs)
|
||||
|
||||
w2_input_scale = torch.nn.Parameter(
|
||||
torch.ones(num_experts, dtype=torch.float32), requires_grad=False
|
||||
)
|
||||
layer.register_parameter("w2_input_scale", w2_input_scale)
|
||||
set_weight_attrs(w2_input_scale, extra_weight_attrs)
|
||||
else:
|
||||
layer.w13_input_scale = None
|
||||
layer.w2_input_scale = None
|
||||
|
||||
def process_weights_after_loading(self, layer: FusedMoE) -> None:
|
||||
# Allow for accessing weights and scales in standard way.
|
||||
w13 = layer.w13_weight
|
||||
w2 = layer.w2_weight
|
||||
w13_scale = layer.w13_weight_scale
|
||||
w2_scale = layer.w2_weight_scale
|
||||
w13_input_scale = layer.w13_input_scale
|
||||
w2_input_scale = layer.w2_input_scale
|
||||
|
||||
# MI300x and MI325x use FNUZ format for FP8. Convert if needed.
|
||||
if current_platform.is_fp8_fnuz():
|
||||
w13, w13_scale, w13_input_scale = normalize_e4m3fn_to_e4m3fnuz(
|
||||
w13, w13_scale, w13_input_scale
|
||||
)
|
||||
w2, w2_scale, w2_input_scale = normalize_e4m3fn_to_e4m3fnuz(
|
||||
w2, w2_scale, w2_input_scale
|
||||
)
|
||||
|
||||
# Per tensor kernels require single activation scale. Use the max.
|
||||
if self.static_input_scales:
|
||||
assert self.input_quant.strategy == QuantizationStrategy.TENSOR
|
||||
assert w13_input_scale is not None and w2_input_scale is not None
|
||||
w13_input_scale, w2_input_scale = process_fp8_input_tensor_strategy_moe(
|
||||
w13_input_scale, w2_input_scale
|
||||
)
|
||||
replace_parameter(layer, "w13_input_scale", w13_input_scale)
|
||||
replace_parameter(layer, "w2_input_scale", w2_input_scale)
|
||||
|
||||
# Per-tensor kernels use a single scale, for W13, but on disk there
|
||||
# is a separate scale for W1 and W3. Requantize with the max scale.
|
||||
if self.weight_quant.strategy == QuantizationStrategy.TENSOR:
|
||||
w13, w13_scale = process_fp8_weight_tensor_strategy_moe(
|
||||
w13,
|
||||
w13_scale,
|
||||
shard_size=layer.intermediate_size_per_partition,
|
||||
num_experts=layer.local_num_experts,
|
||||
is_act_and_mul=self.moe.is_act_and_mul,
|
||||
)
|
||||
|
||||
w13, w2, w13_scale, w2_scale = convert_to_fp8_moe_kernel_format(
|
||||
fp8_backend=self.fp8_backend,
|
||||
layer=layer,
|
||||
w13=w13,
|
||||
w2=w2,
|
||||
w13_scale=w13_scale,
|
||||
w2_scale=w2_scale,
|
||||
w13_input_scale=w13_input_scale,
|
||||
w2_input_scale=w2_input_scale,
|
||||
)
|
||||
|
||||
# Replace parameters with updated versions. Note that this helper
|
||||
# function ensures the replacement is compatible with RL weight reloads.
|
||||
replace_parameter(layer, "w13_weight", w13)
|
||||
replace_parameter(layer, "w2_weight", w2)
|
||||
replace_parameter(layer, "w13_weight_scale", w13_scale)
|
||||
replace_parameter(layer, "w2_weight_scale", w2_scale)
|
||||
|
||||
# Setup modular kernel for TP case and naive DP/EP case.
|
||||
# In non-naive DP/EP case, we will create a ModularKernelMethod.
|
||||
# TODO(rob): unify these so FP8MoEMethod owns the ModularKernel
|
||||
# in both cases.
|
||||
self.moe_quant_config = self.get_fused_moe_quant_config(layer)
|
||||
if self.moe_quant_config:
|
||||
assert self.experts_cls is not None
|
||||
self.moe_kernel = make_fp8_moe_kernel(
|
||||
moe_quant_config=self.moe_quant_config,
|
||||
moe_config=self.moe,
|
||||
fp8_backend=self.fp8_backend,
|
||||
experts_cls=self.experts_cls,
|
||||
routing_tables=layer._maybe_init_expert_routing_tables(),
|
||||
shared_experts=layer.shared_experts,
|
||||
)
|
||||
|
||||
def maybe_make_prepare_finalize(
|
||||
self,
|
||||
routing_tables: tuple[torch.Tensor, torch.Tensor, torch.Tensor] | None = None,
|
||||
) -> mk.FusedMoEPrepareAndFinalizeModular | None:
|
||||
raise ValueError(
|
||||
f"{self.__class__.__name__} uses the new modular kernel initialization "
|
||||
"logic. This function should not be called."
|
||||
)
|
||||
|
||||
def get_fused_moe_quant_config(self, layer: torch.nn.Module) -> FusedMoEQuantConfig:
|
||||
is_per_token = self.input_quant.strategy == QuantizationStrategy.TOKEN
|
||||
return make_fp8_moe_quant_config(
|
||||
fp8_backend=self.fp8_backend,
|
||||
w1_scale=layer.w13_weight_scale,
|
||||
w2_scale=layer.w2_weight_scale,
|
||||
a1_scale=layer.w13_input_scale,
|
||||
a2_scale=layer.w2_input_scale,
|
||||
per_act_token_quant=is_per_token,
|
||||
per_out_ch_quant=is_per_token,
|
||||
block_shape=self.weight_block_size,
|
||||
)
|
||||
|
||||
def apply_monolithic(
|
||||
self,
|
||||
layer: FusedMoE,
|
||||
x: torch.Tensor,
|
||||
router_logits: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
assert self.moe_kernel is not None
|
||||
return self.moe_kernel.apply_monolithic(
|
||||
x,
|
||||
layer.w13_weight,
|
||||
layer.w2_weight,
|
||||
router_logits,
|
||||
activation=layer.activation,
|
||||
global_num_experts=layer.global_num_experts,
|
||||
expert_map=layer.expert_map,
|
||||
apply_router_weight_on_input=layer.apply_router_weight_on_input,
|
||||
num_expert_group=layer.num_expert_group,
|
||||
topk_group=layer.topk_group,
|
||||
e_score_correction_bias=layer.e_score_correction_bias,
|
||||
routed_scaling_factor=layer.routed_scaling_factor,
|
||||
)
|
||||
|
||||
def apply(
|
||||
self,
|
||||
layer: FusedMoE,
|
||||
x: torch.Tensor,
|
||||
topk_weights: torch.Tensor,
|
||||
topk_ids: torch.Tensor,
|
||||
shared_experts_input: torch.Tensor | None,
|
||||
) -> torch.Tensor:
|
||||
assert not self.is_monolithic
|
||||
assert self.moe_kernel is not None
|
||||
return self.moe_kernel.apply(
|
||||
x,
|
||||
layer.w13_weight,
|
||||
layer.w2_weight,
|
||||
topk_weights,
|
||||
topk_ids,
|
||||
activation=layer.activation,
|
||||
global_num_experts=layer.global_num_experts,
|
||||
# TODO(rob): investigate the disable_expert_map introduced by:
|
||||
# https://github.com/vllm-project/vllm/commit/84166fee9770e6fba71a96978b3e7d149392fb28 # noqa: E501
|
||||
expert_map=layer.expert_map,
|
||||
apply_router_weight_on_input=layer.apply_router_weight_on_input,
|
||||
shared_experts_input=shared_experts_input,
|
||||
)
|
||||
|
||||
@property
|
||||
def supports_eplb(self) -> bool:
|
||||
return True
|
||||
+161
@@ -0,0 +1,161 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
|
||||
import torch
|
||||
from compressed_tensors.quantization import (
|
||||
QuantizationArgs,
|
||||
QuantizationStrategy,
|
||||
)
|
||||
|
||||
from vllm.logger import init_logger
|
||||
from vllm.model_executor.layers.fused_moe import (
|
||||
FusedMoE,
|
||||
FusedMoeWeightScaleSupported,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.config import (
|
||||
FusedMoEConfig,
|
||||
FusedMoEQuantConfig,
|
||||
int8_w8a8_moe_quant_config,
|
||||
)
|
||||
from vllm.model_executor.layers.quantization.compressed_tensors.compressed_tensors_moe import ( # noqa E501
|
||||
CompressedTensorsMoEMethod,
|
||||
)
|
||||
from vllm.model_executor.utils import set_weight_attrs
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class CompressedTensorsW8A8Int8MoEMethod(CompressedTensorsMoEMethod):
|
||||
def __init__(
|
||||
self,
|
||||
weight_quant: QuantizationArgs,
|
||||
input_quant: QuantizationArgs,
|
||||
moe: FusedMoEConfig,
|
||||
layer_name: str | None = None,
|
||||
):
|
||||
super().__init__(moe)
|
||||
self.weight_quant = weight_quant
|
||||
self.input_quant = input_quant
|
||||
|
||||
per_channel = (
|
||||
self.weight_quant.strategy == QuantizationStrategy.CHANNEL
|
||||
and self.input_quant.strategy == QuantizationStrategy.TOKEN
|
||||
)
|
||||
if not per_channel:
|
||||
raise ValueError(
|
||||
"For INT8 Fused MoE layers, we require channelwise, "
|
||||
"dynamic per token quantization. Found "
|
||||
f"{self.weight_quant}, {self.input_quant}"
|
||||
)
|
||||
|
||||
self.static_input_scales = not self.input_quant.dynamic
|
||||
if self.static_input_scales:
|
||||
raise ValueError(
|
||||
"For INT8 Fused MoE layers, we require channelwise, "
|
||||
"dynamic per token quantization. Found static input scales."
|
||||
)
|
||||
|
||||
def create_weights(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
num_experts: int,
|
||||
hidden_size: int,
|
||||
intermediate_size_per_partition: int,
|
||||
params_dtype: torch.dtype,
|
||||
**extra_weight_attrs,
|
||||
):
|
||||
params_dtype = torch.int8
|
||||
w13_num_shards = 2 if self.moe.is_act_and_mul else 1
|
||||
|
||||
# WEIGHTS
|
||||
w13_weight = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts,
|
||||
w13_num_shards * intermediate_size_per_partition,
|
||||
hidden_size,
|
||||
dtype=params_dtype,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w13_weight", w13_weight)
|
||||
set_weight_attrs(w13_weight, extra_weight_attrs)
|
||||
|
||||
w2_weight = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts,
|
||||
hidden_size,
|
||||
intermediate_size_per_partition,
|
||||
dtype=params_dtype,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w2_weight", w2_weight)
|
||||
set_weight_attrs(w2_weight, extra_weight_attrs)
|
||||
|
||||
# WEIGHT_SCALES
|
||||
assert self.weight_quant.strategy == QuantizationStrategy.CHANNEL
|
||||
w13_weight_scale = torch.nn.Parameter(
|
||||
torch.ones(
|
||||
num_experts,
|
||||
w13_num_shards * intermediate_size_per_partition,
|
||||
1,
|
||||
dtype=torch.float32,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w13_weight_scale", w13_weight_scale)
|
||||
w2_weight_scale = torch.nn.Parameter(
|
||||
torch.ones(num_experts, hidden_size, 1, dtype=torch.float32),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w2_weight_scale", w2_weight_scale)
|
||||
# Add PER-CHANNEL quantization for FusedMoE.weight_loader.
|
||||
extra_weight_attrs.update(
|
||||
{"quant_method": FusedMoeWeightScaleSupported.CHANNEL.value}
|
||||
)
|
||||
set_weight_attrs(w13_weight_scale, extra_weight_attrs)
|
||||
set_weight_attrs(w2_weight_scale, extra_weight_attrs)
|
||||
|
||||
# INPUT_SCALES
|
||||
assert not self.static_input_scales
|
||||
layer.w13_input_scale = None
|
||||
layer.w2_input_scale = None
|
||||
|
||||
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
|
||||
pass
|
||||
|
||||
def get_fused_moe_quant_config(
|
||||
self, layer: torch.nn.Module
|
||||
) -> FusedMoEQuantConfig | None:
|
||||
return int8_w8a8_moe_quant_config(
|
||||
w1_scale=layer.w13_weight_scale,
|
||||
w2_scale=layer.w2_weight_scale,
|
||||
a1_scale=layer.w13_input_scale,
|
||||
a2_scale=layer.w2_input_scale,
|
||||
per_act_token_quant=True,
|
||||
)
|
||||
|
||||
def apply(
|
||||
self,
|
||||
layer: FusedMoE,
|
||||
x: torch.Tensor,
|
||||
topk_weights: torch.Tensor,
|
||||
topk_ids: torch.Tensor,
|
||||
shared_experts_input: torch.Tensor | None,
|
||||
) -> torch.Tensor:
|
||||
from vllm.model_executor.layers.fused_moe import fused_experts
|
||||
|
||||
return fused_experts(
|
||||
hidden_states=x,
|
||||
w1=layer.w13_weight,
|
||||
w2=layer.w2_weight,
|
||||
topk_weights=topk_weights,
|
||||
topk_ids=topk_ids,
|
||||
inplace=not self.moe.disable_inplace,
|
||||
activation=layer.activation,
|
||||
apply_router_weight_on_input=layer.apply_router_weight_on_input,
|
||||
global_num_experts=layer.global_num_experts,
|
||||
expert_map=layer.expert_map,
|
||||
quant_config=self.moe_quant_config,
|
||||
)
|
||||
+267
@@ -0,0 +1,267 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
|
||||
import torch
|
||||
from compressed_tensors.quantization import (
|
||||
QuantizationArgs,
|
||||
)
|
||||
|
||||
import vllm.model_executor.layers.fused_moe.modular_kernel as mk
|
||||
from vllm.logger import init_logger
|
||||
from vllm.model_executor.layers.fused_moe import (
|
||||
FusedMoE,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.config import (
|
||||
FusedMoEConfig,
|
||||
FusedMoEQuantConfig,
|
||||
int4_w4a16_moe_quant_config,
|
||||
int8_w8a16_moe_quant_config,
|
||||
)
|
||||
from vllm.model_executor.layers.quantization.compressed_tensors.compressed_tensors_moe import ( # noqa E501
|
||||
CompressedTensorsMoEMethod,
|
||||
)
|
||||
from vllm.model_executor.utils import set_weight_attrs
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class CompressedTensorsWNA16MoEMethod(CompressedTensorsMoEMethod):
|
||||
def __init__(
|
||||
self,
|
||||
weight_quant: QuantizationArgs,
|
||||
input_quant: QuantizationArgs | None,
|
||||
moe: FusedMoEConfig,
|
||||
layer_name: str | None = None,
|
||||
):
|
||||
super().__init__(moe)
|
||||
self.weight_quant = weight_quant
|
||||
self.input_quant = input_quant
|
||||
# Extract properties from weight_quant
|
||||
self.num_bits = weight_quant.num_bits
|
||||
self.packed_factor = 32 // weight_quant.num_bits
|
||||
self.strategy = weight_quant.strategy
|
||||
# channelwise is not supported by this kernel
|
||||
assert weight_quant.strategy == "group"
|
||||
self.group_size = weight_quant.group_size
|
||||
# grouped actorder isn't supported by this kernel
|
||||
assert weight_quant.actorder != "group"
|
||||
assert weight_quant.symmetric, (
|
||||
"Only symmetric quantization is supported for MoE"
|
||||
)
|
||||
|
||||
def create_weights(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
num_experts: int,
|
||||
hidden_size: int,
|
||||
intermediate_size_per_partition: int,
|
||||
params_dtype: torch.dtype,
|
||||
**extra_weight_attrs,
|
||||
):
|
||||
# Will transpose the loaded weight along the
|
||||
# intermediate and hidden dim sizes. Will
|
||||
# shard for TP along the transposed dims
|
||||
extra_weight_attrs.update(
|
||||
{"is_transposed": True, "quant_method": self.strategy}
|
||||
)
|
||||
w13_num_shards = 2 if self.moe.is_act_and_mul else 1
|
||||
w13_weight = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts,
|
||||
hidden_size // self.packed_factor,
|
||||
w13_num_shards * intermediate_size_per_partition,
|
||||
dtype=torch.int32,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w13_weight_packed", w13_weight)
|
||||
set_weight_attrs(w13_weight, extra_weight_attrs)
|
||||
|
||||
w2_weight = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts,
|
||||
intermediate_size_per_partition // self.packed_factor,
|
||||
hidden_size,
|
||||
dtype=torch.int32,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w2_weight_packed", w2_weight)
|
||||
set_weight_attrs(w2_weight, extra_weight_attrs)
|
||||
|
||||
w2_scales_size = intermediate_size_per_partition
|
||||
|
||||
if self.strategy == "channel":
|
||||
num_groups_w2 = num_groups_w13 = 1
|
||||
self.group_size = -1
|
||||
else:
|
||||
num_groups_w2 = w2_scales_size // self.group_size
|
||||
num_groups_w13 = hidden_size // self.group_size
|
||||
|
||||
w13_scale = torch.nn.Parameter(
|
||||
torch.ones(
|
||||
num_experts,
|
||||
num_groups_w13,
|
||||
w13_num_shards * intermediate_size_per_partition,
|
||||
dtype=params_dtype,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w13_weight_scale", w13_scale)
|
||||
set_weight_attrs(w13_scale, extra_weight_attrs)
|
||||
|
||||
w2_scale = torch.nn.Parameter(
|
||||
torch.ones(num_experts, num_groups_w2, hidden_size, dtype=params_dtype),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w2_weight_scale", w2_scale)
|
||||
set_weight_attrs(w2_scale, extra_weight_attrs)
|
||||
set_weight_attrs(w2_scale, {"load_full_w2": False})
|
||||
|
||||
w2_weight_shape = torch.nn.Parameter(
|
||||
torch.empty(num_experts, 2), requires_grad=False
|
||||
)
|
||||
layer.register_parameter("w2_weight_shape", w2_weight_shape)
|
||||
set_weight_attrs(w2_weight_shape, extra_weight_attrs)
|
||||
w13_weight_shape = torch.nn.Parameter(
|
||||
torch.empty(num_experts, 2), requires_grad=False
|
||||
)
|
||||
|
||||
layer.register_parameter("w13_weight_shape", w13_weight_shape)
|
||||
set_weight_attrs(w13_weight_shape, extra_weight_attrs)
|
||||
|
||||
w13_g_idx = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts,
|
||||
hidden_size,
|
||||
dtype=torch.int32,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w13_weight_g_idx", w13_g_idx)
|
||||
set_weight_attrs(w13_g_idx, extra_weight_attrs)
|
||||
|
||||
w2_g_idx = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts,
|
||||
intermediate_size_per_partition,
|
||||
dtype=torch.int32,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w2_weight_g_idx", w2_g_idx)
|
||||
set_weight_attrs(w2_g_idx, extra_weight_attrs)
|
||||
|
||||
w13_g_idx_sort_indices = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts,
|
||||
hidden_size,
|
||||
dtype=torch.int32,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w13_g_idx_sort_indices", w13_g_idx_sort_indices)
|
||||
set_weight_attrs(w13_g_idx_sort_indices, extra_weight_attrs)
|
||||
|
||||
w2_g_idx_sort_indices = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts,
|
||||
intermediate_size_per_partition,
|
||||
dtype=torch.int32,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w2_g_idx_sort_indices", w2_g_idx_sort_indices)
|
||||
set_weight_attrs(w2_g_idx_sort_indices, extra_weight_attrs)
|
||||
|
||||
layer.a13_scale = None
|
||||
layer.a2_scale = None
|
||||
|
||||
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
|
||||
# Reconfigure packed weights and scales to match moe_wna16 format
|
||||
layer.w13_weight_packed = torch.nn.Parameter(
|
||||
layer.w13_weight_packed.transpose(1, 2).contiguous().view(torch.uint8),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.w2_weight_packed = torch.nn.Parameter(
|
||||
layer.w2_weight_packed.transpose(1, 2).contiguous().view(torch.uint8),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.w13_weight_scale = torch.nn.Parameter(
|
||||
layer.w13_weight_scale.transpose(1, 2).contiguous(), requires_grad=False
|
||||
)
|
||||
layer.w2_weight_scale = torch.nn.Parameter(
|
||||
layer.w2_weight_scale.transpose(1, 2).contiguous(), requires_grad=False
|
||||
)
|
||||
|
||||
def get_fused_moe_quant_config(
|
||||
self, layer: torch.nn.Module
|
||||
) -> FusedMoEQuantConfig | None:
|
||||
assert self.num_bits == 4 or self.num_bits == 8
|
||||
config_builder = (
|
||||
int4_w4a16_moe_quant_config
|
||||
if self.num_bits == 4
|
||||
else int8_w8a16_moe_quant_config
|
||||
)
|
||||
|
||||
return config_builder(
|
||||
w1_scale=layer.w13_weight_scale,
|
||||
w2_scale=layer.w2_weight_scale,
|
||||
w1_zp=None,
|
||||
w2_zp=None,
|
||||
block_shape=[0, self.group_size],
|
||||
)
|
||||
|
||||
def select_gemm_impl(
|
||||
self,
|
||||
prepare_finalize: mk.FusedMoEPrepareAndFinalizeModular,
|
||||
layer: torch.nn.Module,
|
||||
) -> mk.FusedMoEExpertsModular:
|
||||
if self.moe.is_lora_enabled:
|
||||
assert self.moe_quant_config is not None
|
||||
from vllm.triton_utils import HAS_TRITON
|
||||
|
||||
if HAS_TRITON:
|
||||
from vllm.model_executor.layers.fused_moe import TritonWNA16Experts
|
||||
|
||||
layer.w13_weight = layer.w13_weight_packed
|
||||
layer.w2_weight = layer.w2_weight_packed
|
||||
return TritonWNA16Experts(
|
||||
moe_config=self.moe, quant_config=self.moe_quant_config
|
||||
)
|
||||
else:
|
||||
raise NotImplementedError(
|
||||
"TritonExperts requires Triton. "
|
||||
"Install triton or disable LoRA for MoE."
|
||||
)
|
||||
|
||||
raise NotImplementedError
|
||||
|
||||
def apply(
|
||||
self,
|
||||
layer: FusedMoE,
|
||||
x: torch.Tensor,
|
||||
topk_weights: torch.Tensor,
|
||||
topk_ids: torch.Tensor,
|
||||
shared_experts_input: torch.Tensor | None,
|
||||
) -> torch.Tensor:
|
||||
from vllm.model_executor.layers.fused_moe import fused_experts
|
||||
|
||||
return fused_experts(
|
||||
x,
|
||||
layer.w13_weight_packed,
|
||||
layer.w2_weight_packed,
|
||||
topk_weights=topk_weights,
|
||||
topk_ids=topk_ids,
|
||||
inplace=not self.moe.disable_inplace,
|
||||
activation=layer.activation,
|
||||
apply_router_weight_on_input=layer.apply_router_weight_on_input,
|
||||
global_num_experts=layer.global_num_experts,
|
||||
expert_map=layer.expert_map,
|
||||
quant_config=self.moe_quant_config,
|
||||
)
|
||||
|
||||
@property
|
||||
def supports_eplb(self) -> bool:
|
||||
return True
|
||||
+575
@@ -0,0 +1,575 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
import enum
|
||||
from enum import Enum
|
||||
|
||||
import torch
|
||||
from compressed_tensors.quantization import (
|
||||
QuantizationArgs,
|
||||
)
|
||||
|
||||
import vllm.model_executor.layers.fused_moe.modular_kernel as mk
|
||||
from vllm import _custom_ops as ops
|
||||
from vllm.logger import init_logger
|
||||
from vllm.model_executor.layers.fused_moe import (
|
||||
FusedMoE,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.config import (
|
||||
FusedMoEConfig,
|
||||
FusedMoEQuantConfig,
|
||||
int4_w4a16_moe_quant_config,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.fused_marlin_moe import (
|
||||
BatchedMarlinExperts,
|
||||
MarlinExperts,
|
||||
fused_marlin_moe,
|
||||
)
|
||||
from vllm.model_executor.layers.quantization.compressed_tensors.compressed_tensors_moe import ( # noqa E501
|
||||
CompressedTensorsMoEMethod,
|
||||
)
|
||||
from vllm.model_executor.layers.quantization.compressed_tensors.schemes.compressed_tensors_wNa16 import ( # noqa
|
||||
WNA16_SUPPORTED_TYPES_MAP,
|
||||
)
|
||||
from vllm.model_executor.layers.quantization.utils.flashinfer_mxint4_moe import (
|
||||
flashinfer_trtllm_mxint4_moe,
|
||||
is_flashinfer_mxint4_moe_available,
|
||||
prepare_static_weights_for_trtllm_mxint4_moe,
|
||||
)
|
||||
from vllm.model_executor.layers.quantization.utils.marlin_utils import (
|
||||
get_marlin_input_dtype,
|
||||
marlin_act_int8_process_scales,
|
||||
marlin_make_workspace_new,
|
||||
marlin_moe_permute_scales,
|
||||
)
|
||||
from vllm.model_executor.utils import replace_parameter, set_weight_attrs
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class GPTQMarlinState(Enum):
|
||||
REPACK = enum.auto()
|
||||
READY = enum.auto()
|
||||
|
||||
|
||||
class CompressedTensorsWNA16MarlinMoEMethod(CompressedTensorsMoEMethod):
|
||||
def __init__(
|
||||
self,
|
||||
weight_quant: QuantizationArgs,
|
||||
input_quant: QuantizationArgs | None,
|
||||
moe: FusedMoEConfig,
|
||||
layer_name: str | None = None,
|
||||
):
|
||||
super().__init__(moe)
|
||||
self.weight_quant = weight_quant
|
||||
self.input_quant = input_quant
|
||||
assert weight_quant.symmetric, (
|
||||
"Only symmetric quantization is supported for MoE"
|
||||
)
|
||||
# Extract properties from weight_quant
|
||||
self.num_bits = weight_quant.num_bits
|
||||
self.packed_factor = 32 // weight_quant.num_bits
|
||||
self.strategy = weight_quant.strategy
|
||||
self.group_size = weight_quant.group_size
|
||||
self.actorder = weight_quant.actorder
|
||||
|
||||
self.quant_type = WNA16_SUPPORTED_TYPES_MAP[self.num_bits]
|
||||
|
||||
self.marlin_input_dtype = get_marlin_input_dtype(layer_name)
|
||||
self.use_flashinfer_mxint4_moe = (
|
||||
is_flashinfer_mxint4_moe_available()
|
||||
and self.group_size == 32
|
||||
and weight_quant.num_bits == 4
|
||||
)
|
||||
self.kernel_backend = (
|
||||
"Flashinfer" if self.use_flashinfer_mxint4_moe else "Marlin"
|
||||
)
|
||||
logger.info_once(
|
||||
f"Using {self.kernel_backend} backend for WNA16 MoE "
|
||||
f"(group_size={self.group_size}, num_bits={self.num_bits})",
|
||||
scope="local",
|
||||
)
|
||||
|
||||
def get_weight_shape(
|
||||
self,
|
||||
weight_name: str,
|
||||
num_experts: int,
|
||||
hidden_size: int,
|
||||
intermediate_size_per_partition: int,
|
||||
num_groups_w2: int | None = None,
|
||||
num_groups_w13: int | None = None,
|
||||
) -> tuple[int, int, int]:
|
||||
"""
|
||||
Get the shape of the weight based on the weight name, number of experts
|
||||
hidden size, intermediate size per partition, number of groups for w2,
|
||||
and number of groups for w13. Pass in num_groups_w2 and num_groups_w13
|
||||
for weight scales.
|
||||
"""
|
||||
if weight_name == "w13_scale":
|
||||
assert num_groups_w13 is not None, (
|
||||
"num_groups_w13 must be provided for weight scales"
|
||||
)
|
||||
if weight_name == "w2_scale":
|
||||
assert num_groups_w2 is not None, (
|
||||
"num_groups_w2 must be provided for weight scales"
|
||||
)
|
||||
w13_num_shards = 2 if self.moe.is_act_and_mul else 1
|
||||
shape_map = {
|
||||
"w13_weight": {
|
||||
"Flashinfer": (
|
||||
num_experts,
|
||||
w13_num_shards * intermediate_size_per_partition,
|
||||
hidden_size // self.packed_factor,
|
||||
),
|
||||
"Marlin": (
|
||||
num_experts,
|
||||
hidden_size // self.packed_factor,
|
||||
w13_num_shards * intermediate_size_per_partition,
|
||||
),
|
||||
},
|
||||
"w13_scale": {
|
||||
"Flashinfer": (
|
||||
num_experts,
|
||||
w13_num_shards * intermediate_size_per_partition,
|
||||
num_groups_w13,
|
||||
),
|
||||
"Marlin": (
|
||||
num_experts,
|
||||
num_groups_w13,
|
||||
w13_num_shards * intermediate_size_per_partition,
|
||||
),
|
||||
},
|
||||
"w2_weight": {
|
||||
"Flashinfer": (
|
||||
num_experts,
|
||||
hidden_size,
|
||||
intermediate_size_per_partition // self.packed_factor,
|
||||
),
|
||||
"Marlin": (
|
||||
num_experts,
|
||||
intermediate_size_per_partition // self.packed_factor,
|
||||
hidden_size,
|
||||
),
|
||||
},
|
||||
"w2_scale": {
|
||||
"Flashinfer": (num_experts, hidden_size, num_groups_w2),
|
||||
"Marlin": (num_experts, num_groups_w2, hidden_size),
|
||||
},
|
||||
}
|
||||
return shape_map[weight_name][self.kernel_backend]
|
||||
|
||||
def create_weights(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
num_experts: int,
|
||||
hidden_size: int,
|
||||
intermediate_size_per_partition: int,
|
||||
params_dtype: torch.dtype,
|
||||
**extra_weight_attrs,
|
||||
):
|
||||
intermediate_size_full = extra_weight_attrs.pop("intermediate_size_full")
|
||||
|
||||
# Will transpose the loaded weight along the
|
||||
# intermediate and hidden dim sizes. Will
|
||||
# shard for TP along the transposed dims
|
||||
is_transposed = self.kernel_backend != "Flashinfer"
|
||||
extra_weight_attrs.update(
|
||||
{"is_transposed": is_transposed, "quant_method": self.strategy}
|
||||
)
|
||||
|
||||
w13_weight = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
*self.get_weight_shape(
|
||||
"w13_weight",
|
||||
num_experts,
|
||||
hidden_size,
|
||||
intermediate_size_per_partition,
|
||||
),
|
||||
dtype=torch.int32,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w13_weight_packed", w13_weight)
|
||||
set_weight_attrs(w13_weight, extra_weight_attrs)
|
||||
|
||||
w2_weight = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
*self.get_weight_shape(
|
||||
"w2_weight",
|
||||
num_experts,
|
||||
hidden_size,
|
||||
intermediate_size_per_partition,
|
||||
),
|
||||
dtype=torch.int32,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w2_weight_packed", w2_weight)
|
||||
set_weight_attrs(w2_weight, extra_weight_attrs)
|
||||
|
||||
# In the case where we have actorder/g_idx,
|
||||
# we do not partition the w2 scales
|
||||
load_full_w2 = self.actorder and self.group_size != -1
|
||||
w2_scales_size = (
|
||||
intermediate_size_full if load_full_w2 else intermediate_size_per_partition
|
||||
)
|
||||
|
||||
self.is_k_full = (not self.actorder) or (
|
||||
intermediate_size_per_partition == intermediate_size_full
|
||||
)
|
||||
|
||||
if self.strategy == "channel":
|
||||
num_groups_w2 = num_groups_w13 = 1
|
||||
self.group_size = -1
|
||||
else:
|
||||
num_groups_w2 = w2_scales_size // self.group_size
|
||||
num_groups_w13 = hidden_size // self.group_size
|
||||
|
||||
layer.num_groups_w13 = num_groups_w13
|
||||
layer.num_groups_w2 = num_groups_w2
|
||||
|
||||
w13_scale = torch.nn.Parameter(
|
||||
torch.ones(
|
||||
*self.get_weight_shape(
|
||||
"w13_scale",
|
||||
num_experts,
|
||||
hidden_size,
|
||||
intermediate_size_per_partition,
|
||||
num_groups_w13=num_groups_w13,
|
||||
),
|
||||
dtype=params_dtype,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w13_weight_scale", w13_scale)
|
||||
set_weight_attrs(w13_scale, extra_weight_attrs)
|
||||
|
||||
w2_scale = torch.nn.Parameter(
|
||||
torch.ones(
|
||||
*self.get_weight_shape(
|
||||
"w2_scale",
|
||||
num_experts,
|
||||
hidden_size,
|
||||
intermediate_size_per_partition,
|
||||
num_groups_w2=num_groups_w2,
|
||||
),
|
||||
dtype=params_dtype,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w2_weight_scale", w2_scale)
|
||||
set_weight_attrs(w2_scale, extra_weight_attrs)
|
||||
set_weight_attrs(w2_scale, {"load_full_w2": load_full_w2})
|
||||
|
||||
w2_weight_shape = torch.nn.Parameter(
|
||||
torch.empty(num_experts, 2), requires_grad=False
|
||||
)
|
||||
layer.register_parameter("w2_weight_shape", w2_weight_shape)
|
||||
set_weight_attrs(w2_weight_shape, extra_weight_attrs)
|
||||
w13_weight_shape = torch.nn.Parameter(
|
||||
torch.empty(num_experts, 2), requires_grad=False
|
||||
)
|
||||
|
||||
layer.register_parameter("w13_weight_shape", w13_weight_shape)
|
||||
set_weight_attrs(w13_weight_shape, extra_weight_attrs)
|
||||
|
||||
w13_g_idx = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts,
|
||||
hidden_size,
|
||||
dtype=torch.int32,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w13_weight_g_idx", w13_g_idx)
|
||||
set_weight_attrs(w13_g_idx, extra_weight_attrs)
|
||||
|
||||
w2_g_idx = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts,
|
||||
intermediate_size_per_partition,
|
||||
dtype=torch.int32,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w2_weight_g_idx", w2_g_idx)
|
||||
set_weight_attrs(w2_g_idx, extra_weight_attrs)
|
||||
|
||||
w13_g_idx_sort_indices = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts,
|
||||
hidden_size,
|
||||
dtype=torch.int32,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w13_g_idx_sort_indices", w13_g_idx_sort_indices)
|
||||
set_weight_attrs(w13_g_idx_sort_indices, extra_weight_attrs)
|
||||
|
||||
w2_g_idx_sort_indices = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts,
|
||||
intermediate_size_per_partition,
|
||||
dtype=torch.int32,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w2_g_idx_sort_indices", w2_g_idx_sort_indices)
|
||||
set_weight_attrs(w2_g_idx_sort_indices, extra_weight_attrs)
|
||||
|
||||
layer.a13_scale = None
|
||||
layer.a2_scale = None
|
||||
layer.marlin_state = GPTQMarlinState.REPACK
|
||||
|
||||
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
|
||||
num_experts = layer.w13_weight_g_idx.shape[0]
|
||||
device = layer.w13_weight_g_idx.device
|
||||
if self.kernel_backend == "Flashinfer":
|
||||
dict_weights_mxint4 = prepare_static_weights_for_trtllm_mxint4_moe(
|
||||
layer.w13_weight_packed,
|
||||
layer.w13_weight_scale,
|
||||
layer.w2_weight_packed,
|
||||
layer.w2_weight_scale,
|
||||
)
|
||||
replace_parameter(
|
||||
layer, "w13_weight_packed", dict_weights_mxint4["gemm1_weights"]
|
||||
)
|
||||
replace_parameter(
|
||||
layer, "w13_weight_scale", dict_weights_mxint4["gemm1_scales"]
|
||||
)
|
||||
replace_parameter(
|
||||
layer, "w2_weight_packed", dict_weights_mxint4["gemm2_weights"]
|
||||
)
|
||||
replace_parameter(
|
||||
layer, "w2_weight_scale", dict_weights_mxint4["gemm2_scales"]
|
||||
)
|
||||
return None
|
||||
|
||||
is_a_8bit = (
|
||||
self.marlin_input_dtype is not None
|
||||
and self.marlin_input_dtype.itemsize == 1
|
||||
)
|
||||
|
||||
if self.marlin_input_dtype == torch.float8_e4m3fn:
|
||||
# NOTE: for non-zp quantization format only
|
||||
ops.marlin_int4_fp8_preprocess(layer.w13_weight_packed, inplace=True)
|
||||
ops.marlin_int4_fp8_preprocess(layer.w2_weight_packed, inplace=True)
|
||||
layer.w13_weight_scale.data = layer.w13_weight_scale.data * 512
|
||||
layer.w2_weight_scale.data = layer.w2_weight_scale.data * 512
|
||||
|
||||
# when running models with grouped act order,
|
||||
# resort to g_idx values provided in checkpoint
|
||||
if self.actorder == "group":
|
||||
w13_g_idx_sort_indices = torch.empty_like(layer.w13_weight_g_idx)
|
||||
w2_g_idx_sort_indices = torch.empty_like(layer.w2_weight_g_idx)
|
||||
w13_sorted_g_idx = torch.empty_like(layer.w13_weight_g_idx)
|
||||
w2_sorted_g_idx = torch.empty_like(layer.w2_weight_g_idx)
|
||||
|
||||
for e in range(num_experts):
|
||||
w13_g_idx_sort_indices[e] = torch.argsort(layer.w13_weight_g_idx[e]).to(
|
||||
torch.int32
|
||||
)
|
||||
w2_g_idx_sort_indices[e] = torch.argsort(layer.w2_weight_g_idx[e]).to(
|
||||
torch.int32
|
||||
)
|
||||
w13_sorted_g_idx[e] = layer.w13_weight_g_idx[e][
|
||||
w13_g_idx_sort_indices[e]
|
||||
]
|
||||
w2_sorted_g_idx[e] = layer.w2_weight_g_idx[e][w2_g_idx_sort_indices[e]]
|
||||
|
||||
replace_parameter(layer, "w13_weight_g_idx", w13_sorted_g_idx)
|
||||
replace_parameter(layer, "w2_weight_g_idx", w2_sorted_g_idx)
|
||||
replace_parameter(layer, "w13_g_idx_sort_indices", w13_g_idx_sort_indices)
|
||||
replace_parameter(layer, "w2_g_idx_sort_indices", w2_g_idx_sort_indices)
|
||||
|
||||
else:
|
||||
layer.w13_weight_g_idx = torch.nn.Parameter(
|
||||
torch.empty((num_experts, 0), dtype=torch.int32, device=device),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.w2_weight_g_idx = torch.nn.Parameter(
|
||||
torch.empty((num_experts, 0), dtype=torch.int32, device=device),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.w13_g_idx_sort_indices = torch.nn.Parameter(
|
||||
torch.empty((num_experts, 0), dtype=torch.int32, device=device),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.w2_g_idx_sort_indices = torch.nn.Parameter(
|
||||
torch.empty((num_experts, 0), dtype=torch.int32, device=device),
|
||||
requires_grad=False,
|
||||
)
|
||||
|
||||
marlin_w13_qweight = ops.gptq_marlin_moe_repack(
|
||||
layer.w13_weight_packed,
|
||||
layer.w13_g_idx_sort_indices,
|
||||
layer.w13_weight_packed.shape[1] * self.packed_factor,
|
||||
layer.w13_weight_packed.shape[2],
|
||||
self.num_bits,
|
||||
is_a_8bit=is_a_8bit,
|
||||
)
|
||||
replace_parameter(layer, "w13_weight_packed", marlin_w13_qweight)
|
||||
|
||||
marlin_w2_qweight = ops.gptq_marlin_moe_repack(
|
||||
layer.w2_weight_packed,
|
||||
layer.w2_g_idx_sort_indices,
|
||||
layer.w2_weight_packed.shape[1] * self.packed_factor,
|
||||
layer.w2_weight_packed.shape[2],
|
||||
self.num_bits,
|
||||
is_a_8bit=is_a_8bit,
|
||||
)
|
||||
replace_parameter(layer, "w2_weight_packed", marlin_w2_qweight)
|
||||
|
||||
# Repack scales
|
||||
marlin_w13_scales = marlin_moe_permute_scales(
|
||||
s=layer.w13_weight_scale,
|
||||
size_k=layer.w13_weight_packed.shape[2],
|
||||
size_n=layer.w13_weight_scale.shape[2],
|
||||
group_size=self.group_size,
|
||||
is_a_8bit=is_a_8bit,
|
||||
)
|
||||
if self.marlin_input_dtype == torch.int8 and layer.num_groups_w13 > 1:
|
||||
marlin_w13_scales, w13_input_global_scale = marlin_act_int8_process_scales(
|
||||
marlin_w13_scales
|
||||
)
|
||||
layer.register_parameter(
|
||||
"w13_input_global_scale",
|
||||
torch.nn.Parameter(w13_input_global_scale, requires_grad=False),
|
||||
)
|
||||
replace_parameter(layer, "w13_weight_scale", marlin_w13_scales)
|
||||
|
||||
marlin_w2_scales = marlin_moe_permute_scales(
|
||||
s=layer.w2_weight_scale,
|
||||
size_k=layer.w2_weight_scale.shape[1]
|
||||
* (self.group_size if self.group_size != -1 else self.packed_factor),
|
||||
size_n=layer.w2_weight_scale.shape[2],
|
||||
group_size=self.group_size,
|
||||
is_a_8bit=is_a_8bit,
|
||||
)
|
||||
if self.marlin_input_dtype == torch.int8 and layer.num_groups_w2 > 1:
|
||||
marlin_w2_scales, w2_input_global_scale = marlin_act_int8_process_scales(
|
||||
marlin_w2_scales
|
||||
)
|
||||
layer.register_parameter(
|
||||
"w2_input_global_scale",
|
||||
torch.nn.Parameter(w2_input_global_scale, requires_grad=False),
|
||||
)
|
||||
replace_parameter(layer, "w2_weight_scale", marlin_w2_scales)
|
||||
|
||||
layer.workspace = marlin_make_workspace_new(device, 4)
|
||||
|
||||
def get_fused_moe_quant_config(
|
||||
self, layer: torch.nn.Module
|
||||
) -> FusedMoEQuantConfig | None:
|
||||
if self.num_bits != 4:
|
||||
return None
|
||||
return int4_w4a16_moe_quant_config(
|
||||
w1_scale=layer.w13_weight_scale,
|
||||
w2_scale=layer.w2_weight_scale,
|
||||
w1_zp=None,
|
||||
w2_zp=None,
|
||||
block_shape=[0, self.group_size],
|
||||
)
|
||||
|
||||
def select_gemm_impl(
|
||||
self,
|
||||
prepare_finalize: mk.FusedMoEPrepareAndFinalizeModular,
|
||||
layer: torch.nn.Module,
|
||||
) -> mk.FusedMoEExpertsModular:
|
||||
assert self.num_bits == 4, "only supporting w4"
|
||||
layer.w13_weight = layer.w13_weight_packed
|
||||
layer.w2_weight = layer.w2_weight_packed
|
||||
assert all([w is not None for w in [layer.w13_weight, layer.w2_weight]])
|
||||
assert self.moe_quant_config is not None
|
||||
if (
|
||||
prepare_finalize.activation_format
|
||||
== mk.FusedMoEActivationFormat.BatchedExperts
|
||||
):
|
||||
max_num_tokens_per_rank = prepare_finalize.max_num_tokens_per_rank()
|
||||
assert max_num_tokens_per_rank is not None
|
||||
return BatchedMarlinExperts(
|
||||
max_num_tokens=max_num_tokens_per_rank,
|
||||
num_dispatchers=prepare_finalize.num_dispatchers(),
|
||||
moe_config=self.moe,
|
||||
quant_config=self.moe_quant_config,
|
||||
w13_g_idx=layer.w13_weight_g_idx,
|
||||
w2_g_idx=layer.w2_weight_g_idx,
|
||||
w13_g_idx_sort_indices=layer.w13_g_idx_sort_indices,
|
||||
w2_g_idx_sort_indices=layer.w2_g_idx_sort_indices,
|
||||
is_k_full=self.is_k_full,
|
||||
)
|
||||
else:
|
||||
return MarlinExperts(
|
||||
moe_config=self.moe,
|
||||
quant_config=self.moe_quant_config,
|
||||
w13_g_idx=layer.w13_weight_g_idx,
|
||||
w2_g_idx=layer.w2_weight_g_idx,
|
||||
w13_g_idx_sort_indices=layer.w13_g_idx_sort_indices,
|
||||
w2_g_idx_sort_indices=layer.w2_g_idx_sort_indices,
|
||||
is_k_full=self.is_k_full,
|
||||
)
|
||||
|
||||
@property
|
||||
def is_monolithic(self) -> bool:
|
||||
return self.kernel_backend == "Flashinfer"
|
||||
|
||||
def apply_monolithic(
|
||||
self,
|
||||
layer: FusedMoE,
|
||||
x: torch.Tensor,
|
||||
router_logits: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
assert self.kernel_backend == "Flashinfer"
|
||||
return flashinfer_trtllm_mxint4_moe(
|
||||
x=x,
|
||||
router_logits=router_logits,
|
||||
w13_weight_packed=layer.w13_weight_packed,
|
||||
w13_weight_scale=layer.w13_weight_scale,
|
||||
w2_weight_packed=layer.w2_weight_packed,
|
||||
w2_weight_scale=layer.w2_weight_scale,
|
||||
global_num_experts=layer.global_num_experts,
|
||||
top_k=layer.top_k,
|
||||
intermediate_size_per_partition=layer.intermediate_size_per_partition,
|
||||
local_num_experts=layer.local_num_experts,
|
||||
ep_rank=layer.ep_rank,
|
||||
num_expert_group=layer.num_expert_group,
|
||||
topk_group=layer.topk_group,
|
||||
e_score_correction_bias=layer.e_score_correction_bias,
|
||||
routing_method_type=layer.routing_method_type,
|
||||
)
|
||||
|
||||
def apply(
|
||||
self,
|
||||
layer: FusedMoE,
|
||||
x: torch.Tensor,
|
||||
topk_weights: torch.Tensor,
|
||||
topk_ids: torch.Tensor,
|
||||
shared_experts_input: torch.Tensor | None,
|
||||
) -> torch.Tensor:
|
||||
assert self.kernel_backend == "Marlin"
|
||||
return fused_marlin_moe(
|
||||
x,
|
||||
layer.w13_weight_packed,
|
||||
layer.w2_weight_packed,
|
||||
None,
|
||||
None,
|
||||
layer.w13_weight_scale,
|
||||
layer.w2_weight_scale,
|
||||
topk_weights,
|
||||
topk_ids,
|
||||
input_global_scale1=getattr(layer, "w13_input_global_scale", None),
|
||||
input_global_scale2=getattr(layer, "w2_input_global_scale", None),
|
||||
quant_type_id=self.quant_type.id,
|
||||
apply_router_weight_on_input=layer.apply_router_weight_on_input,
|
||||
global_num_experts=layer.global_num_experts,
|
||||
activation=layer.activation,
|
||||
expert_map=layer.expert_map,
|
||||
g_idx1=layer.w13_weight_g_idx,
|
||||
g_idx2=layer.w2_weight_g_idx,
|
||||
sort_indices1=layer.w13_g_idx_sort_indices,
|
||||
sort_indices2=layer.w2_g_idx_sort_indices,
|
||||
workspace=layer.workspace,
|
||||
input_dtype=self.marlin_input_dtype,
|
||||
is_k_full=self.is_k_full,
|
||||
inplace=not self.moe.disable_inplace,
|
||||
)
|
||||
+23
@@ -5,10 +5,12 @@ from collections.abc import Callable
|
||||
import torch
|
||||
from torch.nn.parameter import Parameter
|
||||
|
||||
from vllm.logger import init_logger
|
||||
from vllm.model_executor.layers.quantization.compressed_tensors.schemes import (
|
||||
CompressedTensorsScheme,
|
||||
)
|
||||
from vllm.model_executor.layers.quantization.utils.nvfp4_utils import (
|
||||
NvFp4LinearBackend,
|
||||
apply_nvfp4_linear,
|
||||
convert_to_nvfp4_linear_kernel_format,
|
||||
select_nvfp4_linear_backend,
|
||||
@@ -19,6 +21,9 @@ from vllm.model_executor.parameter import (
|
||||
PerTensorScaleParameter,
|
||||
)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
__all__ = ["CompressedTensorsW4A4Fp4"]
|
||||
|
||||
|
||||
@@ -27,6 +32,10 @@ class CompressedTensorsW4A4Fp4(CompressedTensorsScheme):
|
||||
self.backend = select_nvfp4_linear_backend()
|
||||
self.group_size = 16
|
||||
|
||||
self.swizzle = None
|
||||
if self.backend == NvFp4LinearBackend.EMULATION:
|
||||
self.swizzle = False
|
||||
|
||||
@classmethod
|
||||
def get_min_capability(cls) -> int:
|
||||
return 75
|
||||
@@ -89,6 +98,19 @@ class CompressedTensorsW4A4Fp4(CompressedTensorsScheme):
|
||||
# Rename CT checkpoint names to standardized names
|
||||
layer.weight = layer.weight_packed
|
||||
del layer.weight_packed
|
||||
|
||||
if (
|
||||
torch.unique(layer.input_global_scale).numel() != 1
|
||||
or torch.unique(layer.weight_global_scale).numel() != 1
|
||||
):
|
||||
logger.warning_once(
|
||||
"In NVFP4 linear, the global scale for input or weight are different"
|
||||
" for parallel layers (e.g. q_proj, k_proj, v_proj). This "
|
||||
" will likely result in reduced accuracy. Please verify the model"
|
||||
" accuracy. Consider using a checkpoint with a shared global NVFP4"
|
||||
" scale for fused layers."
|
||||
)
|
||||
|
||||
# Process global scales (CT stores as divisors, i.e. 1/scale)
|
||||
input_global_scale_inv = layer.input_global_scale.max().to(torch.float32)
|
||||
layer.input_global_scale = Parameter(
|
||||
@@ -121,4 +143,5 @@ class CompressedTensorsW4A4Fp4(CompressedTensorsScheme):
|
||||
layer=layer,
|
||||
x=x,
|
||||
bias=bias,
|
||||
swizzle=self.swizzle,
|
||||
)
|
||||
|
||||
@@ -71,6 +71,7 @@ from vllm.model_executor.layers.quantization.utils.mxfp8_utils import (
|
||||
mxfp8_e4m3_quantize,
|
||||
)
|
||||
from vllm.model_executor.layers.quantization.utils.nvfp4_utils import (
|
||||
NvFp4LinearBackend,
|
||||
apply_nvfp4_linear,
|
||||
convert_to_nvfp4_linear_kernel_format,
|
||||
select_nvfp4_linear_backend,
|
||||
@@ -1074,6 +1075,10 @@ class ModelOptNvFp4LinearMethod(LinearMethodBase):
|
||||
self.marlin_input_dtype = None
|
||||
self.backend = select_nvfp4_linear_backend()
|
||||
|
||||
self.swizzle = None
|
||||
if self.backend == NvFp4LinearBackend.EMULATION:
|
||||
self.swizzle = False
|
||||
|
||||
def create_weights(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
@@ -1149,10 +1154,23 @@ class ModelOptNvFp4LinearMethod(LinearMethodBase):
|
||||
layer.register_parameter("weight_scale", weight_scale)
|
||||
|
||||
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
|
||||
if (
|
||||
torch.unique(layer.input_scale).numel() != 1
|
||||
or torch.unique(layer.weight_scale_2).numel() != 1
|
||||
):
|
||||
logger.warning_once(
|
||||
"In NVFP4 linear, the global scale for input or weight are different"
|
||||
" for parallel layers (e.g. q_proj, k_proj, v_proj). This "
|
||||
" will likely results in reduce accuracy. Please verify the model"
|
||||
" accuracy. Consider using a checkpoint with a shared global NVFP4"
|
||||
" scale for parallel layers."
|
||||
)
|
||||
|
||||
# Rename ModelOpt checkpoint names to standardized names
|
||||
input_global_scale = layer.input_scale.max().to(torch.float32)
|
||||
layer.input_global_scale = Parameter(input_global_scale, requires_grad=False)
|
||||
del layer.input_scale
|
||||
|
||||
weight_global_scale = layer.weight_scale_2.max().to(torch.float32)
|
||||
layer.weight_global_scale = Parameter(weight_global_scale, requires_grad=False)
|
||||
del layer.weight_scale_2
|
||||
@@ -1179,6 +1197,7 @@ class ModelOptNvFp4LinearMethod(LinearMethodBase):
|
||||
layer=layer,
|
||||
x=x,
|
||||
bias=bias,
|
||||
swizzle=self.swizzle,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -267,20 +267,26 @@ class QuarkOCP_MX(QuarkScheme):
|
||||
def get_min_capability(cls) -> int:
|
||||
return 70
|
||||
|
||||
def process_dynamic_mxfp4_weights_after_loading(
|
||||
self, layer: torch.nn.Module
|
||||
) -> None:
|
||||
w_q, w_s = dynamic_mxfp4_quant(layer.weight)
|
||||
layer.weight_scale = torch.nn.Parameter(w_s.T.contiguous(), requires_grad=False)
|
||||
layer.weight = torch.nn.Parameter(w_q, requires_grad=False)
|
||||
|
||||
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
|
||||
layer.weight = torch.nn.Parameter(layer.weight.data, requires_grad=False)
|
||||
|
||||
if self.emulate:
|
||||
layer.weight_scale = torch.nn.Parameter(
|
||||
layer.weight_scale.data, requires_grad=False
|
||||
)
|
||||
if self.dynamic_mxfp4_quant:
|
||||
self.process_dynamic_mxfp4_weights_after_loading(layer)
|
||||
else:
|
||||
layer.weight_scale = torch.nn.Parameter(
|
||||
layer.weight_scale.data, requires_grad=False
|
||||
)
|
||||
else:
|
||||
if self.dynamic_mxfp4_quant:
|
||||
w_q, w_s = dynamic_mxfp4_quant(layer.weight)
|
||||
layer.weight_scale = torch.nn.Parameter(
|
||||
w_s.T.contiguous(), requires_grad=False
|
||||
)
|
||||
layer.weight = torch.nn.Parameter(w_q, requires_grad=False)
|
||||
self.process_dynamic_mxfp4_weights_after_loading(layer)
|
||||
elif self.rocm_use_aiter_fp4_asm_gemm:
|
||||
# shuffle weight scale
|
||||
weight_scale_shuffle = layer.weight_scale.data
|
||||
|
||||
@@ -322,20 +322,23 @@ def _shuffle_deepseek_fp8_moe_weights(
|
||||
block_k = 128
|
||||
num_experts = w13.shape[0]
|
||||
|
||||
w13_shuffled: list[torch.Tensor] = []
|
||||
w2_shuffled: list[torch.Tensor] = []
|
||||
M13, K13 = w13.shape[1], w13.shape[2]
|
||||
M2, K2 = w2.shape[1], w2.shape[2]
|
||||
w13_out = torch.empty(
|
||||
num_experts, K13 // block_k, M13, block_k, dtype=torch.uint8, device=w13.device
|
||||
)
|
||||
w2_out = torch.empty(
|
||||
num_experts, K2 // block_k, M2, block_k, dtype=torch.uint8, device=w2.device
|
||||
)
|
||||
|
||||
for i in range(num_experts):
|
||||
t13 = shuffle_matrix_a(w13[i].view(torch.uint8), epilogue_tile_m)
|
||||
t13 = convert_to_block_layout(t13, block_k)
|
||||
w13_shuffled.append(t13)
|
||||
w13_out[i] = convert_to_block_layout(t13, block_k)
|
||||
|
||||
t2 = shuffle_matrix_a(w2[i].view(torch.uint8), epilogue_tile_m)
|
||||
t2 = convert_to_block_layout(t2, block_k)
|
||||
w2_shuffled.append(t2)
|
||||
w2_out[i] = convert_to_block_layout(t2, block_k)
|
||||
|
||||
w13_out = torch.stack(w13_shuffled).view(torch.float8_e4m3fn)
|
||||
w2_out = torch.stack(w2_shuffled).view(torch.float8_e4m3fn)
|
||||
return w13_out, w2_out
|
||||
return w13_out.view(torch.float8_e4m3fn), w2_out.view(torch.float8_e4m3fn)
|
||||
|
||||
|
||||
def _shuffle_mxfp8_moe_weights(
|
||||
|
||||
@@ -24,7 +24,7 @@ logger = init_logger(__name__)
|
||||
|
||||
|
||||
def is_fp4_marlin_supported():
|
||||
return current_platform.has_device_capability(75)
|
||||
return current_platform.is_cuda() and current_platform.has_device_capability(75)
|
||||
|
||||
|
||||
def _nvfp4_compute_scale_factor(
|
||||
@@ -43,9 +43,9 @@ def _nvfp4_compute_scale_factor(
|
||||
ws_float = marlin_scales.float() * (2**7)
|
||||
nonzero_mask = ws_float > 0
|
||||
if nonzero_mask.any():
|
||||
min_val = ws_float[nonzero_mask].min()
|
||||
if min_val < 2:
|
||||
sf = (2 / min_val).log2().ceil().exp2()
|
||||
max_val = ws_float[nonzero_mask].max()
|
||||
if max_val < 448 * (2**7):
|
||||
sf = (448 * (2**7) / max_val).log2().floor().exp2()
|
||||
return sf.item()
|
||||
return 1.0
|
||||
|
||||
@@ -105,7 +105,9 @@ def nvfp4_marlin_process_scales(
|
||||
if scale_factor > 1.0:
|
||||
marlin_scales = (marlin_scales.float() * scale_factor).to(torch.half)
|
||||
|
||||
marlin_scales = (marlin_scales * (2**7)).view(torch.int16) << 1
|
||||
marlin_scales = marlin_scales * (2**7)
|
||||
marlin_scales[marlin_scales < 2] = 0
|
||||
marlin_scales = marlin_scales.view(torch.int16) << 1
|
||||
marlin_scales = marlin_scales.view(torch.float8_e4m3fn)
|
||||
marlin_scales = marlin_scales[:, 1::2].contiguous()
|
||||
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
from types import SimpleNamespace
|
||||
|
||||
import torch
|
||||
|
||||
from vllm.scalar_type import scalar_types
|
||||
@@ -11,9 +13,10 @@ __all__ = [
|
||||
]
|
||||
|
||||
FLOAT4_E2M1_MAX = scalar_types.float4_e2m1f.max()
|
||||
FLOAT4_E2M1_MAX_RECIPROCAL = 1 / FLOAT4_E2M1_MAX
|
||||
|
||||
kE2M1ToFloat = torch.tensor(
|
||||
[0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0], dtype=torch.float32
|
||||
kE2M1ToFloat_handle = SimpleNamespace(
|
||||
val=torch.tensor([0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0], dtype=torch.float32)
|
||||
)
|
||||
|
||||
|
||||
@@ -29,8 +32,9 @@ def break_fp4_bytes(a, dtype):
|
||||
# Vectorized sign and magnitude extraction
|
||||
signs = (combined & 0x08).to(torch.bool) # Sign bits
|
||||
abs_vals = (combined & 0x07).to(torch.long)
|
||||
|
||||
kE2M1 = kE2M1ToFloat_handle.val
|
||||
# Device-aware lookup and sign application
|
||||
kE2M1 = kE2M1ToFloat.to(device=a.device)
|
||||
values = kE2M1[abs_vals] * torch.where(signs, -1.0, 1.0)
|
||||
# Reshape to final form
|
||||
return values.reshape(m, n * 2).to(dtype=dtype)
|
||||
@@ -47,7 +51,12 @@ def convert_swizzled_to_linear(a_sf_swizzled: torch.Tensor, m, k, block_size):
|
||||
|
||||
|
||||
def dequantize_to_dtype(
|
||||
tensor_fp4, tensor_sf, global_scale, dtype, device, block_size=16
|
||||
tensor_fp4: torch.Tensor,
|
||||
tensor_sf: torch.Tensor,
|
||||
global_scale: torch.Tensor | float,
|
||||
dtype: torch.dtype,
|
||||
block_size: int = 16,
|
||||
swizzle: bool | None = True,
|
||||
):
|
||||
"""Dequantize the fp4 tensor back to high precision."""
|
||||
# Two fp4 values are packed into one uint8.
|
||||
@@ -57,8 +66,10 @@ def dequantize_to_dtype(
|
||||
tensor_f32 = break_fp4_bytes(tensor_fp4, torch.float32)
|
||||
tensor_f32 = tensor_f32.reshape(m, k // block_size, block_size)
|
||||
tensor_sf = tensor_sf.view(torch.float8_e4m3fn)
|
||||
tensor_sf = convert_swizzled_to_linear(tensor_sf, m, k, block_size)
|
||||
tensor_sf_dtype = tensor_sf.to(torch.float32) / global_scale
|
||||
|
||||
if swizzle:
|
||||
tensor_sf = convert_swizzled_to_linear(tensor_sf, m, k, block_size)
|
||||
tensor_sf_dtype = tensor_sf.to(torch.float32) * global_scale
|
||||
|
||||
# scale the tensor
|
||||
out = (tensor_f32 * tensor_sf_dtype.unsqueeze(-1)).reshape(m, k)
|
||||
@@ -67,7 +78,8 @@ def dequantize_to_dtype(
|
||||
|
||||
def get_reciprocal(x):
|
||||
if isinstance(x, torch.Tensor):
|
||||
return torch.where(x == 0, torch.tensor(0.0, dtype=x.dtype), 1.0 / x)
|
||||
# torch.where yields operation not permitted when stream is capturing.
|
||||
return 1.0 / (x + (x == 0) * 1e8)
|
||||
elif isinstance(x, (float, int)):
|
||||
return 0.0 if x == 0 else 1.0 / x
|
||||
else:
|
||||
@@ -94,7 +106,7 @@ def ref_nvfp4_quant(x, global_scale, block_size):
|
||||
m, n = x.shape
|
||||
x = torch.reshape(x, (m, n // block_size, block_size))
|
||||
vec_max = torch.max(torch.abs(x), dim=-1, keepdim=True)[0].to(torch.float32)
|
||||
scale = global_scale * (vec_max * get_reciprocal(FLOAT4_E2M1_MAX))
|
||||
scale = global_scale * (vec_max * FLOAT4_E2M1_MAX_RECIPROCAL)
|
||||
scale = torch.clamp(scale, max=448, min=-448)
|
||||
scale = scale.to(torch.float8_e4m3fn).to(torch.float32)
|
||||
output_scale = get_reciprocal(scale * get_reciprocal(global_scale))
|
||||
@@ -111,6 +123,7 @@ def run_nvfp4_emulations(
|
||||
weight: torch.Tensor,
|
||||
weight_scale_swizzled: torch.Tensor,
|
||||
weight_global_scale: torch.Tensor,
|
||||
swizzle: bool | None = True,
|
||||
):
|
||||
group_size = 16
|
||||
x_m, x_k = x.shape
|
||||
@@ -132,8 +145,8 @@ def run_nvfp4_emulations(
|
||||
weight_scale_swizzled.data,
|
||||
weight_global_scale,
|
||||
output_dtype,
|
||||
x.device,
|
||||
group_size,
|
||||
swizzle=swizzle,
|
||||
)
|
||||
|
||||
# matmul
|
||||
|
||||
@@ -17,31 +17,99 @@ from vllm.model_executor.layers.quantization.utils.marlin_utils_fp4 import (
|
||||
prepare_fp4_layer_for_marlin,
|
||||
)
|
||||
from vllm.model_executor.layers.quantization.utils.nvfp4_emulation_utils import (
|
||||
kE2M1ToFloat_handle,
|
||||
run_nvfp4_emulations,
|
||||
)
|
||||
from vllm.platforms import current_platform
|
||||
from vllm.utils.flashinfer import flashinfer_scaled_fp4_mm, has_flashinfer
|
||||
from vllm.utils.import_utils import has_fbgemm_gpu
|
||||
from vllm.utils.math_utils import round_up
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
# NOTE: This is ordered by preferred backend.
|
||||
# Example: if both are available, FLASHINFER_CUTLASS is preferred to VLLM_CUTLASS.
|
||||
class NvFp4LinearBackend(Enum):
|
||||
VLLM_CUTLASS = "cutlass"
|
||||
FLASHINFER_CUTLASS = "flashinfer-cutlass"
|
||||
VLLM_CUTLASS = "cutlass"
|
||||
MARLIN = "marlin"
|
||||
FLASHINFER_TRTLLM = "flashinfer-trtllm"
|
||||
FLASHINFER_CUDNN = "flashinfer-cudnn"
|
||||
FBGEMM = "fbgemm"
|
||||
MARLIN = "marlin"
|
||||
EMULATION = "emulation"
|
||||
|
||||
|
||||
NVFP4_LINEAR_BACKENDS = list(NvFp4LinearBackend)
|
||||
|
||||
|
||||
def is_backend_supported(backend: NvFp4LinearBackend) -> tuple[bool, str | None]:
|
||||
reason = None
|
||||
supported = True
|
||||
|
||||
if backend == NvFp4LinearBackend.FLASHINFER_CUTLASS:
|
||||
# cutlass_fp4_supported() checks that the vLLM NVFP4 kernels (both
|
||||
# quantization and GEMM) were compiled for the current SM version.
|
||||
# FlashInfer backends still rely on the vLLM quantization kernels,
|
||||
# so we gate them on the same check.
|
||||
supported = (
|
||||
cutlass_fp4_supported()
|
||||
and current_platform.has_device_capability(100)
|
||||
and has_flashinfer()
|
||||
)
|
||||
|
||||
if not supported:
|
||||
reason = "FlashInfer is required, >=sm_100 is required"
|
||||
elif backend == NvFp4LinearBackend.VLLM_CUTLASS:
|
||||
supported = cutlass_fp4_supported()
|
||||
if not supported:
|
||||
reason = "Cutlass is required"
|
||||
elif backend == NvFp4LinearBackend.MARLIN:
|
||||
supported = is_fp4_marlin_supported()
|
||||
if not supported:
|
||||
reason = "Marlin is required"
|
||||
elif backend in [
|
||||
NvFp4LinearBackend.FLASHINFER_TRTLLM,
|
||||
NvFp4LinearBackend.FLASHINFER_CUDNN,
|
||||
]:
|
||||
supported = has_flashinfer()
|
||||
if not supported:
|
||||
reason = "FlashInfer is required"
|
||||
elif backend == NvFp4LinearBackend.FBGEMM:
|
||||
supported = has_fbgemm_gpu()
|
||||
if not supported:
|
||||
reason = "fbgemm_gpu is required"
|
||||
elif backend == NvFp4LinearBackend.EMULATION:
|
||||
# e.g. AMD Instinct does not support native NVFP4.
|
||||
unsupported_reasons = {}
|
||||
for other_backend in NVFP4_LINEAR_BACKENDS:
|
||||
if other_backend == NvFp4LinearBackend.EMULATION:
|
||||
continue
|
||||
other_supported, other_reason = is_backend_supported(other_backend)
|
||||
if not other_supported:
|
||||
unsupported_reasons[other_backend] = other_reason
|
||||
|
||||
if unsupported_reasons:
|
||||
unsupported_reasons_str = "\n - ".join(
|
||||
[f"{b.value}: {r}" for b, r in unsupported_reasons.items()]
|
||||
)
|
||||
logger.warning_once(
|
||||
f"NVFP4 linear falling back to the slow and unoptimized "
|
||||
f"backend=NvFp4LinearBackend.EMULATION as no optimized backend is "
|
||||
f"available (unavailable reasons:\n - {unsupported_reasons_str}\n). "
|
||||
"In case you expect one of these backend to be used, "
|
||||
"please verify your environment."
|
||||
)
|
||||
|
||||
return supported, reason
|
||||
|
||||
|
||||
def select_nvfp4_linear_backend() -> NvFp4LinearBackend:
|
||||
"""
|
||||
Select the best available NVFP4 GEMM backend based on environment
|
||||
configuration and platform capabilities.
|
||||
"""
|
||||
backend: NvFp4LinearBackend | None = None
|
||||
selected_backend: NvFp4LinearBackend | None = None
|
||||
|
||||
if envs.VLLM_USE_FBGEMM:
|
||||
try:
|
||||
@@ -51,51 +119,36 @@ def select_nvfp4_linear_backend() -> NvFp4LinearBackend:
|
||||
"Backend fbgemm requires fbgemm.f4f4bf16 operator, "
|
||||
"Please install with: pip install fbgemm-gpu-genai"
|
||||
) from exc
|
||||
backend = NvFp4LinearBackend.FBGEMM
|
||||
selected_backend = NvFp4LinearBackend.FBGEMM
|
||||
elif envs.VLLM_USE_NVFP4_CT_EMULATIONS:
|
||||
backend = NvFp4LinearBackend.EMULATION
|
||||
selected_backend = NvFp4LinearBackend.EMULATION
|
||||
elif envs.VLLM_NVFP4_GEMM_BACKEND is None:
|
||||
# Auto-select best available backend.
|
||||
# cutlass_fp4_supported() checks that the vLLM NVFP4 kernels (both
|
||||
# quantization and GEMM) were compiled for the current SM version.
|
||||
# FlashInfer backends still rely on the vLLM quantization kernels,
|
||||
# so we gate them on the same check.
|
||||
if (
|
||||
cutlass_fp4_supported()
|
||||
and current_platform.has_device_capability(100)
|
||||
and has_flashinfer()
|
||||
):
|
||||
backend = NvFp4LinearBackend.FLASHINFER_CUTLASS
|
||||
elif cutlass_fp4_supported():
|
||||
backend = NvFp4LinearBackend.VLLM_CUTLASS
|
||||
elif is_fp4_marlin_supported():
|
||||
backend = NvFp4LinearBackend.MARLIN
|
||||
for backend in NVFP4_LINEAR_BACKENDS:
|
||||
supported, reason = is_backend_supported(backend)
|
||||
if supported:
|
||||
selected_backend = backend
|
||||
break
|
||||
else:
|
||||
backend = NvFp4LinearBackend(envs.VLLM_NVFP4_GEMM_BACKEND)
|
||||
selected_backend = NvFp4LinearBackend(envs.VLLM_NVFP4_GEMM_BACKEND)
|
||||
|
||||
# Validate that the backend is supported
|
||||
if backend in (
|
||||
NvFp4LinearBackend.FLASHINFER_CUTLASS,
|
||||
NvFp4LinearBackend.FLASHINFER_TRTLLM,
|
||||
NvFp4LinearBackend.FLASHINFER_CUDNN,
|
||||
):
|
||||
assert has_flashinfer(), f"FlashInfer is required for {backend}"
|
||||
assert cutlass_fp4_supported(), (
|
||||
f"{backend} requires vLLM NVFP4 quantization kernels compiled "
|
||||
f"for the current GPU (SM {current_platform.get_device_capability()})"
|
||||
)
|
||||
elif backend == NvFp4LinearBackend.VLLM_CUTLASS:
|
||||
assert cutlass_fp4_supported(), f"Cutlass is required for {backend}"
|
||||
elif backend == NvFp4LinearBackend.MARLIN:
|
||||
assert is_fp4_marlin_supported(), f"Marlin is required for {backend}"
|
||||
elif backend is None:
|
||||
if selected_backend is None:
|
||||
raise ValueError(
|
||||
f"No NVFP4 GEMM backend selected, "
|
||||
f"available backends: {list(NvFp4LinearBackend)}"
|
||||
f"available backends: {NVFP4_LINEAR_BACKENDS}"
|
||||
)
|
||||
|
||||
logger.info_once(f"Using {backend} for NVFP4 GEMM")
|
||||
return backend
|
||||
supported, reason = is_backend_supported(selected_backend)
|
||||
|
||||
if not supported:
|
||||
raise ValueError(
|
||||
f"The selected backend={selected_backend} is not supported in current "
|
||||
f"environment. Reason: {reason}. Current environment: "
|
||||
f"{envs.VLLM_USE_FBGEMM=}, {envs.VLLM_USE_NVFP4_CT_EMULATIONS=}, "
|
||||
f"{envs.VLLM_NVFP4_GEMM_BACKEND}."
|
||||
)
|
||||
|
||||
logger.info_once(f"Using {selected_backend} for NVFP4 GEMM")
|
||||
return selected_backend
|
||||
|
||||
|
||||
def prepare_weights_for_nvfp4_flashinfer_trtllm(
|
||||
@@ -183,6 +236,10 @@ def convert_to_nvfp4_linear_kernel_format(
|
||||
layer.weight = torch.nn.Parameter(weight, requires_grad=False)
|
||||
layer.weight_scale = torch.nn.Parameter(weight_scale, requires_grad=False)
|
||||
layer.weights_padding_cols = weights_padding_cols
|
||||
elif backend == NvFp4LinearBackend.EMULATION:
|
||||
# We can not call `.to(device)` during cuda graph capture - do it here instead.
|
||||
# (operation not permitted when stream is capturing)
|
||||
kE2M1ToFloat_handle.val = kE2M1ToFloat_handle.val.to(layer.weight.device)
|
||||
|
||||
|
||||
def apply_nvfp4_linear(
|
||||
@@ -190,6 +247,7 @@ def apply_nvfp4_linear(
|
||||
layer: torch.nn.Module,
|
||||
x: torch.Tensor,
|
||||
bias: torch.Tensor | None = None,
|
||||
swizzle: bool | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Apply NVFP4 linear transformation using the specified backend.
|
||||
@@ -220,6 +278,7 @@ def apply_nvfp4_linear(
|
||||
weight=weight,
|
||||
weight_scale_swizzled=weight_scale,
|
||||
weight_global_scale=weight_global_scale,
|
||||
swizzle=swizzle,
|
||||
)
|
||||
if bias is not None:
|
||||
out = out + bias
|
||||
|
||||
@@ -9,6 +9,7 @@ extract_hidden_states speculative decoding method.
|
||||
"""
|
||||
|
||||
from collections.abc import Iterable
|
||||
from dataclasses import replace
|
||||
from typing import ClassVar
|
||||
|
||||
import torch
|
||||
@@ -352,6 +353,10 @@ class ExtractHiddenStatesModel(nn.Module):
|
||||
|
||||
cache_config = vllm_config.cache_config
|
||||
|
||||
# Hidden states dtype should be independent of KV cache dtype.
|
||||
if cache_config is not None and is_quantized_kv_cache(cache_config.cache_dtype):
|
||||
cache_config = replace(cache_config, cache_dtype="auto")
|
||||
|
||||
# Create a single cache-only attention layer
|
||||
# Note: We set num_heads <- self.num_hidden_states
|
||||
# and head_size <- hidden_size so that we can insert
|
||||
|
||||
@@ -50,6 +50,7 @@ from .interfaces import (
|
||||
SupportsPP,
|
||||
)
|
||||
from .utils import (
|
||||
AutoWeightsLoader,
|
||||
PPMissingLayer,
|
||||
is_pp_missing_parameter,
|
||||
make_empty_intermediate_tensors_factory,
|
||||
@@ -495,6 +496,63 @@ class FalconH1Model(nn.Module):
|
||||
hidden_states = self.final_layernorm(hidden_states)
|
||||
return hidden_states
|
||||
|
||||
def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]:
|
||||
stacked_params_mapping = [
|
||||
# (param_name, shard_name, shard_id)
|
||||
("qkv_proj", "q_proj", "q"),
|
||||
("qkv_proj", "k_proj", "k"),
|
||||
("qkv_proj", "v_proj", "v"),
|
||||
("gate_up_proj", "gate_proj", 0),
|
||||
("gate_up_proj", "up_proj", 1),
|
||||
]
|
||||
|
||||
params_dict = dict(self.named_parameters())
|
||||
loaded_params: set[str] = set()
|
||||
for name, loaded_weight in weights:
|
||||
if "rotary_emb.inv_freq" in name:
|
||||
continue
|
||||
|
||||
if "A_log" in name:
|
||||
name = name.replace("A_log", "A")
|
||||
|
||||
if "mamba" in name:
|
||||
name = name.replace("mamba", "mamba.mamba")
|
||||
|
||||
if "scale" in name:
|
||||
# Remapping the name of kv-scale.
|
||||
name = maybe_remap_kv_scale_name(name, params_dict)
|
||||
if name is None:
|
||||
continue
|
||||
|
||||
for param_name, weight_name, shard_id in stacked_params_mapping:
|
||||
if weight_name not in name:
|
||||
continue
|
||||
|
||||
name = name.replace(weight_name, param_name)
|
||||
# Skip loading extra bias for GPTQ models.
|
||||
if name.endswith(".bias") and name not in params_dict:
|
||||
continue
|
||||
# Skip layers on other devices.
|
||||
if is_pp_missing_parameter(name, self):
|
||||
continue
|
||||
param = params_dict[name]
|
||||
weight_loader = param.weight_loader
|
||||
weight_loader(param, loaded_weight, shard_id)
|
||||
break
|
||||
else:
|
||||
# Skip loading extra bias for GPTQ models.
|
||||
if name.endswith(".bias") and name not in params_dict:
|
||||
continue
|
||||
if is_pp_missing_parameter(name, self):
|
||||
continue
|
||||
|
||||
param = params_dict[name]
|
||||
weight_loader = getattr(param, "weight_loader", default_weight_loader)
|
||||
weight_loader(param, loaded_weight)
|
||||
loaded_params.add(name)
|
||||
|
||||
return loaded_params
|
||||
|
||||
|
||||
class FalconH1ForCausalLM(
|
||||
nn.Module,
|
||||
@@ -632,62 +690,8 @@ class FalconH1ForCausalLM(
|
||||
return logits
|
||||
|
||||
def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]:
|
||||
stacked_params_mapping = [
|
||||
# (param_name, shard_name, shard_id)
|
||||
("qkv_proj", "q_proj", "q"),
|
||||
("qkv_proj", "k_proj", "k"),
|
||||
("qkv_proj", "v_proj", "v"),
|
||||
("gate_up_proj", "gate_proj", 0),
|
||||
("gate_up_proj", "up_proj", 1),
|
||||
]
|
||||
|
||||
params_dict = dict(self.named_parameters())
|
||||
loaded_params: set[str] = set()
|
||||
for name, loaded_weight in weights:
|
||||
if "rotary_emb.inv_freq" in name:
|
||||
continue
|
||||
|
||||
if "A_log" in name:
|
||||
name = name.replace("A_log", "A")
|
||||
|
||||
if "mamba" in name:
|
||||
name = name.replace("mamba", "mamba.mamba")
|
||||
|
||||
if "scale" in name:
|
||||
# Remapping the name of kv-scale.
|
||||
name = maybe_remap_kv_scale_name(name, params_dict)
|
||||
if name is None:
|
||||
continue
|
||||
|
||||
for param_name, weight_name, shard_id in stacked_params_mapping:
|
||||
if weight_name not in name:
|
||||
continue
|
||||
|
||||
name = name.replace(weight_name, param_name)
|
||||
# Skip loading extra bias for GPTQ models.
|
||||
if name.endswith(".bias") and name not in params_dict:
|
||||
continue
|
||||
# Skip layers on other devices.
|
||||
if is_pp_missing_parameter(name, self):
|
||||
continue
|
||||
param = params_dict[name]
|
||||
weight_loader = param.weight_loader
|
||||
weight_loader(param, loaded_weight, shard_id)
|
||||
break
|
||||
else:
|
||||
# Skip loading extra bias for GPTQ models.
|
||||
if name.endswith(".bias") and name not in params_dict:
|
||||
continue
|
||||
if is_pp_missing_parameter(name, self):
|
||||
continue
|
||||
if self.tie_word_embeddings and "lm_head" in name:
|
||||
continue
|
||||
|
||||
param = params_dict[name]
|
||||
weight_loader = getattr(param, "weight_loader", default_weight_loader)
|
||||
weight_loader(param, loaded_weight)
|
||||
loaded_params.add(name)
|
||||
|
||||
if self.tie_word_embeddings:
|
||||
loaded_params.add("lm_head.weight")
|
||||
return loaded_params
|
||||
loader = AutoWeightsLoader(
|
||||
self,
|
||||
skip_prefixes=(["lm_head."] if self.tie_word_embeddings else None),
|
||||
)
|
||||
return loader.load_weights(weights)
|
||||
|
||||
@@ -288,6 +288,35 @@ class NanoNemotronVLProcessingInfo(BaseProcessingInfo):
|
||||
max_num_tiles=max_num_tiles,
|
||||
)
|
||||
|
||||
def get_dummy_image_size_and_max_tokens(
|
||||
self, mm_counts: Mapping[str, int]
|
||||
) -> tuple[tuple[int, int], int]:
|
||||
processor = self.get_hf_processor()
|
||||
num_images = mm_counts.get("image", 0)
|
||||
|
||||
if tiler := processor.dynamic_tiler:
|
||||
budget = tiler.max_num_tokens_available(text_prompt_length=num_images)
|
||||
target_width, target_height = (
|
||||
tiler.width_and_height_for_max_num_tokens_available(budget)
|
||||
)
|
||||
return (
|
||||
(target_width, target_height),
|
||||
tiler._get_num_embeddings(target_width, target_height),
|
||||
)
|
||||
|
||||
max_num_tiles = processor.max_num_tiles
|
||||
target_width, target_height = self.get_image_size_with_most_features(
|
||||
max_num_tiles
|
||||
)
|
||||
return (
|
||||
(target_width, target_height),
|
||||
processor.get_num_image_tokens(
|
||||
image_width=target_width,
|
||||
image_height=target_height,
|
||||
max_num_tiles=max_num_tiles,
|
||||
),
|
||||
)
|
||||
|
||||
def get_num_frames_with_most_features(
|
||||
self,
|
||||
seq_len: int,
|
||||
@@ -306,6 +335,26 @@ class NanoNemotronVLProcessingInfo(BaseProcessingInfo):
|
||||
max_frames_per_video = max_tubelets_per_video * T
|
||||
return max(max_frames_per_video, 1)
|
||||
|
||||
def get_mm_max_tokens_per_item(
|
||||
self, seq_len: int, mm_counts: Mapping[str, int]
|
||||
) -> Mapping[str, int]:
|
||||
mm_max_tokens: dict[str, int] = {}
|
||||
|
||||
if mm_counts.get("image", 0) > 0:
|
||||
_, mm_max_tokens["image"] = self.get_dummy_image_size_and_max_tokens(
|
||||
mm_counts
|
||||
)
|
||||
|
||||
if mm_counts.get("video", 0) > 0:
|
||||
assert self.supports_video
|
||||
mm_max_tokens["video"] = seq_len
|
||||
|
||||
if mm_counts.get("audio", 0) > 0:
|
||||
assert self.supports_audio
|
||||
mm_max_tokens["audio"] = seq_len
|
||||
|
||||
return mm_max_tokens
|
||||
|
||||
|
||||
class NanoNemotronVLMultiModalProcessor(
|
||||
BaseMultiModalProcessor[NanoNemotronVLProcessingInfo]
|
||||
@@ -708,17 +757,10 @@ class NanoNemotronVLDummyInputsBuilder(
|
||||
mm_options: Mapping[str, BaseDummyOptions],
|
||||
) -> MultiModalDataDict:
|
||||
num_images = mm_counts.get("image", 0)
|
||||
(target_width, target_height), _ = (
|
||||
self.info.get_dummy_image_size_and_max_tokens(mm_counts)
|
||||
)
|
||||
processor = self.info.get_hf_processor()
|
||||
if tiler := processor.dynamic_tiler:
|
||||
budget = tiler.max_num_tokens_available(text_prompt_length=num_images)
|
||||
target_width, target_height = (
|
||||
tiler.width_and_height_for_max_num_tokens_available(budget)
|
||||
)
|
||||
else:
|
||||
max_num_tiles = 12
|
||||
target_width, target_height = self.info.get_image_size_with_most_features(
|
||||
max_num_tiles
|
||||
)
|
||||
|
||||
image_overrides = mm_options.get("image")
|
||||
|
||||
|
||||
@@ -37,6 +37,7 @@ from vllm.logger import init_logger
|
||||
from vllm.reasoning.abs_reasoning_parsers import ReasoningParser
|
||||
from vllm.tokenizers import TokenizerLike
|
||||
from vllm.tool_parsers.abstract_tool_parser import ToolParser
|
||||
from vllm.tool_parsers.utils import Tool
|
||||
from vllm.utils import random_uuid
|
||||
|
||||
logger = init_logger(__name__)
|
||||
@@ -542,10 +543,10 @@ class _WrappedParser(DelegatingParser):
|
||||
reasoning_parser_cls: type[ReasoningParser] | None = None
|
||||
tool_parser_cls: type[ToolParser] | None = None
|
||||
|
||||
def __init__(self, tokenizer: TokenizerLike):
|
||||
def __init__(self, tokenizer: TokenizerLike, tools: list[Tool] | None = None):
|
||||
super().__init__(tokenizer)
|
||||
# Instantiate the underlying parsers from class attributes
|
||||
if self.__class__.reasoning_parser_cls is not None:
|
||||
self._reasoning_parser = self.__class__.reasoning_parser_cls(tokenizer)
|
||||
if self.__class__.tool_parser_cls is not None:
|
||||
self._tool_parser = self.__class__.tool_parser_cls(tokenizer)
|
||||
self._tool_parser = self.__class__.tool_parser_cls(tokenizer, tools)
|
||||
|
||||
@@ -409,6 +409,7 @@ class RocmPlatform(Platform):
|
||||
"mxfp4",
|
||||
"torchao",
|
||||
"bitsandbytes",
|
||||
"modelopt_fp4",
|
||||
]
|
||||
|
||||
@classmethod
|
||||
|
||||
@@ -2,7 +2,9 @@
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
from typing import Any
|
||||
|
||||
from packaging.version import Version
|
||||
from transformers import PretrainedConfig, WhisperConfig
|
||||
from transformers import __version__ as TRANSFORMERS_VERSION
|
||||
|
||||
from vllm.logger import init_logger
|
||||
|
||||
@@ -134,6 +136,10 @@ def _remap_mistral_yarn_args(config: dict) -> dict:
|
||||
# Cast to remove Transformers > v5 type warnings
|
||||
config["rope_parameters"][new_name] = cast(yarn_config.pop(old_name))
|
||||
|
||||
# Ignore apply_yarn_scaling in Transformers > v5 RoPE validation to remove warnings
|
||||
if Version(TRANSFORMERS_VERSION) >= Version("5.3.0.dev0"):
|
||||
config["ignore_keys_at_rope_validation"] = {"apply_yarn_scaling"}
|
||||
|
||||
assert len(yarn_config) == 0, f"Unparsed yarn config: {yarn_config}"
|
||||
|
||||
return config
|
||||
|
||||
@@ -461,3 +461,8 @@ def has_aiter() -> bool:
|
||||
def has_mori() -> bool:
|
||||
"""Whether the optional `mori` package is available."""
|
||||
return _has_module("mori")
|
||||
|
||||
|
||||
def has_fbgemm_gpu() -> bool:
|
||||
"""Whether the optional `fbgemm_gpu` package is available."""
|
||||
return _has_module("fbgemm_gpu")
|
||||
|
||||
@@ -13,6 +13,7 @@ from vllm.model_executor.layers.attention.mla_attention import (
|
||||
MLACommonImpl,
|
||||
MLACommonMetadata,
|
||||
)
|
||||
from vllm.platforms import current_platform
|
||||
from vllm.platforms.interface import DeviceCapability
|
||||
from vllm.triton_utils import triton
|
||||
from vllm.utils.torch_utils import is_quantized_kv_cache
|
||||
@@ -116,7 +117,7 @@ class TritonMLAImpl(MLACommonImpl[MLACommonMetadata]):
|
||||
if is_quantized_kv_cache(self.kv_cache_dtype):
|
||||
self.supports_quant_query_input = False
|
||||
|
||||
self._sm_count = torch.cuda.get_device_properties(0).multi_processor_count
|
||||
self._sm_count = current_platform.num_compute_units()
|
||||
|
||||
def _flash_attn_varlen_diff_headdims(
|
||||
self, q, k, v, return_softmax_lse=False, softmax_scale=None, **kwargs
|
||||
|
||||
@@ -228,7 +228,6 @@ if current_platform.is_rocm():
|
||||
num_kv_heads,
|
||||
BLOCK_SIZE: tl.constexpr,
|
||||
QUANT: tl.constexpr,
|
||||
IS_FNUZ: tl.constexpr,
|
||||
):
|
||||
tid = tl.program_id(0)
|
||||
head_id = tl.program_id(1)
|
||||
@@ -314,7 +313,6 @@ if current_platform.is_rocm():
|
||||
num_kv_heads,
|
||||
BLOCK_SIZE=head_size,
|
||||
QUANT=QUANT,
|
||||
IS_FNUZ=current_platform.fp8_dtype() == torch.float8_e4m3fnuz,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -93,9 +93,8 @@ class FilterReusedOffloadingManager(OffloadingManager):
|
||||
]
|
||||
|
||||
# Delegate to the backing manager with only the eligible hashes.
|
||||
# Passing an empty list is intentional and safe — both
|
||||
# LRUOffloadingManager and ARCOffloadingManager handle it correctly,
|
||||
# returning a PrepareStoreOutput with empty lists.
|
||||
# Passing an empty list is intentional and safe — CPUOffloadingManager
|
||||
# handles it correctly, returning a PrepareStoreOutput with empty lists.
|
||||
return self._backing.prepare_store(eligible)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
@@ -17,7 +17,6 @@ class AsyncOutput(AsyncModelRunnerOutput):
|
||||
num_sampled_tokens: torch.Tensor,
|
||||
main_stream: torch.cuda.Stream,
|
||||
copy_stream: torch.cuda.Stream,
|
||||
copy_event: torch.cuda.Event,
|
||||
):
|
||||
# NOTE(woosuk): We must retain references to the GPU tensors,
|
||||
# as the copy operations are performed on a different CUDA stream than
|
||||
@@ -25,7 +24,7 @@ class AsyncOutput(AsyncModelRunnerOutput):
|
||||
self.model_runner_output = model_runner_output
|
||||
self.sampler_output = sampler_output
|
||||
self.num_sampled_tokens = num_sampled_tokens
|
||||
self.copy_event = copy_event
|
||||
self.copy_event = torch.cuda.Event()
|
||||
|
||||
with stream(copy_stream, main_stream):
|
||||
copy_stream.wait_stream(main_stream)
|
||||
@@ -78,12 +77,11 @@ class AsyncPoolingOutput(AsyncModelRunnerOutput):
|
||||
is_valid: torch.Tensor | None,
|
||||
main_stream: torch.cuda.Stream,
|
||||
copy_stream: torch.cuda.Stream,
|
||||
copy_event: torch.cuda.Event,
|
||||
):
|
||||
self.model_runner_output = model_runner_output
|
||||
self.pooler_output = pooler_output
|
||||
self.is_valid = is_valid
|
||||
self.copy_event = copy_event
|
||||
self.copy_event = torch.cuda.Event()
|
||||
|
||||
with stream(copy_stream, main_stream):
|
||||
copy_stream.wait_stream(main_stream)
|
||||
|
||||
@@ -130,7 +130,6 @@ class GPUModelRunner(LoRAModelRunnerMixin):
|
||||
|
||||
self.use_async_scheduling = self.scheduler_config.async_scheduling
|
||||
self.output_copy_stream = torch.cuda.Stream(self.device)
|
||||
self.output_copy_event = torch.cuda.Event()
|
||||
|
||||
# Pipeline parallelism.
|
||||
self.use_pp = self.parallel_config.pipeline_parallel_size > 1
|
||||
@@ -1180,7 +1179,6 @@ class GPUModelRunner(LoRAModelRunnerMixin):
|
||||
num_sampled_tokens=num_sampled,
|
||||
main_stream=self.main_stream,
|
||||
copy_stream=self.output_copy_stream,
|
||||
copy_event=self.output_copy_event,
|
||||
)
|
||||
|
||||
mm_inputs: tuple[list[torch.Tensor], torch.Tensor] | None = None
|
||||
@@ -1270,7 +1268,6 @@ class GPUModelRunner(LoRAModelRunnerMixin):
|
||||
is_valid=is_valid,
|
||||
main_stream=self.main_stream,
|
||||
copy_stream=self.output_copy_stream,
|
||||
copy_event=self.output_copy_event,
|
||||
)
|
||||
|
||||
self.postprocess_pool(input_batch)
|
||||
|
||||
Reference in New Issue
Block a user