Compare commits

..
Author SHA1 Message Date
yewentao256 4ac64ec057 deprecate embed&token_classify
Signed-off-by: yewentao256 <zhyanwentao@126.com>
2026-05-25 13:34:43 +00:00
66 changed files with 1243 additions and 8663 deletions
@@ -1,208 +0,0 @@
#!/usr/bin/env python3
"""Aggregate per-step coverage JSON files into a test-selection mapping.
Downloads all coverage_*.json artifacts from the current Buildkite build,
then produces two output files:
1. coverage_map.json — inverted index: {source_file: [step_keys]}
Used by the pipeline generator to determine which steps to trigger.
2. step_coverage.json — forward index: {step_key: [source_files]}
Useful for debugging and understanding test coverage.
Usage:
# Run as a Buildkite step at the end of nightly CI
python3 .buildkite/scripts/coverage/aggregate-coverage.py
# Or locally with downloaded artifacts
python3 .buildkite/scripts/coverage/aggregate-coverage.py --local-dir ./artifacts/
"""
import argparse
import json
import os
import subprocess
import sys
import tempfile
from collections import defaultdict
from pathlib import Path
def download_artifacts(dest_dir: str) -> list[str]:
"""Download all coverage_*.json artifacts from the current build."""
try:
subprocess.run(
["buildkite-agent", "artifact", "download", "coverage_*.json", dest_dir],
check=True,
capture_output=True,
text=True,
)
except FileNotFoundError:
print("buildkite-agent not found, skipping download", file=sys.stderr)
return []
except subprocess.CalledProcessError as e:
print(f"Artifact download failed: {e.stderr}", file=sys.stderr)
return []
return list(Path(dest_dir).glob("coverage_*.json"))
def load_coverage_files(files: list[Path]) -> dict[str, list[str]]:
"""Load coverage JSON files and extract source files per step.
Returns: {step_key: [source_files]}
"""
step_coverage = {}
for filepath in files:
filename = filepath.name
# coverage_<step_key>.json -> step_key
step_key = filename.removeprefix("coverage_").removesuffix(".json")
try:
with open(filepath) as f:
data = json.load(f)
except (json.JSONDecodeError, OSError) as e:
print(f"Warning: skipping {filename}: {e}", file=sys.stderr)
continue
source_files = []
for fpath, fdata in data.get("files", {}).items():
# Skip files with zero executed lines — coverage.py reports
# all files in the source tree, not just those actually run.
# Supports both full format (summary.covered_lines) and
# stripped format (covered_lines directly).
covered = fdata.get("covered_lines") or fdata.get("summary", {}).get("covered_lines", 0)
if covered == 0:
continue
# If function-level data is available, skip import-only files
# (files where only module-level code ran but no named functions
# were actually called).
funcs_called = fdata.get("functions_called")
if funcs_called is not None and funcs_called == 0:
continue
# Normalize paths to be relative to the vllm package root.
# coverage.py may report absolute paths or paths relative to
# the installed package location. We only care about files
# under the vllm/ directory.
normalized = _normalize_path(fpath)
if normalized:
source_files.append(normalized)
if source_files:
step_coverage[step_key] = sorted(set(source_files))
print(f" {step_key}: {len(source_files)} source files")
return step_coverage
def _normalize_path(path: str) -> str | None:
"""Normalize a coverage path to a vllm-relative path.
Returns None for paths outside the vllm package (tests, third-party, etc).
"""
# Strip common prefixes from installed package paths
markers = ["/site-packages/", "/dist-packages/", "/vllm-workspace/src/"]
for marker in markers:
idx = path.find(marker)
if idx != -1:
path = path[idx + len(marker):]
break
# Also handle paths that are already relative
if path.startswith("vllm/"):
return path
# Handle absolute paths that contain /vllm/
idx = path.find("/vllm/")
if idx != -1:
return path[idx + 1:]
return None
def build_inverted_index(
step_coverage: dict[str, list[str]],
) -> dict[str, list[str]]:
"""Build {source_file: [step_keys]} from {step_key: [source_files]}."""
inverted = defaultdict(list)
for step_key, source_files in step_coverage.items():
for src_file in source_files:
inverted[src_file].append(step_key)
# Sort step lists for deterministic output
return {k: sorted(v) for k, v in sorted(inverted.items())}
def main():
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument(
"--local-dir",
help="Directory containing coverage_*.json files (skip artifact download)",
)
parser.add_argument(
"--output-dir",
default=".",
help="Directory to write output files (default: cwd)",
)
args = parser.parse_args()
if args.local_dir:
artifact_dir = args.local_dir
files = list(Path(artifact_dir).glob("coverage_*.json"))
else:
artifact_dir = tempfile.mkdtemp(prefix="coverage_artifacts_")
files = download_artifacts(artifact_dir)
if not files:
print("No coverage files found. Nothing to aggregate.")
sys.exit(0)
print(f"Found {len(files)} coverage files:")
# Build the forward index: step -> source files
step_coverage = load_coverage_files(files)
if not step_coverage:
print("No valid coverage data found.")
sys.exit(0)
# Build the inverted index: source file -> steps
coverage_map = build_inverted_index(step_coverage)
# Write outputs
output_dir = Path(args.output_dir)
output_dir.mkdir(parents=True, exist_ok=True)
step_coverage_path = output_dir / "step_coverage.json"
with open(step_coverage_path, "w") as f:
json.dump(step_coverage, f, indent=2)
print(f"\nWrote {step_coverage_path} ({len(step_coverage)} steps)")
coverage_map_path = output_dir / "coverage_map.json"
with open(coverage_map_path, "w") as f:
json.dump(coverage_map, f, indent=2)
print(f"Wrote {coverage_map_path} ({len(coverage_map)} source files)")
# Summary stats
total_files = len(coverage_map)
total_mappings = sum(len(v) for v in coverage_map.values())
print(f"\nSummary: {total_files} source files mapped to "
f"{len(step_coverage)} steps ({total_mappings} total mappings)")
# Upload aggregated files as artifacts
for output_file in [step_coverage_path, coverage_map_path]:
try:
subprocess.run(
["buildkite-agent", "artifact", "upload", str(output_file)],
check=True,
capture_output=True,
text=True,
)
print(f"Uploaded {output_file}")
except (FileNotFoundError, subprocess.CalledProcessError):
pass # Not in Buildkite or upload failed — that's fine for local runs
if __name__ == "__main__":
main()
@@ -1,42 +0,0 @@
#!/bin/bash
# Upload coverage data for the current Buildkite step.
# Called automatically at the end of each step when COLLECT_COVERAGE=1.
#
# Expects:
# - .coverage.${BUILDKITE_STEP_KEY} data file from coverage run --append
# - BUILDKITE_STEP_KEY, BUILDKITE_BUILD_NUMBER env vars
#
# Produces:
# - coverage_${BUILDKITE_STEP_KEY}.json uploaded as a Buildkite artifact
set -euo pipefail
STEP_KEY="${BUILDKITE_STEP_KEY:-unknown}"
DATA_FILE=".coverage.${STEP_KEY}"
OUTPUT_JSON="coverage_${STEP_KEY}.json"
if [ ! -f "$DATA_FILE" ]; then
echo "~~~ No coverage data file found ($DATA_FILE), skipping upload"
exit 0
fi
echo "~~~ :bar_chart: Exporting coverage data for step: ${STEP_KEY}"
coverage json \
--data-file="$DATA_FILE" \
-o "$OUTPUT_JSON" \
--omit='*/tests/*,*/test_*,*/__pycache__/*' \
2>&1 || {
echo "Warning: coverage json export failed, skipping"
exit 0
}
FILE_COUNT=$(python3 -c "import json; d=json.load(open('$OUTPUT_JSON')); print(len(d.get('files', {})))" 2>/dev/null || echo "?")
echo "Coverage captured ${FILE_COUNT} source files for step ${STEP_KEY}"
buildkite-agent artifact upload "$OUTPUT_JSON" 2>&1 || {
echo "Warning: artifact upload failed"
exit 0
}
echo "Uploaded $OUTPUT_JSON"
-5
View File
@@ -644,11 +644,6 @@ steps:
- pytest -v -s plugins_tests/test_terratorch_io_processor_plugins.py
- pip uninstall prithvi_io_processor_plugin -y
# END: `io_processor` plugins test
# BEGIN: `bge_m3_sparse io_processor` test
- pip install -e ./plugins/bge_m3_sparse_plugin
- pytest -v -s plugins_tests/test_bge_m3_sparse_io_processor_plugins.py
- pip uninstall bge_m3_sparse_plugin -y
# END: `bge_m3_sparse io_processor` test
# BEGIN: `stat_logger` plugins test
- pip install -e ./plugins/vllm_add_dummy_stat_logger
- pytest -v -s plugins_tests/test_stats_logger_plugins.py
-4
View File
@@ -23,10 +23,6 @@ steps:
- pip install -e ./plugins/prithvi_io_processor_plugin
- pytest -v -s plugins_tests/test_terratorch_io_processor_plugins.py
- pip uninstall prithvi_io_processor_plugin -y
# test bge_m3_sparse io_processor plugin
- pip install -e ./plugins/bge_m3_sparse_plugin
- pytest -v -s plugins_tests/test_bge_m3_sparse_io_processor_plugins.py
- pip uninstall bge_m3_sparse_plugin -y
# end io_processor plugins test
# begin stat_logger plugins test
- pip install -e ./plugins/vllm_add_dummy_stat_logger
@@ -1,465 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Benchmark the fused MoE-LoRA fast path (one-shot) vs two-kernel baseline.
The "one_shot" provider goes through `vllm.lora.ops.triton_ops.fused_moe_lora`
which dispatches to the single-kernel one-shot implementation when
fully_sharded=False (the prefill default).
The "two_kernel" provider drives `fused_moe_lora_shrink` + `fused_moe_lora_expand`
directly, bypassing the dispatch and matching the legacy two-kernel path's
work distribution. This isolates the win from kernel fusion.
Run:
.venv/bin/python -m benchmarks.kernels.benchmark_fused_moe_lora_one_shot
.venv/bin/python -m benchmarks.kernels.benchmark_fused_moe_lora_one_shot \\
--model qwen3moe
"""
from __future__ import annotations
import argparse
import os
import random
import torch
from vllm import _custom_ops as ops
from vllm.lora.ops.triton_ops import (
fused_moe_lora,
fused_moe_lora_expand,
fused_moe_lora_shrink,
)
from vllm.triton_utils import triton
DTYPE = torch.bfloat16
DEVICE = "cuda"
# ----- input fabrication -----------------------------------------------------
def _round_up(x: int, base: int) -> int:
return ((x + base - 1) // base) * base
def _ceildiv(x: int, y: int) -> int:
return (x + y - 1) // y
def _assign_loras(num_tokens: int, num_sequences: int, max_loras: int) -> torch.Tensor:
tokens_per_seq = num_tokens // num_sequences
rem = num_tokens % num_sequences
out = torch.empty(num_tokens, dtype=torch.int32)
start = 0
for i in range(num_sequences):
end = start + tokens_per_seq + (1 if i < rem else 0)
out[start:end] = random.randint(0, max_loras - 1)
start = end
return out
def _assign_experts(num_tokens: int, num_experts: int, top_k: int):
expert_indices = torch.empty((num_tokens, top_k), dtype=torch.int32)
for i in range(num_tokens):
expert_indices[i] = torch.randperm(num_experts)[:top_k]
weights = torch.rand((num_tokens, top_k), dtype=torch.float32)
weights = weights / weights.sum(dim=1, keepdim=True)
return expert_indices, weights
def _make_inputs(
M: int,
K: int,
N_per_slice: int,
rank: int,
num_experts: int,
top_k: int,
max_loras: int,
num_slices: int,
block_size_m: int,
):
"""Mirrors the production caller's tensor layout."""
torch.manual_seed(0)
random.seed(0)
num_sequences = max(1, min(M, 8))
topk_ids_cpu, topk_weights_cpu = _assign_experts(M, num_experts, top_k)
token_lora_cpu = _assign_loras(M, num_sequences, max_loras)
lora_ids_cpu = torch.full((max_loras + 1,), -1, dtype=torch.int32)
uniq = torch.unique(token_lora_cpu, sorted=True)
lora_ids_cpu[: uniq.size(0)].copy_(uniq)
topk_ids = topk_ids_cpu.to(DEVICE)
topk_weights = topk_weights_cpu.to(device=DEVICE, dtype=DTYPE)
token_lora_mapping = token_lora_cpu.to(DEVICE)
lora_ids = lora_ids_cpu.to(DEVICE)
adapter_enabled = torch.ones(max_loras + 1, dtype=torch.int32, device=DEVICE)
lora_a = [
torch.randn((max_loras, num_experts, rank, K), dtype=DTYPE, device=DEVICE)
/ max(K, 1) ** 0.5
for _ in range(num_slices)
]
lora_b = [
torch.randn(
(max_loras, num_experts, N_per_slice, rank),
dtype=DTYPE,
device=DEVICE,
)
/ max(rank, 1) ** 0.5
for _ in range(num_slices)
]
hidden = torch.randn((M, K), dtype=DTYPE, device=DEVICE)
out_template = torch.zeros(
(M, top_k, num_slices * N_per_slice), dtype=DTYPE, device=DEVICE
)
# Sorted-path metadata (the prefill default).
max_pad = topk_ids.numel() + num_experts * (block_size_m - 1)
max_pad = _round_up(max_pad, block_size_m)
max_blocks = _ceildiv(max_pad, block_size_m)
sorted_token_ids = torch.empty(
(max_loras * max_pad,), dtype=torch.int32, device=DEVICE
)
expert_ids = torch.empty(
(max_loras * max_blocks,), dtype=torch.int32, device=DEVICE
)
num_post = torch.empty((max_loras,), dtype=torch.int32, device=DEVICE)
ops.moe_lora_align_block_size(
topk_ids,
token_lora_mapping,
num_experts,
block_size_m,
max_loras,
max_pad,
max_blocks,
sorted_token_ids,
expert_ids,
num_post,
adapter_enabled,
lora_ids,
)
expert_ids = expert_ids.view(max_loras, -1).contiguous()
sorted_token_ids = sorted_token_ids.view(max_loras, -1).contiguous()
num_active = torch.tensor([max_loras + 1], dtype=torch.int32, device="cpu")
return dict(
hidden=hidden,
lora_a=lora_a,
lora_b=lora_b,
topk_weights=topk_weights,
sorted_token_ids=sorted_token_ids,
expert_ids=expert_ids,
num_post=num_post,
token_lora_mapping=token_lora_mapping,
lora_ids=lora_ids,
num_active=num_active,
adapter_enabled=adapter_enabled,
out_template=out_template,
# bookkeeping
M=M,
K=K,
N_per_slice=N_per_slice,
rank=rank,
num_experts=num_experts,
top_k=top_k,
max_loras=max_loras,
num_slices=num_slices,
block_size_m=block_size_m,
)
# ----- providers -------------------------------------------------------------
def _run_one_shot(inp: dict):
"""Drive `fused_moe_lora` with fully_sharded=False -> one-shot fast path."""
out = inp["out_template"].clone()
fused_moe_lora(
out,
inp["hidden"],
inp["lora_a"],
inp["lora_b"],
inp["topk_weights"],
inp["sorted_token_ids"],
inp["expert_ids"],
inp["num_post"],
inp["token_lora_mapping"],
inp["rank"],
inp["top_k"],
inp["lora_ids"],
inp["num_active"],
inp["adapter_enabled"],
inp["block_size_m"],
64,
32,
8,
4,
3,
1,
inp["block_size_m"],
64,
32,
8,
4,
3,
1,
False,
False,
0,
)
return out
def _run_two_kernel(inp: dict):
"""Drive `fused_moe_lora_shrink` + `fused_moe_lora_expand` directly,
bypassing the dispatch. Matches the legacy two-kernel work distribution.
"""
M = inp["M"]
top_k = inp["top_k"]
rank = inp["rank"]
num_slices = inp["num_slices"]
N_per_slice = inp["N_per_slice"]
K = inp["K"]
num_experts = inp["num_experts"]
block_m = inp["block_size_m"]
intermediate = torch.zeros((num_slices, M, top_k, rank), dtype=DTYPE, device=DEVICE)
out = inp["out_template"].clone()
EM = inp["sorted_token_ids"].shape[1]
num_tokens = M * top_k
fused_moe_lora_shrink(
intermediate,
inp["hidden"],
inp["lora_a"],
inp["topk_weights"],
inp["sorted_token_ids"],
inp["expert_ids"],
inp["num_post"],
inp["token_lora_mapping"],
top_k,
inp["lora_ids"],
inp["adapter_enabled"],
torch.device(DEVICE),
rank,
M,
EM,
K,
num_tokens,
num_experts,
num_slices,
block_m,
64,
32,
8,
4,
3,
1,
inp["num_active"],
False,
)
fused_moe_lora_expand(
out,
intermediate,
inp["lora_b"],
inp["topk_weights"],
inp["sorted_token_ids"],
inp["expert_ids"],
inp["num_post"],
inp["token_lora_mapping"],
top_k,
inp["lora_ids"],
inp["adapter_enabled"],
torch.device(DEVICE),
rank,
M,
EM,
K,
num_tokens,
num_experts,
num_slices,
rank,
N_per_slice,
block_m,
64,
32,
8,
4,
3,
1,
inp["num_active"],
False,
0,
)
return out
PROVIDER_FNS = {
"one_shot": _run_one_shot,
"two_kernel": _run_two_kernel,
}
# ----- model presets ---------------------------------------------------------
MODEL_PRESETS: dict[str, dict] = {
# Mixtral-8x7B style: E=8, top_k=2, hidden=4096, intermediate=14336
"mixtral": dict(
K=4096,
N_per_slice=7168,
num_experts=8,
top_k=2,
max_loras=4,
num_slices=2,
block_size_m=64,
),
# Qwen3-MoE / DeepSeek-V2 style: E=64, top_k=8, hidden=2048, inter=1408
"qwen3moe": dict(
K=2048,
N_per_slice=1408,
num_experts=64,
top_k=8,
max_loras=4,
num_slices=2,
block_size_m=64,
),
# GLM-5.1 (zai-org/GLM-5.1-FP8): E=256, top_k=8, hidden=6144,
# moe_intermediate=2048
"glm5_1": dict(
K=6144,
N_per_slice=2048,
num_experts=256,
top_k=8,
max_loras=4,
num_slices=2,
block_size_m=64,
),
}
M_RANGE = [16, 64, 256, 1024, 4096, 16384]
RANK_RANGE = [8, 16, 32, 64]
def get_benchmark(model: str, max_loras: int | None = None):
preset = dict(MODEL_PRESETS[model])
if max_loras is not None:
preset["max_loras"] = max_loras
@triton.testing.perf_report(
triton.testing.Benchmark(
x_names=["M", "rank"],
x_vals=[(M, R) for M in M_RANGE for R in RANK_RANGE],
line_arg="provider",
line_vals=list(PROVIDER_FNS.keys()),
line_names=["one_shot (fused)", "two_kernel (legacy)"],
styles=[("red", "-"), ("blue", "-")],
ylabel="ms",
plot_name=f"fused_moe_lora-{model}-loras{preset['max_loras']}",
args={"preset": preset},
)
)
def benchmark(M, rank, provider, preset):
inp = _make_inputs(
M=M,
K=preset["K"],
N_per_slice=preset["N_per_slice"],
rank=rank,
num_experts=preset["num_experts"],
top_k=preset["top_k"],
max_loras=preset["max_loras"],
num_slices=preset["num_slices"],
block_size_m=preset["block_size_m"],
)
fn = PROVIDER_FNS[provider]
quantiles = [0.5, 0.2, 0.8]
ms, min_ms, max_ms = triton.testing.do_bench(
lambda: fn(inp), quantiles=quantiles
)
return ms, max_ms, min_ms
return benchmark
# ----- correctness sanity ---------------------------------------------------
def calculate_diff(model: str, M: int, rank: int, max_loras: int | None = None):
preset = dict(MODEL_PRESETS[model])
if max_loras is not None:
preset["max_loras"] = max_loras
inp = _make_inputs(
M=M,
K=preset["K"],
N_per_slice=preset["N_per_slice"],
rank=rank,
num_experts=preset["num_experts"],
top_k=preset["top_k"],
max_loras=preset["max_loras"],
num_slices=preset["num_slices"],
block_size_m=preset["block_size_m"],
)
out_one = _run_one_shot(inp)
out_two = _run_two_kernel(inp)
max_abs = (out_one.float() - out_two.float()).abs().max().item()
print(
f" model={model:<9} M={M:<6} rank={rank:<3} "
f"max|one_shot - two_kernel|={max_abs:.4g} "
f"ref|max|={out_two.float().abs().max().item():.3g}"
)
if max_abs <= 5e-2:
print(" ✅ outputs match within bf16 tolerance")
else:
print(" ❌ outputs differ beyond expected bf16 noise")
# ----- main ------------------------------------------------------------------
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument(
"--model",
type=str,
default="mixtral",
choices=list(MODEL_PRESETS.keys()),
help="Model preset to sweep",
)
parser.add_argument(
"--save-path",
type=str,
default="./configs/fused_moe_lora_one_shot/",
help="Directory to save benchmark results",
)
parser.add_argument(
"--check-only",
action="store_true",
help="Run correctness sanity check only, no perf sweep",
)
parser.add_argument(
"--max-loras",
type=int,
default=None,
help="Override max_loras in the model preset (number of LoRA adapters "
"active in the batch). Defaults to the preset's value.",
)
args = parser.parse_args()
print(f"Correctness check ({args.model}):")
calculate_diff(args.model, M=256, rank=32, max_loras=args.max_loras)
if args.check_only:
raise SystemExit(0)
effective_max_loras = (
args.max_loras
if args.max_loras is not None
else MODEL_PRESETS[args.model]["max_loras"]
)
print(f"\nGPU: {torch.cuda.get_device_name()}")
print(f"Model preset: {args.model} max_loras={effective_max_loras}\n")
benchmark = get_benchmark(args.model, max_loras=args.max_loras)
os.makedirs(args.save_path, exist_ok=True)
benchmark.run(print_data=True, save_path=args.save_path)
+48 -81
View File
@@ -408,19 +408,9 @@ class AttentionScheduler {
const int64_t cache_size = cpu_utils::get_available_l2_size();
const int32_t max_num_q_per_iter = input.max_num_q_per_iter;
const int32_t kv_len_alignment = input.kv_block_alignment;
bool has_decode_request = false;
bool decode_only_batch = true;
for (int32_t req_id = 0; req_id < input.num_reqs; ++req_id) {
const int32_t q_token_num =
input.query_start_loc[req_id + 1] - input.query_start_loc[req_id];
has_decode_request = has_decode_request || (q_token_num == 1);
decode_only_batch = decode_only_batch && (q_token_num == 1);
}
int32_t q_head_per_kv = input.num_heads_q / input.num_heads_kv;
const bool supports_gqa = q_head_per_kv <= max_num_q_per_iter;
const bool use_gqa_fast_path = supports_gqa && decode_only_batch;
const bool use_gqa_scratchpad = supports_gqa && has_decode_request;
if (!use_gqa_scratchpad) {
const bool use_gqa = (max_num_q_per_iter % q_head_per_kv == 0);
if (!use_gqa) {
q_head_per_kv = 1; // fallback to MHA
}
const int32_t min_split_kv_len =
@@ -690,7 +680,7 @@ class AttentionScheduler {
metadata_ptr->attention_scratchpad_size_per_thread *
metadata_ptr->thread_num +
metadata_ptr->reduction_scratchpad_size_per_kv_head *
(use_gqa_fast_path ? input.num_heads_kv : input.num_heads_q);
(use_gqa ? input.num_heads_kv : input.num_heads_q);
cpu_utils::ScratchPadManager::get_scratchpad_manager()->realloc(
scratchpad_size);
@@ -1419,24 +1409,13 @@ class AttentionMainLoop {
const int32_t q_head_num = input->num_heads;
const int32_t kv_head_num = input->num_kv_heads;
const int32_t q_heads_per_kv = q_head_num / kv_head_num;
AttentionWorkItemGroup* const workitem_groups =
metadata.workitem_groups_ptr;
const int32_t* cu_workitem_num_per_thread =
metadata.cu_workitem_num_per_thread;
ReductionWorkItemGroup* const reduction_items =
metadata.reduction_items_ptr;
const bool supports_gqa = q_heads_per_kv <= max_q_head_num_per_iter;
bool decode_only_batch = true;
for (int32_t i = 0; i < metadata.workitem_group_num; ++i) {
decode_only_batch =
decode_only_batch && (workitem_groups[i].q_token_num == 1);
}
const bool use_gqa_fast_path = supports_gqa && decode_only_batch;
const int32_t actual_kv_head_num =
use_gqa_fast_path ? kv_head_num : q_head_num;
const int32_t actual_q_heads_per_kv =
use_gqa_fast_path ? q_heads_per_kv : 1;
const bool use_gqa =
(max_q_head_num_per_iter % q_heads_per_kv == 0) ? true : false;
const int32_t actual_kv_head_num = use_gqa ? kv_head_num : q_head_num;
const int32_t actual_q_heads_per_kv = use_gqa ? q_heads_per_kv : 1;
TORCH_CHECK_LE(actual_q_heads_per_kv, max_q_head_num_per_iter);
const int32_t max_q_token_num_per_iter =
max_q_head_num_per_iter / actual_q_heads_per_kv;
const int64_t q_token_num_stride = input->query_num_tokens_stride;
const int64_t q_head_num_stride = input->query_num_heads_stride;
const int64_t kv_cache_head_num_stride = input->cache_num_kv_heads_stride;
@@ -1482,6 +1461,15 @@ class AttentionMainLoop {
sizeof(q_buffer_t), sizeof(logits_buffer_t),
sizeof(partial_output_buffer_t), max_q_head_num_per_iter,
max_q_head_num_per_iter);
const int32_t default_q_tile_token_num =
default_tile_size / actual_q_heads_per_kv;
AttentionWorkItemGroup* const workitem_groups =
metadata.workitem_groups_ptr;
const int32_t* cu_workitem_num_per_thread =
metadata.cu_workitem_num_per_thread;
ReductionWorkItemGroup* const reduction_items =
metadata.reduction_items_ptr;
const int32_t effective_thread_num = metadata.effective_thread_num;
const int32_t reduction_item_num = metadata.reduction_item_num;
@@ -1525,6 +1513,8 @@ class AttentionMainLoop {
cu_workitem_num_per_thread[thread_offset + 1] -
cu_workitem_num_per_thread[thread_offset];
const int32_t q_head_start_idx = kv_head_idx * actual_q_heads_per_kv;
for (int32_t workitem_group_idx = 0;
workitem_group_idx < curr_workitem_groups_num;
++workitem_group_idx) {
@@ -1539,21 +1529,6 @@ class AttentionMainLoop {
const int32_t q_token_id_start =
current_workitem_group->q_token_id_start;
const int32_t q_token_num = current_workitem_group->q_token_num;
const bool curr_use_gqa =
use_gqa_fast_path || (supports_gqa && q_token_num == 1);
if (!use_gqa_fast_path && curr_use_gqa &&
kv_head_idx % q_heads_per_kv != 0) {
continue;
}
const int32_t curr_q_heads_per_kv =
curr_use_gqa ? q_heads_per_kv : 1;
const int32_t curr_max_q_token_num_per_iter =
max_q_head_num_per_iter / curr_q_heads_per_kv;
const int32_t curr_default_q_tile_token_num =
default_tile_size / curr_q_heads_per_kv;
const int32_t q_head_start_idx =
use_gqa_fast_path ? (kv_head_idx * q_heads_per_kv)
: kv_head_idx;
// taskgroup general information
const int32_t q_end = input->query_start_loc[current_group_idx + 1];
@@ -1567,7 +1542,7 @@ class AttentionMainLoop {
current_workitem_group->local_split_id == 0);
for (int32_t q_token_offset = 0; q_token_offset < q_token_num;
q_token_offset += curr_default_q_tile_token_num) {
q_token_offset += default_q_tile_token_num) {
bool first_iter_flag[AttentionScheduler::MaxQTileIterNum];
for (int32_t i = 0; i < AttentionScheduler::MaxQTileIterNum;
++i) {
@@ -1577,9 +1552,9 @@ class AttentionMainLoop {
const int32_t q_token_start_idx =
q_start + q_token_offset + q_token_id_start;
const int32_t actual_q_token_num = std::min(
curr_default_q_tile_token_num, q_token_num - q_token_offset);
default_q_tile_token_num, q_token_num - q_token_offset);
const int32_t q_head_tile_size =
actual_q_token_num * curr_q_heads_per_kv;
actual_q_token_num * actual_q_heads_per_kv;
const int32_t rounded_q_head_tile_size =
((q_head_tile_size + max_q_head_num_per_iter - 1) /
max_q_head_num_per_iter) *
@@ -1616,9 +1591,10 @@ class AttentionMainLoop {
AttentionScheduler::align_kv_tile_pos(
kv_tile_start_pos, kv_tile_end_pos, blocksize_alignment);
const int32_t curr_kv_head_idx =
use_gqa_fast_path ? kv_head_idx
: (kv_head_idx / q_heads_per_kv);
int32_t curr_kv_head_idx =
use_gqa ? kv_head_idx
: (kv_head_idx /
q_heads_per_kv); // for GQA disabled case
// std::printf("thread_id: %d, req_id: %d, q_token_start: %d,
// q_token_end: %d, q_head_start: %d, q_head_end: %d, kv_head_idx:
@@ -1653,12 +1629,12 @@ class AttentionMainLoop {
(s_aux != nullptr ? s_aux + q_head_start_idx : nullptr);
// copy the Q tile to q_buffer, the logical layout of q_buffer is
// [actual_q_token_num, curr_q_heads_per_kv, head_dim]
// [actual_q_token_num, actual_q_heads_per_kv, head_dim]
{
attn_impl.copy_q_heads_tile(
q_tile_ptr, q_buffer, actual_q_token_num,
curr_q_heads_per_kv, q_token_num_stride, q_head_num_stride,
scale);
actual_q_heads_per_kv, q_token_num_stride,
q_head_num_stride, scale);
}
if (use_sink) {
@@ -1672,29 +1648,29 @@ class AttentionMainLoop {
float* __restrict__ curr_max_buffer = max_buffer;
for (int32_t token_idx = 0; token_idx < actual_q_token_num;
++token_idx) {
for (int32_t head_idx = 0; head_idx < curr_q_heads_per_kv;
for (int32_t head_idx = 0; head_idx < actual_q_heads_per_kv;
++head_idx) {
curr_sum_buffer[head_idx] = 1.0f;
curr_max_buffer[head_idx] = s_aux_fp32[head_idx];
}
curr_sum_buffer += curr_q_heads_per_kv;
curr_max_buffer += curr_q_heads_per_kv;
curr_sum_buffer += actual_q_heads_per_kv;
curr_max_buffer += actual_q_heads_per_kv;
}
} else {
float* __restrict__ curr_sum_buffer = sum_buffer;
float* __restrict__ curr_max_buffer = max_buffer;
for (int32_t token_idx = 0; token_idx < actual_q_token_num;
++token_idx) {
for (int32_t head_idx = 0; head_idx < curr_q_heads_per_kv;
for (int32_t head_idx = 0; head_idx < actual_q_heads_per_kv;
++head_idx) {
curr_sum_buffer[head_idx] = 0.0f;
curr_max_buffer[head_idx] =
std::numeric_limits<float>::lowest();
}
curr_sum_buffer += curr_q_heads_per_kv;
curr_max_buffer += curr_q_heads_per_kv;
curr_sum_buffer += actual_q_heads_per_kv;
curr_max_buffer += actual_q_heads_per_kv;
}
}
@@ -1707,17 +1683,16 @@ class AttentionMainLoop {
kv_tile_pos_left + kv_tile_size, rounded_kv_tile_end_pos);
for (int32_t q_head_tile_token_offset = 0;
q_head_tile_token_offset < actual_q_token_num;
q_head_tile_token_offset +=
curr_max_q_token_num_per_iter) {
q_head_tile_token_offset += max_q_token_num_per_iter) {
const int32_t q_tile_pos_left =
q_tile_start_pos + q_head_tile_token_offset;
const int32_t q_tile_token_num =
std::min(curr_max_q_token_num_per_iter,
std::min(max_q_token_num_per_iter,
actual_q_token_num - q_head_tile_token_offset);
const int32_t q_tile_head_offset =
q_head_tile_token_offset * curr_q_heads_per_kv;
q_head_tile_token_offset * actual_q_heads_per_kv;
const int32_t q_tile_head_num =
q_tile_token_num * curr_q_heads_per_kv;
q_tile_token_num * actual_q_heads_per_kv;
const int32_t q_tile_pos_right =
q_tile_pos_left + q_tile_token_num;
const auto [actual_kv_tile_pos_left,
@@ -1727,7 +1702,7 @@ class AttentionMainLoop {
q_tile_pos_right, sliding_window_left,
sliding_window_right);
const int32_t q_iter_idx =
q_head_tile_token_offset / curr_max_q_token_num_per_iter;
q_head_tile_token_offset / max_q_token_num_per_iter;
if (actual_kv_tile_pos_right <= actual_kv_tile_pos_left) {
continue;
@@ -1793,7 +1768,7 @@ class AttentionMainLoop {
aligned_actual_kv_tile_pos_left,
aligned_actual_kv_tile_pos_right, actual_kv_token_num,
kv_cache_block_num_stride, q_tile_head_num,
q_tile_token_num, q_tile_pos_left, curr_q_heads_per_kv,
q_tile_token_num, q_tile_pos_left, actual_q_heads_per_kv,
block_size, sliding_window_left, sliding_window_right,
scale, softcap_scale, curr_alibi_slopes,
first_iter_flag[q_iter_idx], use_sink, debug_info);
@@ -1807,11 +1782,11 @@ class AttentionMainLoop {
final_output(partial_q_buffer,
reinterpret_cast<query_t*>(input->output) +
output_buffer_offset,
sum_buffer, curr_q_heads_per_kv,
sum_buffer, actual_q_heads_per_kv,
actual_q_token_num, q_head_num, output_v_scale);
} else {
const int32_t stride =
curr_q_heads_per_kv * split_kv_q_token_num_threshold;
actual_q_heads_per_kv * split_kv_q_token_num_threshold;
buffer_manager.update(kv_head_idx, total_reduction_split_num,
head_dim, stride, sizeof(float));
volatile bool* split_flag_buffer =
@@ -1847,26 +1822,18 @@ class AttentionMainLoop {
const int32_t curr_split_id = curr_workitem_groups->split_start_id;
const int32_t curr_split_num = curr_workitem_groups->split_num;
const int32_t current_group_idx = curr_workitem_groups->req_id;
const bool curr_use_gqa =
use_gqa_fast_path || (supports_gqa && curr_output_token_num == 1);
if (!use_gqa_fast_path && curr_use_gqa &&
kv_head_idx % q_heads_per_kv != 0) {
continue;
}
const int32_t curr_q_heads_per_kv = curr_use_gqa ? q_heads_per_kv : 1;
const int32_t curr_output_head_num =
curr_output_token_num * curr_q_heads_per_kv;
curr_output_token_num * actual_q_heads_per_kv;
const int32_t q_start = input->query_start_loc[current_group_idx];
const int32_t q_token_start_idx = q_start + curr_output_token_idx;
const int32_t q_head_start_idx =
use_gqa_fast_path ? (kv_head_idx * q_heads_per_kv) : kv_head_idx;
const int32_t q_head_start_idx = kv_head_idx * actual_q_heads_per_kv;
size_t output_buffer_offset =
q_token_start_idx * q_head_num * head_dim +
q_head_start_idx * head_dim;
const int32_t stride =
curr_q_heads_per_kv * split_kv_q_token_num_threshold;
actual_q_heads_per_kv * split_kv_q_token_num_threshold;
buffer_manager.update(kv_head_idx, total_reduction_split_num,
head_dim, stride, sizeof(float));
volatile bool* split_flag_buffer =
@@ -1885,7 +1852,7 @@ class AttentionMainLoop {
final_output(
split_output_buffer,
reinterpret_cast<query_t*>(input->output) + output_buffer_offset,
split_sum_buffer, curr_q_heads_per_kv, curr_output_token_num,
split_sum_buffer, actual_q_heads_per_kv, curr_output_token_num,
q_head_num, output_v_scale);
}
}
+1 -1
View File
@@ -12,4 +12,4 @@ ray[data]
setuptools==78.1.0
setuptools-rust>=1.9.0
nixl==0.3.0
tpu-inference==0.20.0
tpu-inference==0.19.0
@@ -192,7 +192,7 @@ def test_preprocess_cmpl_applies_mm_processor_kwargs_to_renderer(
llm.renderer = renderer
monkeypatch.setattr(
"vllm.entrypoints.offline_utils.parse_model_prompt",
"vllm.entrypoints.llm.parse_model_prompt",
lambda _model_config, parsed_prompt: parsed_prompt,
)
@@ -226,7 +226,7 @@ def test_preprocess_cmpl_keeps_prompt_mm_processor_kwargs_when_no_override(
llm.renderer = renderer
monkeypatch.setattr(
"vllm.entrypoints.offline_utils.parse_model_prompt",
"vllm.entrypoints.llm.parse_model_prompt",
lambda _model_config, parsed_prompt: parsed_prompt,
)
@@ -1,199 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import math
import pytest
import torch
import torch.nn.functional as F
from vllm.platforms import current_platform
if not (
current_platform.is_cuda() and current_platform.is_device_capability_family(100)
):
pytest.skip(
reason="GDN CuteDSL prefill requires CUDA SM10x.",
allow_module_level=True,
)
from vllm.model_executor.layers.fla.ops import ( # noqa: E402
chunk_gated_delta_rule,
)
from vllm.model_executor.layers.fla.ops.index import ( # noqa: E402
prepare_chunk_indices,
prepare_chunk_offsets,
)
from vllm.model_executor.layers.mamba.ops.gdn_chunk_cutedsl import ( # noqa: E402
chunk_gated_delta_rule_cutedsl,
prepare_metadata_cutedsl,
)
@pytest.mark.parametrize("num_seqs", [1, 5, 257])
@pytest.mark.parametrize("state_dtype", [torch.bfloat16, torch.float32])
def test_gdn_chunk_cutedsl_correctness(num_seqs: int, state_dtype: torch.dtype):
seq_lens = torch.randint(
1,
130,
(num_seqs,),
dtype=torch.int32,
)
cu_seqlens = torch.zeros(num_seqs + 1, device="cuda", dtype=torch.int32)
cu_seqlens[1:] = seq_lens.to(device="cuda").cumsum(0)
total_tokens = int(cu_seqlens[-1].item())
num_k_heads = 4
num_v_heads = 8
head_k_dim = 128
head_v_dim = 128
dtype = torch.bfloat16
q = torch.randn(
1,
total_tokens,
num_k_heads,
head_k_dim,
device="cuda",
dtype=dtype,
)
k = torch.randn_like(q)
v = torch.randn(
1,
total_tokens,
num_v_heads,
head_v_dim,
device="cuda",
dtype=dtype,
)
q = F.normalize(q.float(), p=2, dim=-1).to(dtype)
k = F.normalize(k.float(), p=2, dim=-1).to(dtype)
a = torch.randn(
1,
total_tokens,
num_v_heads,
device="cuda",
dtype=dtype,
)
b = torch.randn(
1,
total_tokens,
num_v_heads,
device="cuda",
dtype=dtype,
)
# Match upstream FLA GatedDeltaNet synthetic initialization:
# https://github.com/fla-org/flash-linear-attention/blob/main/fla/layers/gated_deltanet.py
A = torch.empty(num_v_heads, device="cuda", dtype=torch.float32).uniform_(0, 16)
A_log = torch.log(A)
dt = torch.exp(
torch.rand(num_v_heads, device="cuda", dtype=torch.float32)
* (math.log(0.1) - math.log(0.001))
+ math.log(0.001)
)
dt = torch.clamp(dt, min=1e-4)
dt_bias = dt + torch.log(-torch.expm1(-dt))
g = -A_log.exp().view(1, 1, num_v_heads) * F.softplus(
a.float() + dt_bias.view(1, 1, num_v_heads)
)
beta = torch.sigmoid(b.float())
initial_state = (
torch.randn(
num_seqs,
num_v_heads,
head_v_dim,
head_k_dim,
device="cuda",
dtype=state_dtype,
)
* 0.05
)
# check metadata kernel
chunk_indices, chunk_offsets = prepare_metadata_cutedsl(cu_seqlens, total_tokens)
torch.accelerator.synchronize()
expected_indices = prepare_chunk_indices(cu_seqlens, 64)
expected_offsets = prepare_chunk_offsets(cu_seqlens, 64)
total_chunks = int(expected_offsets[-1].item())
torch.testing.assert_close(chunk_offsets, expected_offsets.to(torch.int32))
torch.testing.assert_close(
chunk_indices[:total_chunks],
expected_indices,
)
ref_o, ref_state = chunk_gated_delta_rule(
q=q,
k=k,
v=v,
g=g,
beta=beta,
initial_state=initial_state,
output_final_state=True,
cu_seqlens=cu_seqlens,
use_qk_l2norm_in_kernel=False,
)
actual_core_attn_out = torch.empty(
total_tokens,
num_v_heads,
head_v_dim,
device="cuda",
dtype=dtype,
)
actual_o, actual_state = chunk_gated_delta_rule_cutedsl(
q=q,
k=k,
v=v,
g=g,
beta=beta,
initial_state=initial_state,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
chunk_offsets=chunk_offsets,
core_attn_out=actual_core_attn_out,
)
torch.accelerator.synchronize()
# check main kernel
o_error = (actual_o.float() - ref_o.float()).abs()
state_error = (
actual_state.float() - ref_state.to(actual_state.dtype).float()
).abs()
assert o_error.max().item() < 2e-3
assert o_error.mean().item() < 6e-5
assert state_error.max().item() < 2e-2
assert state_error.mean().item() < 6e-4
core_attn_out_error = (
actual_core_attn_out.float() - actual_o.squeeze(0).float()
).abs()
assert core_attn_out_error.max().item() == 0
# check main kernel when core_attn_out is not passed
no_buffer_o, no_buffer_state = chunk_gated_delta_rule_cutedsl(
q=q,
k=k,
v=v,
g=g,
beta=beta,
initial_state=initial_state,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
chunk_offsets=chunk_offsets,
)
torch.accelerator.synchronize()
no_buffer_o_error = (no_buffer_o.float() - ref_o.float()).abs()
no_buffer_state_error = (
no_buffer_state.float() - ref_state.to(no_buffer_state.dtype).float()
).abs()
buffer_o_error = (no_buffer_o.float() - actual_o.float()).abs()
buffer_state_error = (
no_buffer_state.float() - actual_state.to(no_buffer_state.dtype).float()
).abs()
assert no_buffer_o_error.max().item() < 2e-3
assert no_buffer_o_error.mean().item() < 6e-5
assert no_buffer_state_error.max().item() < 2e-2
assert no_buffer_state_error.mean().item() < 6e-4
assert buffer_o_error.max().item() == 0
assert buffer_state_error.max().item() == 0
-495
View File
@@ -141,7 +141,6 @@ def use_fused_moe_lora_kernel(
block_size,
fully_sharded=False,
offset=0,
add_inputs=True,
):
max_num_tokens_padded = topk_ids.numel() + num_experts * (block_size - 1)
max_num_tokens_padded = round_up(max_num_tokens_padded, block_size)
@@ -223,7 +222,6 @@ def use_fused_moe_lora_kernel(
mul_routed_weight,
fully_sharded=fully_sharded,
offset=offset,
add_inputs=add_inputs,
)
@@ -373,7 +371,6 @@ def use_fused_moe_lora_kernel_naive(
block_size,
fully_sharded=False,
offset=0,
add_inputs=True,
):
"""
Test helper for naive_block_assignment path.
@@ -438,7 +435,6 @@ def use_fused_moe_lora_kernel_naive(
mul_routed_weight=mul_routed_weight,
fully_sharded=fully_sharded,
offset=offset,
add_inputs=add_inputs,
)
@@ -749,494 +745,3 @@ def use_fused_moe_lora_kernel_tensor_parallel(
output = tensor_model_parallel_all_reduce(output)
torch.testing.assert_close(output, ref_output, atol=1e-2, rtol=1e-2)
# -- one-shot fast-path coverage --------------------------------------------
# The fused shrink+expand one-shot kernel pads `BLOCK_R` to next_pow2(rank),
# with a floor of 16 (tensor-core minimum). Small ranks (4, 8) exercise the
# rank-dim masking and are not covered by the original tests, which start at
# rank=16. The legacy two-kernel path additionally fails on rank=4 in TMA
# mode because the rank-dim stride (rank * elem_size) is not 16-byte
# aligned; the one-shot fast path takes precedence whenever fully_sharded
# is False so this regression is hidden in normal use, but the test still
# ensures the one-shot logic is correct against the pytorch reference.
@pytest.mark.parametrize("num_tokens", [16, 100])
@pytest.mark.parametrize("top_k_num", [2])
@pytest.mark.parametrize("num_experts", [8, 64])
@pytest.mark.parametrize("max_loras", [4])
@pytest.mark.parametrize("N", [1408])
@pytest.mark.parametrize("K", [2048])
@pytest.mark.parametrize("max_lora_rank", [4, 8])
@pytest.mark.parametrize("block_size", [16, 64])
@pytest.mark.parametrize("num_slices", [1, 2])
@pytest.mark.parametrize("dtype", DTYPES)
@pytest.mark.parametrize("device", DEVICES)
@pytest.mark.parametrize("seed", SEED)
def test_fused_moe_lora_kernel_small_rank(
num_tokens,
top_k_num,
num_experts,
max_loras,
N,
K,
max_lora_rank,
block_size,
num_slices,
dtype,
device,
seed,
):
"""One-shot fast path covering rank<16 (padded to BLOCK_R=16 inside kernel)."""
torch.set_default_device(device)
set_random_seed(seed)
num_sequences = max(1, min(num_tokens, 8))
topk_ids, topk_weights, token_lora_mapping, lora_ids = sample_data(
num_tokens, num_sequences, max_loras, num_experts, top_k_num
)
lora_a_stacked = [
torch.rand(
(max_loras, num_experts, max_lora_rank, K),
dtype=dtype,
)
for _ in range(num_slices)
]
lora_b_stacked = [
torch.rand(
(max_loras, num_experts, N // num_slices, max_lora_rank),
dtype=dtype,
)
for _ in range(num_slices)
]
hidden_states = torch.rand((num_tokens, K), dtype=dtype)
output = torch.zeros((num_tokens, top_k_num, N), dtype=dtype)
use_fused_moe_lora_kernel(
topk_ids,
topk_weights,
token_lora_mapping,
max_lora_rank,
top_k_num,
lora_ids,
lora_a_stacked,
lora_b_stacked,
hidden_states,
output,
max_loras,
num_experts,
block_size,
)
output_ref = use_torch(
hidden_states,
token_lora_mapping,
topk_ids,
lora_a_stacked,
lora_b_stacked,
top_k_num,
num_slices,
)
torch.testing.assert_close(output, output_ref, atol=1e-2, rtol=1e-2)
@pytest.mark.parametrize("num_tokens", [16, 64])
@pytest.mark.parametrize("top_k_num", [2])
@pytest.mark.parametrize("num_experts", [8])
@pytest.mark.parametrize("max_loras", [4])
@pytest.mark.parametrize("N", [2048])
@pytest.mark.parametrize("K", [4096])
@pytest.mark.parametrize("max_lora_rank", [8, 16, 32, 64])
@pytest.mark.parametrize("block_size", [64])
@pytest.mark.parametrize("num_slices", [2])
@pytest.mark.parametrize("dtype", [torch.bfloat16])
@pytest.mark.parametrize("device", DEVICES)
@pytest.mark.parametrize("seed", SEED)
def test_fused_moe_lora_kernel_npid_path(
num_tokens,
top_k_num,
num_experts,
max_loras,
N,
K,
max_lora_rank,
block_size,
num_slices,
dtype,
device,
seed,
):
"""Exercise the small-batch / NPID > 1 branch of the one-shot fast path.
With these sizes the one-shot wrapper computes NPID_FACTOR > 1 (base CTA count
< SM count), so each program covers only an outer chunk of N. The
cross-outer-block write mask is the correctness-critical bit.
"""
torch.set_default_device(device)
set_random_seed(seed)
num_sequences = max(1, min(num_tokens, 4))
topk_ids, topk_weights, token_lora_mapping, lora_ids = sample_data(
num_tokens, num_sequences, max_loras, num_experts, top_k_num
)
lora_a_stacked = [
torch.rand(
(max_loras, num_experts, max_lora_rank, K),
dtype=dtype,
)
for _ in range(num_slices)
]
lora_b_stacked = [
torch.rand(
(max_loras, num_experts, N // num_slices, max_lora_rank),
dtype=dtype,
)
for _ in range(num_slices)
]
hidden_states = torch.rand((num_tokens, K), dtype=dtype)
output = torch.zeros((num_tokens, top_k_num, N), dtype=dtype)
use_fused_moe_lora_kernel(
topk_ids,
topk_weights,
token_lora_mapping,
max_lora_rank,
top_k_num,
lora_ids,
lora_a_stacked,
lora_b_stacked,
hidden_states,
output,
max_loras,
num_experts,
block_size,
)
output_ref = use_torch(
hidden_states,
token_lora_mapping,
topk_ids,
lora_a_stacked,
lora_b_stacked,
top_k_num,
num_slices,
)
torch.testing.assert_close(output, output_ref, atol=2e-2, rtol=2e-2)
# -- one-shot corner-case coverage ------------------------------------------
# Each of the following exercises a path where the kernel is launched but
# every program early-exits, leaving the output unchanged. The contract is
# additive (`output += contribution`), so an empty contribution must leave
# the input residual untouched.
def _build_one_shot_inputs(
num_tokens,
top_k_num,
num_experts,
max_loras,
max_lora_rank,
K,
N,
num_slices,
block_size,
dtype,
):
"""Common scaffolding for the corner-case tests below."""
num_sequences = max(1, min(num_tokens, 4)) if num_tokens > 0 else 1
if num_tokens > 0:
topk_ids, topk_weights, token_lora_mapping, lora_ids = sample_data(
num_tokens, num_sequences, max_loras, num_experts, top_k_num
)
else:
# M=0 path: caller may still hand us empty tensors with the right shape.
topk_ids = torch.empty((0, top_k_num), dtype=torch.int32)
topk_weights = torch.empty((0, top_k_num), dtype=torch.float32)
token_lora_mapping = torch.empty((0,), dtype=torch.int32)
lora_ids = torch.full((max_loras + 1,), -1, dtype=torch.int32)
lora_a_stacked = [
torch.rand((max_loras, num_experts, max_lora_rank, K), dtype=dtype)
for _ in range(num_slices)
]
lora_b_stacked = [
torch.rand(
(max_loras, num_experts, N // num_slices, max_lora_rank), dtype=dtype
)
for _ in range(num_slices)
]
hidden_states = torch.rand((max(num_tokens, 0), K), dtype=dtype)
return (
topk_ids,
topk_weights.to(dtype),
token_lora_mapping,
lora_ids,
lora_a_stacked,
lora_b_stacked,
hidden_states,
)
def _call_one_shot(
output,
hidden_states,
lora_a_stacked,
lora_b_stacked,
topk_weights,
sorted_token_ids,
expert_ids,
num_tokens_post_padded,
token_lora_mapping,
max_lora_rank,
top_k_num,
lora_ids,
num_active_loras,
adapter_enabled,
block_size,
add_inputs=True,
):
"""Direct call into fused_moe_lora with one-shot-routed defaults."""
from vllm.lora.ops.triton_ops import fused_moe_lora as _op
_op(
output,
hidden_states,
lora_a_stacked,
lora_b_stacked,
topk_weights,
sorted_token_ids,
expert_ids,
num_tokens_post_padded,
token_lora_mapping,
max_lora_rank,
top_k_num,
lora_ids,
num_active_loras,
adapter_enabled,
block_size,
32,
64,
1,
4,
3,
1,
block_size,
32,
64,
1,
4,
3,
1,
False,
False,
0,
add_inputs,
)
@pytest.mark.parametrize(
"trigger",
["sorted_lora_ids_neg", "naive_mapping_neg", "naive_all_disabled"],
)
@pytest.mark.parametrize("device", DEVICES)
def test_fused_moe_lora_kernel_one_shot_early_exit(trigger, device):
"""one-shot must leave the residual byte-identical when every program
must early-exit. Three trigger conditions are covered:
- "sorted_lora_ids_neg": sorted path, lora_ids all -1 (lora_id<0 check)
- "naive_mapping_neg": naive path, token_lora_mapping all -1
- "naive_all_disabled": naive path, adapter_enabled all 0
"""
torch.set_default_device(device)
set_random_seed(0)
# Per-trigger shapes: naive_mapping_neg needs the naive dispatch gate
# `num_tokens*top_k*8 <= num_experts*max_loras` to hold, hence the
# larger E/max_loras and smaller num_tokens.
if trigger == "naive_mapping_neg":
num_tokens, top_k, E, max_loras, R = 8, 2, 64, 8, 16
elif trigger == "naive_all_disabled":
num_tokens, top_k, E, max_loras, R = 32, 2, 8, 4, 32
else: # sorted_lora_ids_neg
num_tokens, top_k, E, max_loras, R = 32, 2, 8, 4, 16
K, N = 1024, 1024
block_size, num_slices, dtype = 16, 2, torch.bfloat16
(
topk_ids,
topk_weights,
token_lora_mapping,
lora_ids,
lora_a_stacked,
lora_b_stacked,
hidden_states,
) = _build_one_shot_inputs(
num_tokens, top_k, E, max_loras, R, K, N, num_slices, block_size, dtype
)
adapter_enabled = torch.ones(max_loras + 1, dtype=torch.int32)
num_active_loras = torch.tensor([max_loras + 1], dtype=torch.int32, device="cpu")
if trigger == "sorted_lora_ids_neg":
lora_ids = torch.full((max_loras + 1,), -1, dtype=torch.int32)
max_pad = topk_ids.numel() + E * (block_size - 1)
max_pad = round_up(max_pad, block_size)
max_blocks = CEILDIV(max_pad, block_size)
sorted_token_ids = torch.zeros((max_loras, max_pad), dtype=torch.int32)
expert_ids = torch.full((max_loras, max_blocks), -1, dtype=torch.int32)
num_post = torch.zeros((max_loras,), dtype=torch.int32)
else:
sorted_token_ids = None
expert_ids = topk_ids.reshape(-1).contiguous()
num_post = None
if trigger == "naive_mapping_neg":
token_lora_mapping = torch.full((num_tokens,), -1, dtype=torch.int32)
lora_ids = torch.full((max_loras + 1,), -1, dtype=torch.int32)
else: # naive_all_disabled
adapter_enabled = torch.zeros(max_loras + 1, dtype=torch.int32)
residual = torch.randn((num_tokens, top_k, N), dtype=dtype) * 0.1
output = residual.clone()
_call_one_shot(
output,
hidden_states,
lora_a_stacked,
lora_b_stacked,
topk_weights,
sorted_token_ids,
expert_ids,
num_post,
token_lora_mapping,
R,
top_k,
lora_ids,
num_active_loras,
adapter_enabled,
block_size,
)
torch.testing.assert_close(output, residual, atol=0, rtol=0)
@pytest.mark.parametrize("device", DEVICES)
def test_fused_moe_lora_kernel_zero_grid_no_crash(device):
"""num_active_loras=0 (or num_slices=0) would otherwise launch a grid
with a zero dimension. one-shot wrapper must short-circuit before launch."""
torch.set_default_device(device)
set_random_seed(0)
num_tokens, top_k, E, max_loras, R, K, N = 8, 2, 8, 4, 16, 1024, 1024
block_size, num_slices, dtype = 16, 2, torch.bfloat16
(
topk_ids,
topk_weights,
token_lora_mapping,
lora_ids,
lora_a_stacked,
lora_b_stacked,
hidden_states,
) = _build_one_shot_inputs(
num_tokens,
top_k,
E,
max_loras,
R,
K,
N,
num_slices,
block_size,
dtype,
)
adapter_enabled = torch.ones(max_loras + 1, dtype=torch.int32)
num_active_loras = torch.tensor([0], dtype=torch.int32, device="cpu")
residual = torch.randn((num_tokens, top_k, N), dtype=dtype) * 0.1
output = residual.clone()
# sorted path is the one that uses num_active_loras for grid axis 2
max_pad = topk_ids.numel() + E * (block_size - 1)
max_pad = round_up(max_pad, block_size)
max_blocks = CEILDIV(max_pad, block_size)
sorted_token_ids = torch.zeros((max_loras, max_pad), dtype=torch.int32)
expert_ids = torch.full((max_loras, max_blocks), -1, dtype=torch.int32)
num_post = torch.zeros((max_loras,), dtype=torch.int32)
_call_one_shot(
output,
hidden_states,
lora_a_stacked,
lora_b_stacked,
topk_weights,
sorted_token_ids,
expert_ids,
num_post,
token_lora_mapping,
R,
top_k,
lora_ids,
num_active_loras,
adapter_enabled,
block_size,
)
torch.testing.assert_close(output, residual, atol=0, rtol=0)
@pytest.mark.parametrize("device", DEVICES)
def test_fused_moe_lora_kernel_rejects_bad_block_size_m(device):
"""one-shot must surface a clear assertion when shrink_block_size_m is not
a power of 2 / less than 16, instead of the cryptic Triton compile
failure (`arange's range must be a power of 2`)."""
torch.set_default_device(device)
set_random_seed(0)
num_tokens, top_k, E, max_loras, R, K, N = 32, 2, 8, 4, 16, 1024, 1024
num_slices, dtype = 2, torch.bfloat16
block_size = 24 # NOT a power of 2
(
topk_ids,
topk_weights,
token_lora_mapping,
lora_ids,
lora_a_stacked,
lora_b_stacked,
hidden_states,
) = _build_one_shot_inputs(
num_tokens,
top_k,
E,
max_loras,
R,
K,
N,
num_slices,
16,
dtype,
)
# Build sorted-mode metadata at block_size=16 so shapes are sane,
# but pass block_size=24 to the op (the buggy combination).
max_pad = topk_ids.numel() + E * (16 - 1)
max_pad = round_up(max_pad, 16)
max_blocks = CEILDIV(max_pad, 16)
sorted_token_ids = torch.zeros((max_loras, max_pad), dtype=torch.int32)
expert_ids = torch.full((max_loras, max_blocks), -1, dtype=torch.int32)
num_post = torch.zeros((max_loras,), dtype=torch.int32)
adapter_enabled = torch.ones(max_loras + 1, dtype=torch.int32)
num_active_loras = torch.tensor([max_loras + 1], dtype=torch.int32, device="cpu")
output = torch.zeros((num_tokens, top_k, N), dtype=dtype)
with pytest.raises(AssertionError, match="shrink_block_size_m"):
_call_one_shot(
output,
hidden_states,
lora_a_stacked,
lora_b_stacked,
topk_weights,
sorted_token_ids,
expert_ids,
num_post,
token_lora_mapping,
R,
top_k,
lora_ids,
num_active_loras,
adapter_enabled,
block_size,
)
+2 -2
View File
@@ -8,9 +8,9 @@ import torch
from vllm.models.deepseek_v4.nvidia.model import (
DeepseekV4MegaMoEExperts,
_stage_deepseek_v4_mega_moe_inputs,
make_deepseek_v4_expert_params_mapping,
)
from vllm.models.deepseek_v4.nvidia.ops.prepare_megamoe import prepare_megamoe_inputs
from vllm.platforms import current_platform
pytestmark = pytest.mark.skipif(
@@ -164,7 +164,7 @@ def test_deepseek_v4_mega_moe_fused_input_staging_is_bitwise_exact():
fused_topk_idx = torch.empty_like(ref_topk_idx)
fused_topk_weights = torch.empty_like(ref_topk_weights)
prepare_megamoe_inputs(
_stage_deepseek_v4_mega_moe_inputs(
hidden_states,
topk_weights,
topk_ids,
@@ -1,6 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
def register_bge_m3_sparse_embeddings_processor():
return "bge_m3_sparse_processor.sparse_embeddings_processor.BgeM3SparseEmbeddingsProcessor" # noqa: E501
@@ -1,206 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
from collections.abc import Sequence
from vllm.config import ModelConfig, PoolerConfig, VllmConfig
from vllm.entrypoints.openai.engine.protocol import UsageInfo
from vllm.entrypoints.pooling.base.protocol import EmbedRequestMixin
from vllm.inputs import PromptType
from vllm.outputs import PoolingRequestOutput
from vllm.plugins.io_processors.interface import IOProcessor
from vllm.pooling_params import PoolingParams
from vllm.renderers import BaseRenderer
from vllm.tokenizers.detokenizer_utils import convert_ids_list_to_tokens
from .types import (
EMBED_TASKS,
SparseEmbeddingCompletionRequestMixin,
SparseEmbeddingResponse,
SparseEmbeddingResponseData,
SparseEmbeddingTokenWeight,
)
class BgeM3SparseEmbeddingsProcessor(
IOProcessor[SparseEmbeddingCompletionRequestMixin, SparseEmbeddingResponse]
):
def __init__(self, vllm_config: VllmConfig, renderer: BaseRenderer):
super().__init__(vllm_config, renderer)
self.offline_requests: list[SparseEmbeddingCompletionRequestMixin] = []
self.online_requests: dict[str, SparseEmbeddingCompletionRequestMixin] = {}
self.renderer: BaseRenderer = renderer
self.default_pooling_params = {}
pooler_config: PoolerConfig = vllm_config.model_config.pooler_config
if pooler_config is not None:
for param in ["use_activation", "dimensions"]:
if getattr(pooler_config, param, None) is None:
continue
self.default_pooling_params[param] = getattr(pooler_config, param)
self.embed_dimensions = vllm_config.model_config.embedding_size
self.embed_request_queue: list[EmbedRequestMixin] = []
def __repr__(self) -> str:
return (
f"BgeM3SparseEmbeddingsProcessor("
f"embed_dimensions={self.embed_dimensions}, "
f"default_pooling_params={self.default_pooling_params})"
)
def merge_pooling_params(
self,
params: PoolingParams | None = None,
) -> PoolingParams:
if params is None:
params = PoolingParams()
# refer to PoolingCompletionRequest.to_pooling_params
# set and verify pooling params
params.skip_reading_prefix_cache = True
raw_embed_request = self.embed_request_queue.pop(0)
if raw_embed_request.embed_task not in EMBED_TASKS:
raise ValueError(
f"Unsupported task {raw_embed_request}, "
f"Supported tasks are {EMBED_TASKS}"
)
params.task = "embed&token_classify"
params.use_activation = raw_embed_request.use_activation
if params.use_activation is None:
params.use_activation = True
params.dimensions = raw_embed_request.dimensions
model_config: ModelConfig = self.vllm_config.model_config
for param in self.default_pooling_params:
if getattr(params, param, None) is None:
setattr(params, param, self.default_pooling_params[param])
if params.dimensions is not None:
if not model_config.is_matryoshka:
raise ValueError(
f'Model "{model_config.served_model_name}" does not '
f"support matryoshka representation, "
f"changing output dimensions will lead to poor results."
)
mds = model_config.matryoshka_dimensions
if mds is not None:
if params.dimensions not in mds:
raise ValueError(
f"Model {model_config.served_model_name!r} "
f"only supports {str(mds)} matryoshka dimensions, "
f"use other output dimensions will "
f"lead to poor results."
)
elif params.dimensions < 1:
raise ValueError("Dimensions must be greater than 0")
return params
def parse_request(
self, request_data: object
) -> SparseEmbeddingCompletionRequestMixin:
# for vllm.entrypoints.llm.LLM, offline mode, calls `encode` directly.
if isinstance(request_data, dict):
return SparseEmbeddingCompletionRequestMixin(**request_data)
raise TypeError("request_data should be a dictionary")
def pre_process(
self,
prompt: SparseEmbeddingCompletionRequestMixin,
request_id: str | None = None,
**kwargs,
) -> PromptType | Sequence[PromptType]:
if request_id is not None:
assert request_id not in self.online_requests, "request_id duplicated"
self.online_requests[request_id] = prompt
self.embed_request_queue.extend(prompt.to_embed_requests_online())
else:
self.offline_requests.append(prompt)
self.embed_request_queue.extend(prompt.to_embed_requests_offline())
return prompt.input
def _get_sparse_embedding_request(self, request_id: str | None = None):
if request_id:
return self.online_requests.pop(request_id, None)
return self.offline_requests.pop(0)
def _build_sparse_embedding_token_weights(
self,
sparse_embedding: dict[int, float],
return_tokens: bool = False,
) -> list[SparseEmbeddingTokenWeight]:
token_ids = sparse_embedding.keys()
token_weights = sparse_embedding.values()
tokens = [None] * len(token_ids)
if return_tokens and self.renderer is not None:
tokens = convert_ids_list_to_tokens(
self.renderer.get_tokenizer(), token_ids
)
sparse_embedding_output: list[SparseEmbeddingTokenWeight] = []
for token_id, weight, token in zip(token_ids, token_weights, tokens):
sparse_embedding_output.append(
SparseEmbeddingTokenWeight(
token_id=token_id, weight=weight, token=token
)
)
return sparse_embedding_output
def post_process(
self,
model_output: Sequence[PoolingRequestOutput],
request_id: str | None = None,
**kwargs,
) -> SparseEmbeddingResponse:
num_prompt_tokens = 0
response_data = []
raw_request = self._get_sparse_embedding_request(request_id)
has_dense_embed = raw_request.embed_task in ["dense", "dense&sparse"]
has_sparse_embed = raw_request.embed_task in ["sparse", "dense&sparse"]
embed_dimensions = (
self.embed_dimensions
if raw_request.dimensions is None
else raw_request.dimensions
)
for idx in range(len(model_output)):
mo = model_output[idx]
sparse_embedding_dict: dict[int, float] = {}
num_prompt_tokens += len(mo.prompt_token_ids)
dense_embedding: list[float] | None = None
sparse_embedding: list[SparseEmbeddingTokenWeight] | None = None
if has_dense_embed:
dense_embedding = mo.outputs.data[:embed_dimensions].tolist()
if has_sparse_embed:
sparse_weights = mo.outputs.data[embed_dimensions:].tolist()
if len(mo.prompt_token_ids) != len(sparse_weights):
# this is the case that add_special_tokens is True,
# which means first token and last token are special tokens
mo.prompt_token_ids = mo.prompt_token_ids[1:]
for token_id, weight in zip(mo.prompt_token_ids, sparse_weights):
sparse_embedding_dict[token_id] = max(
weight, sparse_embedding_dict.get(token_id, 0.0)
)
sparse_embedding = self._build_sparse_embedding_token_weights(
sparse_embedding_dict,
raw_request.return_tokens,
)
response_data.append(
SparseEmbeddingResponseData(
index=idx,
object=raw_request.embed_task,
sparse_embedding=sparse_embedding,
dense_embedding=dense_embedding,
)
)
usage = UsageInfo(
prompt_tokens=num_prompt_tokens,
total_tokens=num_prompt_tokens,
)
resp = SparseEmbeddingResponse(
data=response_data,
usage=usage,
)
return resp
@@ -1,59 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
from typing import Literal, get_args
from pydantic import BaseModel, Field
from vllm.entrypoints.openai.engine.protocol import UsageInfo
from vllm.entrypoints.pooling.base.protocol import (
CompletionRequestMixin,
EmbedRequestMixin,
)
EmbedTask = Literal[
"sparse",
"dense",
"dense&sparse",
]
EMBED_TASKS: tuple[EmbedTask, ...] = get_args(EmbedTask)
class SparseEmbeddingCompletionRequestMixin(CompletionRequestMixin, EmbedRequestMixin):
return_tokens: bool | None = Field(
default=None,
description="Whether to return dict shows the mapping of token_id to text."
"`None` or False means not return.",
)
embed_task: EmbedTask = Field(
default="dense&sparse",
description="embed task, can be one of 'sparse', 'dense' , 'dense&sparse', "
"default to 'dense&sparse'",
)
def to_embed_requests_offline(self) -> list[EmbedRequestMixin]:
if isinstance(self.input, list):
return [self] * len(self.input)
return [self]
def to_embed_requests_online(self) -> list[EmbedRequestMixin]:
return [self]
class SparseEmbeddingTokenWeight(BaseModel):
token_id: int
weight: float
token: str | None
class SparseEmbeddingResponseData(BaseModel):
index: int
object: str = "dense&sparse"
sparse_embedding: list[SparseEmbeddingTokenWeight] | None
dense_embedding: list[float] | None
class SparseEmbeddingResponse(BaseModel):
data: list[SparseEmbeddingResponseData]
usage: UsageInfo
@@ -1,15 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
from setuptools import setup
setup(
name="bge-m3-sparse-plugin",
version="0.1",
packages=["bge_m3_sparse_processor"],
entry_points={
"vllm.io_processor_plugins": [
"bge_m3_sparse_plugin = bge_m3_sparse_processor:register_bge_m3_sparse_embeddings_processor", # noqa: E501
]
},
)
@@ -1,235 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import json
import pytest
import requests
# Test configuration for BGE-M3 sparse plugin
from tests.utils import RemoteOpenAIServer
from vllm.entrypoints.pooling.pooling.protocol import IOProcessorResponse
model_config = {
"model_name": "BAAI/bge-m3",
"plugin": "bge_m3_sparse_plugin",
"test_input": "What is the capital of France?",
"hf_overrides": json.dumps(
{"architectures": ["BgeM3EmbeddingModel"], "head_dtype": "float16"}
),
}
dense_embedding_sum = [
-0.7214539647102356, # "What is the capital of France?"
-0.6926871538162231, # "What is the capital of Germany?"
-0.7129564881324768, # "What is the capital of Spain?"
]
def _float_close(expected: object, result: object):
assert isinstance(expected, float) and isinstance(result, float), (
f"{expected=} or {result=} is not float"
)
return (expected - result) < 1e-3 or abs(expected / result - 1) < 1e-3
def _get_attr_or_val(obj: object | dict, key: str):
if isinstance(obj, dict) and key in obj:
return obj[key]
return getattr(obj, key, None)
def _check_dense_embedding(data, index=0):
assert _float_close(sum(data), dense_embedding_sum[index]), (
"dense-embedding result not match"
)
def _check_sparse_embedding(data, check_tokens=False):
expected_weights = [
{"token_id": 32, "weight": 0.0552978515625, "token": "?"},
{"token_id": 70, "weight": 0.09808349609375, "token": "the"},
{"token_id": 83, "weight": 0.08154296875, "token": "is"},
{"token_id": 111, "weight": 0.11810302734375, "token": "of"},
{"token_id": 4865, "weight": 0.1171875, "token": "What"},
{"token_id": 9942, "weight": 0.292236328125, "token": "France"},
{"token_id": 10323, "weight": 0.2802734375, "token": "capital"},
]
expected_embed = {x["token_id"]: x for x in expected_weights}
assert len(data) == len(expected_embed)
for entry in data:
expected_val = expected_embed[_get_attr_or_val(entry, "token_id")]
assert _float_close(
expected_val["weight"], _get_attr_or_val(entry, "weight")
), f"actual embed {entry} not equal to {expected_val}"
if check_tokens:
assert expected_val["token"] == _get_attr_or_val(entry, "token"), (
f"actual embed {entry} not equal to {expected_val}"
)
else:
assert _get_attr_or_val(entry, "token") is None, (
f"{entry} should not return token"
)
@pytest.fixture(scope="function")
def server():
args = [
"--runner",
"pooling",
"--enforce-eager",
"--max-num-seqs",
"32",
"--hf_overrides",
model_config["hf_overrides"],
"--io-processor-plugin",
model_config["plugin"],
]
with RemoteOpenAIServer(model_config["model_name"], args) as remote_server:
yield remote_server
@pytest.mark.asyncio
@pytest.mark.parametrize(
"return_tokens",
[True, False],
)
async def test_bge_m3_sparse_plugin_online(
server: RemoteOpenAIServer, return_tokens: bool
):
"""Test BGE-M3 sparse plugin in online mode via API."""
request_payload = {
"model": model_config["model_name"],
"task": "plugin",
"data": {"input": model_config["test_input"], "return_tokens": return_tokens},
}
ret = requests.post(
server.url_for("pooling"),
json=request_payload,
)
response = ret.json()
# Verify the request response is in the correct format
assert (parsed_response := IOProcessorResponse(**response).data)
# Verify the output is formatted as expected for this plugin
assert _get_attr_or_val(parsed_response, "data")
assert len(_get_attr_or_val(parsed_response, "data")) > 0
data_entry = _get_attr_or_val(parsed_response, "data")[0]
assert _get_attr_or_val(data_entry, "object") == "dense&sparse"
assert _get_attr_or_val(data_entry, "sparse_embedding")
# Verify sparse embedding format
sparse_embedding = _get_attr_or_val(data_entry, "sparse_embedding")
assert isinstance(sparse_embedding, list)
_check_sparse_embedding(sparse_embedding, return_tokens)
# Verify dense embedding format
dense_embedding = _get_attr_or_val(data_entry, "dense_embedding")
assert isinstance(dense_embedding, list)
_check_dense_embedding(dense_embedding)
# Verify usage information
usage = _get_attr_or_val(parsed_response, "usage")
assert usage, f"usage not found for {parsed_response}"
assert _get_attr_or_val(usage, "prompt_tokens") > 0
assert _get_attr_or_val(usage, "total_tokens") == _get_attr_or_val(
usage, "prompt_tokens"
)
@pytest.mark.parametrize(
"return_tokens",
[True, False],
)
def test_bge_m3_sparse_plugin_offline(vllm_runner, return_tokens: bool):
"""Test BGE-M3 sparse plugin in offline mode."""
prompt = {
"data": {
"input": model_config["test_input"],
"return_tokens": return_tokens,
}
}
with vllm_runner(
model_config["model_name"],
runner="pooling",
enforce_eager=True,
max_num_seqs=32,
io_processor_plugin=model_config["plugin"],
hf_overrides=json.loads(model_config["hf_overrides"]),
default_torch_num_threads=1,
) as llm_runner:
llm = llm_runner.get_llm()
pooler_output = llm.encode(prompt, pooling_task="plugin")
outputs = pooler_output[0]
# Verify output structure
assert hasattr(outputs, "outputs")
response = outputs.outputs
assert hasattr(response, "data")
assert len(response.data) == 1
# Verify response data
for i, output in enumerate(response.data):
# Each output should have sparse embeddings
sparse_embedding = output.sparse_embedding
assert isinstance(sparse_embedding, list)
_check_sparse_embedding(sparse_embedding, return_tokens)
dense_embedding = output.dense_embedding
assert isinstance(dense_embedding, list)
_check_dense_embedding(dense_embedding)
# Verify usage
assert response.usage.prompt_tokens > 0
assert response.usage.total_tokens == response.usage.prompt_tokens
def test_bge_m3_sparse_plugin_offline_multiple_inputs(vllm_runner):
"""Test BGE-M3 sparse plugin with multiple inputs in offline mode."""
prompts = {
"data": {
"input": [
"What is the capital of France?",
"What is the capital of Germany?",
"What is the capital of Spain?",
],
"return_tokens": True,
}
}
with vllm_runner(
model_config["model_name"],
runner="pooling",
enforce_eager=True,
max_num_seqs=32,
io_processor_plugin=model_config["plugin"],
hf_overrides=json.loads(model_config["hf_overrides"]),
default_torch_num_threads=1,
) as llm_runner:
llm = llm_runner.get_llm()
pooler_output = llm.encode(prompts, pooling_task="plugin")
outputs = pooler_output[0]
# Verify output structure
assert hasattr(outputs, "outputs")
response = outputs.outputs
assert hasattr(response, "data")
assert len(response.data) == 3
for i, output in enumerate(response.data):
# Each output should have sparse embeddings
sparse_embedding = output.sparse_embedding
assert isinstance(sparse_embedding, list)
dense_embedding = output.dense_embedding
assert isinstance(dense_embedding, list)
_check_dense_embedding(dense_embedding, i)
# Verify usage
assert response.usage.prompt_tokens > 0
assert response.usage.total_tokens == response.usage.prompt_tokens
@@ -108,14 +108,14 @@ def test_register_quantization_config(caplog_vllm):
assert get_quantization_config("custom_quant") == CustomQuantConfig
# The quantization method `custom_quant` is already exists,
# should raise a debug message when re-registering it.
with caplog_vllm.at_level(logging.DEBUG, logger="vllm"):
# should raise a warning when re-registering it.
with caplog_vllm.at_level(logging.WARNING):
register_quantization_config("custom_quant")(CustomQuantConfig)
assert any(
"The quantization method 'custom_quant' already exists" in message
for message in caplog_vllm.messages
), "Expected a debug message when re-registering custom_quant"
), "Expected a warning when re-registering custom_quant"
@pytest.mark.parametrize(
@@ -103,18 +103,14 @@ def test_worker_methods_delegate_to_store_worker():
worker = mock_worker_cls.return_value
worker.get_finished.return_value = ({"req-1"}, {"req-2"})
worker.get_block_ids_with_load_errors.return_value = {3, 4}
connector.bind_connector_metadata(metadata)
connector.register_kv_caches(kv_caches)
result = connector.get_finished(finished_req_ids)
invalid_block_ids = connector.get_block_ids_with_load_errors()
worker.register_kv_caches.assert_called_once_with(kv_caches)
worker.get_finished.assert_called_once_with(finished_req_ids, metadata)
assert result == ({"req-1"}, {"req-2"})
worker.get_block_ids_with_load_errors.assert_called_once_with()
assert invalid_block_ids == {3, 4}
def test_get_kv_connector_kv_cache_events_returns_none_when_empty():
@@ -300,85 +300,3 @@ def test_store_mask_fast_path_single_attention_group():
assert len(coord.attention_groups) == 1
masks = coord.store_mask(64)
assert masks == ([True] * 4, [True] * 4)
# ----- Eagle / MTP interaction with load_mask -----
def test_lookup_with_eagle_pops_last_full_attention_block():
"""Sanity: with use_eagle, find_longest_cache_hit drops the last block.
Pairs with the load_mask test below to lock the round-trip contract."""
groups = [KVCacheGroupSpec(["L0"], _full(16))]
coord = _make_coord(groups, hash_block_size=16, use_eagle=True)
hs = _hashes(4)
cmap = ExternalCachedBlockPool({(0, bytes(h)) for h in hs})
_masks, hit = coord.find_longest_cache_hit(
hs, max_length=64, cached_block_pool=cmap
)
# 4 blocks present, eagle pops 1 → 3 blocks = 48 tokens.
assert hit == 48
def test_load_mask_with_eagle_does_not_double_prune_full_attention():
"""Regression for silent KV corruption with MTP/EAGLE-3.
The recv side calls ``load_mask(block_hashes, token_len)`` where
``token_len`` is already the eagle-pruned hit length from ``lookup``.
A second eagle pop here used to shorten the mask by one extra block;
``process_tokens`` then yielded a chunk past the mask, which the worker
silently skipped — leaving the trailing block of the loaded prefix
uninitialized in local KV.
"""
groups = [KVCacheGroupSpec(["L0"], _full(16))]
coord = _make_coord(groups, hash_block_size=16, use_eagle=True)
hs = _hashes(4)
cmap = ExternalCachedBlockPool({(0, bytes(h)) for h in hs})
_masks, hit = coord.find_longest_cache_hit(
hs, max_length=64, cached_block_pool=cmap
)
assert hit == 48 # eagle popped 1 block
masks = coord.load_mask(hs, token_len=hit)
# Every chunk that process_tokens(token_len=48, ...) would yield must
# have a corresponding mask slot. process_tokens emits chunk_id 0..2
# (start=0, 16, 32), so the mask must be length 3, all True.
assert masks[0] == [True, True, True]
def test_load_mask_with_eagle_hybrid_full_plus_swa():
"""Hybrid (FullAttn + SWA) with eagle: load_mask must cover every chunk
in [0, token_len) for the FullAttn group; SWA group keeps its
tail-window mask."""
groups = [
KVCacheGroupSpec(["L0"], _full(16)),
KVCacheGroupSpec(["L1"], _swa(16, 32)),
]
coord = _make_coord(groups, hash_block_size=16, use_eagle=True)
hs = _hashes(4)
exists = {(g, bytes(h)) for g in (0, 1) for h in hs}
cmap = ExternalCachedBlockPool(exists)
_masks, hit = coord.find_longest_cache_hit(
hs, max_length=64, cached_block_pool=cmap
)
# FullAttn dictates the convergence; eagle pops one block off it.
assert hit == 48
masks = coord.load_mask(hs, token_len=hit)
# FullAttn: all chunks populated locally.
assert masks[0] == [True, True, True]
# SWA: tail-window only (ceil((32-1)/16) = 2 trailing blocks).
assert masks[1][-2:] == [True, True]
def test_load_mask_without_eagle_unchanged():
"""Sanity: when eagle is off, load_mask is identical to the pre-fix path."""
groups = [KVCacheGroupSpec(["L0"], _full(16))]
coord = _make_coord(groups, hash_block_size=16, use_eagle=False)
hs = _hashes(4)
cmap = ExternalCachedBlockPool({(0, bytes(h)) for h in hs})
_masks, hit = coord.find_longest_cache_hit(
hs, max_length=64, cached_block_pool=cmap
)
assert hit == 64
masks = coord.load_mask(hs, token_len=hit)
assert masks[0] == [True, True, True, True]
@@ -77,7 +77,6 @@ def _make_store_sending_thread(
def _make_store_recving_thread(
store: MagicMock,
*,
tp_rank: int = 0,
disk_offload_buffer_budget_bytes: int | None = None,
) -> mooncake_store_worker.KVCacheStoreRecvingThread:
from vllm.v1.kv_cache_interface import FullAttentionSpec, KVCacheGroupSpec
@@ -97,7 +96,7 @@ def _make_store_recving_thread(
store=store,
token_databases=[token_database],
block_size=16,
tp_rank=tp_rank,
tp_rank=0,
ready_event=threading.Event(),
coord=coord,
disk_offload_buffer_budget_bytes=disk_offload_buffer_budget_bytes,
@@ -461,67 +460,6 @@ def test_store_sending_thread_only_skips_on_no_available_handle():
assert store.batch_put_from_multi_buffers.call_count == 2
def test_store_recving_thread_reports_failed_block_ids():
store = MagicMock()
store.batch_get_into_multi_buffers.return_value = [256, -5, -7]
thread = _make_store_recving_thread(store)
thread._handle_request(
_make_load_req(
"req-a",
[b"a0", b"a1", b"a2"],
token_len=48,
)
)
assert thread.get_and_clear_finished_requests() == {"req-a"}
assert thread.get_and_clear_block_ids_with_load_errors() == {1, 2}
assert thread.get_and_clear_block_ids_with_load_errors() == set()
def test_store_recving_thread_reports_failed_block_ids_after_rotation():
store = MagicMock()
store.batch_get_into_multi_buffers.return_value = [256, -5, 256]
thread = _make_store_recving_thread(store, tp_rank=1)
thread._handle_request(
_make_load_req(
"req-a",
[b"a0", b"a1", b"a2"],
token_len=48,
)
)
assert thread.get_and_clear_block_ids_with_load_errors() == {2}
def test_store_recving_thread_reports_all_attempted_blocks_on_exception():
store = MagicMock()
store.batch_get_into_multi_buffers.side_effect = RuntimeError("boom")
thread = _make_store_recving_thread(store)
thread._handle_request(
_make_load_req(
"req-a",
[b"a0", b"a1", b"a2"],
token_len=48,
)
)
assert thread.get_and_clear_finished_requests() == {"req-a"}
assert thread.get_and_clear_block_ids_with_load_errors() == {0, 1, 2}
def test_store_worker_get_block_ids_with_load_errors_delegates_to_recv_thread():
recv_thread = MagicMock()
recv_thread.get_and_clear_block_ids_with_load_errors.return_value = {3, 4}
w = _make_bare_worker()
w.kv_recv_thread = recv_thread
assert w.get_block_ids_with_load_errors() == {3, 4}
recv_thread.get_and_clear_block_ids_with_load_errors.assert_called_once_with()
def test_store_sending_thread_passes_replicate_config_when_preferred_segment_set():
store = MagicMock()
store.batch_is_exist.side_effect = lambda keys: [0] * len(keys)
@@ -744,7 +682,6 @@ def test_recv_thread_stops_after_first_failing_disk_offload_sub_batch():
thread._handle_request(req)
assert store.batch_get_into_multi_buffers.call_count == 1
assert thread.get_and_clear_block_ids_with_load_errors() == {0, 1}
def test_recv_thread_skips_split_when_budget_holds_all_keys():
@@ -774,26 +711,21 @@ def test_recv_thread_skips_split_when_budget_holds_all_keys():
def test_recv_thread_reports_unsplittable_key_larger_than_budget():
# Non-zero tp_rank exercises the rotation: the oversized key may not sit
# at the original request's first block. Every block must still be marked
# invalid, since none are loaded.
store = MagicMock()
thread = _make_store_recving_thread(
store,
tp_rank=2,
disk_offload_buffer_budget_bytes=_DISK_OFFLOAD_BUDGET_TOO_SMALL,
)
req = _make_load_req(
"req-a",
[b"a0", b"a1", b"a2"],
token_len=48,
[b"a0"],
token_len=16,
)
thread._handle_request(req)
assert store.batch_get_into_multi_buffers.call_count == 0
assert thread.get_and_clear_block_ids_with_load_errors() == {0, 1, 2}
def test_requester_worker_init_uses_positional_setup(tmp_path, monkeypatch):
-1
View File
@@ -1540,7 +1540,6 @@ class ModelConfig:
return "token_classify"
priority: list[PoolingTask] = [
"embed&token_classify",
"embed",
"classify",
"token_embed",
-134
View File
@@ -1,134 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
from cutlass import BFloat16, Float32, Int64, Uint32, cute
from cutlass._mlir import ir
from cutlass._mlir.dialects import llvm, vector
from cutlass.cute.nvgpu import cpasync
from cutlass.cutlass_dsl import T, dsl_user_op
# https://github.com/NVIDIA/cutlass/blob/v4.3.2/include/cute/arch/copy_sm90_desc.hpp#L193-L197
EVICT_NORMAL = Int64(0x1000000000000000)
EVICT_FIRST = Int64(0x12F0000000000000)
EVICT_LAST = Int64(0x14F0000000000000)
@dsl_user_op
def recast_val(x, dtype, *, loc=None, ip=None):
return dtype(llvm.bitcast(dtype.mlir_type, x.ir_value(loc=loc, ip=ip)))
def simple_tma_copy(atom, src, dst, mbar=None, cache_policy=None):
"""A simple helper that wraps group_modes() and tma_partition()
NOTE: this should be called WITHOUT cute.elect_one()
"""
if isinstance(atom.op, cpasync.CopyBulkTensorTileG2SOp):
gmem = src
smem = dst
elif isinstance(atom.op, cpasync.CopyBulkTensorTileS2GOp):
smem = src
gmem = dst
else:
raise ValueError
s_part, g_part = cpasync.tma_partition(
atom,
0,
cute.make_layout(1),
cute.group_modes(smem, 0),
cute.group_modes(gmem, 0),
)
if isinstance(atom.op, cpasync.CopyBulkTensorTileG2SOp):
cute.copy(atom, g_part, s_part, tma_bar_ptr=mbar, cache_policy=cache_policy)
elif isinstance(atom.op, cpasync.CopyBulkTensorTileS2GOp):
cute.copy(atom, s_part, g_part, cache_policy=cache_policy)
else:
raise ValueError
# can't find the equivalent in nvvm
@dsl_user_op
def fence_before_tma_store(*, loc=None, ip=None):
llvm.inline_asm(
T.i32(),
[],
"mov.u32 $0, 0;\n\t"
"fence.proxy.async::generic.release.sync_restrict::shared::cta.cluster;",
"=r",
has_side_effects=True,
is_align_stack=False,
loc=loc,
ip=ip,
)
@dsl_user_op
def mma_bf16(
a: cute.TensorSSA, b: cute.TensorSSA, c: cute.TensorSSA, *, loc=None, ip=None
):
if a.element_type == BFloat16:
a = cute.recast_tensor(a, Uint32)
if b.element_type == BFloat16:
b = cute.recast_tensor(b, Uint32)
mlir_ty = Float32.mlir_type
out = llvm.inline_asm(
llvm.StructType.get_literal([mlir_ty] * 4),
[a[i].ir_value(loc=loc, ip=ip) for i in range(4)]
+ [b[i].ir_value(loc=loc, ip=ip) for i in range(2)]
+ [c[i].ir_value(loc=loc, ip=ip) for i in range(4)],
"mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 "
"{$0, $1, $2, $3}, {$4, $5, $6, $7}, {$8, $9}, "
"{$10, $11, $12, $13};",
"=f,=f,=f,=f,r,r,r,r,r,r,f,f,f,f",
has_side_effects=False,
is_align_stack=False,
loc=loc,
ip=ip,
)
vec = vector.from_elements(
ir.VectorType.get([4], mlir_ty, loc=loc),
[llvm.extractvalue(mlir_ty, out, [i], loc=loc, ip=ip) for i in range(4)],
loc=loc,
ip=ip,
)
return cute.TensorSSA(vec, 4, Float32)
@dsl_user_op
def _bf16x2_abs(a: Uint32, *, loc=None, ip=None) -> Uint32:
out = llvm.inline_asm(
T.i32(),
[a.ir_value(loc=loc, ip=ip)],
"abs.bf16x2 $0, $1;",
"=r,r",
has_side_effects=False,
is_align_stack=False,
)
return Uint32(out)
@dsl_user_op
def _bf16x2_max(a: Uint32, b: Uint32, *, loc=None, ip=None) -> Uint32:
out = llvm.inline_asm(
T.i32(),
[a.ir_value(loc=loc, ip=ip), b.ir_value(loc=loc, ip=ip)],
"max.bf16x2 $0, $1, $2;",
"=r,r,r",
has_side_effects=False,
is_align_stack=False,
)
return Uint32(out)
@dsl_user_op
def _bf16x2_mul(a: Uint32, b: Uint32, *, loc=None, ip=None) -> Uint32:
out = llvm.inline_asm(
T.i32(),
[a.ir_value(loc=loc, ip=ip), b.ir_value(loc=loc, ip=ip)],
"mul.rn.bf16x2 $0, $1, $2;",
"=r,r,r",
has_side_effects=False,
is_align_stack=False,
)
return Uint32(out)
-219
View File
@@ -1,219 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
# this module is named _tcgen05 to avoid name collision with cute.nvgpu.tcgen05
import cutlass
from cutlass import Boolean, Float32, Int32, Uint32, Uint64, cute
from cutlass._mlir import ir
from cutlass._mlir.dialects import llvm, nvvm, vector
from cutlass.cutlass_dsl import dsl_user_op
NVVM_CTA_GROUP_MAP = [
None,
nvvm.Tcgen05GroupKind.CTA_1,
nvvm.Tcgen05GroupKind.CTA_2,
]
LDST_MAP = {
"32x32b": (nvvm.Tcgen05LdStShape.SHAPE_32X32B, 1),
"16x128b": (nvvm.Tcgen05LdStShape.SHAPE_16X128B, 2),
"16x256b": (nvvm.Tcgen05LdStShape.SHAPE_16X256B, 4),
}
def _make_tmem_llvm_ptr(addr, *, loc=None, ip=None):
ptr_ty = llvm.PointerType.get(cute.AddressSpace.tmem.value)
val = Int32(addr).ir_value(loc=loc, ip=ip)
return llvm.inttoptr(ptr_ty, val, loc=loc, ip=ip)
@dsl_user_op
def alloc(
taddr: cute.Pointer,
cta_group: int = 1,
*,
loc=None,
ip=None,
) -> None:
nvvm.tcgen05_alloc(
taddr.to_llvm_ptr(loc=loc, ip=ip),
Uint32(512).ir_value(loc=loc, ip=ip),
group=NVVM_CTA_GROUP_MAP[cta_group],
loc=loc,
ip=ip,
)
@dsl_user_op
def dealloc(cta_group: int = 1, *, loc=None, ip=None) -> None:
nvvm.tcgen05_dealloc(
_make_tmem_llvm_ptr(0, loc=loc, ip=ip),
Int32(512).ir_value(loc=loc, ip=ip),
group=NVVM_CTA_GROUP_MAP[cta_group],
loc=loc,
ip=ip,
)
def make_bf16_idesc(
MMA_M: int,
MMA_N: int,
*,
negate_A: bool = False,
negate_B: bool = False,
transpose_A: bool = False,
transpose_B: bool = False,
):
idesc = Uint32(
(1 << 4) | (1 << 7) | (1 << 10) | ((MMA_N >> 3) << 17) | ((MMA_M >> 4) << 24)
)
idesc |= Uint32(negate_A) << 13
idesc |= Uint32(negate_B) << 14
idesc |= Uint32(transpose_A) << 15
idesc |= Uint32(transpose_B) << 16
return idesc
def make_sdesc_128B_swizzle(LBO: int):
SBO = 8 * 128
return Uint64((LBO >> 4 << 16) | (SBO >> 4 << 32) | (1 << 46) | (2 << 61))
@dsl_user_op
def mma_f16(
d_tmem,
a_desc,
b_desc,
idesc,
enable_input_d,
cta_group: int = 1,
*,
loc=None,
ip=None,
) -> None:
nvvm.tcgen05_mma(
nvvm.Tcgen05MMAKind.F16,
NVVM_CTA_GROUP_MAP[cta_group],
_make_tmem_llvm_ptr(d_tmem, loc=loc, ip=ip),
Uint64(a_desc).ir_value(loc=loc, ip=ip),
Uint64(b_desc).ir_value(loc=loc, ip=ip),
Int32(idesc).ir_value(loc=loc, ip=ip),
Boolean(enable_input_d).ir_value(loc=loc, ip=ip),
loc=loc,
ip=ip,
)
@dsl_user_op
def mma_ts_f16(
d_tmem,
a_tmem,
b_desc,
idesc,
enable_input_d,
cta_group: int = 1,
*,
loc=None,
ip=None,
) -> None:
nvvm.tcgen05_mma(
nvvm.Tcgen05MMAKind.F16,
NVVM_CTA_GROUP_MAP[cta_group],
_make_tmem_llvm_ptr(d_tmem, loc=loc, ip=ip),
_make_tmem_llvm_ptr(a_tmem, loc=loc, ip=ip),
Uint64(b_desc).ir_value(loc=loc, ip=ip),
Int32(idesc).ir_value(loc=loc, ip=ip),
Boolean(enable_input_d).ir_value(loc=loc, ip=ip),
loc=loc,
ip=ip,
)
@dsl_user_op
def commit(mbar, cta_mask=None, cta_group: int = 1, *, loc=None, ip=None):
mbar_llvm = mbar.to_llvm_ptr(loc=loc, ip=ip)
group = NVVM_CTA_GROUP_MAP[cta_group]
if cutlass.const_expr(cta_mask is not None):
nvvm.tcgen05_commit_arrive(
mbar_llvm,
multicast_mask=cta_mask.ir_value(loc=loc, ip=ip),
group=group,
loc=loc,
ip=ip,
)
else:
nvvm.tcgen05_commit_arrive(mbar_llvm, group=group, loc=loc, ip=ip)
@dsl_user_op
def ld(row, col, shape: str, num: int, *, loc=None, ip=None):
nvvm_shape, regs_per_num = LDST_MAP[shape]
num_regs = regs_per_num * num
tmem = (Int32(row) << Int32(16)) | Int32(col)
tmem_ptr = _make_tmem_llvm_ptr(tmem, loc=loc, ip=ip)
if num_regs == 1:
reg = nvvm.tcgen05_ld(Int32.mlir_type, nvvm_shape, tmem_ptr, loc=loc, ip=ip)
reg_f32 = llvm.bitcast(Float32.mlir_type, reg, loc=loc, ip=ip)
return Float32(reg_f32)
else:
vec_i32_ty = ir.VectorType.get([num_regs], Int32.mlir_type, loc=loc)
vec_f32_ty = ir.VectorType.get([num_regs], Float32.mlir_type, loc=loc)
regs = nvvm.tcgen05_ld(vec_i32_ty, nvvm_shape, tmem_ptr, loc=loc, ip=ip)
regs_f32 = llvm.bitcast(vec_f32_ty, regs, loc=loc, ip=ip)
return cute.TensorSSA(regs_f32, (num_regs,), Float32)
@dsl_user_op
def st(row, col, shape: str, num: int, vals, *, loc=None, ip=None) -> None:
# if input is TensorSSA, convert to Tensor so we can bitcast
if isinstance(vals, cute.TensorSSA):
vals_ = cute.make_rmem_tensor_like(vals)
vals_.store(vals)
vals = vals_
# bitcast to Int32
vals = cute.recast_tensor(vals, Int32)
nvvm_shape, regs_per_num = LDST_MAP[shape]
num_regs = regs_per_num * num
tmem = (Int32(row) << Int32(16)) | Int32(col)
tmem_ptr = _make_tmem_llvm_ptr(tmem, loc=loc, ip=ip)
if num_regs == 1:
nvvm.tcgen05_st(
nvvm_shape,
tmem_ptr,
vals[0].ir_value(loc=loc, ip=ip),
loc=loc,
ip=ip,
)
else:
vec_i32_ty = ir.VectorType.get([num_regs], Int32.mlir_type, loc=loc)
val_vec = vector.from_elements(
vec_i32_ty,
[vals[i].ir_value(loc=loc, ip=ip) for i in range(num_regs)],
loc=loc,
ip=ip,
)
nvvm.tcgen05_st(nvvm_shape, tmem_ptr, val_vec, loc=loc, ip=ip)
@dsl_user_op
def fence_after_thread_sync(*, loc=None, ip=None):
nvvm.tcgen05_fence(nvvm.Tcgen05FenceKind.AFTER_THREAD_SYNC, loc=loc, ip=ip)
@dsl_user_op
def fence_before_thread_sync(*, loc=None, ip=None):
nvvm.tcgen05_fence(nvvm.Tcgen05FenceKind.BEFORE_THREAD_SYNC, loc=loc, ip=ip)
@dsl_user_op
def wait_ld(*, loc=None, ip=None):
nvvm.tcgen05_wait(nvvm.Tcgen05WaitKind.LOAD, loc=loc, ip=ip)
@dsl_user_op
def wait_st(*, loc=None, ip=None):
nvvm.tcgen05_wait(nvvm.Tcgen05WaitKind.STORE, loc=loc, ip=ip)
@@ -270,10 +270,6 @@ class MooncakeStoreConnector(KVConnectorBase_V1, SupportsHMA):
assert isinstance(metadata, MooncakeStoreConnectorMetadata)
return self.connector_worker.get_finished(finished_req_ids, metadata)
def get_block_ids_with_load_errors(self) -> set[int]:
assert self.connector_worker is not None
return self.connector_worker.get_block_ids_with_load_errors()
def get_kv_connector_kv_cache_events(
self,
) -> MooncakeStoreKVEvents | None:
@@ -115,21 +115,12 @@ class MooncakeStoreCoordinator:
block_hashes: list[BlockHash],
max_length: int,
cached_block_pool: ExternalCachedBlockPool,
*,
apply_eagle: bool = True,
) -> tuple[tuple[list[bool], ...], int]:
"""Returns ``(load_mask_per_group, hit_length)``. ``mask[g][i]`` is True iff
group ``g`` populates chunk ``i`` locally (e.g. SWA and Mamba tail-only);
recv-side callers skip False slots.
``apply_eagle`` controls whether the per-spec ``use_eagle`` last-block
pop is applied. Lookup callers want it (the drafter requires recomputing
the last block); per-chunk mask callers must not, because ``token_len``
already reflects the eagle-pruned hit length and a second pop would
leave the trailing block unloaded.
"""
recv-side callers skip False slots."""
blocks_per_group, hit_length = self._find_hit_blocks(
block_hashes, max_length, cached_block_pool, apply_eagle=apply_eagle
block_hashes, max_length, cached_block_pool
)
masks = tuple(
[blk is not cached_block_pool.null_block for blk in blocks]
@@ -146,17 +137,8 @@ class MooncakeStoreCoordinator:
spec would populate chunk ``i`` locally at length ``token_len``
(e.g. SWA / Mamba tail-only).
"""
# ``apply_eagle=False`` because ``token_len`` is already the
# eagle-pruned hit length returned by ``client.lookup``. Re-applying
# the pop here would shorten the mask by one extra block; the recv
# thread would then silently skip the trailing chunk yielded by
# ``db.process_tokens`` and leave that block uninitialized in the
# local KV pool.
masks, _ = self.find_longest_cache_hit(
block_hashes,
token_len,
ExternalCachedBlockPool(),
apply_eagle=False,
block_hashes, token_len, ExternalCachedBlockPool()
)
return masks
@@ -213,17 +195,10 @@ class MooncakeStoreCoordinator:
block_hashes: list[BlockHash],
max_length: int,
cached_block_pool: ExternalCachedBlockPool,
*,
apply_eagle: bool = True,
) -> tuple[tuple[list[KVCacheBlock], ...], int]:
"""Mirrors HybridKVCacheCoordinator.find_longest_cache_hit but
dispatches via spec_manager_map (we don't allocate managers).
When ``apply_eagle`` is False, ignore ``eagle_attn_group_indices`` —
used by ``load_mask`` to avoid popping a second block on top of the
one already removed by the lookup.
"""
eagle_indices = self.eagle_attn_group_indices if apply_eagle else set()
if len(self.attention_groups) == 1:
spec, group_ids, manager_cls = self.attention_groups[0]
hashes = self.block_hashes_for_spec(block_hashes, spec)
@@ -233,7 +208,7 @@ class MooncakeStoreCoordinator:
kv_cache_group_ids=group_ids,
block_pool=cast(BlockPool, cached_block_pool),
kv_cache_spec=spec,
use_eagle=(0 in eagle_indices),
use_eagle=(0 in self.eagle_attn_group_indices),
alignment_tokens=spec.block_size,
)
num_groups = len(self.kv_cache_groups)
@@ -262,7 +237,9 @@ class MooncakeStoreCoordinator:
)
continue
use_eagle = idx in eagle_indices and idx not in eagle_verified
use_eagle = (
idx in self.eagle_attn_group_indices and idx not in eagle_verified
)
_max_length = curr_hit_length
if use_eagle:
_max_length = min(curr_hit_length + spec.block_size, max_length)
@@ -20,7 +20,7 @@ import time
from collections import defaultdict
from collections.abc import Callable
from dataclasses import dataclass
from typing import Any, Literal, TypeVar
from typing import Any, Literal
import regex as re
import torch
@@ -68,12 +68,6 @@ DEFAULT_GLOBAL_SEGMENT_SIZE = 4 * 1024 * 1024 * 1024 # 4 GiB
DEFAULT_LOCAL_BUFFER_SIZE = 4 * 1024 * 1024 * 1024 # 4 GiB
MOONCAKE_NO_AVAILABLE_HANDLE = -200
_T = TypeVar("_T")
def _rotate_list(values: list[_T], offset: int) -> list[_T]:
return values[offset:] + values[:offset]
# Mirrors FileStorageConfig::local_buffer_size in Mooncake C++.
DEFAULT_MOONCAKE_DISK_STAGING_BUFFER_BYTES = 1280 * 1024 * 1024
@@ -717,9 +711,6 @@ class KVCacheStoreRecvingThread(KVTransferThread):
name="KVCacheStoreRecvingThread",
record_operation=record_operation,
)
# _invalid_block_ids can be access by both the Worker and RecvingThread
self._invalid_block_ids_lock = threading.Lock()
self._invalid_block_ids: set[int] = set()
self.disk_offload_buffer_budget_bytes = disk_offload_buffer_budget_bytes
self.usable_disk_offload_buffer_budget_bytes = (
None
@@ -730,16 +721,6 @@ class KVCacheStoreRecvingThread(KVTransferThread):
)
self.coord = coord
def _add_load_error_block_ids(self, block_ids: list[int]) -> None:
with self._invalid_block_ids_lock:
self._invalid_block_ids.update(block_ids)
def get_and_clear_block_ids_with_load_errors(self) -> set[int]:
with self._invalid_block_ids_lock:
invalid_block_ids = self._invalid_block_ids.copy()
self._invalid_block_ids.clear()
return invalid_block_ids
def _handle_request(self, req_meta: ReqMeta):
token_len = req_meta.load_spec.token_len # type: ignore[union-attr]
req_id = req_meta.req_id
@@ -756,7 +737,6 @@ class KVCacheStoreRecvingThread(KVTransferThread):
addr_list: list[list[int]] = []
size_list: list[list[int]] = []
key_list: list[str] = []
block_id_list: list[int] = []
for g_idx, db in enumerate(self.token_databases):
mask = load_mask_per_group[g_idx]
for start, end, key in db.process_tokens(
@@ -765,29 +745,25 @@ class KVCacheStoreRecvingThread(KVTransferThread):
chunk_idx = start // db.block_size
if chunk_idx >= len(mask) or not mask[chunk_idx]:
continue
addr, size, block_id = db.prepare_value(
start, end, req_meta.block_ids[g_idx]
)
addr, size, _ = db.prepare_value(start, end, req_meta.block_ids[g_idx])
key_list.append(key.to_string())
addr_list.append(addr)
size_list.append(size)
block_id_list.append(block_id)
# Rotate aligned lists by tp_rank for load balancing.
# Rotate lists by tp_rank for load balancing
rotation = self.tp_rank % len(key_list)
key_list_c = _rotate_list(key_list, rotation)
addr_list_c = _rotate_list(addr_list, rotation)
size_list_c = _rotate_list(size_list, rotation)
block_id_list_c = _rotate_list(block_id_list, rotation)
key_list_c = key_list[rotation:] + key_list[:rotation]
addr_list_c = addr_list[rotation:] + addr_list[:rotation]
size_list_c = size_list[rotation:] + size_list[:rotation]
load_batches = [(key_list_c, addr_list_c, size_list_c, block_id_list_c)]
load_batches = [(key_list_c, addr_list_c, size_list_c)]
if self.usable_disk_offload_buffer_budget_bytes is not None:
total_staging_bytes = sum(
_estimate_disk_offload_staging_bytes(size) for size in size_list_c
)
if total_staging_bytes > self.usable_disk_offload_buffer_budget_bytes:
assert self.disk_offload_buffer_budget_bytes is not None
split_batches, oversized_key = _split_disk_offload_load_batches(
load_batches, oversized_key = _split_disk_offload_load_batches(
key_list_c,
addr_list_c,
size_list_c,
@@ -796,10 +772,6 @@ class KVCacheStoreRecvingThread(KVTransferThread):
)
if oversized_key is not None:
oversized_key_index = key_list_c.index(oversized_key)
# Mark every block: we skip the whole request, and the
# tp_rank rotation means oversized_key isn't necessarily
# the first block in the request's original order.
self._add_load_error_block_ids(block_id_list_c)
oversized_key_bytes = _estimate_disk_offload_staging_bytes(
size_list_c[oversized_key_index]
)
@@ -814,25 +786,12 @@ class KVCacheStoreRecvingThread(KVTransferThread):
self.set_finished_request(req_id)
self.request_queue.task_done()
return
load_batches = []
block_id_offset = 0
for batch_keys, batch_addrs, batch_sizes in split_batches:
next_block_id_offset = block_id_offset + len(batch_keys)
batch_block_ids = block_id_list_c[
block_id_offset:next_block_id_offset
]
load_batches.append(
(batch_keys, batch_addrs, batch_sizes, batch_block_ids)
)
block_id_offset = next_block_id_offset
current_batch_keys: list[str] = key_list_c
current_batch_block_ids: list[int] = block_id_list_c
batch_bytes = 0
try:
for batch_keys, batch_addrs, batch_sizes, batch_block_ids in load_batches:
for batch_keys, batch_addrs, batch_sizes in load_batches:
current_batch_keys = batch_keys
current_batch_block_ids = batch_block_ids
batch_bytes = _sum_batch_bytes(batch_sizes)
tiers_by_key: dict[str, str] | None = None
if envs.VLLM_MOONCAKE_STORE_TIER_LOG:
@@ -847,10 +806,8 @@ class KVCacheStoreRecvingThread(KVTransferThread):
req_id, batch_keys, res, tiers_by_key
)
failed = [
(key, value, block_id)
for key, value, block_id in zip(
batch_keys, res, batch_block_ids, strict=True
)
(key, value)
for key, value in zip(batch_keys, res, strict=True)
if value < 0
]
self._record_operation(
@@ -862,19 +819,15 @@ class KVCacheStoreRecvingThread(KVTransferThread):
num_failed_keys=len(failed),
)
if failed:
self._add_load_error_block_ids(
[block_id for _, _, block_id in failed]
)
logger.warning(
"Failed to get %d Mooncake keys from sub-batch "
"(batch_keys=%d, first_failures=%s)",
len(failed),
len(batch_keys),
[(key, value) for key, value, _ in failed[:3]],
failed[:3],
)
break
except Exception as e:
self._add_load_error_block_ids(current_batch_block_ids)
self._record_operation(
"load_get",
load_get_start,
@@ -1286,11 +1239,6 @@ class MooncakeStoreWorker:
)
return done_sending, done_recving
def get_block_ids_with_load_errors(self) -> set[int]:
if self.kv_recv_thread is None:
return set()
return self.kv_recv_thread.get_and_clear_block_ids_with_load_errors()
def _record_kv_connector_operation(
self,
operation: str,
+2 -2
View File
@@ -712,7 +712,7 @@ class EngineArgs:
)
fail_on_environ_validation: bool = False
gdn_prefill_backend: Literal["flashinfer", "triton", "cutedsl"] | None = None
gdn_prefill_backend: Literal["flashinfer", "triton"] | None = None
def __post_init__(self):
# support `EngineArgs(compilation_config={...})`
@@ -1527,7 +1527,7 @@ class EngineArgs:
parser.add_argument(
"--gdn-prefill-backend",
dest="gdn_prefill_backend",
choices=["flashinfer", "triton", "cutedsl"],
choices=["flashinfer", "triton"],
default=None,
help="Select GDN prefill backend.",
)
@@ -2,13 +2,17 @@
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import itertools
from abc import ABC, abstractmethod
from collections.abc import Callable, Iterable, Sequence
from typing import Any
from tqdm import tqdm
from vllm import RequestOutput, TextPrompt, TokensPrompt
from vllm.entrypoints.offline_utils import OfflineInferenceMixin
from vllm import PromptType, RequestOutput, TextPrompt, TokensPrompt
from vllm.inputs import EngineInput
from vllm.logger import init_logger
from vllm.lora.request import LoRARequest
from vllm.renderers import BaseRenderer
from vllm.sampling_params import BeamSearchParams, SamplingParams
from .utils import (
@@ -21,9 +25,11 @@ from .utils import (
logger = init_logger(__name__)
class BeamSearchOfflineMixin(OfflineInferenceMixin):
class BeamSearchOfflineMixin(ABC):
"""Offline inference for beam search"""
renderer: BaseRenderer
def beam_search(
self,
prompts: list[TokensPrompt | TextPrompt],
@@ -180,3 +186,41 @@ class BeamSearchOfflineMixin(OfflineInferenceMixin):
outputs.append(BeamSearchOutput(sequences=best_beams))
return outputs
@abstractmethod
def _preprocess_cmpl(
self,
prompts: Sequence[PromptType],
tokenization_kwargs: dict[str, Any] | None = None,
mm_processor_kwargs: dict[str, Any] | None = None,
) -> Sequence[EngineInput]:
raise NotImplementedError
@abstractmethod
def _lora_request_to_seq(
self,
lora_request: LoRARequest | None | Sequence[LoRARequest | None],
num_requests: int,
) -> Sequence[LoRARequest | None]:
raise NotImplementedError
@abstractmethod
def _params_to_seq(
self,
params: SamplingParams,
num_requests: int,
) -> Sequence[SamplingParams]:
raise NotImplementedError
@abstractmethod
def _render_and_run_requests(
self,
prompts: Iterable[EngineInput],
params: Sequence[SamplingParams],
output_type: type[RequestOutput],
*,
lora_requests: Sequence[LoRARequest | None] | None = None,
priorities: Sequence[int] | None = None,
use_tqdm: bool | Callable[..., tqdm] = True,
):
raise NotImplementedError
+596 -7
View File
@@ -1,7 +1,7 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
from collections.abc import Callable, Sequence
from collections.abc import Callable, Iterable, Sequence
from pathlib import Path
from typing import TYPE_CHECKING, Any
@@ -9,7 +9,7 @@ import cloudpickle
import torch.nn as nn
from pydantic import ValidationError
from tqdm.auto import tqdm
from typing_extensions import overload
from typing_extensions import TypeVar, overload
from vllm.config import (
AttentionConfig,
@@ -41,29 +41,47 @@ from vllm.entrypoints.chat_utils import (
from vllm.entrypoints.generate.beam_search.offline import BeamSearchOfflineMixin
from vllm.entrypoints.pooling.offline import PoolingOfflineMixin
from vllm.entrypoints.utils import log_non_default_args
from vllm.inputs import PromptType
from vllm.inputs import (
EngineInput,
PromptType,
)
from vllm.logger import init_logger
from vllm.lora.request import LoRARequest
from vllm.model_executor.layers.quantization import QuantizationMethods
from vllm.outputs import PoolingRequestOutput, RequestOutput
from vllm.platforms import current_platform
from vllm.sampling_params import SamplingParams
from vllm.pooling_params import PoolingParams
from vllm.renderers import ChatParams, merge_kwargs
from vllm.renderers.inputs.preprocess import (
conversation_to_seq,
parse_model_prompt,
prompt_to_seq,
)
from vllm.sampling_params import RequestOutputKind, SamplingParams
from vllm.tokenizers import TokenizerLike
from vllm.usage.usage_lib import UsageContext
from vllm.utils.counter import Counter
from vllm.utils.mistral import is_mistral_tokenizer
from vllm.utils.tqdm_utils import maybe_tqdm
from vllm.v1.engine import PauseMode
from vllm.v1.engine.llm_engine import LLMEngine
from vllm.v1.sample.logits_processor import LogitsProcessor
from .offline_utils import _O, _R, OfflineInferenceMixin
if TYPE_CHECKING:
from vllm.v1.metrics.reader import Metric
logger = init_logger(__name__)
_O = TypeVar(
"_O",
bound=RequestOutput | PoolingRequestOutput,
default=RequestOutput | PoolingRequestOutput,
)
_P = TypeVar("_P", bound=SamplingParams | PoolingParams | None)
_R = TypeVar("_R", default=Any)
class LLM(BeamSearchOfflineMixin, PoolingOfflineMixin, OfflineInferenceMixin):
class LLM(BeamSearchOfflineMixin, PoolingOfflineMixin):
"""An LLM for generating texts from given prompts and sampling parameters.
This class includes a tokenizer, a language model (possibly distributed
@@ -567,6 +585,58 @@ class LLM(BeamSearchOfflineMixin, PoolingOfflineMixin, OfflineInferenceMixin):
return self._run_engine(output_type, use_tqdm=use_tqdm)
def _resolve_mm_lora(
self,
prompt: EngineInput,
lora_request: LoRARequest | None,
) -> LoRARequest | None:
if prompt["type"] != "multimodal":
return lora_request
lora_config = self.llm_engine.vllm_config.lora_config
default_mm_loras = None if lora_config is None else lora_config.default_mm_loras
if not default_mm_loras:
return lora_request
prompt_modalities = prompt["mm_placeholders"].keys()
intersection = set(prompt_modalities).intersection(default_mm_loras.keys())
if not intersection:
return lora_request
if len(intersection) > 1:
# TODO: Would be nice to be able to have multiple loras per prompt
logger.warning(
"Multiple modality specific loras were registered and would be "
"used by a single prompt consuming several modalities; "
"currently we only support one lora per request; as such, "
"lora(s) registered with modalities: %s will be skipped",
intersection,
)
return lora_request
# Build the LoRA request; the ID of the default mm lora is the
# index of the modality name sorted alphabetically + 1.
modality_name = intersection.pop()
modality_lora_path = default_mm_loras[modality_name]
modality_lora_id = sorted(default_mm_loras).index(modality_name) + 1
# If we have a collision, warn if there is a collision,
# but always send the explicitly provided request.
if lora_request:
if lora_request.lora_int_id != modality_lora_id:
logger.warning(
"A modality with a registered lora and a lora_request "
"with a different ID were provided; falling back to the "
"lora_request as we only apply one LoRARequest per prompt"
)
return lora_request
return LoRARequest(
modality_name,
modality_lora_id,
modality_lora_path,
)
def collective_rpc(
self,
method: str | Callable[..., _R],
@@ -612,6 +682,139 @@ class LLM(BeamSearchOfflineMixin, PoolingOfflineMixin, OfflineInferenceMixin):
"""
return self.llm_engine.apply_model(func)
def _preprocess_cmpl(
self,
prompts: Sequence[PromptType],
tokenization_kwargs: dict[str, Any] | None = None,
mm_processor_kwargs: dict[str, Any] | None = None,
) -> Sequence[EngineInput]:
"""
Convert prompt inputs from LLM APIs (other than [LLM.chat][]) into
a format that can be passed to `_add_request`.
Refer to [LLM.generate][] for a complete description of the arguments.
Returns:
A list of `EngineInput` objects ready to be passed into LLMEngine.
"""
renderer = self.renderer
model_config = self.model_config
parsed_prompts = [
parse_model_prompt(model_config, prompt) for prompt in prompts
]
tok_params = renderer.default_cmpl_tok_params.with_kwargs(
**(tokenization_kwargs or {})
)
prompt_extras = (
None
if mm_processor_kwargs is None
else {"mm_processor_kwargs": mm_processor_kwargs}
)
return renderer.render_cmpl(
parsed_prompts,
tok_params,
prompt_extras=prompt_extras,
)
def _preprocess_cmpl_one(
self,
prompt: PromptType,
tokenization_kwargs: dict[str, Any] | None = None,
mm_processor_kwargs: dict[str, Any] | None = None,
) -> EngineInput:
(engine_input,) = self._preprocess_cmpl(
[prompt],
tokenization_kwargs,
mm_processor_kwargs=mm_processor_kwargs,
)
return engine_input
def _preprocess_chat(
self,
conversations: Sequence[list[ChatCompletionMessageParam]],
chat_template: str | None = None,
chat_template_content_format: ChatTemplateContentFormatOption = "auto",
chat_template_kwargs: dict[str, Any] | None = None,
add_generation_prompt: bool = True,
continue_final_message: bool = False,
tools: list[dict[str, Any]] | None = None,
tokenization_kwargs: dict[str, Any] | None = None,
mm_processor_kwargs: dict[str, Any] | None = None,
) -> Sequence[EngineInput]:
"""
Convert a list of conversations into prompts so that they can then
be used as input for other LLM APIs.
Refer to [LLM.chat][] for a complete description of the arguments.
Returns:
A list of `EngineInput` objects ready to be passed into LLMEngine.
"""
renderer = self.renderer
chat_params = ChatParams(
chat_template=chat_template,
chat_template_content_format=chat_template_content_format,
chat_template_kwargs=merge_kwargs(
chat_template_kwargs,
dict(
add_generation_prompt=add_generation_prompt,
continue_final_message=continue_final_message,
tools=tools,
tokenize=(
is_mistral_tokenizer(renderer.tokenizer)
or self.model_config.enable_prompt_embeds
),
),
),
mm_processor_kwargs=mm_processor_kwargs,
)
tok_params = renderer.default_chat_tok_params.with_kwargs(
**(tokenization_kwargs or {})
)
prompt_extras = (
None
if mm_processor_kwargs is None
else {"mm_processor_kwargs": mm_processor_kwargs}
)
_, engine_inputs = renderer.render_chat(
conversations,
chat_params,
tok_params,
prompt_extras=prompt_extras,
)
return engine_inputs
def _preprocess_chat_one(
self,
conversation: list[ChatCompletionMessageParam],
chat_template: str | None = None,
chat_template_content_format: ChatTemplateContentFormatOption = "auto",
chat_template_kwargs: dict[str, Any] | None = None,
add_generation_prompt: bool = True,
continue_final_message: bool = False,
tools: list[dict[str, Any]] | None = None,
tokenization_kwargs: dict[str, Any] | None = None,
mm_processor_kwargs: dict[str, Any] | None = None,
) -> EngineInput:
(engine_input,) = self._preprocess_chat(
[conversation],
chat_template=chat_template,
chat_template_content_format=chat_template_content_format,
chat_template_kwargs=chat_template_kwargs,
add_generation_prompt=add_generation_prompt,
continue_final_message=continue_final_message,
tools=tools,
tokenization_kwargs=tokenization_kwargs,
mm_processor_kwargs=mm_processor_kwargs,
)
return engine_input
def chat(
self,
messages: list[ChatCompletionMessageParam]
@@ -855,6 +1058,392 @@ class LLM(BeamSearchOfflineMixin, PoolingOfflineMixin, OfflineInferenceMixin):
"""
return self.llm_engine.get_metrics()
def _params_to_seq(
self,
params: _P | Sequence[_P],
num_requests: int,
) -> Sequence[_P]:
if isinstance(params, Sequence):
if len(params) != num_requests:
raise ValueError(
f"The lengths of prompts ({num_requests}) "
f"and params ({len(params)}) must be the same."
)
return params
return [params] * num_requests
def _lora_request_to_seq(
self,
lora_request: LoRARequest | None | Sequence[LoRARequest | None],
num_requests: int,
) -> Sequence[LoRARequest | None]:
if isinstance(lora_request, Sequence):
if len(lora_request) != num_requests:
raise ValueError(
f"The lengths of prompts ({num_requests}) "
f"and lora_request ({len(lora_request)}) must be the same."
)
return lora_request
return [lora_request] * num_requests
def _priority_to_seq(
self,
priority: list[int] | None,
num_requests: int,
) -> Sequence[int]:
if priority is not None:
if len(priority) != num_requests:
raise ValueError(
f"The lengths of prompts ({num_requests}) "
f"and priority ({len(priority)}) must be the same."
)
return priority
return [0] * num_requests
def _add_completion_requests(
self,
prompts: PromptType | Sequence[PromptType],
params: SamplingParams
| PoolingParams
| Sequence[SamplingParams | PoolingParams],
*,
use_tqdm: bool | Callable[..., tqdm] = True,
lora_request: Sequence[LoRARequest] | LoRARequest | None = None,
priority: list[int] | None = None,
tokenization_kwargs: dict[str, Any] | None = None,
mm_processor_kwargs: dict[str, Any] | None = None,
) -> list[str]:
seq_prompts = prompt_to_seq(prompts)
seq_params = self._params_to_seq(params, len(seq_prompts))
seq_lora_requests = self._lora_request_to_seq(lora_request, len(seq_prompts))
seq_priority = self._priority_to_seq(priority, len(seq_prompts))
return self._render_and_add_requests(
prompts=(
self._preprocess_cmpl_one(
prompt,
tokenization_kwargs,
mm_processor_kwargs=mm_processor_kwargs,
)
for prompt in maybe_tqdm(
seq_prompts,
use_tqdm=use_tqdm,
desc="Rendering prompts",
)
),
params=seq_params,
lora_requests=seq_lora_requests,
priorities=seq_priority,
)
def _run_completion(
self,
prompts: PromptType | Sequence[PromptType],
params: SamplingParams
| PoolingParams
| Sequence[SamplingParams | PoolingParams],
output_type: type[_O],
*,
use_tqdm: bool | Callable[..., tqdm] = True,
lora_request: Sequence[LoRARequest] | LoRARequest | None = None,
priority: list[int] | None = None,
tokenization_kwargs: dict[str, Any] | None = None,
mm_processor_kwargs: dict[str, Any] | None = None,
):
self._add_completion_requests(
prompts=prompts,
params=params,
use_tqdm=use_tqdm,
lora_request=lora_request,
priority=priority,
tokenization_kwargs=tokenization_kwargs,
mm_processor_kwargs=mm_processor_kwargs,
)
return self._run_engine(use_tqdm=use_tqdm, output_type=output_type)
def _run_chat(
self,
messages: list[ChatCompletionMessageParam]
| Sequence[list[ChatCompletionMessageParam]],
params: SamplingParams
| PoolingParams
| Sequence[SamplingParams | PoolingParams],
output_type: type[_O],
*,
use_tqdm: bool | Callable[..., tqdm] = True,
lora_request: Sequence[LoRARequest] | LoRARequest | None = None,
chat_template: str | None = None,
chat_template_content_format: ChatTemplateContentFormatOption = "auto",
add_generation_prompt: bool = True,
continue_final_message: bool = False,
tools: list[dict[str, Any]] | None = None,
chat_template_kwargs: dict[str, Any] | None = None,
tokenization_kwargs: dict[str, Any] | None = None,
mm_processor_kwargs: dict[str, Any] | None = None,
):
self._add_chat_requests(
messages=messages,
params=params,
use_tqdm=use_tqdm,
lora_request=lora_request,
chat_template=chat_template,
chat_template_content_format=chat_template_content_format,
chat_template_kwargs=chat_template_kwargs,
add_generation_prompt=add_generation_prompt,
continue_final_message=continue_final_message,
tools=tools,
tokenization_kwargs=tokenization_kwargs,
mm_processor_kwargs=mm_processor_kwargs,
)
return self._run_engine(output_type=output_type, use_tqdm=use_tqdm)
def _add_chat_requests(
self,
messages: list[ChatCompletionMessageParam]
| Sequence[list[ChatCompletionMessageParam]],
params: SamplingParams
| PoolingParams
| Sequence[SamplingParams | PoolingParams],
*,
use_tqdm: bool | Callable[..., tqdm] = True,
lora_request: Sequence[LoRARequest] | LoRARequest | None = None,
priority: list[int] | None = None,
chat_template: str | None = None,
chat_template_content_format: ChatTemplateContentFormatOption = "auto",
add_generation_prompt: bool = True,
continue_final_message: bool = False,
tools: list[dict[str, Any]] | None = None,
chat_template_kwargs: dict[str, Any] | None = None,
tokenization_kwargs: dict[str, Any] | None = None,
mm_processor_kwargs: dict[str, Any] | None = None,
) -> list[str]:
seq_convs = conversation_to_seq(messages)
seq_params = self._params_to_seq(params, len(seq_convs))
seq_lora_requests = self._lora_request_to_seq(lora_request, len(seq_convs))
seq_priority = self._priority_to_seq(priority, len(seq_convs))
# When thinking is enabled or tools are provided, and the model
# uses special tokens for structured output (e.g. Gemma4's
# <|channel>, <|tool_call>, <|"|>), automatically set
# skip_special_tokens=False so these tokens are preserved in
# output.text for downstream parsing.
needs_parsing = (
chat_template_kwargs and chat_template_kwargs.get("enable_thinking")
) or tools
if needs_parsing:
self._adjust_params_for_parsing(seq_params)
return self._render_and_add_requests(
prompts=(
self._preprocess_chat_one(
conversation,
chat_template=chat_template,
chat_template_content_format=chat_template_content_format,
chat_template_kwargs=chat_template_kwargs,
add_generation_prompt=add_generation_prompt,
continue_final_message=continue_final_message,
tools=tools,
tokenization_kwargs=tokenization_kwargs,
mm_processor_kwargs=mm_processor_kwargs,
)
for conversation in maybe_tqdm(
seq_convs,
use_tqdm=use_tqdm,
desc="Rendering conversations",
)
),
params=seq_params,
lora_requests=seq_lora_requests,
priorities=seq_priority,
)
def _adjust_params_for_parsing(
self, params: Sequence[SamplingParams | PoolingParams]
) -> None:
"""Set ``skip_special_tokens=False`` when the model encodes
structured output syntax as special tokens.
Models like Gemma4 register thinking delimiters
(``<|channel>``/``<channel|>``) and tool call tokens
(``<|tool_call>``/``<tool_call|>``/``<|"|>``) as special tokens.
The default ``skip_special_tokens=True`` strips them from
``output.text``, breaking parsing of both reasoning blocks and
tool calls.
This is a no-op for models whose structured tokens are regular
text tokens (e.g. DeepSeek's ``<think>``/``</think>``).
"""
# The offline API currently lacks a unified rendering pipeline.
# Until the planned Renderer refactor is complete, we hardcode
# this token preservation logic specifically for Gemma4 models
# to avoid regressions on other models.
hf_config = getattr(self.model_config, "hf_config", None)
architectures = getattr(hf_config, "architectures", [])
if any("Gemma4" in arch for arch in architectures):
tokenizer = self.renderer.get_tokenizer()
vocab = tokenizer.get_vocab()
special_ids = set(getattr(tokenizer, "all_special_ids", []))
# Tokens used for thinking delimiters and tool call syntax
# that some models (Gemma4) register as special tokens.
structured_tokens = (
"<|channel>",
"<channel|>", # thinking delimiters
"<|tool_call>",
"<tool_call|>", # tool call delimiters
'<|"|>', # string quoting in tool args
)
needs_special = any(
vocab.get(tok) in special_ids
for tok in structured_tokens
if tok in vocab
)
if needs_special:
for sp in params:
if isinstance(sp, SamplingParams) and sp.skip_special_tokens:
sp.skip_special_tokens = False
def _render_and_run_requests(
self,
prompts: Iterable[EngineInput],
params: Sequence[SamplingParams | PoolingParams],
output_type: type[_O],
*,
lora_requests: Sequence[LoRARequest | None] | None = None,
priorities: Sequence[int] | None = None,
use_tqdm: bool | Callable[..., tqdm] = True,
):
if isinstance(prompts, (list, tuple)):
logger.warning_once(
"Rendering all prompts before adding them to the engine "
"is less efficient than performing both on the same prompt "
"before processing the next prompt. You should instead pass "
"a generator that renders one prompt per iteration, as that allows "
"engine execution to begin for the first prompt while processing "
"the next prompt."
)
self._render_and_add_requests(
prompts=prompts,
params=params,
lora_requests=lora_requests,
priorities=priorities,
)
return self._run_engine(output_type, use_tqdm=use_tqdm)
def _render_and_add_requests(
self,
prompts: Iterable[EngineInput],
params: Sequence[SamplingParams | PoolingParams],
*,
lora_requests: Sequence[LoRARequest | None] | None = None,
priorities: Sequence[int] | None = None,
) -> list[str]:
added_request_ids: list[str] = []
try:
for i, prompt in enumerate(prompts):
request_id = self._add_request(
prompt,
params[i],
lora_request=self._resolve_mm_lora(
prompt,
None if lora_requests is None else lora_requests[i],
),
priority=0 if priorities is None else priorities[i],
)
added_request_ids.append(request_id)
except Exception as e:
if added_request_ids:
self.llm_engine.abort_request(added_request_ids, internal=True)
raise e
return added_request_ids
def _add_request(
self,
prompt: EngineInput,
params: SamplingParams | PoolingParams,
lora_request: LoRARequest | None = None,
priority: int = 0,
) -> str:
if isinstance(params, SamplingParams):
# We only care about the final output
params.output_kind = RequestOutputKind.FINAL_ONLY
request_id = str(next(self.request_counter))
return self.llm_engine.add_request(
request_id,
prompt,
params,
lora_request=lora_request,
priority=priority,
)
def _run_engine(
self,
output_type: type[_O] | tuple[type[_O], ...],
*,
use_tqdm: bool | Callable[..., tqdm] = True,
) -> list[_O]:
# Initialize tqdm.
if use_tqdm:
num_requests = self.llm_engine.get_num_unfinished_requests()
tqdm_func = use_tqdm if callable(use_tqdm) else tqdm
pbar = tqdm_func(
total=num_requests,
desc="Processed prompts",
dynamic_ncols=True,
postfix=(f"est. speed input: {0:.2f} toks/s, output: {0:.2f} toks/s"),
)
# Run the engine.
outputs: list[_O] = []
total_in_toks = 0
total_out_toks = 0
while self.llm_engine.has_unfinished_requests():
step_outputs = self.llm_engine.step()
for output in step_outputs:
assert isinstance(output, output_type)
if output.finished:
outputs.append(output) # type: ignore[arg-type]
if use_tqdm:
if isinstance(output, RequestOutput):
# Calculate tokens only for RequestOutput
n = len(output.outputs)
assert output.prompt_token_ids is not None
total_in_toks += len(output.prompt_token_ids) * n
in_spd = total_in_toks / pbar.format_dict["elapsed"]
total_out_toks += sum(
len(stp.token_ids) for stp in output.outputs
)
out_spd = total_out_toks / pbar.format_dict["elapsed"]
pbar.postfix = (
f"est. speed input: {in_spd:.2f} toks/s, "
f"output: {out_spd:.2f} toks/s"
)
pbar.update(n)
else:
pbar.update(1)
if pbar.n == num_requests:
pbar.refresh()
if use_tqdm:
pbar.close()
# Sort the outputs by request ID.
# This is necessary because some requests may be finished earlier than
# its previous requests.
return sorted(outputs, key=lambda x: int(x.request_id))
def init_weight_transfer_engine(
self, request: WeightTransferInitRequest | dict
) -> None:
-626
View File
@@ -1,626 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
from collections.abc import Callable, Iterable, Sequence
from typing import Any
from tqdm import tqdm
from typing_extensions import TypeVar
from vllm import (
PoolingParams,
PoolingRequestOutput,
PromptType,
RequestOutput,
SamplingParams,
)
from vllm.config import ModelConfig
from vllm.entrypoints.chat_utils import (
ChatCompletionMessageParam,
ChatTemplateContentFormatOption,
)
from vllm.inputs import EngineInput
from vllm.logger import init_logger
from vllm.lora.request import LoRARequest
from vllm.renderers import BaseRenderer, ChatParams, merge_kwargs
from vllm.renderers.inputs.preprocess import (
conversation_to_seq,
parse_model_prompt,
prompt_to_seq,
)
from vllm.sampling_params import RequestOutputKind
from vllm.utils.counter import Counter
from vllm.utils.mistral import is_mistral_tokenizer
from vllm.utils.tqdm_utils import maybe_tqdm
from vllm.v1.engine.llm_engine import LLMEngine
logger = init_logger(__name__)
_P = TypeVar("_P", bound=SamplingParams | PoolingParams | None)
_O = TypeVar(
"_O",
bound=RequestOutput | PoolingRequestOutput,
default=RequestOutput | PoolingRequestOutput,
)
_R = TypeVar("_R", default=Any)
class OfflineInferenceMixin:
"""Offline inference utils"""
request_counter: Counter
renderer: BaseRenderer
llm_engine: "LLMEngine"
model_config: ModelConfig
def _resolve_mm_lora(
self,
prompt: EngineInput,
lora_request: LoRARequest | None,
) -> LoRARequest | None:
if prompt["type"] != "multimodal":
return lora_request
lora_config = self.llm_engine.vllm_config.lora_config
default_mm_loras = None if lora_config is None else lora_config.default_mm_loras
if not default_mm_loras:
return lora_request
prompt_modalities = prompt["mm_placeholders"].keys()
intersection = set(prompt_modalities).intersection(default_mm_loras.keys())
if not intersection:
return lora_request
if len(intersection) > 1:
# TODO: Would be nice to be able to have multiple loras per prompt
logger.warning(
"Multiple modality specific loras were registered and would be "
"used by a single prompt consuming several modalities; "
"currently we only support one lora per request; as such, "
"lora(s) registered with modalities: %s will be skipped",
intersection,
)
return lora_request
# Build the LoRA request; the ID of the default mm lora is the
# index of the modality name sorted alphabetically + 1.
modality_name = intersection.pop()
modality_lora_path = default_mm_loras[modality_name]
modality_lora_id = sorted(default_mm_loras).index(modality_name) + 1
# If we have a collision, warn if there is a collision,
# but always send the explicitly provided request.
if lora_request:
if lora_request.lora_int_id != modality_lora_id:
logger.warning(
"A modality with a registered lora and a lora_request "
"with a different ID were provided; falling back to the "
"lora_request as we only apply one LoRARequest per prompt"
)
return lora_request
return LoRARequest(
modality_name,
modality_lora_id,
modality_lora_path,
)
def _preprocess_cmpl(
self,
prompts: Sequence[PromptType],
tokenization_kwargs: dict[str, Any] | None = None,
mm_processor_kwargs: dict[str, Any] | None = None,
) -> Sequence[EngineInput]:
"""
Convert prompt inputs from LLM APIs (other than [LLM.chat][]) into
a format that can be passed to `_add_request`.
Refer to [LLM.generate][] for a complete description of the arguments.
Returns:
A list of `EngineInput` objects ready to be passed into LLMEngine.
"""
renderer = self.renderer
model_config = self.model_config
parsed_prompts = [
parse_model_prompt(model_config, prompt) for prompt in prompts
]
tok_params = renderer.default_cmpl_tok_params.with_kwargs(
**(tokenization_kwargs or {})
)
prompt_extras = (
None
if mm_processor_kwargs is None
else {"mm_processor_kwargs": mm_processor_kwargs}
)
return renderer.render_cmpl(
parsed_prompts,
tok_params,
prompt_extras=prompt_extras,
)
def _preprocess_cmpl_one(
self,
prompt: PromptType,
tokenization_kwargs: dict[str, Any] | None = None,
mm_processor_kwargs: dict[str, Any] | None = None,
) -> EngineInput:
(engine_input,) = self._preprocess_cmpl(
[prompt],
tokenization_kwargs,
mm_processor_kwargs=mm_processor_kwargs,
)
return engine_input
def _preprocess_chat(
self,
conversations: Sequence[list[ChatCompletionMessageParam]],
chat_template: str | None = None,
chat_template_content_format: ChatTemplateContentFormatOption = "auto",
chat_template_kwargs: dict[str, Any] | None = None,
add_generation_prompt: bool = True,
continue_final_message: bool = False,
tools: list[dict[str, Any]] | None = None,
tokenization_kwargs: dict[str, Any] | None = None,
mm_processor_kwargs: dict[str, Any] | None = None,
) -> Sequence[EngineInput]:
"""
Convert a list of conversations into prompts so that they can then
be used as input for other LLM APIs.
Refer to [LLM.chat][] for a complete description of the arguments.
Returns:
A list of `EngineInput` objects ready to be passed into LLMEngine.
"""
renderer = self.renderer
chat_params = ChatParams(
chat_template=chat_template,
chat_template_content_format=chat_template_content_format,
chat_template_kwargs=merge_kwargs(
chat_template_kwargs,
dict(
add_generation_prompt=add_generation_prompt,
continue_final_message=continue_final_message,
tools=tools,
tokenize=(
is_mistral_tokenizer(renderer.tokenizer)
or self.model_config.enable_prompt_embeds
),
),
),
mm_processor_kwargs=mm_processor_kwargs,
)
tok_params = renderer.default_chat_tok_params.with_kwargs(
**(tokenization_kwargs or {})
)
prompt_extras = (
None
if mm_processor_kwargs is None
else {"mm_processor_kwargs": mm_processor_kwargs}
)
_, engine_inputs = renderer.render_chat(
conversations,
chat_params,
tok_params,
prompt_extras=prompt_extras,
)
return engine_inputs
def _preprocess_chat_one(
self,
conversation: list[ChatCompletionMessageParam],
chat_template: str | None = None,
chat_template_content_format: ChatTemplateContentFormatOption = "auto",
chat_template_kwargs: dict[str, Any] | None = None,
add_generation_prompt: bool = True,
continue_final_message: bool = False,
tools: list[dict[str, Any]] | None = None,
tokenization_kwargs: dict[str, Any] | None = None,
mm_processor_kwargs: dict[str, Any] | None = None,
) -> EngineInput:
(engine_input,) = self._preprocess_chat(
[conversation],
chat_template=chat_template,
chat_template_content_format=chat_template_content_format,
chat_template_kwargs=chat_template_kwargs,
add_generation_prompt=add_generation_prompt,
continue_final_message=continue_final_message,
tools=tools,
tokenization_kwargs=tokenization_kwargs,
mm_processor_kwargs=mm_processor_kwargs,
)
return engine_input
def _params_to_seq(
self,
params: _P | Sequence[_P],
num_requests: int,
) -> Sequence[_P]:
if isinstance(params, Sequence):
if len(params) != num_requests:
raise ValueError(
f"The lengths of prompts ({num_requests}) "
f"and params ({len(params)}) must be the same."
)
return params
return [params] * num_requests
def _lora_request_to_seq(
self,
lora_request: LoRARequest | None | Sequence[LoRARequest | None],
num_requests: int,
) -> Sequence[LoRARequest | None]:
if isinstance(lora_request, Sequence):
if len(lora_request) != num_requests:
raise ValueError(
f"The lengths of prompts ({num_requests}) "
f"and lora_request ({len(lora_request)}) must be the same."
)
return lora_request
return [lora_request] * num_requests
def _priority_to_seq(
self,
priority: list[int] | None,
num_requests: int,
) -> Sequence[int]:
if priority is not None:
if len(priority) != num_requests:
raise ValueError(
f"The lengths of prompts ({num_requests}) "
f"and priority ({len(priority)}) must be the same."
)
return priority
return [0] * num_requests
def _add_completion_requests(
self,
prompts: PromptType | Sequence[PromptType],
params: SamplingParams
| PoolingParams
| Sequence[SamplingParams | PoolingParams],
*,
use_tqdm: bool | Callable[..., tqdm] = True,
lora_request: Sequence[LoRARequest] | LoRARequest | None = None,
priority: list[int] | None = None,
tokenization_kwargs: dict[str, Any] | None = None,
mm_processor_kwargs: dict[str, Any] | None = None,
) -> list[str]:
seq_prompts = prompt_to_seq(prompts)
seq_params = self._params_to_seq(params, len(seq_prompts))
seq_lora_requests = self._lora_request_to_seq(lora_request, len(seq_prompts))
seq_priority = self._priority_to_seq(priority, len(seq_prompts))
return self._render_and_add_requests(
prompts=(
self._preprocess_cmpl_one(
prompt,
tokenization_kwargs,
mm_processor_kwargs=mm_processor_kwargs,
)
for prompt in maybe_tqdm(
seq_prompts,
use_tqdm=use_tqdm,
desc="Rendering prompts",
)
),
params=seq_params,
lora_requests=seq_lora_requests,
priorities=seq_priority,
)
def _run_completion(
self,
prompts: PromptType | Sequence[PromptType],
params: SamplingParams
| PoolingParams
| Sequence[SamplingParams | PoolingParams],
output_type: type[_O],
*,
use_tqdm: bool | Callable[..., tqdm] = True,
lora_request: Sequence[LoRARequest] | LoRARequest | None = None,
priority: list[int] | None = None,
tokenization_kwargs: dict[str, Any] | None = None,
mm_processor_kwargs: dict[str, Any] | None = None,
):
self._add_completion_requests(
prompts=prompts,
params=params,
use_tqdm=use_tqdm,
lora_request=lora_request,
priority=priority,
tokenization_kwargs=tokenization_kwargs,
mm_processor_kwargs=mm_processor_kwargs,
)
return self._run_engine(use_tqdm=use_tqdm, output_type=output_type)
def _run_chat(
self,
messages: list[ChatCompletionMessageParam]
| Sequence[list[ChatCompletionMessageParam]],
params: SamplingParams
| PoolingParams
| Sequence[SamplingParams | PoolingParams],
output_type: type[_O],
*,
use_tqdm: bool | Callable[..., tqdm] = True,
lora_request: Sequence[LoRARequest] | LoRARequest | None = None,
chat_template: str | None = None,
chat_template_content_format: ChatTemplateContentFormatOption = "auto",
add_generation_prompt: bool = True,
continue_final_message: bool = False,
tools: list[dict[str, Any]] | None = None,
chat_template_kwargs: dict[str, Any] | None = None,
tokenization_kwargs: dict[str, Any] | None = None,
mm_processor_kwargs: dict[str, Any] | None = None,
):
self._add_chat_requests(
messages=messages,
params=params,
use_tqdm=use_tqdm,
lora_request=lora_request,
chat_template=chat_template,
chat_template_content_format=chat_template_content_format,
chat_template_kwargs=chat_template_kwargs,
add_generation_prompt=add_generation_prompt,
continue_final_message=continue_final_message,
tools=tools,
tokenization_kwargs=tokenization_kwargs,
mm_processor_kwargs=mm_processor_kwargs,
)
return self._run_engine(output_type=output_type, use_tqdm=use_tqdm)
def _add_chat_requests(
self,
messages: list[ChatCompletionMessageParam]
| Sequence[list[ChatCompletionMessageParam]],
params: SamplingParams
| PoolingParams
| Sequence[SamplingParams | PoolingParams],
*,
use_tqdm: bool | Callable[..., tqdm] = True,
lora_request: Sequence[LoRARequest] | LoRARequest | None = None,
priority: list[int] | None = None,
chat_template: str | None = None,
chat_template_content_format: ChatTemplateContentFormatOption = "auto",
add_generation_prompt: bool = True,
continue_final_message: bool = False,
tools: list[dict[str, Any]] | None = None,
chat_template_kwargs: dict[str, Any] | None = None,
tokenization_kwargs: dict[str, Any] | None = None,
mm_processor_kwargs: dict[str, Any] | None = None,
) -> list[str]:
seq_convs = conversation_to_seq(messages)
seq_params = self._params_to_seq(params, len(seq_convs))
seq_lora_requests = self._lora_request_to_seq(lora_request, len(seq_convs))
seq_priority = self._priority_to_seq(priority, len(seq_convs))
# When thinking is enabled or tools are provided, and the model
# uses special tokens for structured output (e.g. Gemma4's
# <|channel>, <|tool_call>, <|"|>), automatically set
# skip_special_tokens=False so these tokens are preserved in
# output.text for downstream parsing.
needs_parsing = (
chat_template_kwargs and chat_template_kwargs.get("enable_thinking")
) or tools
if needs_parsing:
self._adjust_params_for_parsing(seq_params)
return self._render_and_add_requests(
prompts=(
self._preprocess_chat_one(
conversation,
chat_template=chat_template,
chat_template_content_format=chat_template_content_format,
chat_template_kwargs=chat_template_kwargs,
add_generation_prompt=add_generation_prompt,
continue_final_message=continue_final_message,
tools=tools,
tokenization_kwargs=tokenization_kwargs,
mm_processor_kwargs=mm_processor_kwargs,
)
for conversation in maybe_tqdm(
seq_convs,
use_tqdm=use_tqdm,
desc="Rendering conversations",
)
),
params=seq_params,
lora_requests=seq_lora_requests,
priorities=seq_priority,
)
def _adjust_params_for_parsing(
self, params: Sequence[SamplingParams | PoolingParams]
) -> None:
"""Set ``skip_special_tokens=False`` when the model encodes
structured output syntax as special tokens.
Models like Gemma4 register thinking delimiters
(``<|channel>``/``<channel|>``) and tool call tokens
(``<|tool_call>``/``<tool_call|>``/``<|"|>``) as special tokens.
The default ``skip_special_tokens=True`` strips them from
``output.text``, breaking parsing of both reasoning blocks and
tool calls.
This is a no-op for models whose structured tokens are regular
text tokens (e.g. DeepSeek's ``<think>``/``</think>``).
"""
# The offline API currently lacks a unified rendering pipeline.
# Until the planned Renderer refactor is complete, we hardcode
# this token preservation logic specifically for Gemma4 models
# to avoid regressions on other models.
hf_config = getattr(self.model_config, "hf_config", None)
architectures = getattr(hf_config, "architectures", [])
if any("Gemma4" in arch for arch in architectures):
tokenizer = self.renderer.get_tokenizer()
vocab = tokenizer.get_vocab()
special_ids = set(getattr(tokenizer, "all_special_ids", []))
# Tokens used for thinking delimiters and tool call syntax
# that some models (Gemma4) register as special tokens.
structured_tokens = (
"<|channel>",
"<channel|>", # thinking delimiters
"<|tool_call>",
"<tool_call|>", # tool call delimiters
'<|"|>', # string quoting in tool args
)
needs_special = any(
vocab.get(tok) in special_ids
for tok in structured_tokens
if tok in vocab
)
if needs_special:
for sp in params:
if isinstance(sp, SamplingParams) and sp.skip_special_tokens:
sp.skip_special_tokens = False
def _render_and_run_requests(
self,
prompts: Iterable[EngineInput],
params: Sequence[SamplingParams | PoolingParams],
output_type: type[_O],
*,
lora_requests: Sequence[LoRARequest | None] | None = None,
priorities: Sequence[int] | None = None,
use_tqdm: bool | Callable[..., tqdm] = True,
):
if isinstance(prompts, (list, tuple)):
logger.warning_once(
"Rendering all prompts before adding them to the engine "
"is less efficient than performing both on the same prompt "
"before processing the next prompt. You should instead pass "
"a generator that renders one prompt per iteration, as that allows "
"engine execution to begin for the first prompt while processing "
"the next prompt."
)
self._render_and_add_requests(
prompts=prompts,
params=params,
lora_requests=lora_requests,
priorities=priorities,
)
return self._run_engine(output_type, use_tqdm=use_tqdm)
def _render_and_add_requests(
self,
prompts: Iterable[EngineInput],
params: Sequence[SamplingParams | PoolingParams],
*,
lora_requests: Sequence[LoRARequest | None] | None = None,
priorities: Sequence[int] | None = None,
) -> list[str]:
added_request_ids: list[str] = []
try:
for i, prompt in enumerate(prompts):
request_id = self._add_request(
prompt,
params[i],
lora_request=self._resolve_mm_lora(
prompt,
None if lora_requests is None else lora_requests[i],
),
priority=0 if priorities is None else priorities[i],
)
added_request_ids.append(request_id)
except Exception as e:
if added_request_ids:
self.llm_engine.abort_request(added_request_ids, internal=True)
raise e
return added_request_ids
def _add_request(
self,
prompt: EngineInput,
params: SamplingParams | PoolingParams,
lora_request: LoRARequest | None = None,
priority: int = 0,
) -> str:
if isinstance(params, SamplingParams):
# We only care about the final output
params.output_kind = RequestOutputKind.FINAL_ONLY
request_id = str(next(self.request_counter))
return self.llm_engine.add_request(
request_id,
prompt,
params,
lora_request=lora_request,
priority=priority,
)
def _run_engine(
self,
output_type: type[_O] | tuple[type[_O], ...],
*,
use_tqdm: bool | Callable[..., tqdm] = True,
) -> list[_O]:
# Initialize tqdm.
if use_tqdm:
num_requests = self.llm_engine.get_num_unfinished_requests()
tqdm_func = use_tqdm if callable(use_tqdm) else tqdm
pbar = tqdm_func(
total=num_requests,
desc="Processed prompts",
dynamic_ncols=True,
postfix=(f"est. speed input: {0:.2f} toks/s, output: {0:.2f} toks/s"),
)
# Run the engine.
outputs: list[_O] = []
total_in_toks = 0
total_out_toks = 0
while self.llm_engine.has_unfinished_requests():
step_outputs = self.llm_engine.step()
for output in step_outputs:
assert isinstance(output, output_type)
if output.finished:
outputs.append(output) # type: ignore[arg-type]
if use_tqdm:
if isinstance(output, RequestOutput):
# Calculate tokens only for RequestOutput
n = len(output.outputs)
assert output.prompt_token_ids is not None
total_in_toks += len(output.prompt_token_ids) * n
in_spd = total_in_toks / pbar.format_dict["elapsed"]
total_out_toks += sum(
len(stp.token_ids) for stp in output.outputs
)
out_spd = total_out_toks / pbar.format_dict["elapsed"]
pbar.postfix = (
f"est. speed input: {in_spd:.2f} toks/s, "
f"output: {out_spd:.2f} toks/s"
)
pbar.update(n)
else:
pbar.update(1)
if pbar.n == num_requests:
pbar.refresh()
if use_tqdm:
pbar.close()
# Sort the outputs by request ID.
# This is necessary because some requests may be finished earlier than
# its previous requests.
return sorted(outputs, key=lambda x: int(x.request_id))
+68 -4
View File
@@ -1,24 +1,34 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
from collections.abc import Callable, Sequence
from abc import ABC, abstractmethod
from collections.abc import Callable, Iterable, Sequence
from typing import Any
from tqdm.auto import tqdm
from typing_extensions import TypeVar
from vllm.config import ModelConfig
from vllm.entrypoints.chat_utils import ChatTemplateConfig
from vllm.entrypoints.offline_utils import OfflineInferenceMixin
from vllm.inputs import DataPrompt, PromptType
from vllm.inputs import (
DataPrompt,
EngineInput,
PromptType,
)
from vllm.logger import init_logger
from vllm.lora.request import LoRARequest
from vllm.outputs import (
ClassificationRequestOutput,
EmbeddingRequestOutput,
PoolingRequestOutput,
RequestOutput,
ScoringRequestOutput,
)
from vllm.pooling_params import PoolingParams
from vllm.renderers import BaseRenderer
from vllm.sampling_params import SamplingParams
from vllm.tasks import SCORE_TYPE_MAP, PoolingTask, SupportedTask
from vllm.v1.engine.llm_engine import LLMEngine
from .factories import init_pooling_io_processors
from .scoring.io_processor import ScoringIOProcessor
@@ -27,10 +37,20 @@ from .typing import OfflineInputsContext, OfflineOutputsContext
logger = init_logger(__name__)
_P = TypeVar("_P", bound=SamplingParams | PoolingParams | None)
_O = TypeVar(
"_O",
bound=RequestOutput | PoolingRequestOutput,
default=RequestOutput | PoolingRequestOutput,
)
class PoolingOfflineMixin(OfflineInferenceMixin):
class PoolingOfflineMixin(ABC):
"""Offline inference for pooling models"""
renderer: BaseRenderer
llm_engine: "LLMEngine"
model_config: ModelConfig
runner_type: str
chat_template: str | None
supported_tasks: tuple[SupportedTask, ...]
@@ -444,3 +464,47 @@ class PoolingOfflineMixin(OfflineInferenceMixin):
)
return [ScoringRequestOutput.from_base(item) for item in outputs]
@abstractmethod
def _params_to_seq(
self,
params: _P | Sequence[_P],
num_requests: int,
) -> Sequence[_P]:
raise NotImplementedError
@abstractmethod
def _lora_request_to_seq(
self,
lora_request: LoRARequest | None | Sequence[LoRARequest | None],
num_requests: int,
) -> Sequence[LoRARequest | None]:
raise NotImplementedError
@abstractmethod
def _priority_to_seq(
self,
priority: list[int] | None,
num_requests: int,
) -> Sequence[int]:
raise NotImplementedError
@abstractmethod
def _render_and_add_requests(
self,
prompts: Iterable[EngineInput],
params: Sequence[SamplingParams | PoolingParams],
*,
lora_requests: Sequence[LoRARequest | None] | None = None,
priorities: Sequence[int] | None = None,
) -> list[str]:
raise NotImplementedError
@abstractmethod
def _run_engine(
self,
output_type: type[_O] | tuple[type[_O], ...],
*,
use_tqdm: bool | Callable[..., tqdm] = True,
) -> list[_O]:
raise NotImplementedError
-2
View File
@@ -1971,8 +1971,6 @@ environment_variables: dict[str, Callable[[], Any]] = {
int(os.getenv("VLLM_USE_SIMPLE_KV_OFFLOAD", "0"))
),
# Whether to enable dual cuda streams for LoRA computation
# (used by both BaseLinearLayerWithLoRA and FusedMoEWithLoRA to
# overlap the base layer compute with the LoRA fast path).
"VLLM_LORA_ENABLE_DUAL_STREAM": lambda: bool(
int(os.getenv("VLLM_LORA_ENABLE_DUAL_STREAM", "0"))
),
+8 -1
View File
@@ -25,9 +25,16 @@ from vllm.utils.multi_stream_utils import maybe_execute_in_parallel
from vllm.utils.torch_utils import direct_register_custom_op
from .base import BaseLayerWithLoRA
from .utils import _get_lora_aux_cuda_stream, _get_lora_device
from .utils import _get_lora_device
if envs.VLLM_LORA_ENABLE_DUAL_STREAM:
_lora_aux_cuda_stream: torch.cuda.Stream | None = None
def _get_lora_aux_cuda_stream() -> torch.cuda.Stream | None:
global _lora_aux_cuda_stream
if _lora_aux_cuda_stream is None and current_platform.is_cuda_alike():
_lora_aux_cuda_stream = torch.cuda.Stream()
return _lora_aux_cuda_stream
def lora_linear_async(
layer_name: str,
+1 -25
View File
@@ -19,9 +19,8 @@ from vllm.model_executor.layers.fused_moe.modular_kernel import FusedMoEKernel
from vllm.model_executor.layers.fused_moe.prepare_finalize import (
MoEPrepareAndFinalizeNoDPEPModular,
)
from vllm.platforms import current_platform
from .utils import _get_lora_aux_cuda_stream, _get_lora_device
from .utils import _get_lora_device
class FusedMoEWithLoRA(BaseLayerWithLoRA):
@@ -35,9 +34,6 @@ class FusedMoEWithLoRA(BaseLayerWithLoRA):
self.tp_size = self.base_layer.tp_size
self.tp_rank = self.base_layer.tp_rank
self.device = _get_lora_device(base_layer)
self._enable_aux_cuda_stream = envs.VLLM_LORA_ENABLE_DUAL_STREAM
self._init_lora_stream_context()
# For non-gated MoE (is_act_and_mul=False), only 1 slice is needed
# since there's only up_proj (w1), not gate_proj + up_proj (w1 + w3)
self._w13_slices = 2 if base_layer.moe_config.is_act_and_mul else 1
@@ -69,25 +65,7 @@ class FusedMoEWithLoRA(BaseLayerWithLoRA):
FusedMoEModularMethod(self.base_layer.quant_method, moe_kernel)
)
def _init_lora_stream_context(self) -> None:
self._lora_stream: torch.cuda.Stream | None = None
self._events: tuple[torch.cuda.Event, ...] | None = None
if not self._enable_aux_cuda_stream:
return
if not current_platform.is_cuda_alike():
return
self._lora_stream = _get_lora_aux_cuda_stream()
# 4 events: 2 per (base GEMM, LoRA) pair so w13 and w2 don't reuse
# the same event objects; reuse-within-a-pair is fine because the
# second pair starts only after intermediate_cache1.add_() has joined.
self._events = tuple(torch.cuda.Event() for _ in range(4))
def _build_lora_context(self):
use_dual_stream = (
self._enable_aux_cuda_stream
and not self.fully_sharded
and self._lora_stream is not None
)
return MoELoRAContext(
w13_lora_a_stacked=self.w13_lora_a_stacked,
w13_lora_b_stacked=self.w13_lora_b_stacked,
@@ -103,8 +81,6 @@ class FusedMoEWithLoRA(BaseLayerWithLoRA):
local_num_experts=self.base_layer.local_num_experts,
punica_wrapper=self.punica_wrapper,
use_tuned_config=bool(envs.VLLM_TUNED_CONFIG_FOLDER),
aux_stream=self._lora_stream if use_dual_stream else None,
events=self._events if use_dual_stream else None,
)
def _create_lora_a_weights(
-13
View File
@@ -7,22 +7,9 @@ from enum import Enum
import torch
import torch.nn as nn
from vllm import envs
from vllm.model_executor.layers.fused_moe.fused_moe import try_get_optimal_moe_config
from vllm.platforms import current_platform
from vllm.utils.math_utils import next_power_of_2
_lora_aux_cuda_stream: torch.cuda.Stream | None = None
def _get_lora_aux_cuda_stream() -> torch.cuda.Stream | None:
if not envs.VLLM_LORA_ENABLE_DUAL_STREAM:
return None
global _lora_aux_cuda_stream
if _lora_aux_cuda_stream is None and current_platform.is_cuda_alike():
_lora_aux_cuda_stream = torch.cuda.Stream()
return _lora_aux_cuda_stream
class LoRAMappingType(Enum):
LANGUAGE = 1
@@ -105,742 +105,6 @@ def _get_c_ptrs(
_LORA_PTR_DICT: dict[tuple[int, ...], torch.tensor] = {}
# ---------------------------------------------------------------------------
# Fully-fused MoE-LoRA kernel (one-shot): shrink + expand combined into a single
# launch with the rank-dim intermediate kept in registers. Used by the fast
# path of `_fused_moe_lora` for `fully_sharded=False`. The legacy two-kernel
# path (`_fused_moe_lora_kernel` above) is retained for `fully_sharded=True`
# because that path needs to materialise the intermediate cache for an
# all_reduce / all_gather between shrink and expand.
# ---------------------------------------------------------------------------
@triton.heuristics({"EVEN_K": lambda args: args["K"] % args["BLOCK_K"] == 0})
@triton.jit
def _fused_moe_lora_one_shot_kernel(
# ---- pointers ----
x_ptr,
A_ptrs,
B_ptrs,
out_ptr,
topk_weights_ptr,
sorted_token_ids_ptr,
expert_ids_ptr,
num_tokens_post_padded_ptr,
token_lora_mapping_ptr,
lora_ids_ptr,
adapter_enabled_ptr,
# ---- dims ----
N,
K,
num_valid_tokens,
top_k_num,
max_loras,
# ---- strides ----
stride_xm,
stride_xk,
stride_A_lora,
stride_A_expert,
stride_A_r,
stride_A_k,
stride_B_lora,
stride_B_expert,
stride_B_n,
stride_B_r,
stride_om,
stride_on,
stride_tl_,
stride_el,
# ---- scalar ----
slice_n_offset,
# ---- constexpr (set per call) ----
token_mapping_factor: tl.constexpr,
naive_block_assignment: tl.constexpr,
MUL_ROUTED_WEIGHT: tl.constexpr,
BLOCK_M: tl.constexpr,
BLOCK_R: tl.constexpr,
actual_rank: tl.constexpr,
NPID_FACTOR: tl.constexpr,
BLOCK_N: tl.constexpr,
BLOCK_K: tl.constexpr,
EVEN_K: tl.constexpr,
ADD_INPUTS: tl.constexpr,
):
pid_full = tl.program_id(axis=0)
pid_m = pid_full // NPID_FACTOR
pid_n_outer = pid_full % NPID_FACTOR
slice_id = tl.program_id(axis=1)
lora_idx = tl.program_id(axis=2)
# Resolve lora_id.
if naive_block_assignment:
token_idx_for_lora = pid_m // top_k_num
lora_id = tl.load(token_lora_mapping_ptr + token_idx_for_lora)
else:
lora_id = tl.load(lora_ids_ptr + lora_idx)
if lora_id < 0:
return
if lora_id >= max_loras:
return
enabled = tl.load(adapter_enabled_ptr + lora_id)
if enabled == 0:
return
if not naive_block_assignment:
ntpp = tl.load(num_tokens_post_padded_ptr + lora_id)
if pid_m * BLOCK_M >= ntpp:
return
# Resolve expert_id.
if naive_block_assignment:
expert_id = tl.load(expert_ids_ptr + pid_m)
else:
ind = lora_id * stride_el + pid_m
expert_id = tl.load(
expert_ids_ptr + ind, mask=ind < max_loras * stride_el, other=-1
)
if expert_id < 0:
return
# Compute offs_token (flat token ids).
offs = tl.arange(0, BLOCK_M).to(tl.int64)
if naive_block_assignment:
offs_token = tl.where(offs == 0, pid_m, num_valid_tokens)
else:
offs_token_id = pid_m * BLOCK_M + offs
token_ind = stride_tl_ * lora_id + offs_token_id
offs_token = tl.load(
sorted_token_ids_ptr + token_ind,
mask=token_ind < max_loras * stride_tl_,
other=num_valid_tokens,
)
token_mask = offs_token < num_valid_tokens
# N range owned by this program. Splitting [0, N) into NPID_FACTOR
# contiguous outer blocks lets us scale parallelism for small batches.
n_per_outer = tl.cdiv(N, NPID_FACTOR)
n_lo = pid_n_outer * n_per_outer
n_hi = tl.minimum((pid_n_outer + 1) * n_per_outer, N)
if n_lo >= N:
return
# Slice pointers.
cur_A_ptr = tl.load(A_ptrs + slice_id).to(tl.pointer_type(out_ptr.dtype.element_ty))
cur_B_ptr = tl.load(B_ptrs + slice_id).to(tl.pointer_type(out_ptr.dtype.element_ty))
A_base = cur_A_ptr + lora_id * stride_A_lora + expert_id * stride_A_expert
B_base = cur_B_ptr + lora_id * stride_B_lora + expert_id * stride_B_expert
# SHRINK: tmp[BLOCK_M, BLOCK_R] = x @ A^T, accumulated in fp32 registers.
offs_r = tl.arange(0, BLOCK_R)
rank_mask = offs_r < actual_rank
# Clamp rank offsets so OOB rows of A / B map to address 0; the mask
# zeros the loaded values. Required when BLOCK_R > actual_rank
# (e.g. rank=4 padded to 16) -- without clamping, tl.load would address
# the next expert's memory.
safe_offs_r = tl.where(rank_mask, offs_r, 0)
offs_k = tl.arange(0, BLOCK_K)
offs_x_row = offs_token // token_mapping_factor
x_ptrs = x_ptr + offs_x_row[:, None] * stride_xm + offs_k[None, :] * stride_xk
a_ptrs = A_base + offs_k[:, None] * stride_A_k + safe_offs_r[None, :] * stride_A_r
tmp = tl.zeros((BLOCK_M, BLOCK_R), dtype=tl.float32)
if EVEN_K:
for _ in range(0, K, BLOCK_K):
x = tl.load(x_ptrs, mask=token_mask[:, None], other=0.0)
a = tl.load(a_ptrs, mask=rank_mask[None, :], other=0.0)
tmp += tl.dot(x, a)
x_ptrs += BLOCK_K * stride_xk
a_ptrs += BLOCK_K * stride_A_k
else:
for kb in range(0, K, BLOCK_K):
k_remain = K - kb
k_mask = offs_k < k_remain
x = tl.load(x_ptrs, mask=token_mask[:, None] & k_mask[None, :], other=0.0)
a = tl.load(a_ptrs, mask=k_mask[:, None] & rank_mask[None, :], other=0.0)
tmp += tl.dot(x, a)
x_ptrs += BLOCK_K * stride_xk
a_ptrs += BLOCK_K * stride_A_k
tmp_typed = tmp.to(out_ptr.dtype.element_ty)
# EXPAND: out[tokens, n] += tmp @ B^T, looped over BLOCK_N tiles within
# this program's [n_lo, n_hi). The (offs_n < n_hi) mask is required
# whenever BLOCK_N > n_per_outer to keep adjacent outer blocks from
# writing into each other's columns.
if MUL_ROUTED_WEIGHT:
moe_w = tl.load(topk_weights_ptr + offs_token, mask=token_mask, other=0.0).to(
tl.float32
)
out_slice_base = out_ptr + slice_id * slice_n_offset
for n_start in range(n_lo, n_hi, BLOCK_N):
offs_n = n_start + tl.arange(0, BLOCK_N)
n_mask = (offs_n < N) & (offs_n < n_hi)
b_ptrs = (
B_base + safe_offs_r[:, None] * stride_B_r + offs_n[None, :] * stride_B_n
)
b = tl.load(b_ptrs, mask=rank_mask[:, None] & n_mask[None, :], other=0.0)
acc = tl.dot(tmp_typed, b) # (BLOCK_M, BLOCK_N) fp32
if MUL_ROUTED_WEIGHT:
acc = acc * moe_w[:, None]
out_ptrs = (
out_slice_base
+ offs_token[:, None] * stride_om
+ offs_n[None, :] * stride_on
)
out_mask = token_mask[:, None] & n_mask[None, :]
if ADD_INPUTS:
prev = tl.load(out_ptrs, mask=out_mask, other=0.0)
tl.store(out_ptrs, prev + acc.to(out_ptr.dtype.element_ty), mask=out_mask)
else:
tl.store(out_ptrs, acc.to(out_ptr.dtype.element_ty), mask=out_mask)
def _run_fused_moe_lora_one_shot(
output: torch.Tensor,
qcurr_hidden_states: torch.Tensor,
lora_a_stacked: list[torch.Tensor],
lora_b_stacked: list[torch.Tensor],
topk_weights: torch.Tensor,
sorted_token_ids: torch.Tensor | None,
expert_ids: torch.Tensor,
num_tokens_post_padded: torch.Tensor | None,
token_lora_mapping: torch.Tensor,
max_lora_rank: int,
top_k_num: int,
lora_ids: torch.Tensor,
num_active_loras: torch.Tensor,
adapter_enabled: torch.Tensor,
mul_routed_weight: bool,
block_size_m: int,
add_inputs: bool = True,
) -> None:
"""Fast-path wrapper: launches one fused shrink+expand kernel.
The shape contract matches `_fused_moe_lora`. `output` has shape
`(num_tokens, top_k_num, num_slices * N_per_slice)`. When
`add_inputs=True` (default) the kernel reads-modifies-writes `output`
in place; when `add_inputs=False` the kernel overwrites `output` with
the LoRA delta only. The latter is used by the dual-stream path that
sums LoRA into the base output on a separate stream.
"""
num_slices = len(lora_a_stacked)
device = qcurr_hidden_states.device
A0 = lora_a_stacked[0]
B0 = lora_b_stacked[0]
max_loras_w = A0.shape[0]
rank = A0.shape[2]
K = A0.shape[3]
N_per_slice = B0.shape[2]
# rank padding is to next pow2 with a floor of 16 (tensor-core minimum
# K-dim). Beyond 128 the (BLOCK_M, BLOCK_R) accumulator outgrows the
# register file; rank tiling would be needed but is out of scope for
# this kernel. Tried floor=32 to double MMA density per K-step but it
# regressed across all M (+8 to +40%): the (64,32) fp32 accumulator +
# widened B tile pushed register count past spill threshold, lowering
# occupancy by more than the MMA gain saved.
assert rank <= 128, (
f"fused_moe_lora_one_shot supports max_lora_rank<=128; got rank={rank}"
)
BLOCK_R = max(triton.next_power_of_2(rank), 16)
num_experts = A0.shape[1]
naive = sorted_token_ids is None
if sorted_token_ids is None:
EM_grid = topk_weights.numel()
BLOCK_M = 16
stride_tl_ = 0
stride_el = 0
grid_lora_dim = 1
else:
EM_grid = sorted_token_ids.shape[1]
# BLOCK_M must equal moe_lora_align_block_size's block_size. The
# caller passes that explicitly; deriving it from tensor shapes is
# unsafe because sorted_token_ids.shape[1] is the raw padded length
# (not necessarily a multiple of block_size — e.g. OLMoE prefill
# produces sorted=139200 with expert_ids=1088 and block_size=128).
# tl.arange and tl.dot need block_size_m to be a power of 2 and at
# least 16. The Python-side assertion gives a clearer error than
# the cryptic Triton compile failure.
assert block_size_m >= 16 and (block_size_m & (block_size_m - 1)) == 0, (
f"shrink_block_size_m must be a power of 2 and >=16; got {block_size_m}"
)
BLOCK_M = block_size_m
stride_tl_ = sorted_token_ids.stride(0)
stride_el = expert_ids.stride(0)
grid_lora_dim = int(num_active_loras.item())
# Empty-work guards: the grid would otherwise have a zero dimension,
# which Triton rejects. None of these is a hot path in production -- a
# batch with zero tokens, an EM_grid of zero, or zero active LoRAs all
# mean there's nothing to add to `output`.
if EM_grid == 0 or grid_lora_dim == 0 or num_slices == 0:
return
token_mapping_factor = 1 if mul_routed_weight else top_k_num
A_ptrs = _get_ptr(lora_a_stacked, device)
B_ptrs = _get_ptr(lora_b_stacked, device)
# Flatten (num_tokens, top_k) → flat_token axis. The kernel addresses
# output via offs_token * stride_om, which is correct iff the dim-0 /
# dim-1 strides collapse cleanly: stride(0) == top_k * stride(1). All
# production callers pass contiguous output, so this always holds; the
# explicit check guards against future regressions where a non-trivial
# view (e.g. permute) would silently break in-place accumulation.
assert output.dim() == 3, f"output must be 3-D, got {output.shape}"
assert output.stride(0) == output.shape[1] * output.stride(1), (
"fused_moe_lora_one_shot requires output.stride(0) == top_k*stride(1); "
f"got shape={output.shape} strides={output.stride()}"
)
out_view = output.view(-1, output.shape[-1])
M_blocks = triton.cdiv(EM_grid, BLOCK_M) if not naive else EM_grid
# NPID_FACTOR heuristic: scale N-axis parallelism when base CTA count is
# short of saturating the SM array. Cap by the cost of redundant shrink.
sm_count = torch.cuda.get_device_properties(device).multi_processor_count
base_programs = max(M_blocks * num_slices * grid_lora_dim, 1)
shrink_ratio = K / max(K + N_per_slice, 1)
max_npid_by_budget = max(1, int(1.5 / max(shrink_ratio, 1e-3)) + 1)
target = 2 * sm_count
if base_programs >= int(1.5 * sm_count):
npid = 1
else:
npid_occ = max(1, min(16, (target + base_programs - 1) // base_programs))
npid = min(npid_occ, max_npid_by_budget)
npid = max(1, min(npid, max(1, N_per_slice // 128)))
# Robust defaults across the prefill regime (H100/H200/B200, bf16/fp16).
# NPID > 1 is the small-M / under-saturated path -- more warps help
# amortise the inner-N expand loop. ns=3 instead of 4: GB200 ncu showed
# the 4-stage pipeline pushed register count to 168/thread and capped
# achieved occupancy at ~17% (3 blocks/SM, register-bound); ns=3 frees
# ~30 regs/thread which keeps a 4th block resident on small grids.
# Tried BLOCK_N=64 for w13 (N=192) to avoid the half-wasted second
# tile: regressed 11-29% because the "waste" was just masked stores
# (cheap) and the extra iteration added load + index overhead.
if npid > 1:
block_n, nw, ns = 128, 8, 3
else:
block_n, nw, ns = 128, 4, 3
# BLOCK_K choice: for hidden-sized K (≥256, i.e. the K=hidden_size
# shrink input on w13) force BLOCK_K=128 -- the wider tile halves the
# K-loop trip count and removes the scoreboard stalls that dominated
# M=16-64 on GB200 (kernel time -13% to -37% vs the work_per_expert
# heuristic which picked 64 for low-tokens-per-expert ratios). For
# small-K shapes (e.g. w2 with K=192 where the down-proj reads the
# MoE intermediate) keep the work_per_expert heuristic: BLOCK_K=128
# would force the EVEN_K=False masked path and add no K-loop savings
# (K/64=3 vs K/128=2 masked) while inflating per-program startup.
if K >= 256:
block_k = 128
else:
work_per_expert = topk_weights.numel() / max(num_experts, 1)
block_k = 128 if work_per_expert >= 16 else 64
grid = (M_blocks * npid, num_slices, grid_lora_dim)
_fused_moe_lora_one_shot_kernel[grid](
qcurr_hidden_states,
A_ptrs,
B_ptrs,
out_view,
topk_weights,
sorted_token_ids,
expert_ids,
num_tokens_post_padded,
token_lora_mapping,
lora_ids,
adapter_enabled,
N_per_slice,
K,
topk_weights.numel(),
top_k_num,
max_loras_w,
qcurr_hidden_states.stride(0),
qcurr_hidden_states.stride(1),
A0.stride(0),
A0.stride(1),
A0.stride(2),
A0.stride(3),
B0.stride(0),
B0.stride(1),
B0.stride(2),
B0.stride(3),
out_view.stride(0),
out_view.stride(1),
stride_tl_,
stride_el,
N_per_slice,
token_mapping_factor=token_mapping_factor,
naive_block_assignment=naive,
MUL_ROUTED_WEIGHT=mul_routed_weight,
BLOCK_M=BLOCK_M,
BLOCK_R=BLOCK_R,
actual_rank=rank,
NPID_FACTOR=npid,
BLOCK_N=block_n,
BLOCK_K=block_k,
ADD_INPUTS=add_inputs,
num_warps=nw,
num_stages=ns,
)
# ---------------------------------------------------------------------------
# Small-batch (decode-style) fused MoE-LoRA kernel — sub-path of the
# one_shot fast path.
# ---------------------------------------------------------------------------
@triton.heuristics({"EVEN_K": lambda args: args["K"] % args["BLOCK_K"] == 0})
@triton.jit
def _fused_moe_lora_small_batch_kernel(
# ---- pointers ----
x_ptr,
A_ptrs,
B_ptrs,
out_ptr,
topk_weights_ptr,
expert_ids_ptr, # (num_tokens * top_k_num,)
token_lora_mapping_ptr, # (num_tokens,)
adapter_enabled_ptr,
# ---- dims ----
N,
K,
top_k_num,
max_loras,
work_total, # = pair_slices * n_chunks_per_pair_slice
pair_slices, # = num_tokens * top_k_num * NUM_SLICES
# ---- strides ----
stride_xm,
stride_xk,
stride_A_lora,
stride_A_expert,
stride_A_r,
stride_A_k,
stride_B_lora,
stride_B_expert,
stride_B_n,
stride_B_r,
stride_om,
stride_on,
# ---- scalar (runtime ints, NOT constexpr) ----
# n_tiles_per_program / n_chunks_per_pair_slice are deliberately
# runtime: each distinct value would otherwise trigger a fresh Triton
# compile -> fresh kernel binary -> fresh CUDA graph instance per
# batch size. Production traces showed that variant explosion adding
# ~5.9k graph instantiations on top of legacy. Runtime args mean one
# shared binary across all chunk sizes.
slice_n_offset,
n_tiles_per_program,
n_chunks_per_pair_slice,
# ---- constexpr ----
token_mapping_factor: tl.constexpr,
MUL_ROUTED_WEIGHT: tl.constexpr,
ADD_INPUTS: tl.constexpr,
BLOCK_R: tl.constexpr,
actual_rank: tl.constexpr,
BLOCK_N: tl.constexpr,
BLOCK_K: tl.constexpr,
NUM_SLICES: tl.constexpr,
EVEN_K: tl.constexpr,
):
"""Persistent fused MoE-LoRA kernel for naive_block_assignment inputs.
Each program owns one (pair × slice × n_chunk) work item. A "chunk"
covers `n_tiles_per_program` consecutive output-N tiles, all of which
share a single shrink so the rank-vector is computed once per
program and the A weights for that (lora, expert, slice) are loaded
once instead of n_tiles_per_program times.
The wrapper picks `n_tiles_per_program` to keep the grid close to
2*SM_count: at very small batch (work_total SM_count) the chunk
size collapses to 1 and behaviour matches a per-tile GEMV; as batch
grows the chunk grows so we trade some N-axis parallelism for shrink
reuse. When `work_total` exceeds the launched grid, the outer stride
loop drains the leftover work units serially.
"""
pid = tl.program_id(axis=0)
num_programs = tl.num_programs(axis=0)
offs_r = tl.arange(0, BLOCK_R)
rank_mask = offs_r < actual_rank
# Clamp OOB rank lanes so they address row 0 of A/B; the mask zeros
# the loaded values. Required when BLOCK_R > actual_rank (e.g. rank=4
# padded to 16) -- without clamping, tl.load would address the next
# expert's memory.
safe_offs_r = tl.where(rank_mask, offs_r, 0)
offs_k = tl.arange(0, BLOCK_K)
# Persistent stride loop: when grid < work_total each program walks
# multiple work items. When grid == work_total the loop runs exactly
# once and the kernel degenerates to the per-tile GEMV.
for work_id in range(pid, work_total, num_programs):
n_chunk_idx = work_id % n_chunks_per_pair_slice
pair_slice_idx = work_id // n_chunks_per_pair_slice
# NUM_SLICES is constexpr (typ. 1 or 2) so divmod folds.
pair_idx = pair_slice_idx // NUM_SLICES
slice_id = pair_slice_idx % NUM_SLICES
# Resolve lora_id / expert_id; skip the body for inactive lanes.
# Using a single `valid` flag instead of early `return` keeps the
# outer stride loop alive — `return` would exit the whole program
# and skip later work items assigned to this SM.
token_idx = pair_idx // top_k_num
lora_id = tl.load(token_lora_mapping_ptr + token_idx)
valid = (lora_id >= 0) & (lora_id < max_loras)
enabled = tl.load(adapter_enabled_ptr + tl.where(valid, lora_id, 0))
valid = valid & (enabled != 0)
expert_id = tl.load(expert_ids_ptr + pair_idx)
valid = valid & (expert_id >= 0)
if valid:
cur_A_ptr = tl.load(A_ptrs + slice_id).to(
tl.pointer_type(out_ptr.dtype.element_ty)
)
cur_B_ptr = tl.load(B_ptrs + slice_id).to(
tl.pointer_type(out_ptr.dtype.element_ty)
)
A_base = cur_A_ptr + lora_id * stride_A_lora + expert_id * stride_A_expert
B_base = cur_B_ptr + lora_id * stride_B_lora + expert_id * stride_B_expert
x_row = pair_idx // token_mapping_factor
x_row_ptr = x_ptr + x_row * stride_xm
# SHRINK GEMV (once per program; reused across n_tiles_per_program
# expand tiles below). Sum-reduction over BLOCK_K with fp32
# accumulator — same precision path as the one_shot kernel.
rank_vec = tl.zeros((BLOCK_R,), dtype=tl.float32)
if EVEN_K:
for kb in range(0, K, BLOCK_K):
cur_k = kb + offs_k
x_tile = tl.load(x_row_ptr + cur_k * stride_xk).to(tl.float32)
a_tile = tl.load(
A_base
+ safe_offs_r[:, None] * stride_A_r
+ cur_k[None, :] * stride_A_k,
mask=rank_mask[:, None],
other=0.0,
).to(tl.float32)
rank_vec += tl.sum(a_tile * x_tile[None, :], axis=1)
else:
for kb in range(0, K, BLOCK_K):
cur_k = kb + offs_k
k_mask = cur_k < K
x_tile = tl.load(
x_row_ptr + cur_k * stride_xk, mask=k_mask, other=0.0
).to(tl.float32)
a_tile = tl.load(
A_base
+ safe_offs_r[:, None] * stride_A_r
+ cur_k[None, :] * stride_A_k,
mask=rank_mask[:, None] & k_mask[None, :],
other=0.0,
).to(tl.float32)
rank_vec += tl.sum(a_tile * x_tile[None, :], axis=1)
# EXPAND: walk n_tiles_per_program consecutive output-N tiles
# using the same rank_vec. The loop is a runtime range (not
# tl.static_range) so a single compiled kernel handles every
# chunk size — see the note on the kernel signature.
n_tile_start = n_chunk_idx * n_tiles_per_program
out_row_ptr = out_ptr + slice_id * slice_n_offset + pair_idx * stride_om
if MUL_ROUTED_WEIGHT:
moe_w = tl.load(topk_weights_ptr + pair_idx).to(tl.float32)
for nt in range(n_tiles_per_program):
n_lo = (n_tile_start + nt) * BLOCK_N
if n_lo < N:
offs_n = n_lo + tl.arange(0, BLOCK_N)
n_mask = offs_n < N
b_tile = tl.load(
B_base
+ offs_n[:, None] * stride_B_n
+ safe_offs_r[None, :] * stride_B_r,
mask=n_mask[:, None] & rank_mask[None, :],
other=0.0,
).to(tl.float32)
out_tile = tl.sum(b_tile * rank_vec[None, :], axis=1)
if MUL_ROUTED_WEIGHT:
out_tile = out_tile * moe_w
out_ptrs = out_row_ptr + offs_n * stride_on
if ADD_INPUTS:
prev = tl.load(out_ptrs, mask=n_mask, other=0.0).to(tl.float32)
tl.store(
out_ptrs,
(prev + out_tile).to(out_ptr.dtype.element_ty),
mask=n_mask,
)
else:
tl.store(
out_ptrs,
out_tile.to(out_ptr.dtype.element_ty),
mask=n_mask,
)
def _pick_small_batch_chunk(pair_slices: int, N_tiles: int, sm_count: int) -> int:
"""Pick `n_tiles_per_program` so the launched grid stays near
2*SM_count.
Sizes for occupancy first (more programs in flight better latency
hiding for the K-loop A/x loads). Once the per-tile grid already
exceeds 2*SM_count we increase the chunk size to amortise the shrink
cost at that point the GPU is saturated by per-program work and
packing more tiles per program lets the rank_vec be reused.
"""
target_grid = max(1, 2 * sm_count)
total_work = pair_slices * N_tiles
if total_work <= target_grid:
return 1
ntpp = (total_work + target_grid - 1) // target_grid
return min(ntpp, N_tiles)
def _run_fused_moe_lora_small_batch(
output: torch.Tensor,
qcurr_hidden_states: torch.Tensor,
lora_a_stacked: list[torch.Tensor],
lora_b_stacked: list[torch.Tensor],
topk_weights: torch.Tensor,
expert_ids_flat: torch.Tensor, # (num_tokens * top_k_num,)
token_lora_mapping: torch.Tensor,
top_k_num: int,
adapter_enabled: torch.Tensor,
mul_routed_weight: bool,
add_inputs: bool = True,
) -> None:
"""Small-batch GEMV-style wrapper. Naive-block-assignment inputs only.
Shape contract matches `_run_fused_moe_lora_one_shot`: `output` is
`(num_tokens, top_k_num, num_slices * N_per_slice)` with
contiguous-style strides, `expert_ids_flat` is the flattened
`topk_ids` of shape `(num_tokens * top_k_num,)`, and the
rank-padded LoRA weights live in `lora_a_stacked` /
`lora_b_stacked`.
The kernel is persistent over (pair × slice × n_chunk) work items
each program does one shrink and reuses the rank vector across
`n_tiles_per_program` expand tiles. The chunk size scales with the
pair-slice count so very small batches keep per-tile parallelism
while medium batches cut redundant shrinks.
"""
num_slices = len(lora_a_stacked)
device = qcurr_hidden_states.device
A0 = lora_a_stacked[0]
B0 = lora_b_stacked[0]
max_loras_w = A0.shape[0]
rank = A0.shape[2]
K = A0.shape[3]
N_per_slice = B0.shape[2]
# Rank padding: floor 16 (tensor-core min K), ceil to next pow2. The
# ≤64 cap is set conservatively for the prototype: at rank 64 the
# per-program register footprint is rank_vec(64 fp32) + b_tile(BLOCK_N
# × 64 fp32) = e.g. 128*64*4 = 32 KiB, comfortably within the 64 KiB
# register file even with num_warps=8. Doubling to 128 would push us
# against the limit and require shared-memory staging.
assert rank <= 64, f"fused_moe_lora_small_batch supports rank<=64; got rank={rank}"
BLOCK_R = max(triton.next_power_of_2(rank), 16)
num_tokens = topk_weights.shape[0]
M_grid = num_tokens * top_k_num
if M_grid == 0 or num_slices == 0:
return
token_mapping_factor = 1 if mul_routed_weight else top_k_num
A_ptrs = _get_ptr(lora_a_stacked, device)
B_ptrs = _get_ptr(lora_b_stacked, device)
assert output.dim() == 3, f"output must be 3-D, got {output.shape}"
assert output.stride(0) == output.shape[1] * output.stride(1), (
"fused_moe_lora_small_batch requires output.stride(0) == "
f"top_k*stride(1); got shape={output.shape} strides={output.stride()}"
)
out_view = output.view(-1, output.shape[-1])
# Block sizes. BLOCK_N=128 matches the one_shot's expand tile and gives
# 6-24 N tiles for typical N ∈ [768, 3072], enough to saturate the SM
# array once M_grid * num_slices reaches ~SM_count. BLOCK_K=128 halves
# the K-loop trip count vs 64 and pays for itself once K ≥ 1024 (the
# only regime we care about — hidden sizes are always large here).
BLOCK_N = 128
BLOCK_K = 128
nw = 4
ns = 3
N_tiles = triton.cdiv(N_per_slice, BLOCK_N)
pair_slices = M_grid * num_slices
sm_count = torch.cuda.get_device_properties(device).multi_processor_count
n_tiles_per_program = _pick_small_batch_chunk(pair_slices, N_tiles, sm_count)
n_chunks = triton.cdiv(N_tiles, n_tiles_per_program)
work_total = pair_slices * n_chunks
# Grid sizing: keep parallelism uncapped when work_total is small (so
# very small batches still spread across SMs); cap at 2*SM_count once
# we have plenty of work, letting the in-kernel stride loop drain the
# remainder.
grid_size = min(work_total, max(1, 2 * sm_count))
grid = (grid_size,)
_fused_moe_lora_small_batch_kernel[grid](
qcurr_hidden_states,
A_ptrs,
B_ptrs,
out_view,
topk_weights,
expert_ids_flat,
token_lora_mapping,
adapter_enabled,
N_per_slice,
K,
top_k_num,
max_loras_w,
work_total,
pair_slices,
qcurr_hidden_states.stride(0),
qcurr_hidden_states.stride(1),
A0.stride(0),
A0.stride(1),
A0.stride(2),
A0.stride(3),
B0.stride(0),
B0.stride(1),
B0.stride(2),
B0.stride(3),
out_view.stride(0),
out_view.stride(1),
N_per_slice,
n_tiles_per_program,
n_chunks,
token_mapping_factor=token_mapping_factor,
MUL_ROUTED_WEIGHT=mul_routed_weight,
ADD_INPUTS=add_inputs,
BLOCK_R=BLOCK_R,
actual_rank=rank,
BLOCK_N=BLOCK_N,
BLOCK_K=BLOCK_K,
NUM_SLICES=num_slices,
num_warps=nw,
num_stages=ns,
)
def _get_ptr(lora_weights: list[torch.Tensor], device: torch.device):
"""
`_LORA_PTR_DICT` collects the required information during `profile_run`,
@@ -1442,7 +706,6 @@ def _fused_moe_lora(
mul_routed_weight: bool = False,
fully_sharded: bool = False,
offset: int = 0,
add_inputs: bool = True,
) -> None:
assert len(lora_a_stacked) == len(lora_b_stacked) > 0
assert topk_weights.dim() == qcurr_hidden_states.dim() == 2
@@ -1465,59 +728,6 @@ def _fused_moe_lora(
)
assert output.shape[0] == topk_weights.shape[0]
assert top_k_num == topk_weights.shape[1]
# Fast path: single fused kernel
if not fully_sharded:
M_pairs = topk_weights.numel()
if (
sorted_token_ids is None
and max_lora_rank <= 64
and M_pairs * max_lora_rank <= 1024
):
_run_fused_moe_lora_small_batch(
output,
qcurr_hidden_states,
lora_a_stacked,
lora_b_stacked,
topk_weights,
expert_ids,
token_lora_mapping,
top_k_num,
adapter_enabled,
mul_routed_weight,
add_inputs=add_inputs,
)
return
# shrink/expand BLOCK_SIZE_M must match the block_size that
# moe_lora_align_block_size used; both shrink and expand pass the
# same value (asserted by `shrink_block_size_m == expand_block_size_m`
# below).
_run_fused_moe_lora_one_shot(
output,
qcurr_hidden_states,
lora_a_stacked,
lora_b_stacked,
topk_weights,
sorted_token_ids,
expert_ids,
num_tokens_post_padded,
token_lora_mapping,
max_lora_rank,
top_k_num,
lora_ids,
num_active_loras,
adapter_enabled,
mul_routed_weight,
shrink_block_size_m,
add_inputs=add_inputs,
)
return
assert add_inputs, (
"fused_moe_lora(add_inputs=False) is only supported on the "
"fully_sharded=False fast path"
)
device = qcurr_hidden_states.device
num_slices = len(lora_a_stacked)
w1_lora_b_stacked = lora_b_stacked[0]
@@ -1684,7 +894,6 @@ def _fused_moe_lora_fake(
mul_routed_weight: bool = False,
fully_sharded: bool = False,
offset: int = 0,
add_inputs: bool = True,
) -> None:
return
-2
View File
@@ -514,7 +514,6 @@ class PunicaWrapperBase(PunicaWrapperABC):
num_slices: int,
fully_sharded: bool,
use_tuned_config: bool,
add_inputs: bool = True,
token_lora_mapping: torch.Tensor | None = None,
) -> tuple[
torch.Tensor | None,
@@ -555,7 +554,6 @@ class PunicaWrapperBase(PunicaWrapperABC):
fully_sharded: bool,
tp_rank: int,
use_tuned_config: bool,
add_inputs: bool = True,
) -> None:
"""Apply w2 LoRA to y (intermediate_cache3) in-place before moe_sum.
-6
View File
@@ -435,7 +435,6 @@ class PunicaWrapperGPU(PunicaWrapperBase):
fully_sharded: bool = False,
offset: int = 0,
token_lora_mapping: torch.Tensor | None = None,
add_inputs: bool = True,
):
"""
Performs a fused forward computation for LoRA of Mixture-of-Experts (MoE) layer.
@@ -485,7 +484,6 @@ class PunicaWrapperGPU(PunicaWrapperBase):
mul_routed_weight,
fully_sharded,
offset,
add_inputs,
)
def add_lora_w13(
@@ -508,7 +506,6 @@ class PunicaWrapperGPU(PunicaWrapperBase):
num_slices: int,
fully_sharded: bool,
use_tuned_config: bool,
add_inputs: bool = True,
token_lora_mapping: torch.Tensor | None = None,
) -> tuple[
torch.Tensor | None,
@@ -613,7 +610,6 @@ class PunicaWrapperGPU(PunicaWrapperBase):
adapter_enabled,
fully_sharded=fully_sharded,
token_lora_mapping=token_lora_mapping,
add_inputs=add_inputs,
)
return (
@@ -644,7 +640,6 @@ class PunicaWrapperGPU(PunicaWrapperBase):
fully_sharded: bool,
tp_rank: int,
use_tuned_config: bool,
add_inputs: bool = True,
) -> None:
import functools
@@ -727,5 +722,4 @@ class PunicaWrapperGPU(PunicaWrapperBase):
fully_sharded=fully_sharded,
offset=offset,
token_lora_mapping=token_lora_mapping,
add_inputs=add_inputs,
)
-4
View File
@@ -390,7 +390,6 @@ class PunicaWrapperXPU(PunicaWrapperBase):
fully_sharded: bool = False,
offset: int = 0,
token_lora_mapping: torch.Tensor | None = None,
add_inputs: bool = True,
):
"""
Performs a fused forward computation for LoRA of Mixture-of-Experts (MoE) layer.
@@ -440,7 +439,6 @@ class PunicaWrapperXPU(PunicaWrapperBase):
mul_routed_weight,
fully_sharded,
offset,
add_inputs,
)
def add_lora_w13(
@@ -463,7 +461,6 @@ class PunicaWrapperXPU(PunicaWrapperBase):
num_slices: int,
fully_sharded: bool,
use_tuned_config: bool,
add_inputs: bool = True,
token_lora_mapping: torch.Tensor | None = None,
) -> tuple[
torch.Tensor | None,
@@ -597,7 +594,6 @@ class PunicaWrapperXPU(PunicaWrapperBase):
fully_sharded: bool,
tp_rank: int,
use_tuned_config: bool,
add_inputs: bool = True,
) -> None:
import functools
@@ -43,16 +43,6 @@ class MoELoRAContext:
# try_get_optimal_moe_lora_config for Triton kernel tile configs.
use_tuned_config: bool
# Optional dual-stream support for overlapping each (base GEMM, LoRA)
# pair. When aux_stream is None, the experts.apply() path runs the
# original sequential schedule. When set, base GEMM runs on the default
# stream and the LoRA fast-path writes the delta into a fresh buffer on
# aux_stream, which the default stream sums in afterwards.
# Events are paired one-per-overlap-pair: events[0,1] for w13,
# events[2,3] for w2, so the two pairs do not race on the same event.
aux_stream: torch.cuda.Stream | None = None
events: tuple[torch.cuda.Event, ...] | None = None
# Per-rank token→LoRA mapping after EP dispatch. Set by
# FusedMoEPrepareAndFinalizeModular.prepare() when EP+LoRA is active, read
# by LoRAExpertsMixin helpers in place of punica_wrapper's global mapping.
@@ -45,7 +45,6 @@ class LoRAExpertsMixin:
w2: torch.Tensor,
num_tokens: int,
top_k_num: int,
add_inputs: bool = True,
) -> tuple[
torch.Tensor | None,
torch.Tensor | None,
@@ -71,7 +70,6 @@ class LoRAExpertsMixin:
lora_context.w13_num_slices,
lora_context.fully_sharded,
lora_context.use_tuned_config,
add_inputs=add_inputs,
token_lora_mapping=lora_context.local_token_lora_mapping,
)
@@ -90,7 +88,6 @@ class LoRAExpertsMixin:
w1: torch.Tensor,
w2: torch.Tensor,
top_k_num: int,
add_inputs: bool = True,
) -> None:
lora_context.punica_wrapper.add_lora_w2(
y,
@@ -112,5 +109,4 @@ class LoRAExpertsMixin:
lora_context.fully_sharded,
lora_context.tp_rank,
lora_context.use_tuned_config,
add_inputs=add_inputs,
)
@@ -48,7 +48,6 @@ from vllm.model_executor.layers.quantization.utils.quant_utils import (
)
from vllm.platforms import current_platform
from vllm.triton_utils import tl
from vllm.utils.multi_stream_utils import maybe_execute_in_parallel
class TritonExperts(LoRAExpertsMixin, mk.FusedMoEExpertsModular):
@@ -248,99 +247,55 @@ class TritonExperts(LoRAExpertsMixin, mk.FusedMoEExpertsModular):
)
)
# LoRA w13: applied to intermediate_cache1 before activation. When
# the LoRA layer requested a dual-stream schedule, we run base w13
# GEMM on the default stream and the LoRA fast-path on aux_stream;
# the LoRA writes its delta into a fresh zero buffer (add_inputs=
# False) and we sum it into intermediate_cache1 after both finish.
invoke_fused_moe_triton_kernel(
hidden_states,
w1,
intermediate_cache1,
a1q_scale if a1q_scale is not None else self.a1_scale,
self.w1_scale,
None, # topk_weights
sorted_token_ids,
expert_ids,
num_tokens_post_padded,
False, # mul_routed_weights
top_k_num,
config,
compute_type=compute_type,
use_fp8_w8a8=self.quant_config.use_fp8_w8a8,
use_int8_w8a8=self.quant_config.use_int8_w8a8,
use_int8_w8a16=self.quant_config.use_int8_w8a16,
use_int4_w4a16=self.quant_config.use_int4_w4a16,
per_channel_quant=self.per_act_token_quant,
block_shape=self.block_shape,
B_bias=self.w1_bias,
)
# LoRA w13: applied to intermediate_cache1 before activation, using
# hidden_states as the lora_a input. moe_lora_align_block_size is
# called once here and results reused for the w2 LoRA below.
sorted_token_ids_lora = None
expert_ids_lora = None
num_tokens_post_padded_lora = None
token_lora_mapping = None
lora_context = self._lora_context
def _base_w13_fn():
invoke_fused_moe_triton_kernel(
hidden_states,
w1,
intermediate_cache1,
a1q_scale if a1q_scale is not None else self.a1_scale,
self.w1_scale,
None, # topk_weights
sorted_token_ids,
expert_ids,
num_tokens_post_padded,
False, # mul_routed_weights
top_k_num,
config,
compute_type=compute_type,
use_fp8_w8a8=self.quant_config.use_fp8_w8a8,
use_int8_w8a8=self.quant_config.use_int8_w8a8,
use_int8_w8a16=self.quant_config.use_int8_w8a16,
use_int4_w4a16=self.quant_config.use_int4_w4a16,
per_channel_quant=self.per_act_token_quant,
block_shape=self.block_shape,
B_bias=self.w1_bias,
)
if lora_context is not None and lora_context.aux_stream is not None:
# add_inputs=False: kernel overwrites lora_delta_w13. zeros (not
# empty) so untouched rows -- e.g. blocks where every program
# early-exits because lora_id<0 -- stay at zero and the trailing
# add_() is a no-op there.
lora_delta_w13 = torch.zeros_like(intermediate_cache1)
def _lora_w13_fn():
return self.apply_w13_lora(
lora_context,
y=lora_delta_w13,
x=hidden_states,
topk_ids=topk_ids,
topk_weights=topk_weights,
expert_map=expert_map,
w1=w1,
w2=w2,
num_tokens=num_tokens,
top_k_num=top_k_num,
add_inputs=False,
)
assert lora_context.events is not None
_, lora_meta = maybe_execute_in_parallel(
_base_w13_fn,
_lora_w13_fn,
lora_context.events[0],
lora_context.events[1],
lora_context.aux_stream,
)
if lora_context is not None:
(
sorted_token_ids_lora,
expert_ids_lora,
num_tokens_post_padded_lora,
token_lora_mapping,
) = lora_meta
intermediate_cache1.add_(lora_delta_w13)
else:
_base_w13_fn()
if lora_context is not None:
(
sorted_token_ids_lora,
expert_ids_lora,
num_tokens_post_padded_lora,
token_lora_mapping,
) = self.apply_w13_lora(
lora_context,
y=intermediate_cache1,
x=hidden_states,
topk_ids=topk_ids,
topk_weights=topk_weights,
expert_map=expert_map,
w1=w1,
w2=w2,
num_tokens=num_tokens,
top_k_num=top_k_num,
)
) = self.apply_w13_lora(
lora_context,
y=intermediate_cache1,
x=hidden_states,
topk_ids=topk_ids,
topk_weights=topk_weights,
expert_map=expert_map,
w1=w1,
w2=w2,
num_tokens=num_tokens,
top_k_num=top_k_num,
)
a2q_scale: torch.Tensor | None = None
@@ -373,82 +328,48 @@ class TritonExperts(LoRAExpertsMixin, mk.FusedMoEExpertsModular):
quantization_emulation=self.quantization_emulation,
)
invoke_fused_moe_triton_kernel(
qintermediate_cache2,
w2,
intermediate_cache3,
a2q_scale,
self.w2_scale,
topk_weights,
sorted_token_ids,
expert_ids,
num_tokens_post_padded,
not apply_router_weight_on_input,
1,
config,
compute_type=compute_type,
use_fp8_w8a8=self.quant_config.use_fp8_w8a8,
use_int8_w8a8=self.quant_config.use_int8_w8a8,
use_int8_w8a16=self.quant_config.use_int8_w8a16,
use_int4_w4a16=self.quant_config.use_int4_w4a16,
per_channel_quant=self.per_act_token_quant,
block_shape=self.block_shape,
B_bias=self.w2_bias,
)
# LoRA w2: applied to intermediate_cache3 before moe_sum, using the
# unquantized intermediate_cache2 as the lora_a input. Reuses the
# sorted_token_ids_lora computed above. Same dual-stream pattern as
# the w13 pair: base GEMM on default stream, LoRA delta on aux,
# join via .add_() into intermediate_cache3.
def _base_w2_fn():
invoke_fused_moe_triton_kernel(
qintermediate_cache2,
w2,
intermediate_cache3,
a2q_scale,
self.w2_scale,
topk_weights,
sorted_token_ids,
expert_ids,
num_tokens_post_padded,
not apply_router_weight_on_input,
1,
config,
compute_type=compute_type,
use_fp8_w8a8=self.quant_config.use_fp8_w8a8,
use_int8_w8a8=self.quant_config.use_int8_w8a8,
use_int8_w8a16=self.quant_config.use_int8_w8a16,
use_int4_w4a16=self.quant_config.use_int4_w4a16,
per_channel_quant=self.per_act_token_quant,
block_shape=self.block_shape,
B_bias=self.w2_bias,
# sorted_token_ids_lora computed above.
if lora_context is not None:
self.apply_w2_lora(
lora_context,
y=intermediate_cache3,
x=intermediate_cache2,
topk_weights=topk_weights,
sorted_token_ids_lora=sorted_token_ids_lora,
expert_ids_lora=expert_ids_lora,
num_tokens_post_padded_lora=num_tokens_post_padded_lora,
token_lora_mapping=token_lora_mapping,
num_tokens=num_tokens,
w1=w1,
w2=w2,
top_k_num=top_k_num,
)
if lora_context is not None and lora_context.aux_stream is not None:
lora_delta_w2 = torch.zeros_like(intermediate_cache3)
def _lora_w2_fn():
self.apply_w2_lora(
lora_context,
y=lora_delta_w2,
x=intermediate_cache2,
topk_weights=topk_weights,
sorted_token_ids_lora=sorted_token_ids_lora,
expert_ids_lora=expert_ids_lora,
num_tokens_post_padded_lora=num_tokens_post_padded_lora,
token_lora_mapping=token_lora_mapping,
num_tokens=num_tokens,
w1=w1,
w2=w2,
top_k_num=top_k_num,
add_inputs=False,
)
assert lora_context.events is not None
maybe_execute_in_parallel(
_base_w2_fn,
_lora_w2_fn,
lora_context.events[2],
lora_context.events[3],
lora_context.aux_stream,
)
intermediate_cache3.add_(lora_delta_w2)
else:
_base_w2_fn()
if lora_context is not None:
self.apply_w2_lora(
lora_context,
y=intermediate_cache3,
x=intermediate_cache2,
topk_weights=topk_weights,
sorted_token_ids_lora=sorted_token_ids_lora,
expert_ids_lora=expert_ids_lora,
num_tokens_post_padded_lora=num_tokens_post_padded_lora,
token_lora_mapping=token_lora_mapping,
num_tokens=num_tokens,
w1=w1,
w2=w2,
top_k_num=top_k_num,
)
# separate function is required for MoE + LoRA
self.moe_sum(intermediate_cache3, output)
@@ -164,8 +164,6 @@ def _moe_forward_shared_fake(
return shared_out, fused_out
# NOTE: `moe_forward` and `moe_forward_shared` being opaque custom ops is a
# load-bearing assumption for the MoE-LoRA dual-stream path.
direct_register_custom_op(
op_name="moe_forward",
op_func=_moe_forward,
@@ -3,7 +3,6 @@
"""Inference-only Qwen3-Next/Qwen3.5 model."""
import functools
from typing import Literal
import torch
from einops import rearrange
@@ -84,7 +83,7 @@ logger = init_logger(__name__)
# TODO(arpera): remove ``_is_libs_cu13_install_intact`` and its caller in
# ``_resolve_gdn_prefill_backend`` once the upstream packaging bug is
# ``_should_use_flashinfer_gdn_prefill`` once the upstream packaging bug is
# fixed and the broken wheels are yanked / superseded on PyPI:
# https://github.com/NVIDIA/cutlass/issues/3170
# https://github.com/NVIDIA/cutlass/issues/3259
@@ -147,12 +146,11 @@ def _is_libs_cu13_install_intact() -> bool:
return True
def _resolve_gdn_prefill_backend(
vllm_config: VllmConfig,
) -> tuple[str, Literal["triton", "flashinfer", "cutedsl"]]:
"""Resolve GDN prefill backend.
def _should_use_flashinfer_gdn_prefill(backend: str, head_k_dim: int | None) -> bool:
"""Whether to use FlashInfer's GDN prefill kernel instead of the
Triton/FLA fallback.
FlashInfer's GDN prefill kernel is chosen when:
Requirements:
* ``requested in ["flashinfer", "auto"]``;
* ``platform == cuda``;
* one of the following:
@@ -160,78 +158,47 @@ def _resolve_gdn_prefill_backend(
- Blackwell (SM10.x) with ``head_k_dim == 128``, ``cuda_runtime >= 13``,
and an intact ``nvidia-cutlass-dsl-libs-cu13`` install on disk
(see :func:`_is_libs_cu13_install_intact`).
In-tree CuteDSL GDN prefill kernel is chosen when:
* "cutedsl" is requested; (opt-in only)
* Blackwell (SM10.x) with ``head_k_dim == 128``;
"""
additional_config = vllm_config.additional_config
backend_cfg = (
additional_config.get("gdn_prefill_backend", "auto")
if isinstance(additional_config, dict)
else "auto"
)
backend = str(backend_cfg).strip().lower()
if backend not in ["flashinfer", "auto"]:
return False
if not current_platform.is_cuda():
return backend, "triton"
head_k_dim = getattr(
vllm_config.model_config.hf_config, "linear_key_head_dim", None
)
supports_flashinfer = False
supports_cutedsl = False
return False
if current_platform.is_device_capability(90):
supports_flashinfer = True
elif (
current_platform.is_device_capability_family(100)
and head_k_dim == 128
and current_platform.get_cuda_runtime_major() >= 13
):
supports_flashinfer = _is_libs_cu13_install_intact()
supports_cutedsl = True
if not supports_flashinfer:
logger.warning_once(
"FlashInfer Blackwell GDN requires an intact nvidia-cutlass-dsl"
"-libs-cu13 install, but some on-disk files do not match the "
"SHA-256 declared in its RECORD (install-order race in "
"nvidia-cutlass-dsl packaging -- see "
"https://github.com/NVIDIA/cutlass/issues/3170 and "
"https://github.com/NVIDIA/cutlass/issues/3259). Falling back "
"to Triton/FLA. Repair with: pip install --force-reinstall "
"--no-deps nvidia-cutlass-dsl-libs-cu13"
)
if backend in ["flashinfer", "auto"] and supports_flashinfer:
return backend, "flashinfer"
if backend == "cutedsl" and supports_cutedsl:
return backend, "cutedsl"
return backend, "triton"
return True # Hopper — no further constraints.
if not current_platform.is_device_capability_family(100):
return False # Neither Hopper nor Blackwell.
if head_k_dim != 128:
return False
if current_platform.get_cuda_runtime_major() < 13:
return False
if not _is_libs_cu13_install_intact():
logger.warning_once(
"FlashInfer Blackwell GDN requires an intact nvidia-cutlass-dsl"
"-libs-cu13 install, but some on-disk files do not match the "
"SHA-256 declared in its RECORD (install-order race in "
"nvidia-cutlass-dsl packaging — see "
"https://github.com/NVIDIA/cutlass/issues/3170 and "
"https://github.com/NVIDIA/cutlass/issues/3259). Falling back "
"to Triton/FLA. Repair with: pip install --force-reinstall "
"--no-deps nvidia-cutlass-dsl-libs-cu13"
)
return False
return True
def _log_gdn_backend_decision(
vllm_config: VllmConfig,
requested_backend: str,
active_backend: str,
backend: str, head_k_dim: int | None, use_flashinfer: bool
) -> None:
"""Log the GDN prefill backend choice in the attention-selector style."""
head_k_dim = getattr(
vllm_config.model_config.hf_config, "linear_key_head_dim", None
)
chosen = {
"flashinfer": "FlashInfer",
"cutedsl": "CuteDSL",
"triton": "Triton/FLA",
}[active_backend]
chosen = "FlashInfer" if use_flashinfer else "Triton/FLA"
logger.info_once(
"Using %s GDN prefill kernel (requested=%s, head_k_dim=%s).",
chosen,
requested_backend,
backend,
head_k_dim,
)
if active_backend == "flashinfer" and current_platform.is_device_capability(90):
# JIT-compiled cutlass path is only used on SM90 (Hopper).
if use_flashinfer and current_platform.is_device_capability(90):
logger.warning_once(
"FlashInfer GDN prefill is JIT-compiled; first run may take a "
"while. Set --gdn-prefill-backend triton to skip JIT.",
@@ -289,26 +256,25 @@ def fi_chunk_gated_delta_rule(
@CustomOp.register("chunk_gated_delta_rule")
class ChunkGatedDeltaRule(CustomOp):
def __init__(self) -> None:
def __init__(self, head_k_dim: int | None = None) -> None:
super().__init__()
vllm_config = get_current_vllm_config()
backend, active_backend = _resolve_gdn_prefill_backend(vllm_config)
self.gdn_prefill_backend = active_backend
additional_config = get_current_vllm_config().additional_config
assert isinstance(additional_config, dict)
backend_cfg = additional_config.get("gdn_prefill_backend", "auto")
backend = str(backend_cfg).strip().lower()
if backend in ("flashinfer", "cutedsl") and active_backend != backend:
use_flashinfer = _should_use_flashinfer_gdn_prefill(backend, head_k_dim)
if backend == "flashinfer" and not use_flashinfer:
logger.warning_once(
"GDN prefill backend '%s' is selected but cannot use this "
"kernel on the current platform. Falling back to Triton/FLA.",
backend,
"GDN prefill backend 'flashinfer' is selected but "
"cannot use this kernel on the current platform. "
"Falling back to Triton/FLA."
)
_log_gdn_backend_decision(vllm_config, backend, active_backend)
_log_gdn_backend_decision(backend, head_k_dim, use_flashinfer)
if active_backend == "flashinfer":
self._forward_method = self.forward_cuda
elif active_backend == "cutedsl":
self._forward_method = self.forward_cutedsl
else:
self._forward_method = self.forward_native
self._forward_method = (
self.forward_cuda if use_flashinfer else self.forward_native
)
def forward_cuda(
self,
@@ -372,49 +338,6 @@ class ChunkGatedDeltaRule(CustomOp):
core_attn_out=core_attn_out,
)
def forward_cutedsl(
self,
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
initial_state: torch.Tensor,
output_final_state: bool,
cu_seqlens: torch.Tensor | None = None,
chunk_indices: torch.Tensor | None = None,
chunk_offsets: torch.Tensor | None = None,
use_qk_l2norm_in_kernel: bool = True,
core_attn_out: torch.Tensor | None = None,
):
from vllm.model_executor.layers.mamba.ops.gdn_chunk_cutedsl import (
chunk_gated_delta_rule_cutedsl,
)
if use_qk_l2norm_in_kernel:
q = l2norm_fwd(q)
k = l2norm_fwd(k)
assert cu_seqlens is not None
assert chunk_indices is not None
assert chunk_offsets is not None
o, final_state = chunk_gated_delta_rule_cutedsl(
q=q,
k=k,
v=v,
g=g,
beta=beta,
initial_state=initial_state,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
chunk_offsets=chunk_offsets,
core_attn_out=core_attn_out,
)
if not output_final_state:
final_state = None
return o, final_state
@PluggableLayer.register("qwen_gated_delta_net_attention")
class QwenGatedDeltaNetAttention(GatedDeltaNetAttention):
@@ -551,8 +474,7 @@ class QwenGatedDeltaNetAttention(GatedDeltaNetAttention):
prefix=f"{prefix}.out_proj",
)
self.chunk_gated_delta_rule = ChunkGatedDeltaRule()
self.gdn_prefill_backend = self.chunk_gated_delta_rule.gdn_prefill_backend
self.chunk_gated_delta_rule = ChunkGatedDeltaRule(head_k_dim=self.head_k_dim)
self._prefill_kernels_warmed_up = False
self.enable_packed_recurrent_decode = (
envs.VLLM_ENABLE_FLA_PACKED_RECURRENT_DECODE
@@ -1138,16 +1060,6 @@ class QwenGatedDeltaNetAttention(GatedDeltaNetAttention):
)
cu_seqlens = torch.tensor([0, T], device=device, dtype=torch.int32)
# CuteDSL kernels require metadata
chunk_indices = None
chunk_offsets = None
if self.gdn_prefill_backend == "cutedsl":
from vllm.model_executor.layers.mamba.ops.gdn_chunk_cutedsl import (
prepare_metadata_cutedsl,
)
chunk_indices, chunk_offsets = prepare_metadata_cutedsl(cu_seqlens, T)
try:
self.chunk_gated_delta_rule(
q=q,
@@ -1158,8 +1070,6 @@ class QwenGatedDeltaNetAttention(GatedDeltaNetAttention):
initial_state=state,
output_final_state=True,
cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices,
chunk_offsets=chunk_offsets,
use_qk_l2norm_in_kernel=False,
)
except Exception:
@@ -1178,20 +1088,7 @@ class QwenGatedDeltaNetAttention(GatedDeltaNetAttention):
self.prefix,
)
finally:
del (
dummy_mixed_qkv,
q,
k,
v,
dummy_a,
dummy_b,
g,
beta,
state,
cu_seqlens,
chunk_indices,
chunk_offsets,
)
del dummy_mixed_qkv, q, k, v, dummy_a, dummy_b, g, beta, state, cu_seqlens
torch.accelerator.empty_cache()
@@ -1,251 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
from functools import cache
import cutlass
import torch
from cuda.bindings.driver import CUstream
from cutlass import Int32, cute
from quack.compile_utils import make_fake_tensor
from vllm.triton_utils import triton
from .kernel_h import h_cutedsl
from .kernel_kkt_inv_uw import kkt_inv_uw_cutedsl
from .kernel_o import o_cutedsl
class PrepMetaKernel:
def __init__(self, BT: int) -> None:
self.BT = BT
self.num_warps = 8
@cute.jit
def __call__(
self,
cu_seqlens: cute.Tensor,
chunk_indices: cute.Tensor,
chunk_offsets: cute.Tensor,
stream: CUstream,
):
block = (self.num_warps * 32, 1, 1)
self.kernel(
cu_seqlens,
chunk_indices,
chunk_offsets,
).launch(grid=(1, 1, 1), block=block, stream=stream)
@cute.kernel
def kernel(
self,
cu_seqlens: cute.Tensor,
chunk_indices: cute.Tensor,
chunk_offsets: cute.Tensor,
):
tid, _, _ = cute.arch.thread_idx()
warp_id = cute.arch.make_warp_uniform(tid // 32)
lane_id = tid % 32
num_seqs = cu_seqlens.shape[0] - 1
num_warps = self.num_warps
tb_size = num_warps * 32
if tid == 0:
chunk_offsets[0] = 0
coarsen = cute.ceil_div(num_seqs, tb_size)
seq_start = tid * coarsen
num_iters = cutlass.min(seq_start + coarsen, num_seqs) - seq_start
# First pass: compute this thread's total chunk count.
thread_sum = Int32(0)
for i in range(num_iters):
seq_id = seq_start + i
seqlen = cu_seqlens[seq_id + 1] - cu_seqlens[seq_id]
thread_sum += cute.ceil_div(seqlen, self.BT)
# warp parallel scan
cu_num_chunks = thread_sum
for i in cutlass.range_constexpr(5):
offset = cutlass.const_expr(1 << i)
lower = cute.arch.shuffle_sync_up(
cu_num_chunks, offset=offset, mask_and_clamp=0
)
if lane_id >= offset:
cu_num_chunks += lower
# cross-warp cumsum (CTA-wide)
smem = cutlass.utils.SmemAllocator()
warp_num_chunks = smem.allocate_array(Int32, num_warps)
if lane_id == 31:
warp_num_chunks[warp_id] = cu_num_chunks
cute.arch.sync_threads()
for i in cutlass.range_constexpr(1, num_warps):
if warp_id >= i:
cu_num_chunks += warp_num_chunks[i - 1]
chunk_start = cu_num_chunks - thread_sum
# Second pass: recompute per-sequence chunk counts and write results.
for i in range(num_iters):
seq_id = seq_start + i
seqlen = cu_seqlens[seq_id + 1] - cu_seqlens[seq_id]
num_chunks = cute.ceil_div(seqlen, self.BT)
chunk_end = chunk_start + num_chunks
chunk_offsets[seq_id + 1] = chunk_end
for chunk_id in range(num_chunks):
chunk_indices[chunk_start + chunk_id, 0] = seq_id
chunk_indices[chunk_start + chunk_id, 1] = chunk_id
chunk_start = chunk_end
@cache
@staticmethod
def compile(BT: int):
cu_entries = cute.sym_int()
upper_bound_chunks = cute.sym_int()
cu_seqlens = make_fake_tensor(Int32, (cu_entries,), divisibility=1)
chunk_indices = make_fake_tensor(Int32, (upper_bound_chunks, 2), divisibility=2)
chunk_offsets = make_fake_tensor(Int32, (cu_entries,), divisibility=1)
kernel = PrepMetaKernel(BT)
stream = cute.runtime.make_fake_stream(use_tvm_ffi_env_stream=True)
return cute.compile(
kernel,
cu_seqlens,
chunk_indices,
chunk_offsets,
stream,
options="--enable-tvm-ffi",
)
def _upper_bound_chunks(num_seqs: int, total_tokens: int, chunk_size: int) -> int:
return (num_seqs - 1) + triton.cdiv(total_tokens - (num_seqs - 1), chunk_size)
def prepare_metadata_cutedsl(
cu_seqlens: torch.Tensor,
total_tokens: int,
chunk_size: int = 64,
) -> tuple[torch.Tensor, torch.Tensor]:
num_seqs = cu_seqlens.numel() - 1
upper_bound_chunks = _upper_bound_chunks(num_seqs, total_tokens, chunk_size)
chunk_offsets = cu_seqlens.new_empty(num_seqs + 1, dtype=torch.int32)
chunk_indices = cu_seqlens.new_empty((upper_bound_chunks, 2), dtype=torch.int32)
PrepMetaKernel.compile(chunk_size)(cu_seqlens, chunk_indices, chunk_offsets)
return chunk_indices, chunk_offsets
def chunk_gated_delta_rule_cutedsl(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
initial_state: torch.Tensor,
cu_seqlens: torch.Tensor,
chunk_indices: torch.Tensor,
chunk_offsets: torch.Tensor,
core_attn_out: torch.Tensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Run the GDN chunk CuteDSL prefill kernels.
Args:
q: Query tensor with shape ``[1, T, H, K]``.
k: Key tensor with shape ``[1, T, H, K]``.
v: Value tensor with shape ``[1, T, Hv, V]``.
g: Log-space decay tensor with shape ``[1, T, Hv]``.
beta: Delta-rule beta tensor with shape ``[1, T, Hv]``.
initial_state: Recurrent state with shape ``[N, Hv, V, K]``.
cu_seqlens: Cumulative sequence lengths with shape ``[N + 1]``.
chunk_indices: Chunk index metadata with shape ``[NT, 2]``.
chunk_offsets: Cumulative chunk offsets with shape ``[N + 1]``.
core_attn_out: Optional output buffer with shape ``[T, Hv, V]``.
Returns:
A tuple ``(output, final_state)`` where ``output`` has shape
``[1, T, Hv, V]`` and ``final_state`` has shape ``[N, Hv, V, K]``.
When ``core_attn_out`` is provided, ``output`` is an unsqueezed view of
that buffer.
"""
q_3d = q.squeeze(0)
k_3d = k.squeeze(0)
v_3d = v.squeeze(0)
g_2d = g.squeeze(0)
beta_2d = beta.squeeze(0)
_, _, head_k_dim = k_3d.shape
_, num_v_heads, head_v_dim = v_3d.shape
chunk_size = 64
upper_bound_chunks = chunk_indices.shape[0]
pad_t = upper_bound_chunks * chunk_size
total_chunks_ptr = chunk_offsets[-1:]
g_cu = torch.empty_like(g_2d, dtype=torch.float32)
u = q_3d.new_empty(pad_t, num_v_heads, head_v_dim)
w = q_3d.new_empty(pad_t, num_v_heads, head_k_dim)
num_sms = torch.cuda.get_device_properties(q.device).multi_processor_count
kkt_inv_uw_cutedsl(
k_3d,
v_3d,
u,
w,
g_2d,
beta_2d,
g_cu,
cu_seqlens,
chunk_indices,
total_chunks_ptr,
num_sms=num_sms,
)
h = k_3d.new_empty(
upper_bound_chunks,
num_v_heads,
head_v_dim,
head_k_dim,
)
v_new = q_3d.new_empty(pad_t, num_v_heads, head_v_dim)
final_state = torch.empty_like(initial_state)
h_cutedsl(
k_3d,
u,
w,
v_new,
g_cu,
h,
initial_state,
final_state,
cu_seqlens,
chunk_offsets,
)
output = core_attn_out if core_attn_out is not None else torch.empty_like(v_3d)
scale = head_k_dim**-0.5
o_cutedsl(
q_3d,
k_3d,
v_new.view(upper_bound_chunks, chunk_size, num_v_heads, head_v_dim),
h,
g_cu,
output,
cu_seqlens,
chunk_indices,
total_chunks_ptr,
scale,
num_sms=num_sms,
)
return output.unsqueeze(0), final_state
__all__ = [
"chunk_gated_delta_rule_cutedsl",
"prepare_metadata_cutedsl",
]
@@ -1,753 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
from functools import cache
import cutlass
import torch
from cuda.bindings.driver import CUstream
from cutlass import BFloat16, Float32, Int32, Int64, Uint32, cute
from cutlass.cute.nvgpu import cpasync, warp
from quack.compile_utils import make_fake_tensor
from vllm.cute_utils import (
EVICT_FIRST,
_tcgen05,
cvt,
fence_before_tma_store,
simple_tma_copy,
)
class Sm100ChunkHKernel:
"""For each sequence, compute the chunk recurrent update.
The input V tile is the U output from the KKT/UW kernel. For each chunk:
V_new = U - W @ H.T
(we actually do V_new.T = U.T - H @ W.T instead)
H_scaled = H * exp(g_last)
V_scaled = V_new * exp(g_last - g)
H_new = H_scaled + V_scaled.T @ K
"""
def __init__(
self,
H: int,
Hv: int,
K_dim: int,
V_dim: int,
h_dtype: cutlass.Numeric = Float32,
BT: int = 64,
num_stages: int = 2,
) -> None:
assert Hv % H == 0
assert K_dim == V_dim == 128
assert BT == 64
self.H = H
self.Hv = Hv
self.K_dim = K_dim
self.V_dim = V_dim
self.h_dtype = h_dtype
self.BT = BT
self.num_stages = num_stages
self.num_warps = 10
@cute.jit
def _make_bf16_tma_args(
self,
tensor: cute.Tensor,
dim: cutlass.Constexpr[int],
op: cpasync.TmaCopyOp,
stages: cutlass.Constexpr[int],
):
swizzle_128B = cute.make_swizzle(3, 4, 3)
slayout = cute.make_layout(
(self.BT, 1, (64, dim // 64), stages),
stride=(64, 0, (1, self.BT * 64), self.BT * dim),
)
slayout = cute.make_composed_layout(swizzle_128B, 0, slayout)
atom, tma_tensor = cpasync.make_tiled_tma_atom(
op,
cute.logical_divide(tensor, (None, None, 64)),
slayout,
cta_tiler=(self.BT, 1, dim),
)
return atom, tma_tensor, slayout
@cute.jit
def _make_h_tma_args(self, tensor: cute.Tensor, op: cpasync.TmaCopyOp):
# number of elements to fill 128B
num_elems = 128 // (tensor.element_type.width // 8)
swizzle_128B = cute.make_swizzle(3, 4, 3)
slayout = cute.make_layout(
(1, 1, self.V_dim, (num_elems, self.K_dim // num_elems)),
stride=(0, 0, num_elems, (1, self.V_dim * num_elems)),
)
slayout = cute.make_composed_layout(swizzle_128B, 0, slayout)
atom, tma_tensor = cpasync.make_tiled_tma_atom(
op,
cute.logical_divide(tensor, (None, None, None, num_elems)),
slayout,
cta_tiler=(1, 1, self.V_dim, self.K_dim),
)
return atom, tma_tensor, slayout
@cute.jit
def __call__(
self,
K: cute.Tensor,
V: cute.Tensor,
W: cute.Tensor,
V_new: cute.Tensor,
g_cu: cute.Tensor,
h: cute.Tensor,
h0: cute.Tensor,
ht: cute.Tensor,
cu_seqlens: cute.Tensor,
chunk_offsets: cute.Tensor,
stream: CUstream,
):
tma_g2s = cpasync.CopyBulkTensorTileG2SOp()
tma_s2g = cpasync.CopyBulkTensorTileS2GOp()
K_args = self._make_bf16_tma_args(K, self.K_dim, tma_g2s, self.num_stages)
V_args = self._make_bf16_tma_args(V, self.V_dim, tma_g2s, self.num_stages)
W_args = self._make_bf16_tma_args(W, self.K_dim, tma_g2s, self.num_stages)
V_new_args = self._make_bf16_tma_args(V_new, self.V_dim, tma_s2g, 1)
H0_args = self._make_h_tma_args(h0, tma_g2s)
HT_args = self._make_h_tma_args(ht, tma_s2g)
H_args = self._make_h_tma_args(h, tma_s2g)
grid = (self.Hv, h0.shape[0], 1)
block = (self.num_warps * 32, 1, 1)
self.kernel(
K_args,
V_args,
W_args,
V_new_args,
H0_args,
HT_args,
H_args,
g_cu,
cu_seqlens,
chunk_offsets,
).launch(grid=grid, block=block, stream=stream)
@cute.kernel
def kernel(
self,
K_args: tuple[cute.CopyAtom, cute.Tensor, cute.ComposedLayout],
V_args: tuple[cute.CopyAtom, cute.Tensor, cute.ComposedLayout],
W_args: tuple[cute.CopyAtom, cute.Tensor, cute.ComposedLayout],
V_new_args: tuple[cute.CopyAtom, cute.Tensor, cute.ComposedLayout],
H0_args: tuple[cute.CopyAtom, cute.Tensor, cute.ComposedLayout],
HT_args: tuple[cute.CopyAtom, cute.Tensor, cute.ComposedLayout],
H_args: tuple[cute.CopyAtom, cute.Tensor, cute.ComposedLayout],
g_cu: cute.Tensor,
cu_seqlens: cute.Tensor,
chunk_offsets: cute.Tensor,
):
tid, _, _ = cute.arch.thread_idx()
head_id, seq_id, _ = cute.arch.block_idx()
warp_id = cute.arch.make_warp_uniform(tid // 32)
lane_id = tid % 32
BT = self.BT
V_dim = self.V_dim
K_dim = self.K_dim
num_stages = self.num_stages
is_f32 = self.h_dtype == Float32
K_tma_atom, tmaK, sK_layout = K_args
V_tma_atom, tmaV, sV_layout = V_args
W_tma_atom, tmaW, sW_layout = W_args
V_new_tma_atom, tmaV_new, sV_new_layout = V_new_args
H0_tma_atom, tmaH0, sH0_layout = H0_args
HT_tma_atom, tmaHT, _ = HT_args
H_tma_atom, tmaH, sH_layout = H_args
def allocate_tensor(smem, dtype, layout):
return smem.allocate_tensor(
dtype, layout.outer, byte_alignment=128, swizzle=layout.inner
)
smem = cutlass.utils.SmemAllocator()
# remove size=1 modes
sW = allocate_tensor(smem, BFloat16, sW_layout)[None, 0, None, None]
sV = allocate_tensor(smem, BFloat16, sV_layout)[None, 0, None, None]
sK = allocate_tensor(smem, BFloat16, sK_layout)[None, 0, None, None]
sH0 = allocate_tensor(smem, self.h_dtype, sH0_layout)[0, 0, None, None]
sH = allocate_tensor(smem, BFloat16, sH_layout)[0, 0, None, None]
sV_new = allocate_tensor(smem, BFloat16, sV_new_layout)[None, 0, None, 0]
s_v_scale = smem.allocate_array(Float32, BT)
tma_mbar = smem.allocate_array(Int64, num_stages)
wh_in_mbar = smem.allocate_array(Int64, num_stages)
wh_done_mbar = smem.allocate_array(Int64, num_stages)
vk_in_mbar = smem.allocate_array(Int64, num_stages)
vk_done_mbar = smem.allocate_array(Int64, num_stages)
h0_mbar = smem.allocate_array(Int64, 1)
taddr = smem.allocate(Int32, 4)
wh_tmem = 0
vk_tmem = wh_tmem + BT
h_tmem_base = vk_tmem + K_dim
v_tmem_base = h_tmem_base + K_dim // 2
if warp_id == 0:
with cute.arch.elect_one():
for i in cutlass.range_constexpr(num_stages):
cute.arch.mbarrier_init(tma_mbar + i, 1)
cute.arch.mbarrier_init(wh_in_mbar + i, 256)
cute.arch.mbarrier_init(wh_done_mbar + i, 1)
cute.arch.mbarrier_init(vk_in_mbar + i, 256)
cute.arch.mbarrier_init(vk_done_mbar + i, 1)
cute.arch.mbarrier_init(h0_mbar, 1)
cute.arch.mbarrier_init_fence()
elif warp_id == 1:
cpasync.prefetch_descriptor(H0_tma_atom)
cpasync.prefetch_descriptor(W_tma_atom)
cpasync.prefetch_descriptor(V_tma_atom)
cpasync.prefetch_descriptor(K_tma_atom)
cpasync.prefetch_descriptor(HT_tma_atom)
cpasync.prefetch_descriptor(H_tma_atom)
cpasync.prefetch_descriptor(V_new_tma_atom)
cute.arch.sync_threads()
bos = cu_seqlens[seq_id]
eos = cu_seqlens[seq_id + 1]
seqlen = eos - bos
num_chunks = cute.ceil_div(seqlen, BT)
if warp_id == 9:
# TMA warp
stage_id = 0
parity = 1
k_head_id = head_id // (self.Hv // self.H)
chunk_offset = chunk_offsets[seq_id]
# load H0
with cute.arch.elect_one():
H0_size = V_dim * K_dim * self.h_dtype.width // 8
cute.arch.mbarrier_arrive_and_expect_tx(h0_mbar, H0_size)
simple_tma_copy(
H0_tma_atom, tmaH0[seq_id, head_id, None, None], sH0, h0_mbar
)
# shape: ((BT, num_BT_tiles), (64, 2))
gW_tiles = cute.logical_divide(tmaW[None, head_id, None], (BT, None))
gV_tiles = cute.logical_divide(tmaV[None, head_id, None], (BT, None))
gK_tiles = cute.logical_divide(
cute.domain_offset((bos, 0), tmaK[None, k_head_id, None]),
(BT, None),
)
for chunk_id in range(num_chunks):
mbar = tma_mbar + stage_id
gW = gW_tiles[(None, chunk_offset + chunk_id), None]
gV = gV_tiles[(None, chunk_offset + chunk_id), None]
gK = gK_tiles[(None, chunk_id), None]
# wait for MMA to release the buffer
cute.arch.mbarrier_wait(vk_done_mbar + stage_id, parity)
# load W, V (i.e. U), and K
with cute.arch.elect_one():
STAGE_SIZE = BT * (K_dim + V_dim + K_dim) * 2
cute.arch.mbarrier_arrive_and_expect_tx(mbar, STAGE_SIZE)
simple_tma_copy(
W_tma_atom, gW, sW[None, None, stage_id], mbar, EVICT_FIRST
)
simple_tma_copy(
V_tma_atom, gV, sV[None, None, stage_id], mbar, EVICT_FIRST
)
simple_tma_copy(K_tma_atom, gK, sK[None, None, stage_id], mbar)
stage_id = (stage_id + 1) % num_stages
if stage_id == 0:
parity ^= 1
elif warp_id == 8:
# MMA warp
_tcgen05.alloc(taddr)
stage_id = 0
parity = 0
wh_idesc = _tcgen05.make_bf16_idesc(V_dim, BT, negate_A=True)
vk_idesc = _tcgen05.make_bf16_idesc(V_dim, K_dim, transpose_B=True)
# LBO=BT*128 is ignored for K-major
sdesc_template = _tcgen05.make_sdesc_128B_swizzle(BT * 128)
# when using BF16 state, H is read from smem for the 1st iteration
# variable names in this conditional branch can't be the same as those
# in the mainloop below due to CuteDSL restrictions.
if cutlass.const_expr(not is_f32):
##### 1st MMA: V_new.T = V.T - H @ W.T #####
Haddr0 = sH0[None, None].iterator.toint()
Waddr0 = sW[None, None, stage_id].iterator.toint()
hdesc0_base = sdesc_template | (Haddr0 >> 4)
wdesc0_base = sdesc_template | (Waddr0 >> 4)
cute.arch.mbarrier_wait(tma_mbar + stage_id, parity)
cute.arch.mbarrier_wait(wh_in_mbar + stage_id, parity)
_tcgen05.fence_after_thread_sync()
with cute.arch.elect_one():
for i in cutlass.range_constexpr(K_dim // 64):
for j in cutlass.range_constexpr(64 // 16):
hdesc0 = hdesc0_base | ((i * V_dim * 128 + j * 32) >> 4)
wdesc0 = wdesc0_base | ((i * BT * 128 + j * 32) >> 4)
_tcgen05.mma_f16(wh_tmem, hdesc0, wdesc0, wh_idesc, True)
_tcgen05.commit(wh_done_mbar + stage_id)
##### 2nd MMA: H_new = H + V_new.T @ K #####
Kaddr0 = sK[None, None, stage_id].iterator.toint()
kdesc0_base = sdesc_template | (Kaddr0 >> 4)
cute.arch.mbarrier_wait(vk_in_mbar + stage_id, parity)
_tcgen05.fence_after_thread_sync()
with cute.arch.elect_one():
for k in cutlass.range_constexpr(BT // 16):
vtmem0 = v_tmem_base + k * 8
kdesc0 = kdesc0_base | ((k * 16 * 128) >> 4)
_tcgen05.mma_ts_f16(vk_tmem, vtmem0, kdesc0, vk_idesc, True)
_tcgen05.commit(vk_done_mbar + stage_id)
stage_id = (stage_id + 1) % num_stages
if stage_id == 0:
parity ^= 1
num_iters = num_chunks - int(not is_f32)
for _ in range(num_iters):
##### 1st MMA: V_new.T = V.T - H @ W.T #####
Waddr = sW[None, None, stage_id].iterator.toint()
wdesc_base = sdesc_template | (Waddr >> 4)
cute.arch.mbarrier_wait(tma_mbar + stage_id, parity)
cute.arch.mbarrier_wait(wh_in_mbar + stage_id, parity)
_tcgen05.fence_after_thread_sync()
with cute.arch.elect_one():
for i in cutlass.range_constexpr(K_dim // 64):
for j in cutlass.range_constexpr(64 // 16):
htmem = h_tmem_base + i * 32 + j * 8
wdesc = wdesc_base | ((i * BT * 128 + j * 32) >> 4)
_tcgen05.mma_ts_f16(wh_tmem, htmem, wdesc, wh_idesc, True)
_tcgen05.commit(wh_done_mbar + stage_id)
##### 2nd MMA: H_new = H + V_new.T @ K #####
Kaddr = sK[None, None, stage_id].iterator.toint()
kdesc_base = sdesc_template | (Kaddr >> 4)
cute.arch.mbarrier_wait(vk_in_mbar + stage_id, parity)
_tcgen05.fence_after_thread_sync()
with cute.arch.elect_one():
for k in cutlass.range_constexpr(BT // 16):
vtmem = v_tmem_base + k * 8
kdesc = kdesc_base | ((k * 16 * 128) >> 4)
_tcgen05.mma_ts_f16(vk_tmem, vtmem, kdesc, vk_idesc, True)
_tcgen05.commit(vk_done_mbar + stage_id)
stage_id = (stage_id + 1) % num_stages
if stage_id == 0:
parity ^= 1
elif warp_id >= 4:
# H warps
tid_ = tid % 128
warp_id_ = warp_id % 4
chunk_offset = chunk_offsets[seq_id]
stage_id = 0
vk_stage_id = 0
vk_parity = 0
op = cute.nvgpu.CopyUniversalOp()
cp_16B = cute.make_copy_atom(op, Float32, num_bits_per_copy=128)
##### chunk_id = 0 #####
if True:
chunk_id = 0
end_t = min(bos + (chunk_id + 1) * BT, eos)
last_idx = end_t - 1
h_scale = cute.math.exp(g_cu[last_idx, head_id], fastmath=True)
# for 1st chunk, wait for H0 transfer from gmem
if warp_id_ == 0:
cute.arch.mbarrier_wait(h0_mbar, 0)
cute.arch.barrier(barrier_id=1, number_of_threads=128)
# when H0 is FP32, we need to pack it to BF16
# also store to smem for TMA store later.
if cutlass.const_expr(is_f32):
for i in cutlass.range_constexpr(K_dim // 32):
# H0 smem layout: (V_dim, (32, K_dim/32))
h_f32 = cute.make_rmem_tensor(32, Float32)
cute.copy(cp_16B, sH0[tid_, (None, i)], h_f32)
h_bf16 = cute.make_rmem_tensor(32, BFloat16)
h_bf16.store(h_f32.load().to(BFloat16))
_tcgen05.st(
warp_id_ * 32, h_tmem_base + i * 16, "32x32b", 16, h_bf16
)
# H smem layout: (V_dim, (64, K_dim/64))
dst = cute.local_tile(sH[tid_, None], (32,), (i,))
cute.copy(cp_16B, h_bf16, dst)
_tcgen05.wait_st()
_tcgen05.fence_before_thread_sync()
cute.arch.mbarrier_arrive(wh_in_mbar + stage_id)
# scale H for 2nd MMA
for i in cutlass.range_constexpr(K_dim // 32):
h_f32 = cute.make_rmem_tensor(32, Float32)
if cutlass.const_expr(is_f32):
cute.copy(cp_16B, sH0[tid_, (None, i)], h_f32)
else:
h_bf16 = cute.make_rmem_tensor(32, BFloat16)
sH_src = cute.local_tile(sH0[tid_, None], (32,), (i,))
cute.copy(cp_16B, sH_src, h_bf16)
h_f32.store(
cvt.bf16x2_to_fp32x2(
cute.recast_tensor(h_bf16, Uint32)
).load()
)
for j in cutlass.range_constexpr(32):
h_f32[j] *= h_scale
_tcgen05.st(warp_id_ * 32, vk_tmem + i * 32, "32x32b", 32, h_f32)
_tcgen05.wait_st()
_tcgen05.fence_before_thread_sync()
cute.arch.mbarrier_arrive(vk_in_mbar + stage_id)
# for BF16 H0, we issue TMA store from H0 smem
# for FP32 H0, we issue TMA store from H smem (after packing)
cute.arch.barrier(barrier_id=1, number_of_threads=128)
fence_before_tma_store()
if warp_id_ == 3:
h_src = sH if cutlass.const_expr(is_f32) else sH0
h_dst = tmaH[chunk_offset + chunk_id, head_id, None, None]
simple_tma_copy(H_tma_atom, h_src, h_dst)
with cute.arch.elect_one():
cute.arch.cp_async_bulk_commit_group()
# When H0 is BF16, and there is only 1 chunk, storing
# the final state to sH0 can race before this store
# has finished. hence, we need to wait here.
if cutlass.const_expr(not is_f32):
cute.arch.cp_async_bulk_wait_group(0, read=True)
stage_id = (stage_id + 1) % num_stages
##### subsequent chunks #####
for chunk_id in range(1, num_chunks):
end_t = min(bos + (chunk_id + 1) * BT, eos)
last_idx = end_t - 1
h_scale = cute.math.exp(g_cu[last_idx, head_id], fastmath=True)
# wait for H from previous vk MMA
if warp_id_ == 0:
cute.arch.mbarrier_wait(vk_done_mbar + vk_stage_id, vk_parity)
vk_stage_id = (vk_stage_id + 1) % num_stages
if vk_stage_id == 0:
vk_parity ^= 1
elif warp_id_ == 3:
with cute.arch.elect_one():
cute.arch.cp_async_bulk_wait_group(0, read=True)
cute.arch.barrier(barrier_id=1, number_of_threads=128)
_tcgen05.fence_after_thread_sync()
# load FP32 H from tmem, convert to BF16, store to tmem for 1st MMA,
# store to smem for TMA store later.
for i in cutlass.range_constexpr(K_dim // 32):
h_f32 = _tcgen05.ld(warp_id_ * 32, vk_tmem + i * 32, "32x32b", 32)
h_bf16 = cute.make_rmem_tensor(32, BFloat16)
h_bf16.store(h_f32.to(BFloat16))
_tcgen05.st(
warp_id_ * 32, h_tmem_base + i * 16, "32x32b", 16, h_bf16
)
# H smem layout: (V_dim, (64, K_dim/64))
dst = cute.local_tile(sH[tid_, None], (32,), (i,))
cute.copy(cp_16B, h_bf16, dst)
_tcgen05.wait_st()
_tcgen05.fence_before_thread_sync()
cute.arch.mbarrier_arrive(wh_in_mbar + stage_id)
# scale H for 2nd MMA
for i in cutlass.range_constexpr(K_dim // 32):
h_f32 = cute.make_rmem_tensor(32, Float32)
h_f32.store(
_tcgen05.ld(warp_id_ * 32, vk_tmem + i * 32, "32x32b", 32)
)
for j in cutlass.range_constexpr(32):
h_f32[j] *= h_scale
_tcgen05.st(warp_id_ * 32, vk_tmem + i * 32, "32x32b", 32, h_f32)
_tcgen05.wait_st()
_tcgen05.fence_before_thread_sync()
cute.arch.mbarrier_arrive(vk_in_mbar + stage_id)
# issue TMA store for O kernel
cute.arch.barrier(barrier_id=1, number_of_threads=128)
fence_before_tma_store()
if warp_id_ == 3:
h_dst = tmaH[chunk_offset + chunk_id, head_id, None, None]
simple_tma_copy(H_tma_atom, sH, h_dst)
with cute.arch.elect_one():
cute.arch.cp_async_bulk_commit_group()
stage_id = (stage_id + 1) % num_stages
# handle final state. reuse H0 smem.
if warp_id_ == 0:
cute.arch.mbarrier_wait(vk_done_mbar + vk_stage_id, vk_parity)
cute.arch.barrier(barrier_id=1, number_of_threads=128)
_tcgen05.fence_after_thread_sync()
for i in cutlass.range_constexpr(K_dim // 32):
h_f32 = cute.make_rmem_tensor(32, Float32)
h_f32.store(_tcgen05.ld(warp_id_ * 32, vk_tmem + i * 32, "32x32b", 32))
if cutlass.const_expr(is_f32):
cute.copy(cp_16B, h_f32, sH0[tid_, (None, i)])
else:
h_bf16 = cute.make_rmem_tensor(32, BFloat16)
h_bf16.store(h_f32.load().to(BFloat16))
sH0_dst = cute.local_tile(sH0[tid_, None], (32,), (i,))
cute.copy(cp_16B, h_bf16, sH0_dst)
cute.arch.barrier(barrier_id=1, number_of_threads=128)
if warp_id_ == 0:
ht_dst = tmaHT[seq_id, head_id, None, None]
simple_tma_copy(HT_tma_atom, sH0, ht_dst)
with cute.arch.elect_one():
cute.arch.cp_async_bulk_commit_group()
if warp_id_ == 1:
_tcgen05.dealloc()
else:
# V warps
stage_id = 0
parity = 0
chunk_offset = chunk_offsets[seq_id]
ldsm_trans_op = warp.LdMatrix8x8x16bOp(num_matrices=4, transpose=True)
stsm_trans_op = warp.StMatrix8x8x16bOp(num_matrices=4, transpose=True)
ldsm_trans_atom = cute.make_copy_atom(ldsm_trans_op, BFloat16)
stsm_trans_atom = cute.make_copy_atom(stsm_trans_op, BFloat16)
# ((BT, num_BT_tiles), V_dim)
gV_new_tiles = cute.logical_divide(
tmaV_new[None, head_id, None], (BT, None)
)
# sV shape: [BT, (64, V_dim/64), num_stages]
# sV_view shape: [BT, (8, (8,2)), num_stages]
sV_view = cute.logical_divide(sV, (None, 8, None))
sV_new_view = cute.logical_divide(sV_new, (None, 8))
# [BT, 8, num_stages]
s_col = warp_id * 4 + (lane_id // 8)
sV_view = sV_view[None, (None, s_col), None]
sV_new_view = sV_new_view[None, (None, s_col)]
for chunk_id in range(num_chunks):
# wait for V to arrive
if warp_id == 0:
cute.arch.mbarrier_wait(tma_mbar + stage_id, parity)
cute.arch.barrier(barrier_id=2, number_of_threads=128)
# unpack V BF16->FP32, then store to tmem for 1st MMA
# V smem layout: [BT, (64, V_dim/64)] / [BT, V_dim]
# each iteration, CTA loads [8, V_dim] tile
# (warp loads [8, 32] tile)
for i in cutlass.range_constexpr(BT // 8):
s_row = i * 8 + (lane_id % 8)
v_bf16 = cute.make_rmem_tensor(8, BFloat16)
cute.copy(ldsm_trans_atom, sV_view[s_row, None, stage_id], v_bf16)
v_fp32 = cvt.bf16x2_to_fp32x2(cute.recast_tensor(v_bf16, Uint32))
v_fp32 = cute.logical_divide(v_fp32, 4) # (4, 2)
tcol = wh_tmem + i * 8
_tcgen05.st(warp_id * 32 + 0, tcol, "16x256b", 1, v_fp32[None, 0])
_tcgen05.st(warp_id * 32 + 16, tcol, "16x256b", 1, v_fp32[None, 1])
_tcgen05.wait_st()
_tcgen05.fence_before_thread_sync()
cute.arch.mbarrier_arrive(wh_in_mbar + stage_id)
# load g_cu for scaling
if tid < BT:
end_t = min(bos + (chunk_id + 1) * BT, eos)
last_idx = end_t - 1
t = bos + chunk_id * BT + tid
val = Float32(0.0)
if t < eos:
val = cute.math.exp(
g_cu[last_idx, head_id] - g_cu[t, head_id],
fastmath=True,
)
s_v_scale[tid] = val
# wait for 1st MMA to finish
if warp_id == 2:
cute.arch.mbarrier_wait(wh_done_mbar + stage_id, parity)
elif warp_id == 3:
with cute.arch.elect_one():
cute.arch.cp_async_bulk_wait_group(0, read=True)
cute.arch.barrier(barrier_id=2, number_of_threads=128)
_tcgen05.fence_after_thread_sync()
for i in cutlass.range_constexpr(BT // 8):
v_new = cute.make_rmem_tensor((4, 2), Float32)
tcol = wh_tmem + i * 8
v_new[None, 0].store(
_tcgen05.ld(warp_id * 32 + 0, tcol, "16x256b", 1)
)
v_new[None, 1].store(
_tcgen05.ld(warp_id * 32 + 16, tcol, "16x256b", 1)
)
v_new_bf16 = cute.make_rmem_tensor(8, BFloat16)
v_new_bf16.store(v_new.load().to(BFloat16))
# scale V_new for 2nd MMA
scale0 = s_v_scale[i * 8 + (lane_id % 4) * 2 + 0]
scale1 = s_v_scale[i * 8 + (lane_id % 4) * 2 + 1]
v_scaled = cute.make_rmem_tensor(8, Float32)
for k in cutlass.range_constexpr(4):
v_scaled[k * 2] = v_new[k * 2] * scale0
v_scaled[k * 2 + 1] = v_new[k * 2 + 1] * scale1
v_scaled_bf16 = v_scaled.load().to(BFloat16).reshape((4, 2))
# store V_new BF16 for O kernel
s_row = i * 8 + (lane_id % 8)
cute.copy(stsm_trans_atom, v_new_bf16, sV_new_view[s_row, None])
# store to tmem
tcol = v_tmem_base + i * 4
_tcgen05.st(
warp_id * 32 + 0, tcol, "16x128b", 1, v_scaled_bf16[None, 0]
)
_tcgen05.st(
warp_id * 32 + 16, tcol, "16x128b", 1, v_scaled_bf16[None, 1]
)
_tcgen05.wait_st()
_tcgen05.fence_before_thread_sync()
cute.arch.mbarrier_arrive(vk_in_mbar + stage_id)
# issue TMA store for V_new
cute.arch.barrier(barrier_id=2, number_of_threads=128)
fence_before_tma_store()
if warp_id == 3:
gV = gV_new_tiles[(None, chunk_offset + chunk_id), None]
simple_tma_copy(V_new_tma_atom, sV_new, gV)
with cute.arch.elect_one():
cute.arch.cp_async_bulk_commit_group()
stage_id = (stage_id + 1) % num_stages
if stage_id == 0:
parity ^= 1
@cache
@staticmethod
def compile(
H: int,
Hv: int,
K_dim: int,
V_dim: int,
h_dtype: cutlass.Numeric = Float32,
BT: int = 64,
num_stages: int = 2,
):
total_t = cute.sym_int()
pad_t = cute.sym_int()
total_chunks_n = cute.sym_int()
num_sequences = cute.sym_int()
cu_entries = cute.sym_int()
K = make_fake_tensor(BFloat16, (total_t, H, K_dim), divisibility=16)
V = make_fake_tensor(BFloat16, (pad_t, Hv, V_dim), divisibility=16)
W = make_fake_tensor(BFloat16, (pad_t, Hv, K_dim), divisibility=16)
V_new = make_fake_tensor(BFloat16, (pad_t, Hv, V_dim), divisibility=16)
g_cu = make_fake_tensor(Float32, (total_t, Hv), divisibility=4)
h = make_fake_tensor(
BFloat16, (total_chunks_n, Hv, V_dim, K_dim), divisibility=16
)
h0 = make_fake_tensor(
h_dtype, (num_sequences, Hv, V_dim, K_dim), divisibility=16
)
ht = make_fake_tensor(
h_dtype, (num_sequences, Hv, V_dim, K_dim), divisibility=16
)
cu_seqlens = make_fake_tensor(Int32, (cu_entries,), divisibility=1)
chunk_offsets = make_fake_tensor(Int32, (cu_entries,), divisibility=1)
kernel = Sm100ChunkHKernel(H, Hv, K_dim, V_dim, h_dtype, BT, num_stages)
stream = cute.runtime.make_fake_stream(use_tvm_ffi_env_stream=True)
return cute.compile(
kernel,
K,
V,
W,
V_new,
g_cu,
h,
h0,
ht,
cu_seqlens,
chunk_offsets,
stream,
options="--enable-tvm-ffi",
)
def h_cutedsl(
K: torch.Tensor,
V: torch.Tensor,
W: torch.Tensor,
V_new: torch.Tensor,
g_cu: torch.Tensor,
h: torch.Tensor,
h0: torch.Tensor,
ht: torch.Tensor,
cu_seqlens: torch.Tensor,
chunk_offsets: torch.Tensor,
BT: int = 64,
num_stages: int = 2,
) -> None:
"""Compute H/V_new with the same argument order as the CUDA wrapper."""
_, H, K_dim = K.shape
_, Hv, V_dim = V.shape
h_dtype = {
torch.bfloat16: BFloat16,
torch.float32: Float32,
}[h0.dtype]
Sm100ChunkHKernel.compile(H, Hv, K_dim, V_dim, h_dtype, BT, num_stages)(
K,
V,
W,
V_new,
g_cu,
h,
h0,
ht,
cu_seqlens,
chunk_offsets,
)
h_v2b_cutedsl = h_cutedsl
@@ -1,832 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
from functools import cache
import cutlass
import torch
from cuda.bindings.driver import CUstream
from cutlass import BFloat16, Float32, Int32, Int64, Uint32, cute
from cutlass.cute.nvgpu import cpasync, warp
from quack.compile_utils import make_fake_tensor
from vllm.cute_utils import (
EVICT_FIRST,
_tcgen05,
cvt,
fence_before_tma_store,
mma_bf16,
simple_tma_copy,
)
class Sm100ChunkUWKernel:
"""Compute per-chunk KKT inverse preprocessing and U/W tiles.
Gamma[i,j] = exp(g_cu[i] - g_cu[j])
A = strictLower(beta * (K @ K.T) * Gamma)
Ai = inverse(I + A)
U = (Ai * beta) @ V
W = (Ai * beta * exp(g_cu)) @ K
"""
def __init__(
self,
H: int,
Hv: int,
K_dim: int,
V_dim: int,
num_stages: int = 2,
) -> None:
assert Hv % H == 0
assert K_dim == V_dim == 128
self.H = H
self.Hv = Hv
self.K_dim = K_dim
self.V_dim = V_dim
self.num_stages = num_stages
# hard-code
self.BT = 64
self.num_warps = 2 + 4 + 4
@cute.jit
def _make_tma_args(
self,
tensor: cute.Tensor,
dim: cutlass.Constexpr[int],
num_stages: int,
op: cpasync.TmaCopyOp,
):
# logical layout: [BT, dim]
# permute for TMA: [dim/64, BT, 64] with swizzling
swizzle_128B = cute.make_swizzle(3, 4, 3)
slayout = cute.make_layout(
(self.BT, 1, (64, dim // 64), num_stages),
stride=(64, 0, (1, self.BT * 64), self.BT * dim),
)
slayout = cute.make_composed_layout(swizzle_128B, 0, slayout)
# we need to convert gmem layout to (T, H, (64, D/64)) for make_tiled_tma_atom()
# to emit a single 4D TMA. otherwise, it will emit (D/64)x 3D TMA.
atom, tma_tensor = cpasync.make_tiled_tma_atom(
op,
cute.logical_divide(tensor, (None, None, 64)),
slayout,
cta_tiler=(self.BT, 1, dim),
)
return atom, tma_tensor, slayout
@cute.jit
def __call__(
self,
K: cute.Tensor,
V: cute.Tensor,
U: cute.Tensor,
W: cute.Tensor,
g: cute.Tensor,
beta: cute.Tensor,
g_cu: cute.Tensor,
cu_seqlens: cute.Tensor,
chunk_indices: cute.Tensor,
total_chunks: cute.Tensor,
num_sms: Int32,
stream: CUstream,
):
tma_g2s = cpasync.CopyBulkTensorTileG2SOp()
tma_s2g = cpasync.CopyBulkTensorTileS2GOp()
K_args = self._make_tma_args(K, self.K_dim, self.num_stages, tma_g2s)
V_args = self._make_tma_args(V, self.V_dim, self.num_stages, tma_g2s)
U_args = self._make_tma_args(U, self.V_dim, 1, tma_s2g)
W_args = self._make_tma_args(W, self.K_dim, 1, tma_s2g)
grid = (num_sms // self.Hv, self.Hv, 1)
block = (self.num_warps * 32, 1, 1)
self.kernel(
K_args,
V_args,
U_args,
W_args,
g,
beta,
g_cu,
cu_seqlens,
chunk_indices,
total_chunks,
).launch(grid=grid, block=block, stream=stream)
@cute.kernel
def kernel(
self,
K_args: tuple[cute.CopyAtom, cute.Tensor, cute.ComposedLayout],
V_args: tuple[cute.CopyAtom, cute.Tensor, cute.ComposedLayout],
U_args: tuple[cute.CopyAtom, cute.Tensor, cute.ComposedLayout],
W_args: tuple[cute.CopyAtom, cute.Tensor, cute.ComposedLayout],
g: cute.Tensor,
beta: cute.Tensor,
g_cu: cute.Tensor,
cu_seqlens: cute.Tensor,
chunk_indices: cute.Tensor,
total_chunks: cute.Tensor,
):
tid, _, _ = cute.arch.thread_idx()
bid, head_id, _ = cute.arch.block_idx()
grid_x, _, _ = cute.arch.grid_dim()
warp_id = cute.arch.make_warp_uniform(tid // 32)
lane_id = tid % 32
k_head_id = head_id // (self.Hv // self.H)
BT = self.BT
K_dim = self.K_dim
V_dim = self.V_dim
num_stages = self.num_stages
K_tma_atom, tmaK, sK_layout = K_args
V_tma_atom, tmaV, sV_layout = V_args
U_tma_atom, tmaU, sU_layout = U_args
W_tma_atom, tmaW, sW_layout = W_args
def allocate_tensor(smem, dtype, layout):
return smem.allocate_tensor(
dtype, layout.outer, byte_alignment=128, swizzle=layout.inner
)
smem = cutlass.utils.SmemAllocator()
sK = allocate_tensor(smem, BFloat16, sK_layout)[None, 0, None, None]
sV = allocate_tensor(smem, BFloat16, sV_layout)[None, 0, None, None]
sU = allocate_tensor(smem, BFloat16, sU_layout)[None, 0, None, 0]
sW = allocate_tensor(smem, BFloat16, sW_layout)[None, 0, None, 0]
swizzle_128B = cute.make_swizzle(3, 4, 3)
sA_layout = cute.make_layout((BT, (64, 1)), stride=(64, (1, BT * 64)))
sA_layout = cute.make_composed_layout(swizzle_128B, 0, sA_layout)
sA = allocate_tensor(smem, BFloat16, sA_layout)
sAi = allocate_tensor(smem, BFloat16, sA_layout)
s_beta = smem.allocate_array(Float32, BT)
s_g_cu_exp = smem.allocate_array(Float32, BT)
s_g_cu = smem.allocate_array(Float32, BT)
tma_mbar = smem.allocate_array(Int64, num_stages)
mma_kkt_mbar = smem.allocate_array(Int64, num_stages)
inv_mbar = smem.allocate_array(Int64, num_stages)
mma_u_mbar = smem.allocate_array(Int64, num_stages)
mma_w_mbar = smem.allocate_array(Int64, num_stages)
epi_mbar = smem.allocate_array(Int64, num_stages)
taddr = smem.allocate(Int32, 4)
kkt_tmem = 0
U_tmem_base = kkt_tmem + BT
Ab_tmem_base = U_tmem_base + V_dim * num_stages
assert Ab_tmem_base + (BT // 2) * num_stages <= 512
# prepare ldmatrix/stmatrix ops
ldsm_op = warp.LdMatrix8x8x16bOp(num_matrices=4)
stsm_op = warp.StMatrix8x8x16bOp(num_matrices=4)
ldsm_trans_op = warp.LdMatrix8x8x16bOp(num_matrices=4, transpose=True)
ldsm_atom = cute.make_copy_atom(ldsm_op, BFloat16)
stsm_atom = cute.make_copy_atom(stsm_op, BFloat16)
ldsm_trans_atom = cute.make_copy_atom(ldsm_trans_op, BFloat16)
if warp_id == 0:
with cute.arch.elect_one():
for i in cutlass.range_constexpr(num_stages):
cute.arch.mbarrier_init(tma_mbar + i, 1)
cute.arch.mbarrier_init(mma_kkt_mbar + i, 1)
cute.arch.mbarrier_init(inv_mbar + i, 128)
cute.arch.mbarrier_init(mma_u_mbar + i, 1)
cute.arch.mbarrier_init(mma_w_mbar + i, 1)
cute.arch.mbarrier_init(epi_mbar + i, 128)
cute.arch.mbarrier_init_fence()
elif warp_id == 1:
cpasync.prefetch_descriptor(K_tma_atom)
cpasync.prefetch_descriptor(V_tma_atom)
cpasync.prefetch_descriptor(U_tma_atom)
cpasync.prefetch_descriptor(W_tma_atom)
cute.arch.sync_threads()
num_global_chunks = total_chunks[0]
if warp_id == 9:
# TMA warp
stage_id = 0
parity = 1
for global_chunk_id in range(bid, num_global_chunks, grid_x):
seq_id = chunk_indices[global_chunk_id, 0]
chunk_id = chunk_indices[global_chunk_id, 1]
bos = cu_seqlens[seq_id]
# since off_t is not a multiple of BT, we need to use
# domain_offset() to shift the pointer first.
mbar = tma_mbar + stage_id
gK = cute.local_tile(
cute.domain_offset((bos, 0), tmaK[None, k_head_id, None]),
tiler=(BT, K_dim),
coord=(chunk_id, 0),
)
gV = cute.local_tile(
cute.domain_offset((bos, 0), tmaV[None, head_id, None]),
tiler=(BT, V_dim),
coord=(chunk_id, 0),
)
# when UW MMA is done, K and V TMA buffers are released
cute.arch.mbarrier_wait(mma_u_mbar + stage_id, parity)
with cute.arch.elect_one():
STAGE_SIZE = BT * (K_dim + V_dim) * 2
cute.arch.mbarrier_arrive_and_expect_tx(mbar, STAGE_SIZE)
simple_tma_copy(K_tma_atom, gK, sK[None, None, stage_id], mbar)
simple_tma_copy(
V_tma_atom, gV, sV[None, None, stage_id], mbar, EVICT_FIRST
)
stage_id = (stage_id + 1) % num_stages
if stage_id == 0:
parity ^= 1
elif warp_id == 8:
# MMA warp
_tcgen05.alloc(taddr)
stage_id = 0
parity = 0
kkt_idesc = _tcgen05.make_bf16_idesc(BT, BT)
u_idesc = _tcgen05.make_bf16_idesc(BT, V_dim, transpose_B=True)
w_idesc = _tcgen05.make_bf16_idesc(BT, K_dim, transpose_B=True)
# LBO=BT*128 is ignored for K-major
sdesc_template = _tcgen05.make_sdesc_128B_swizzle(BT * 128)
for global_chunk_id in range(bid, num_global_chunks, grid_x):
U_tmem = U_tmem_base + V_dim * stage_id
W_tmem = U_tmem | (16 << 16)
Ab_tmem = Ab_tmem_base + (BT // 2) * stage_id
Abg_tmem = Ab_tmem | (16 << 16)
##### KKT MMA: KKT = K @ K.T #####
kaddr = sK[None, None, stage_id].iterator.toint()
kdesc_base = sdesc_template | (kaddr >> 4)
# wait for TMA data to arrive
# kkt tmem is guaranteed to be free as this is issued
# after the previous kkt's consumer (inv warps)
cute.arch.mbarrier_wait(tma_mbar + stage_id, parity)
_tcgen05.fence_after_thread_sync()
with cute.arch.elect_one():
for i in cutlass.range_constexpr(K_dim // 64):
for j in cutlass.range_constexpr(64 // 16):
kdesc = kdesc_base | ((i * BT * 128 + j * 32) >> 4)
_tcgen05.mma_f16(
kkt_tmem,
kdesc,
kdesc,
kkt_idesc,
(i > 0) or (j > 0),
)
_tcgen05.commit(mma_kkt_mbar + stage_id)
##### U/W MMA: U = Ab @ V, W = Abg @ K #####
vaddr = sV[None, None, stage_id].iterator.toint()
vdesc = sdesc_template | (vaddr >> 4)
kdesc = sdesc_template | (kaddr >> 4)
# wait for epilogue to release tmem buffer
cute.arch.mbarrier_wait(epi_mbar + stage_id, parity ^ 1)
cute.arch.mbarrier_wait(inv_mbar + stage_id, parity)
_tcgen05.fence_after_thread_sync()
with cute.arch.elect_one():
for i in cutlass.range_constexpr(BT // 16):
_tcgen05.mma_ts_f16(
W_tmem, Abg_tmem + i * 8, kdesc, w_idesc, i > 0
)
kdesc += (16 * 128) >> 4
_tcgen05.commit(mma_w_mbar + stage_id)
for i in cutlass.range_constexpr(BT // 16):
_tcgen05.mma_ts_f16(
U_tmem, Ab_tmem + i * 8, vdesc, u_idesc, i > 0
)
vdesc += (16 * 128) >> 4
_tcgen05.commit(mma_u_mbar + stage_id)
stage_id = (stage_id + 1) % num_stages
if stage_id == 0:
parity ^= 1
cute.arch.mbarrier_wait(epi_mbar + stage_id, parity ^ 1)
_tcgen05.dealloc()
elif warp_id >= 4:
# inv warps
tid_ = tid % 128
warp_id_ = warp_id % 4
stage_id = 0
parity = 0
# view into (16,16) sub-tiles, then ldmatrix layout
sA_ldsm = cute.logical_divide(sA, (16, cute.make_layout((8, 2))))
sAi_ldsm = cute.logical_divide(sAi, (16, cute.make_layout((8, 2))))
sA_ldsm = sA_ldsm[(lane_id % 16, None), ((None, lane_id // 16), None)]
sAi_ldsm = sAi_ldsm[(lane_id % 16, None), ((None, lane_id // 16), None)]
# init Ai smem buffer with zeros (only the first 48 rows)
for i in cutlass.range_constexpr((BT // 4 * 3) * BT // 128):
idx = i * 128 + tid_
sAi[idx // BT, idx % BT] = BFloat16(0.0)
# indices for ldmatrix layout later
row_indices = cute.make_rmem_tensor((1, 2, 1), Int32)
row_indices[0, 0, 0] = warp_id_ * 16 + (lane_id // 4)
row_indices[0, 1, 0] = warp_id_ * 16 + (lane_id // 4) + 8
row_indices = row_indices.load()
col_indices = cute.make_rmem_tensor((2, 1, 2), Int32)
col_indices[0, 0, 0] = (lane_id % 4) * 2 + 0
col_indices[1, 0, 0] = (lane_id % 4) * 2 + 1
col_indices[0, 0, 1] = (lane_id % 4) * 2 + 8
col_indices[1, 0, 1] = (lane_id % 4) * 2 + 9
col_indices = col_indices.load()
for global_chunk_id in range(bid, num_global_chunks, grid_x):
seq_id = chunk_indices[global_chunk_id, 0]
chunk_id = chunk_indices[global_chunk_id, 1]
bos = cu_seqlens[seq_id]
eos = cu_seqlens[seq_id + 1]
off_t = bos + chunk_id * BT
t = off_t + tid_
##### Phase 1: load g and beta #####
if tid_ < BT:
in_bounds = t < eos
beta_val = beta[t, head_id] if in_bounds else Float32(0.0)
g_val = g[t, head_id] if in_bounds else Float32(0.0)
s_beta[tid_] = beta_val
# compute cumsum(g)
# parallel scan within a warp
for i in cutlass.range_constexpr(5):
offset = cutlass.const_expr(1 << i)
lower = cute.arch.shuffle_sync_up(
g_val, offset, mask_and_clamp=0
)
if lane_id >= offset:
g_val += lower
# store warp sum
if lane_id == 31:
s_g_cu[warp_id_] = g_val
cute.arch.barrier(barrier_id=3, number_of_threads=BT)
# add warp sum from lower warps
for i in cutlass.range_constexpr(1, BT // 32):
if warp_id_ >= i:
g_val += s_g_cu[i - 1]
cute.arch.barrier(barrier_id=3, number_of_threads=BT)
# store g_cu to gmem for H and O kernels
if in_bounds:
g_cu[t, head_id] = g_val
# store g and g_cu to smem for later
s_g_cu[tid_] = g_val
s_g_cu_exp[tid_] = cute.math.exp(g_val) if in_bounds else 0.0
##### Phase 2: A = strictLower(beta * kkt * Gamma) #####
if warp_id_ == 0:
cute.arch.mbarrier_wait(mma_kkt_mbar + stage_id, parity)
cute.arch.barrier(barrier_id=1, number_of_threads=128)
_tcgen05.fence_after_thread_sync()
# tmem 16x256b layout / ldmatrix layout
# mode0 is 8 rows together
# mode1 is top and bottom 8 rows
# mode2 is groups of 16 rows
row_coord = (lane_id // 4, None, warp_id_)
s_beta_view = cute.make_tensor(s_beta, (8, 2, 4))
beta_row = s_beta_view[row_coord].load().reshape((1, 2, 1))
s_g_cu_view = cute.make_tensor(s_g_cu, (8, 2, 4))
g_cu_row = s_g_cu_view[row_coord].load().reshape((1, 2, 1))
# mode0 is 2 consecutive elems
# mode1 is top and bottom 8 rows
# mode2 is next 8 columns
# mode3 is repeating that 16x16 tile pattern
kkt = _tcgen05.ld(kkt_tmem, 0, "16x256b", BT // 8)
kkt = kkt.reshape((2, 2, 2, BT // 16))
for i in cutlass.range_constexpr(BT // 16):
# mode0 is 2 elems next to each other
# mode1 is 4 pairs of elems on 1 row
# mode2 is top and bottom 8 rows
# mode3 is next 16 columns
col_coord = (None, lane_id % 4, None, i)
s_g_cu_view = cute.make_tensor(s_g_cu, (2, 4, 2, BT // 16))
g_cu_col = s_g_cu_view[col_coord].load().reshape((2, 1, 2))
Gamma = cute.math.exp(g_cu_row - g_cu_col, fastmath=True)
A = kkt[None, None, None, i] * beta_row * Gamma
# strict lower mask
# NOTE: for OOB t position, s_beta is filled with zeros.
# hence, we don't need to apply bounds check for columns.
A_masked = cute.where(row_indices > col_indices + i * 16, A, 0.0)
# pack to BF16
# CuteDSL doesn't generate cvt.bf16x2.f32 here for some reasons
packed = cute.make_rmem_tensor(4, Uint32)
packed[0] = cvt.fp32x2_to_bf16x2(
A_masked[0, 0, 0], A_masked[1, 0, 0]
)
packed[1] = cvt.fp32x2_to_bf16x2(
A_masked[0, 1, 0], A_masked[1, 1, 0]
)
packed[2] = cvt.fp32x2_to_bf16x2(
A_masked[0, 0, 1], A_masked[1, 0, 1]
)
packed[3] = cvt.fp32x2_to_bf16x2(
A_masked[0, 1, 1], A_masked[1, 1, 1]
)
# store to smem
cute.copy(
stsm_atom,
cute.recast_tensor(packed, BFloat16),
sA_ldsm[warp_id_, None, i],
)
cute.arch.barrier(barrier_id=1, number_of_threads=128)
##### Phase 3: matrix inverse #####
# we use Newton-Schulz iterations to compute the inverse
# of the four 16x16 diagonal blocks.
# Ai_new = 2 Ai - Ai @ M @ Ai
# where M = I + A
#
# we do this with 2 MMAs:
# 1. -AiM = Ai @ (-M)
# 2. Ai_new = 2 Ai + (-AiM) @ Ai
zeros_f32 = cute.make_rmem_tensor(4, Float32)
zeros_f32.fill(0.0)
def set_diagonal(A: cute.Tensor, lane_id: Int32):
"Set the diagonal to 1s"
if lane_id % 9 == 0:
A[0] = (A[0] & Uint32(0xFFFF0000)) | Uint32(0x00003F80)
A[3] = (A[3] & Uint32(0xFFFF0000)) | Uint32(0x00003F80)
elif lane_id % 9 == 4:
A[0] = (A[0] & Uint32(0x0000FFFF)) | Uint32(0x3F800000)
A[3] = (A[3] & Uint32(0x0000FFFF)) | Uint32(0x3F800000)
Ai_bf16 = cute.make_rmem_tensor(8, BFloat16)
mma_B_bf16 = cute.make_rmem_tensor(8, BFloat16)
M_bf16 = cute.make_rmem_tensor(8, BFloat16)
acc = cute.make_rmem_tensor((4, 2), Float32)
# share the same storage
Ai = cute.recast_tensor(Ai_bf16, Uint32)
mma_B = cute.logical_divide(cute.recast_tensor(mma_B_bf16, Uint32), 2)
M = cute.logical_divide(cute.recast_tensor(M_bf16, Uint32), 2)
# initial guess: Ai = I-A
cute.copy(ldsm_atom, sA_ldsm[warp_id_, None, warp_id_], Ai_bf16)
for i in cutlass.range_constexpr(4):
Ai[i] ^= Uint32(0x80008000) # negate A
set_diagonal(Ai, lane_id)
# (4, 2)
Ai_f32 = cute.logical_divide(cvt.bf16x2_to_fp32x2(Ai), 4)
# M is holding -(I+A), stay constant throughout the iterations
cute.copy(ldsm_trans_atom, sA_ldsm[warp_id_, None, warp_id_], M_bf16)
set_diagonal(M, lane_id)
for i in cutlass.range_constexpr(4):
M[i] ^= Uint32(0x80008000)
# 3 rounds of Newton-Schulz
for _ in cutlass.range_constexpr(3):
# First MMA: -AiM = Ai @ (-M)
cute.copy(stsm_atom, Ai_bf16, sA_ldsm[warp_id_, None, warp_id_])
cute.arch.sync_warp()
acc[None, 0] = mma_bf16(Ai, M[None, 0], zeros_f32)
acc[None, 1] = mma_bf16(Ai, M[None, 1], zeros_f32)
Ai_bf16.store(acc.load().to(BFloat16))
# Second MMA: Ai_new = 2Ai + (-AiM) @ Ai
for j in cutlass.range_constexpr(8):
Ai_f32[j] *= 2.0
cute.copy(
ldsm_trans_atom,
sA_ldsm[warp_id_, None, warp_id_],
mma_B_bf16,
)
Ai_f32[None, 0] = mma_bf16(Ai, mma_B[None, 0], Ai_f32[None, 0])
Ai_f32[None, 1] = mma_bf16(Ai, mma_B[None, 1], Ai_f32[None, 1])
Ai_bf16.store(Ai_f32.load().to(BFloat16))
cute.copy(stsm_atom, Ai_bf16, sAi_ldsm[warp_id_, None, warp_id_])
cute.arch.barrier(barrier_id=1, number_of_threads=128)
# off-diagonal by 1
# given
# [ Ai00 ]
# [ A10 Ai11 ]
# [ A20 A21 Ai22 ]
# [ A30 A31 A32 Ai33]
# warp1: Ai10 = -Ai11 @ A10 @ Ai00
# warp2: Ai21 = -Ai22 @ A21 @ Ai11
# warp3: Ai32 = -Ai33 @ A32 @ Ai22
if warp_id_ > 0:
neg_Ai = cute.make_rmem_tensor(4, Uint32)
for i in cutlass.range_constexpr(4):
neg_Ai[i] = Ai[i] ^ Uint32(0x80008000)
cute.copy(
ldsm_trans_atom,
sA_ldsm[warp_id_, None, warp_id_ - 1],
mma_B_bf16,
)
acc[None, 0] = mma_bf16(neg_Ai, mma_B[None, 0], zeros_f32)
acc[None, 1] = mma_bf16(neg_Ai, mma_B[None, 1], zeros_f32)
Ai_bf16.store(acc.load().to(BFloat16))
cute.copy(
ldsm_trans_atom,
sAi_ldsm[warp_id_ - 1, None, warp_id_ - 1],
mma_B_bf16,
)
acc[None, 0] = mma_bf16(Ai, mma_B[None, 0], zeros_f32)
acc[None, 1] = mma_bf16(Ai, mma_B[None, 1], zeros_f32)
Ai_bf16.store(acc.load().to(BFloat16))
cute.copy(
stsm_atom,
Ai_bf16,
sAi_ldsm[warp_id_, None, warp_id_ - 1],
)
cute.arch.barrier(barrier_id=1, number_of_threads=128)
# off-diagonal by 2
# warp0: Ai20 = -Ai22 @ (A20 @ Ai00 + A21 @ Ai10)
# warp1: Ai31 = -Ai33 @ (A31 @ Ai11 + A32 @ Ai21)
if warp_id_ < 2:
cute.copy(
ldsm_atom,
sA_ldsm[warp_id_ + 2, None, warp_id_],
Ai_bf16,
)
cute.copy(
ldsm_trans_atom,
sAi_ldsm[warp_id_, None, warp_id_],
mma_B_bf16,
)
acc[None, 0] = mma_bf16(Ai, mma_B[None, 0], zeros_f32)
acc[None, 1] = mma_bf16(Ai, mma_B[None, 1], zeros_f32)
cute.copy(
ldsm_atom,
sA_ldsm[warp_id_ + 2, None, warp_id_ + 1],
Ai_bf16,
)
cute.copy(
ldsm_trans_atom,
sAi_ldsm[warp_id_ + 1, None, warp_id_],
mma_B_bf16,
)
acc[None, 0] = mma_bf16(Ai, mma_B[None, 0], acc[None, 0])
acc[None, 1] = mma_bf16(Ai, mma_B[None, 1], acc[None, 1])
tmp = cute.make_rmem_tensor(8, BFloat16)
tmp.store(acc.load().to(BFloat16))
cute.copy(stsm_atom, tmp, sAi_ldsm[warp_id_ + 2, None, warp_id_])
cute.arch.sync_warp()
cute.copy(
ldsm_atom, sAi_ldsm[warp_id_ + 2, None, warp_id_ + 2], Ai_bf16
)
for i in cutlass.range_constexpr(4):
Ai[i] ^= Uint32(0x80008000)
cute.copy(
ldsm_trans_atom,
sAi_ldsm[warp_id_ + 2, None, warp_id_],
mma_B_bf16,
)
acc[None, 0] = mma_bf16(Ai, mma_B[None, 0], zeros_f32)
acc[None, 1] = mma_bf16(Ai, mma_B[None, 1], zeros_f32)
tmp.store(acc.load().to(BFloat16))
cute.copy(stsm_atom, tmp, sAi_ldsm[warp_id_ + 2, None, warp_id_])
cute.arch.barrier(barrier_id=1, number_of_threads=128)
# off-diagonal by 3
# warp0: Ai30 = -Ai33 @ (A30 @ Ai00 + A31 @ Ai10 + A32 @ Ai20)
if warp_id_ == 0:
cute.copy(ldsm_atom, sA_ldsm[3, None, 0], Ai_bf16)
cute.copy(ldsm_trans_atom, sAi_ldsm[0, None, 0], mma_B_bf16)
acc[None, 0] = mma_bf16(Ai, mma_B[None, 0], zeros_f32)
acc[None, 1] = mma_bf16(Ai, mma_B[None, 1], zeros_f32)
for i in cutlass.range_constexpr(1, 3):
cute.copy(ldsm_atom, sA_ldsm[3, None, i], Ai_bf16)
cute.copy(ldsm_trans_atom, sAi_ldsm[i, None, 0], mma_B_bf16)
acc[None, 0] = mma_bf16(Ai, mma_B[None, 0], acc[None, 0])
acc[None, 1] = mma_bf16(Ai, mma_B[None, 1], acc[None, 1])
tmp = cute.make_rmem_tensor(8, BFloat16)
tmp.store(acc.load().to(BFloat16))
cute.copy(stsm_atom, tmp, sAi_ldsm[3, None, 0])
cute.arch.sync_warp()
cute.copy(ldsm_atom, sAi_ldsm[3, None, 3], Ai_bf16)
for i in cutlass.range_constexpr(4):
Ai[i] ^= Uint32(0x80008000)
cute.copy(ldsm_trans_atom, sAi_ldsm[3, None, 0], mma_B_bf16)
acc[None, 0] = mma_bf16(Ai, mma_B[None, 0], zeros_f32)
acc[None, 1] = mma_bf16(Ai, mma_B[None, 1], zeros_f32)
tmp.store(acc.load().to(BFloat16))
cute.copy(stsm_atom, tmp, sAi_ldsm[3, None, 0])
##### Phase 4: compute Ab, Abg #####
if warp_id_ == 3:
cute.arch.mbarrier_wait(mma_u_mbar + stage_id, parity ^ 1)
cute.arch.barrier(barrier_id=1, number_of_threads=128)
for i in cutlass.range_constexpr(BT // 16):
cute.copy(ldsm_atom, sAi_ldsm[warp_id_, None, i], Ai_bf16)
col_coord = (None, lane_id % 4, None, i)
s_beta_view = cute.make_tensor(s_beta, (2, 4, 2, BT // 16))
beta_col = s_beta_view[col_coord].load().reshape((2, 1, 2))
s_g_cu_view = cute.make_tensor(s_g_cu_exp, (2, 4, 2, BT // 16))
g_cu_col = s_g_cu_view[col_coord].load().reshape((2, 1, 2))
Ai_f32 = cvt.bf16x2_to_fp32x2(Ai).load().reshape((2, 2, 2))
Ab_f32 = Ai_f32 * beta_col
Ab = Ab_f32.to(BFloat16)
Ab_tmem = Ab_tmem_base + (BT // 2) * stage_id + i * 8
_tcgen05.st(warp_id_ * 32, Ab_tmem, "16x128b", 2, Ab)
Abg_f32 = Ab_f32 * g_cu_col
Abg = Abg_f32.to(BFloat16)
_tcgen05.st(warp_id_ * 32 + 16, Ab_tmem, "16x128b", 2, Abg)
_tcgen05.wait_st()
_tcgen05.fence_before_thread_sync()
cute.arch.mbarrier_arrive(inv_mbar + stage_id)
stage_id = (stage_id + 1) % num_stages
if stage_id == 0:
parity ^= 1
elif warp_id < 4:
# epi warps
stage_id = 0
parity = 0
# ((BT, num_global_chunks), V_dim)
gU_tiles = cute.logical_divide(tmaU[None, head_id, None], (BT, None))
gW_tiles = cute.logical_divide(tmaW[None, head_id, None], (BT, None))
# sW shape: [BT, (64, K_dim/64)]
# sW_view shape: [(8, 2), (4, K_dim/64)]
s_row = warp_id * 16 + lane_id % 16 # select the rows of [16,16] tile
sW_view = cute.zipped_divide(
sW[s_row, None],
tiler=cute.make_layout((8, 2)),
)
sU_view = cute.zipped_divide(
sU[s_row, None],
tiler=cute.make_layout((8, 2)),
)
# select the 8 columns within [16,16] tile
sW_view = sW_view[(None, lane_id // 16), None]
sU_view = sU_view[(None, lane_id // 16), None]
for global_chunk_id in range(bid, num_global_chunks, grid_x):
# wait for W MMA + previous TMA store to finish
U_tmem = U_tmem_base + V_dim * stage_id
if warp_id == 0:
cute.arch.mbarrier_wait(mma_w_mbar + stage_id, parity)
elif warp_id == 1:
with cute.arch.elect_one():
cute.arch.cp_async_bulk_wait_group(0, read=True)
cute.arch.barrier(barrier_id=2, number_of_threads=128)
_tcgen05.fence_after_thread_sync()
w_f32 = _tcgen05.ld(warp_id * 32 + 16, U_tmem, "16x256b", K_dim // 8)
_tcgen05.wait_ld()
w_bf16 = cute.make_rmem_tensor((8, K_dim // 16), BFloat16)
w_bf16.store(w_f32.to(BFloat16))
cute.copy(stsm_atom, w_bf16, sW_view)
# wait for U MMA + issue W TMA store
cute.arch.barrier(barrier_id=2, number_of_threads=128)
fence_before_tma_store()
if warp_id == 0:
cute.arch.mbarrier_wait(mma_u_mbar + stage_id, parity)
elif warp_id == 1:
# don't need to commit
simple_tma_copy(
W_tma_atom, sW, gW_tiles[(None, global_chunk_id), None]
)
cute.arch.barrier(barrier_id=2, number_of_threads=128)
_tcgen05.fence_after_thread_sync()
u_f32 = _tcgen05.ld(warp_id * 32, U_tmem, "16x256b", V_dim // 8)
_tcgen05.wait_ld()
_tcgen05.fence_before_thread_sync()
cute.arch.mbarrier_arrive(epi_mbar + stage_id)
u_bf16 = cute.make_rmem_tensor((8, V_dim // 16), BFloat16)
u_bf16.store(u_f32.to(BFloat16))
cute.copy(stsm_atom, u_bf16, sU_view)
cute.arch.barrier(barrier_id=2, number_of_threads=128)
fence_before_tma_store()
if warp_id == 1:
simple_tma_copy(
U_tma_atom, sU, gU_tiles[(None, global_chunk_id), None]
)
with cute.arch.elect_one():
cute.arch.cp_async_bulk_commit_group()
stage_id = (stage_id + 1) % num_stages
if stage_id == 0:
parity ^= 1
@cache
@staticmethod
def compile(H: int, Hv: int, K_dim: int, V_dim: int, num_stages: int = 2):
total_t = cute.sym_int()
pad_t = cute.sym_int()
total_chunks_n = cute.sym_int()
num_sequences = cute.sym_int()
K = make_fake_tensor(BFloat16, (total_t, H, K_dim), divisibility=16)
V = make_fake_tensor(BFloat16, (total_t, Hv, V_dim), divisibility=16)
U = make_fake_tensor(BFloat16, (pad_t, Hv, V_dim), divisibility=16)
W = make_fake_tensor(BFloat16, (pad_t, Hv, K_dim), divisibility=16)
g = make_fake_tensor(Float32, (total_t, Hv), divisibility=4)
beta = make_fake_tensor(Float32, (total_t, Hv), divisibility=4)
g_cu = make_fake_tensor(Float32, (total_t, Hv), divisibility=4)
cu_seqlens = make_fake_tensor(Int32, (num_sequences,), divisibility=1)
chunk_indices = make_fake_tensor(Int32, (total_chunks_n, 2), divisibility=2)
total_chunks = make_fake_tensor(Int32, (1,), divisibility=1)
kernel = Sm100ChunkUWKernel(H, Hv, K_dim, V_dim, num_stages)
stream = cute.runtime.make_fake_stream(use_tvm_ffi_env_stream=True)
return cute.compile(
kernel,
K,
V,
U,
W,
g,
beta,
g_cu,
cu_seqlens,
chunk_indices,
total_chunks,
Int32(148),
stream,
options="--enable-tvm-ffi",
)
def kkt_inv_uw_cutedsl(
K: torch.Tensor,
V: torch.Tensor,
U: torch.Tensor,
W: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
g_cu: torch.Tensor,
cu_seqlens: torch.Tensor,
chunk_indices: torch.Tensor,
total_chunks: torch.Tensor,
num_sms: int = 148,
) -> None:
_, Hv, V_dim = V.shape
_, H, K_dim = K.shape
Sm100ChunkUWKernel.compile(H, Hv, K_dim, V_dim)(
K,
V,
U,
W,
g,
beta,
g_cu,
cu_seqlens,
chunk_indices,
total_chunks,
num_sms,
)
@@ -1,630 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
from functools import cache
import cutlass
import torch
from cuda.bindings.driver import CUstream
from cutlass import BFloat16, Float32, Int32, Int64, Uint32, cute
from cutlass.cute.nvgpu import cpasync, warp
from quack.compile_utils import make_fake_tensor
from vllm.cute_utils import (
EVICT_FIRST,
_tcgen05,
cvt,
fence_before_tma_store,
simple_tma_copy,
)
class Sm100ChunkOKernel:
"""Compute per-token output from recurrent and intra-chunk terms.
Gamma[i,j] = exp(g_cu[i] - g_cu[j])
P = mask((Q @ K.T) * Gamma)
O = scale * (exp(g_cu) * (Q @ H.T) + P @ V)
"""
def __init__(
self,
H: int,
Hv: int,
K_dim: int,
V_dim: int,
BT: int = 64,
num_stages: int = 2,
) -> None:
assert Hv % H == 0
assert K_dim == 128
assert V_dim == 128
assert BT == 64
self.H = H
self.Hv = Hv
self.K_dim = K_dim
self.V_dim = V_dim
self.BT = BT
self.num_stages = num_stages
self.num_warps = 10
@cute.jit
def _make_bf16_tma_args(
self,
tensor: cute.Tensor,
dim: cutlass.Constexpr[int],
op: cpasync.TmaCopyOp,
stages: cutlass.Constexpr[int],
):
swizzle_128B = cute.make_swizzle(3, 4, 3)
slayout = cute.make_layout(
(self.BT, 1, (64, dim // 64), stages),
stride=(64, 0, (1, self.BT * 64), self.BT * dim),
)
slayout = cute.make_composed_layout(swizzle_128B, 0, slayout)
atom, tma_tensor = cpasync.make_tiled_tma_atom(
op,
cute.logical_divide(tensor, (None, None, 64)),
slayout,
cta_tiler=(self.BT, 1, dim),
)
return atom, tma_tensor, slayout
@cute.jit
def _make_h_tma_args(
self,
tensor: cute.Tensor,
op: cpasync.TmaCopyOp,
stages: cutlass.Constexpr[int],
):
num_elems = 128 // (tensor.element_type.width // 8)
swizzle_128B = cute.make_swizzle(3, 4, 3)
slayout = cute.make_layout(
(1, self.V_dim, (num_elems, self.K_dim // num_elems), stages),
stride=(0, num_elems, (1, self.V_dim * num_elems), self.V_dim * self.K_dim),
)
slayout = cute.make_composed_layout(swizzle_128B, 0, slayout)
atom, tma_tensor = cpasync.make_tiled_tma_atom(
op,
cute.logical_divide(tensor, (None, None, num_elems)),
slayout,
cta_tiler=(1, self.V_dim, self.K_dim),
)
return atom, tma_tensor, slayout
@cute.jit
def __call__(
self,
q: cute.Tensor,
k: cute.Tensor,
v_new_chunks: cute.Tensor,
h: cute.Tensor,
g_cu: cute.Tensor,
o: cute.Tensor,
cu_seqlens: cute.Tensor,
chunk_indices: cute.Tensor,
total_chunks: cute.Tensor,
scale: Float32,
num_sms: Int32,
stream: CUstream,
):
grid = (num_sms // self.Hv, self.Hv, 1)
block = (self.num_warps * 32, 1, 1)
tma_g2s = cpasync.CopyBulkTensorTileG2SOp()
tma_s2g = cpasync.CopyBulkTensorTileS2GOp()
Q_args = self._make_bf16_tma_args(q, self.K_dim, tma_g2s, self.num_stages)
K_args = self._make_bf16_tma_args(k, self.K_dim, tma_g2s, self.num_stages)
V_args = self._make_bf16_tma_args(
v_new_chunks, self.V_dim, tma_g2s, self.num_stages
)
H_args = self._make_h_tma_args(h, tma_g2s, self.num_stages)
O_args = self._make_bf16_tma_args(o, self.V_dim, tma_s2g, 1)
self.kernel(
Q_args,
K_args,
V_args,
H_args,
O_args,
g_cu,
o,
cu_seqlens,
chunk_indices,
total_chunks,
scale,
).launch(grid=grid, block=block, stream=stream)
@cute.kernel
def kernel(
self,
Q_args: tuple[cute.CopyAtom, cute.Tensor, cute.ComposedLayout],
K_args: tuple[cute.CopyAtom, cute.Tensor, cute.ComposedLayout],
V_args: tuple[cute.CopyAtom, cute.Tensor, cute.ComposedLayout],
H_args: tuple[cute.CopyAtom, cute.Tensor, cute.ComposedLayout],
O_args: tuple[cute.CopyAtom, cute.Tensor, cute.ComposedLayout],
g_cu: cute.Tensor,
o: cute.Tensor,
cu_seqlens: cute.Tensor,
chunk_indices: cute.Tensor,
total_chunks: cute.Tensor,
scale: Float32,
):
tid, _, _ = cute.arch.thread_idx()
bid, v_head_id, _ = cute.arch.block_idx()
grid_x, _, _ = cute.arch.grid_dim()
warp_id = cute.arch.make_warp_uniform(tid // 32)
lane_id = tid % 32
BT = self.BT
K_dim = self.K_dim
V_dim = self.V_dim
num_stages = self.num_stages
heads_per_qk = self.Hv // self.H
k_head_id = v_head_id // heads_per_qk
num_global_chunks = total_chunks[0]
Q_tma_atom, tmaQ, sQ_layout = Q_args
K_tma_atom, tmaK, sK_layout = K_args
V_tma_atom, tmaV, sV_layout = V_args
H_tma_atom, tmaH, sH_layout = H_args
O_tma_atom, tmaO, sO_layout = O_args
def allocate_tensor(smem, dtype, layout):
return smem.allocate_tensor(
dtype, layout.outer, byte_alignment=128, swizzle=layout.inner
)
smem = cutlass.utils.SmemAllocator()
sQ = allocate_tensor(smem, BFloat16, sQ_layout)[None, 0, None, None]
sK = allocate_tensor(smem, BFloat16, sK_layout)[None, 0, None, None]
sV = allocate_tensor(smem, BFloat16, sV_layout)[None, 0, None, None]
sH = allocate_tensor(smem, BFloat16, sH_layout)[0, None, None, None]
sO = allocate_tensor(smem, BFloat16, sO_layout)[None, 0, None, 0]
s_g_cu = smem.allocate_array(Float32, BT)
qk_full_mbar = smem.allocate_array(Int64, num_stages)
hv_full_mbar = smem.allocate_array(Int64, num_stages)
qk_empty_mbar = smem.allocate_array(Int64, num_stages)
pv_mma_mbar = smem.allocate_array(Int64, num_stages)
qk_mbar = smem.allocate_array(Int64, 1)
mask_mbar = smem.allocate_array(Int64, 1)
epi_mbar = smem.allocate_array(Int64, 1)
taddr = smem.allocate(Int32, 4)
qk_tmem = 0
p_tmem = 64
out_tmem = 128
qh_tmem = 256
if warp_id == 0:
with cute.arch.elect_one():
for i in cutlass.range_constexpr(num_stages):
cute.arch.mbarrier_init(qk_full_mbar + i, 1)
cute.arch.mbarrier_init(qk_empty_mbar + i, 1)
cute.arch.mbarrier_init(hv_full_mbar + i, 1)
cute.arch.mbarrier_init(pv_mma_mbar + i, 1)
cute.arch.mbarrier_init(qk_mbar, 1)
cute.arch.mbarrier_init(mask_mbar, 128)
cute.arch.mbarrier_init(epi_mbar, 128)
cute.arch.mbarrier_init_fence()
elif warp_id == 9:
cpasync.prefetch_descriptor(Q_tma_atom)
cpasync.prefetch_descriptor(K_tma_atom)
cpasync.prefetch_descriptor(V_tma_atom)
cpasync.prefetch_descriptor(H_tma_atom)
cute.arch.sync_threads()
if warp_id == 9:
# TMA warp
stage_id = 0
parity = 1
for global_chunk_id in range(bid, num_global_chunks, grid_x):
seq_id = chunk_indices[global_chunk_id, 0]
chunk_id = chunk_indices[global_chunk_id, 1]
bos = cu_seqlens[seq_id]
# copy Q and K
q_tile = cute.local_tile(
cute.domain_offset((bos, 0), tmaQ[None, k_head_id, None]),
tiler=(BT, K_dim),
coord=(chunk_id, 0),
)
k_tile = cute.local_tile(
cute.domain_offset((bos, 0), tmaK[None, k_head_id, None]),
tiler=(BT, K_dim),
coord=(chunk_id, 0),
)
mbar = qk_full_mbar + stage_id
cute.arch.mbarrier_wait(qk_empty_mbar + stage_id, parity)
with cute.arch.elect_one():
STAGE_SIZE = BT * (K_dim + K_dim) * 2
cute.arch.mbarrier_arrive_and_expect_tx(mbar, STAGE_SIZE)
simple_tma_copy(Q_tma_atom, q_tile, sQ[None, None, stage_id], mbar)
simple_tma_copy(K_tma_atom, k_tile, sK[None, None, stage_id], mbar)
# copy H and V
gH = tmaH[global_chunk_id * self.Hv + v_head_id, None, None]
gV = cute.local_tile(
tmaV[None, v_head_id, None],
tiler=(BT, V_dim),
coord=(global_chunk_id, 0),
)
mbar = hv_full_mbar + stage_id
cute.arch.mbarrier_wait(pv_mma_mbar + stage_id, parity)
with cute.arch.elect_one():
H_STAGE_SIZE = V_dim * K_dim * 2
V_STAGE_SIZE = BT * V_dim * 2
cute.arch.mbarrier_arrive_and_expect_tx(
mbar, H_STAGE_SIZE + V_STAGE_SIZE
)
simple_tma_copy(
H_tma_atom, gH, sH[None, None, stage_id], mbar, EVICT_FIRST
)
simple_tma_copy(
V_tma_atom, gV, sV[None, None, stage_id], mbar, EVICT_FIRST
)
stage_id = (stage_id + 1) % num_stages
if stage_id == 0:
parity ^= 1
elif warp_id == 8:
# MMA warp
_tcgen05.alloc(taddr)
# LBO=BT*128 is ignored for K-major
sdesc_template = _tcgen05.make_sdesc_128B_swizzle(BT * 128)
qk_idesc = _tcgen05.make_bf16_idesc(BT, BT)
qh_idesc = _tcgen05.make_bf16_idesc(BT, V_dim)
pv_idesc = _tcgen05.make_bf16_idesc(BT, V_dim, transpose_B=True)
stage_id = 0
tma_parity = 0
mask_parity = 0
for global_chunk_id in range(bid, num_global_chunks, grid_x):
qaddr = sQ[None, None, stage_id].iterator.toint()
kaddr = sK[None, None, stage_id].iterator.toint()
haddr = sH[None, None, stage_id].iterator.toint()
vaddr = sV[None, None, stage_id].iterator.toint()
qdesc_base = sdesc_template | (qaddr >> 4)
kdesc_base = sdesc_template | (kaddr >> 4)
hdesc_base = sdesc_template | (haddr >> 4)
vdesc_base = sdesc_template | (vaddr >> 4)
##### 1st MMA: Q @ K.T #####
# do this first to unblock mask(QK)
cute.arch.mbarrier_wait(epi_mbar, mask_parity ^ 1)
cute.arch.mbarrier_wait(qk_full_mbar + stage_id, tma_parity)
_tcgen05.fence_after_thread_sync()
with cute.arch.elect_one():
for i in cutlass.range_constexpr(K_dim // BT):
for j in cutlass.range_constexpr(BT // 16):
qdesc = qdesc_base | ((i * BT * 128 + j * 32) >> 4)
kdesc = kdesc_base | ((i * BT * 128 + j * 32) >> 4)
_tcgen05.mma_f16(
qk_tmem, qdesc, kdesc, qk_idesc, (i > 0) or (j > 0)
)
_tcgen05.commit(qk_mbar)
##### 2nd MMA: Q @ H.T #####
cute.arch.mbarrier_wait(hv_full_mbar + stage_id, tma_parity)
_tcgen05.fence_after_thread_sync()
with cute.arch.elect_one():
for i in cutlass.range_constexpr(K_dim // BT):
for j in cutlass.range_constexpr(BT // 16):
qdesc = qdesc_base | ((i * BT * 128 + j * 32) >> 4)
hdesc = hdesc_base | ((i * V_dim * 128 + j * 32) >> 4)
_tcgen05.mma_f16(
qh_tmem, qdesc, hdesc, qh_idesc, (i > 0) or (j > 0)
)
_tcgen05.commit(qk_empty_mbar + stage_id)
##### 3rd MMA: P @ V #####
# stalled by mask(QK)
cute.arch.mbarrier_wait(mask_mbar, mask_parity)
_tcgen05.fence_after_thread_sync()
with cute.arch.elect_one():
for i in cutlass.range_constexpr(BT // 16):
vdesc = vdesc_base | ((i * 16 * 128) >> 4)
_tcgen05.mma_ts_f16(
out_tmem, p_tmem + i * 8, vdesc, pv_idesc, i > 0
)
_tcgen05.commit(pv_mma_mbar + stage_id)
stage_id = (stage_id + 1) % num_stages
if stage_id == 0:
tma_parity ^= 1
mask_parity ^= 1
# wait for epilogue to finish for deallocation
cute.arch.mbarrier_wait(epi_mbar, mask_parity ^ 1)
_tcgen05.dealloc()
elif warp_id >= 4:
# masking warps
warp_id_ = warp_id % 4
tid_ = tid % 128
row0 = warp_id_ * 16 + lane_id // 4
row1 = row0 + 8
parity = 0
# for ldmatrix layout later
row_indices = cute.make_rmem_tensor(2, Int32)
row_indices[0] = warp_id_ * 16 + lane_id // 4
row_indices[1] = warp_id_ * 16 + lane_id // 4 + 8
row_indices = row_indices.load().reshape((1, 2))
col_indices = cute.make_rmem_tensor(2, Int32)
col_indices[0] = (lane_id % 4) * 2
col_indices[1] = (lane_id % 4) * 2 + 1
col_indices = col_indices.load().reshape((2, 1))
for global_chunk_id in range(bid, num_global_chunks, grid_x):
if tid_ < BT:
seq_id = chunk_indices[global_chunk_id, 0]
chunk_id = chunk_indices[global_chunk_id, 1]
bos = cu_seqlens[seq_id]
eos = cu_seqlens[seq_id + 1]
t_ = bos + chunk_id * BT + tid_
s_g_cu[tid_] = g_cu[t_, v_head_id] if t_ < eos else Float32(0.0)
# wait for QK MMA
if warp_id_ == 0:
cute.arch.mbarrier_wait(qk_mbar, parity)
cute.arch.barrier(barrier_id=1, number_of_threads=128)
_tcgen05.fence_after_thread_sync()
qk = _tcgen05.ld(warp_id_ * 32, qk_tmem, "16x256b", BT // 8)
qk = qk.reshape((2, 2, BT // 8))
_tcgen05.wait_ld()
g_cu_rows = cute.make_rmem_tensor(2, Float32)
g_cu_rows[0] = s_g_cu[row0]
g_cu_rows[1] = s_g_cu[row1]
g_cu_rows = g_cu_rows.load().reshape((1, 2))
for i in cutlass.range_constexpr(BT // 8):
col = i * 8 + (lane_id % 4) * 2
g_cu_cols = cute.make_rmem_tensor(2, Float32)
g_cu_cols[0] = s_g_cu[col]
g_cu_cols[1] = s_g_cu[col + 1]
g_cu_cols = g_cu_cols.load().reshape((2, 1))
# apply gamma and causal mask
Gamma = cute.math.exp(g_cu_rows - g_cu_cols, fastmath=True)
tmp = qk[None, None, i] * Gamma
tmp = cute.where(row_indices >= col_indices + i * 8, tmp, 0.0)
# CuteDSL can't emit cvt.bf16x2.f32 here
attn_lo = cute.make_rmem_tensor(2, Uint32)
attn_lo[0] = cvt.fp32x2_to_bf16x2(tmp[0, 0], tmp[1, 0])
attn_lo[1] = cvt.fp32x2_to_bf16x2(tmp[0, 1], tmp[1, 1])
_tcgen05.st(warp_id_ * 32, p_tmem + i * 4, "16x128b", 1, attn_lo)
_tcgen05.wait_st()
_tcgen05.fence_before_thread_sync()
cute.arch.mbarrier_arrive(mask_mbar)
parity ^= 1
else:
# epilogue warps
# for ldmatrix layout later
row0 = warp_id * 16 + lane_id // 4
row1 = row0 + 8
stage_id = 0
mma_parity = 0
op = cute.nvgpu.CopyUniversalOp()
cp_4B = cute.make_copy_atom(op, BFloat16, num_bits_per_copy=32)
stsm_op = warp.StMatrix8x8x16bOp(num_matrices=4, transpose=False)
stsm_atom = cute.make_copy_atom(stsm_op, BFloat16)
# ldmatrix layout
# [total_seq_len, ((2, 4, WIDTH/8), V_DIM/WIDTH)]
WIDTH = 64
o_view = cute.logical_divide(
o[None, v_head_id, None],
(None, cute.make_layout((2, 4, WIDTH // 8))),
)
# select lane: [total_seq_len, 2, WIDTH/8, V_DIM/WIDTH]
o_view = o_view[None, ((None, lane_id % 4, None), None)]
for global_chunk_id in range(bid, num_global_chunks, grid_x):
seq_id = chunk_indices[global_chunk_id, 0]
chunk_id = chunk_indices[global_chunk_id, 1]
bos = cu_seqlens[seq_id]
eos = cu_seqlens[seq_id + 1]
chunk_start = bos + chunk_id * BT
full_chunk = chunk_start + BT <= eos
g_cu_rows = cute.make_rmem_tensor(2, Float32)
g_cu_rows.fill(0.0)
# load g_cu
if chunk_start + row0 < eos:
g_cu_rows[0] = cute.math.exp(
g_cu[chunk_start + row0, v_head_id], fastmath=True
)
if chunk_start + row1 < eos:
g_cu_rows[1] = cute.math.exp(
g_cu[chunk_start + row1, v_head_id], fastmath=True
)
g_cu_rows = g_cu_rows.load().reshape((1, 2, 1))
if warp_id == 0:
cute.arch.mbarrier_wait(pv_mma_mbar + stage_id, mma_parity)
elif warp_id == 3 and full_chunk:
cute.arch.cp_async_bulk_wait_group(0, read=True)
cute.arch.barrier(barrier_id=2, number_of_threads=128)
_tcgen05.fence_after_thread_sync()
if full_chunk:
# use TMA store: tmem->rmem->smem->gmem
for i in cutlass.range_constexpr(V_dim // WIDTH):
qh = _tcgen05.ld(
warp_id * 32, qh_tmem + i * WIDTH, "16x256b", WIDTH // 8
)
pv = _tcgen05.ld(
warp_id * 32, out_tmem + i * WIDTH, "16x256b", WIDTH // 8
)
_tcgen05.wait_ld()
if i == V_dim // WIDTH - 1:
_tcgen05.fence_before_thread_sync()
cute.arch.mbarrier_arrive(epi_mbar)
qh = qh.reshape((2, 2, WIDTH // 8))
pv = pv.reshape((2, 2, WIDTH // 8))
out_f32 = scale * (g_cu_rows * qh + pv)
out_bf16 = cute.make_rmem_tensor((8, WIDTH // 16), BFloat16)
out_bf16.store(out_f32.to(BFloat16).reshape((8, WIDTH // 16)))
# TODO: issue single cute.copy()
for j in cutlass.range_constexpr(WIDTH // 16):
s_row = warp_id * 16 + lane_id % 16
s_col = i * (WIDTH // 8) + j * 2 + lane_id // 16
sO_tile = cute.local_tile(sO[s_row, None], (8,), (s_col,))
cute.copy(stsm_atom, out_bf16[None, j], sO_tile)
cute.arch.barrier(barrier_id=2, number_of_threads=128)
fence_before_tma_store()
if warp_id == 3:
gO = cute.local_tile(
cute.domain_offset((bos, 0), tmaO[None, v_head_id, None]),
tiler=(BT, V_dim),
coord=(chunk_id, 0),
)
simple_tma_copy(O_tma_atom, sO, gO)
with cute.arch.elect_one():
cute.arch.cp_async_bulk_commit_group()
else:
# direct gmem store
# TODO: explore doing multiple 1D TMAs
for i in cutlass.range_constexpr(V_dim // WIDTH):
qh = _tcgen05.ld(
warp_id * 32, qh_tmem + i * WIDTH, "16x256b", WIDTH // 8
)
pv = _tcgen05.ld(
warp_id * 32, out_tmem + i * WIDTH, "16x256b", WIDTH // 8
)
_tcgen05.wait_ld()
if i == V_dim // WIDTH - 1:
_tcgen05.fence_before_thread_sync()
cute.arch.mbarrier_arrive(epi_mbar)
qh = qh.reshape((2, 2, WIDTH // 8))
pv = pv.reshape((2, 2, WIDTH // 8))
out_f32 = scale * (g_cu_rows * qh + pv)
out_bf16 = cute.make_rmem_tensor((2, 2, WIDTH // 8), BFloat16)
out_bf16.store(out_f32.to(BFloat16))
if chunk_start + row0 < eos:
cute.copy(
cp_4B,
out_bf16[None, 0, None],
o_view[chunk_start + row0, None, None, i],
)
if chunk_start + row1 < eos:
cute.copy(
cp_4B,
out_bf16[None, 1, None],
o_view[chunk_start + row1, None, None, i],
)
stage_id = (stage_id + 1) % num_stages
if stage_id == 0:
mma_parity ^= 1
@cache
@staticmethod
def compile(
H: int,
Hv: int,
K_dim: int,
V_dim: int,
BT: int = 64,
num_stages: int = 2,
):
total_t = cute.sym_int()
pad_t = cute.sym_int()
total_chunks_n = cute.sym_int()
h_outer_n = cute.sym_int()
cu_entries = cute.sym_int()
q = make_fake_tensor(BFloat16, (total_t, H, K_dim), divisibility=16)
k = make_fake_tensor(BFloat16, (total_t, H, K_dim), divisibility=16)
v_new = make_fake_tensor(BFloat16, (pad_t, Hv, V_dim), divisibility=16)
h_flat = make_fake_tensor(BFloat16, (h_outer_n, V_dim, K_dim), divisibility=16)
g_cu = make_fake_tensor(Float32, (total_t, Hv), divisibility=4)
o = make_fake_tensor(BFloat16, (total_t, Hv, V_dim), divisibility=16)
cu_seqlens = make_fake_tensor(Int32, (cu_entries,), divisibility=1)
chunk_indices = make_fake_tensor(Int32, (total_chunks_n, 2), divisibility=2)
total_chunks = make_fake_tensor(Int32, (1,), divisibility=1)
kernel = Sm100ChunkOKernel(
H,
Hv,
K_dim,
V_dim,
BT,
num_stages,
)
stream = cute.runtime.make_fake_stream(use_tvm_ffi_env_stream=True)
return cute.compile(
kernel,
q,
k,
v_new,
h_flat,
g_cu,
o,
cu_seqlens,
chunk_indices,
total_chunks,
Float32(1.0),
Int32(148),
stream,
options="--enable-tvm-ffi",
)
def o_cutedsl(
q: torch.Tensor,
k: torch.Tensor,
v_new_chunks: torch.Tensor,
h: torch.Tensor,
g_cu: torch.Tensor,
o: torch.Tensor,
cu_seqlens: torch.Tensor,
chunk_indices: torch.Tensor,
total_chunks: torch.Tensor,
scale: float,
num_sms: int = 148,
) -> None:
_, H, K_dim = q.shape
_, Hv, V_dim = o.shape
Sm100ChunkOKernel.compile(H, Hv, K_dim, V_dim)(
q,
k,
v_new_chunks.view(-1, Hv, V_dim),
h.view(-1, V_dim, K_dim),
g_cu,
o,
cu_seqlens,
chunk_indices,
total_chunks,
float(scale),
num_sms,
)
+1 -39
View File
@@ -196,42 +196,4 @@ class BOSEOSFilter(Pooler):
return pooled_outputs
class BgeM3Pooler(Pooler):
def __init__(self, token_classify_pooler: Pooler, embed_pooler: Pooler) -> None:
super().__init__()
self.token_classify_pooler = token_classify_pooler
self.embed_pooler = embed_pooler
def forward(
self, hidden_states: torch.Tensor, pooling_metadata: PoolingMetadata
) -> PoolerOutput:
embed_outputs = self.embed_pooler(hidden_states, pooling_metadata)
token_classify_outputs = self.token_classify_pooler(
hidden_states, pooling_metadata
)
pooler_outputs: list[torch.Tensor] = []
for embed_output, token_classify_output in zip(
embed_outputs, token_classify_outputs
):
pooler_outputs.append(
torch.cat(
[embed_output.view(-1), token_classify_output.view(-1)], dim=-1
)
)
return pooler_outputs
def get_supported_tasks(self) -> Set[PoolingTask]:
return {"embed&token_classify"}
def get_pooling_updates(self, task: PoolingTask) -> PoolingParamsUpdate:
return self.embed_pooler.get_pooling_updates(
"embed"
) | self.token_classify_pooler.get_pooling_updates("token_classify")
def extra_repr(self) -> str:
s = f"supported_task={self.get_supported_tasks()}"
return s
__all__ = ["BOSEOSFilter", "DispatchPooler", "IdentityPooler", "BgeM3Pooler"]
__all__ = ["BOSEOSFilter", "DispatchPooler", "IdentityPooler"]
@@ -84,7 +84,7 @@ def register_quantization_config(quantization: str):
def _wrapper(quant_config_cls):
if quantization in QUANTIZATION_METHODS:
logger.debug(
logger.warning(
"The quantization method '%s' already exists and will be "
"overwritten by the quantization config %s.",
quantization,
@@ -407,12 +407,6 @@ class Eagle3DeepseekV2ForCausalLM(DeepseekV2ForCausalLM):
hidden_states: torch.Tensor,
) -> torch.Tensor:
# Combine multiple auxiliary hidden states returned by Eagle3
if self.model.fc_norm is not None:
chunks = hidden_states.chunk(self.model.num_aux_hidden_states, dim=-1)
hidden_states = torch.cat(
[norm(chunk) for norm, chunk in zip(self.model.fc_norm, chunks)],
dim=-1,
)
return self.model.fc(hidden_states)
def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]):
+2 -7
View File
@@ -339,7 +339,7 @@ class GPT2ForSequenceClassification(nn.Module, SupportsCrossEncoding):
super().__init__()
config = vllm_config.model_config.hf_config
self.transformer = GPT2Model(
vllm_config=vllm_config, prefix=maybe_prefix(prefix, "transformer")
vllm_config=vllm_config, prefix=maybe_prefix(prefix, "gpt2")
)
self.score = nn.Linear(
config.n_embd,
@@ -358,7 +358,6 @@ class GPT2ForSequenceClassification(nn.Module, SupportsCrossEncoding):
def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]):
loader = AutoWeightsLoader(self)
weights = _add_transformer_prefix(weights)
return loader.load_weights(weights)
def forward(
@@ -381,10 +380,6 @@ def _add_transformer_prefix(
weights: Iterable[tuple[str, torch.Tensor]],
) -> Iterable[tuple[str, torch.Tensor]]:
for name, tensor in weights:
if (
not name.startswith("transformer.")
and not name.startswith("lm_head.")
and not name.startswith("score.")
):
if not name.startswith("transformer.") and not name.startswith("lm_head"):
name = "transformer." + name
yield name, tensor
+1 -3
View File
@@ -378,9 +378,7 @@ class Resampler4_5(Resampler2_5):
) # D
pos_embed_2d.append(
self.pos_embed[:tgt_h, :tgt_w, :]
.reshape((tgt_h * tgt_w, -1))
.to(device=device, dtype=dtype)
self.pos_embed[:tgt_h, :tgt_w, :].reshape((tgt_h * tgt_w, -1)).to(dtype)
) # patches * D
key_padding_mask[i, patch_len[i] :] = True
+1 -1
View File
@@ -980,7 +980,7 @@ class _ModelRegistry:
raise TypeError(msg)
if model_arch in self.models:
logger.debug(
logger.warning(
"Model architecture %s is already registered, and will be "
"overwritten by the new model class %s.",
model_arch,
-4
View File
@@ -10,7 +10,6 @@ from transformers import RobertaConfig
from vllm.config import ModelConfig, PoolerConfig, VllmConfig
from vllm.model_executor.layers.pooler import (
BgeM3Pooler,
BOSEOSFilter,
DispatchPooler,
Pooler,
@@ -238,9 +237,6 @@ class BgeM3EmbeddingModel(RobertaEmbeddingModel):
# for some reason m3 only filters the bos for colbert vectors
),
"token_classify": token_classify_pooler,
"embed&token_classify": BgeM3Pooler(
token_classify_pooler, embed_pooler
),
}
)
@@ -11,12 +11,6 @@ Three specialized kernels:
- _fused_kv_compress_norm_rope_insert_indexer_mxfp4_attn:
head=128, MXFP4 (block=32), 4 ue8m0 bytes
Additional cutedsl kernels:
- _compress_kv_sparse_attn_cutedsl / _norm_rope_insert_sparse_attn_cutedsl:
CuTe DSL split kernels for C128
- _fused_kv_compress_norm_rope_insert_sparse_attn_cutedsl:
CuTe DSL fused kernels for C4
RoPE is register-based via tl.reshape -> tl.split -> tl.interleave (or the
even/odd halves are consumed directly for MXFP4, no interleave needed).
FP8 UE8M0 quant uses tl.reshape to tile [N_QUANT_BLOCKS, QUANT_BLOCK] for
@@ -25,43 +19,11 @@ even/odd halves, producing (N_QUANT_BLOCKS, MXFP4_BLOCK/2) packed nibbles
and N_QUANT_BLOCKS ue8m0 bytes.
"""
from functools import cache
from vllm.triton_utils import tl, triton
from .fused_indexer_q import _fp32x2_to_fp4x2
@cache
def _get_sparse_attn_cutedsl_impls():
from .sparse_attn_compress_cutedsl import (
_compress_kv_sparse_attn_cutedsl,
_fused_kv_compress_norm_rope_insert_sparse_attn_cutedsl,
_norm_rope_insert_sparse_attn_cutedsl,
)
return (
_compress_kv_sparse_attn_cutedsl,
_norm_rope_insert_sparse_attn_cutedsl,
_fused_kv_compress_norm_rope_insert_sparse_attn_cutedsl,
)
def _compress_kv_sparse_attn_cutedsl(*args, **kwargs):
"""CuTe DSL sparse-attention compress wrapper."""
return _get_sparse_attn_cutedsl_impls()[0](*args, **kwargs)
def _norm_rope_insert_sparse_attn_cutedsl(*args, **kwargs):
"""CuTe DSL RMSNorm/RoPE/FP8-store wrapper."""
return _get_sparse_attn_cutedsl_impls()[1](*args, **kwargs)
def _fused_kv_compress_norm_rope_insert_sparse_attn_cutedsl(*args, **kwargs):
"""CuTe DSL fused C4 sparse-attention compressor wrapper."""
return _get_sparse_attn_cutedsl_impls()[2](*args, **kwargs)
# =============================================================================
# DeepseekV4 Attention path (head=512, nope=448 FP8 + rope=64 bf16)
# =============================================================================
File diff suppressed because it is too large Load Diff
+39 -141
View File
@@ -13,11 +13,9 @@ from vllm.model_executor.layers.attention_layer_base import AttentionLayerBase
from vllm.model_executor.layers.layernorm import RMSNorm
from vllm.model_executor.layers.linear import MergedColumnParallelLinear
from vllm.models.deepseek_v4.common.ops.fused_compress_quant_cache import (
_compress_kv_sparse_attn_cutedsl,
_fused_kv_compress_norm_rope_insert_indexer_attn,
_fused_kv_compress_norm_rope_insert_indexer_mxfp4_attn,
_fused_kv_compress_norm_rope_insert_sparse_attn_cutedsl,
_norm_rope_insert_sparse_attn_cutedsl,
_fused_kv_compress_norm_rope_insert_sparse_attn,
)
from vllm.models.deepseek_v4.common.ops.fused_indexer_q import MXFP4_BLOCK_SIZE
from vllm.platforms import current_platform
@@ -173,33 +171,6 @@ class CompressorStateCache(torch.nn.Module, AttentionLayerBase):
class DeepseekCompressor(nn.Module):
_compressed_kv_buffers: ClassVar[dict[tuple[str, int, int], torch.Tensor]] = {}
@classmethod
def _get_compressed_kv_buffer(
cls,
device: str,
max_num_tokens: int,
head_dim: int,
) -> torch.Tensor:
if device == "cuda" and torch.accelerator.is_available():
device_key = f"cuda:{torch.accelerator.current_device_index()}"
alloc_device = torch.device(device_key)
else:
device_key = str(device)
alloc_device = torch.device(device)
key = (device_key, max_num_tokens, head_dim)
buffer = cls._compressed_kv_buffers.get(key)
if buffer is None:
buffer = torch.empty(
(max_num_tokens, head_dim),
dtype=torch.float32,
device=alloc_device,
)
cls._compressed_kv_buffers[key] = buffer
return buffer
def __init__(
self,
vllm_config: VllmConfig,
@@ -269,24 +240,12 @@ class DeepseekCompressor(nn.Module):
assert not use_fp4_cache, (
"MXFP4 cache is only supported for indexer (head=128)"
)
self._use_cutedsl_sparse_compressor = True
self._use_cutedsl_fused_sparse_compressor = self.compress_ratio == 4
self._compress_kernel = _compress_kv_sparse_attn_cutedsl
self._norm_rope_store_kernel = _norm_rope_insert_sparse_attn_cutedsl
self._fused_sparse_kernel = (
_fused_kv_compress_norm_rope_insert_sparse_attn_cutedsl
)
self._compressed_kv_buffer = self._get_compressed_kv_buffer(
self.device,
vllm_config.scheduler_config.max_num_batched_tokens,
self.head_dim,
)
self._fused_kernel = _fused_kv_compress_norm_rope_insert_sparse_attn
self._quant_block = 64
self._token_stride = self.nope_head_dim + self.rope_head_dim * 2
self._scale_dim = self.nope_head_dim // 64 + 1 # 7 real + 1 pad
self._num_warps = 4
elif self.head_dim == 128:
self._use_cutedsl_sparse_compressor = False
if use_fp4_cache:
self._fused_kernel = (
_fused_kv_compress_norm_rope_insert_indexer_mxfp4_attn
@@ -380,104 +339,43 @@ class DeepseekCompressor(nn.Module):
k_cache_metadata = cast(Any, attn_metadata[self.k_cache_prefix])
kv_cache = self._static_forward_context[self.k_cache_prefix].kv_cache
if self._use_cutedsl_sparse_compressor:
if self._use_cutedsl_fused_sparse_compressor:
self._fused_sparse_kernel(
state_cache,
token_to_req_indices,
positions,
slot_mapping,
block_table,
block_size,
self.norm.weight,
self.rms_norm_eps,
cos_sin_cache,
kv_cache,
k_cache_metadata.slot_mapping,
kv_cache.shape[1], # paged KV cache block size
kv_cache.stride(0),
head_size=self.head_dim,
state_width=state_width,
rope_head_dim=self.rope_head_dim,
fp8_max=448.0,
quant_block=self._quant_block,
token_stride=self._token_stride,
scale_dim=self._scale_dim,
compress_ratio=self.compress_ratio,
overlap=self.overlap,
)
else:
compressed_kv = self._compressed_kv_buffer[:num_actual]
self._compress_kernel(
state_cache,
token_to_req_indices,
positions,
slot_mapping,
block_table,
block_size,
compressed_kv,
head_size=self.head_dim,
state_width=state_width,
compress_ratio=self.compress_ratio,
overlap=self.overlap,
)
self._norm_rope_store_kernel(
compressed_kv,
positions,
slot_mapping,
self.norm.weight,
self.rms_norm_eps,
cos_sin_cache,
kv_cache,
k_cache_metadata.slot_mapping,
kv_cache.shape[1], # paged KV cache block size
kv_cache.stride(0),
head_size=self.head_dim,
rope_head_dim=self.rope_head_dim,
fp8_max=448.0,
quant_block=self._quant_block,
token_stride=self._token_stride,
scale_dim=self._scale_dim,
compress_ratio=self.compress_ratio,
)
else:
self._fused_kernel[(num_actual,)](
# state cache
state_cache,
state_cache.stride(0),
state_cache.stride(1),
# metadata
token_to_req_indices,
positions,
slot_mapping,
block_table,
block_table.stride(0),
block_size,
# RMSNorm
self.norm.weight,
self.rms_norm_eps,
# RoPE
cos_sin_cache,
cos_sin_cache.stride(0),
# KV cache
kv_cache,
k_cache_metadata.slot_mapping,
kv_cache.shape[1], # paged KV cache block size (tokens per block)
# constexprs
HEAD_SIZE=self.head_dim,
TRITON_BLOCK_SIZE=triton.next_power_of_2(self.head_dim),
STATE_WIDTH=state_width,
COMPRESS_RATIO=self.compress_ratio,
OVERLAP=self.overlap,
ROPE_HEAD_DIM=self.rope_head_dim,
FP8_MAX=448.0,
QUANT_BLOCK=self._quant_block,
TOKEN_STRIDE=self._token_stride,
SCALE_DIM=self._scale_dim,
KV_BLOCK_STRIDE=kv_cache.stride(0),
num_warps=self._num_warps,
**pdl_kwargs,
)
self._fused_kernel[(num_actual,)](
# state cache
state_cache,
state_cache.stride(0),
state_cache.stride(1),
# metadata
token_to_req_indices,
positions,
slot_mapping,
block_table,
block_table.stride(0),
block_size,
# RMSNorm
self.norm.weight,
self.rms_norm_eps,
# RoPE
cos_sin_cache,
cos_sin_cache.stride(0),
# KV cache
kv_cache,
k_cache_metadata.slot_mapping,
kv_cache.shape[1], # paged KV cache block size (tokens per block)
# constexprs
HEAD_SIZE=self.head_dim,
TRITON_BLOCK_SIZE=triton.next_power_of_2(self.head_dim),
STATE_WIDTH=state_width,
COMPRESS_RATIO=self.compress_ratio,
OVERLAP=self.overlap,
ROPE_HEAD_DIM=self.rope_head_dim,
FP8_MAX=448.0,
QUANT_BLOCK=self._quant_block,
TOKEN_STRIDE=self._token_stride,
SCALE_DIM=self._scale_dim,
KV_BLOCK_STRIDE=kv_cache.stride(0),
num_warps=self._num_warps,
**pdl_kwargs,
)
@triton.jit
+163 -2
View File
@@ -59,9 +59,9 @@ from vllm.models.deepseek_v4.nvidia.ops.attention import (
DeepseekV4MLAModules,
DeepseekV4MultiHeadLatentAttentionWrapper,
)
from vllm.models.deepseek_v4.nvidia.ops.prepare_megamoe import prepare_megamoe_inputs
from vllm.platforms import current_platform
from vllm.sequence import IntermediateTensors
from vllm.triton_utils import tl, triton
from vllm.utils.torch_utils import direct_register_custom_op
@@ -116,6 +116,167 @@ class DeepseekV4MLP(nn.Module):
return x
@triton.jit
def _deepseek_v4_stage_mega_moe_inputs_kernel(
hidden_states,
x_fp8,
x_sf,
topk_ids,
topk_weights,
topk_idx_out,
topk_weights_out,
hidden_stride_m: tl.constexpr,
hidden_stride_k: tl.constexpr,
x_stride_m: tl.constexpr,
x_stride_k: tl.constexpr,
x_sf_stride_m: tl.constexpr,
x_sf_stride_k: tl.constexpr,
topk_ids_stride_m: tl.constexpr,
topk_ids_stride_k: tl.constexpr,
topk_weights_stride_m: tl.constexpr,
topk_weights_stride_k: tl.constexpr,
topk_idx_stride_m: tl.constexpr,
topk_idx_stride_k: tl.constexpr,
topk_weights_out_stride_m: tl.constexpr,
topk_weights_out_stride_k: tl.constexpr,
hidden_size: tl.constexpr,
top_k: tl.constexpr,
BLOCK_K: tl.constexpr,
GROUP_K: tl.constexpr,
BLOCK_TOPK: tl.constexpr,
) -> None:
token_id = tl.program_id(0)
k_block_id = tl.program_id(1)
k_offsets = k_block_id * BLOCK_K + tl.arange(0, BLOCK_K)
k_mask = k_offsets < hidden_size
hidden = tl.load(
hidden_states + token_id * hidden_stride_m + k_offsets * hidden_stride_k,
mask=k_mask,
other=0.0,
).to(tl.float32)
num_groups: tl.constexpr = BLOCK_K // GROUP_K
hidden_groups = tl.reshape(tl.abs(hidden), [num_groups, GROUP_K])
amax = tl.max(hidden_groups, axis=1)
amax = tl.maximum(amax, 1.0e-4)
scale = amax / 448.0
scale_bits = scale.to(tl.uint32, bitcast=True)
scale_exp = ((scale_bits >> 23) & 0xFF) + ((scale_bits & 0x7FFFFF) != 0).to(
tl.uint32
)
scale_exp = tl.minimum(tl.maximum(scale_exp, 1), 254)
rounded_scale = (scale_exp << 23).to(tl.float32, bitcast=True)
hidden_groups = tl.reshape(hidden, [num_groups, GROUP_K])
scaled = hidden_groups * (1.0 / rounded_scale)[:, None]
scaled = tl.reshape(scaled, [BLOCK_K])
fp8 = scaled.to(tl.float8e4nv)
tl.store(
x_fp8 + token_id * x_stride_m + k_offsets * x_stride_k,
fp8,
mask=k_mask,
)
scale_offsets = tl.arange(0, num_groups)
packed_scale = tl.sum(scale_exp << (scale_offsets * 8), axis=0).to(tl.int32)
tl.store(
x_sf + token_id * x_sf_stride_m + k_block_id * x_sf_stride_k,
packed_scale,
)
if k_block_id == 0:
topk_offsets = tl.arange(0, BLOCK_TOPK)
topk_mask = topk_offsets < top_k
ids = tl.load(
topk_ids + token_id * topk_ids_stride_m + topk_offsets * topk_ids_stride_k,
mask=topk_mask,
other=0,
).to(tl.int64)
tl.store(
topk_idx_out
+ token_id * topk_idx_stride_m
+ topk_offsets * topk_idx_stride_k,
ids,
mask=topk_mask,
)
weights = tl.load(
topk_weights
+ token_id * topk_weights_stride_m
+ topk_offsets * topk_weights_stride_k,
mask=topk_mask,
other=0.0,
)
tl.store(
topk_weights_out
+ token_id * topk_weights_out_stride_m
+ topk_offsets * topk_weights_out_stride_k,
weights,
mask=topk_mask,
)
def _stage_deepseek_v4_mega_moe_inputs(
hidden_states: torch.Tensor,
topk_weights: torch.Tensor,
topk_ids: torch.Tensor,
x_fp8: torch.Tensor,
x_sf: torch.Tensor,
topk_idx_out: torch.Tensor,
topk_weights_out: torch.Tensor,
) -> None:
num_tokens, hidden_size = hidden_states.shape
if num_tokens == 0:
return
if hidden_size % 128 != 0:
raise ValueError(
"DeepSeek V4 MegaMoE input staging requires hidden_size to be "
"a multiple of 128."
)
top_k = topk_ids.shape[1]
if topk_weights.shape != topk_ids.shape:
raise ValueError(
"DeepSeek V4 MegaMoE input staging requires topk_weights and "
"topk_ids to have the same shape."
)
block_k = 128
grid = (num_tokens, triton.cdiv(hidden_size, block_k))
block_topk = triton.next_power_of_2(top_k)
_deepseek_v4_stage_mega_moe_inputs_kernel[grid](
hidden_states,
x_fp8,
x_sf,
topk_ids,
topk_weights,
topk_idx_out,
topk_weights_out,
hidden_states.stride(0),
hidden_states.stride(1),
x_fp8.stride(0),
x_fp8.stride(1),
x_sf.stride(0),
x_sf.stride(1),
topk_ids.stride(0),
topk_ids.stride(1),
topk_weights.stride(0),
topk_weights.stride(1),
topk_idx_out.stride(0),
topk_idx_out.stride(1),
topk_weights_out.stride(0),
topk_weights_out.stride(1),
hidden_size,
top_k,
BLOCK_K=block_k,
GROUP_K=32,
BLOCK_TOPK=block_topk,
num_warps=4,
)
def make_deepseek_v4_expert_params_mapping(
num_experts: int,
) -> list[tuple[str, str, int, str]]:
@@ -381,7 +542,7 @@ class DeepseekV4MegaMoEExperts(nn.Module):
symm_buffer = self.get_symm_buffer()
num_tokens = hidden_states.shape[0]
prepare_megamoe_inputs(
_stage_deepseek_v4_mega_moe_inputs(
hidden_states,
topk_weights,
topk_ids,
@@ -1,13 +1,21 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
from cutlass import Constexpr, Float32, Uint32, cute
import cutlass
import cutlass.cute as cute
from cutlass import Float32, Uint32
from cutlass._mlir import ir
from cutlass._mlir.dialects import llvm, vector
from cutlass.cutlass_dsl import T, dsl_user_op
@dsl_user_op
def fp32x2_to_bf16x2(a: Float32, b: Float32, *, loc=None, ip=None) -> Uint32:
def _recast_val(x, dtype, *, loc=None, ip=None):
return dtype(llvm.bitcast(dtype.mlir_type, x.ir_value(loc=loc, ip=ip)))
@dsl_user_op
def _fp32x2_to_bf16x2(a: Float32, b: Float32, *, loc=None, ip=None) -> Uint32:
out = llvm.inline_asm(
T.i32(),
[a.ir_value(loc=loc, ip=ip), b.ir_value(loc=loc, ip=ip)],
@@ -20,37 +28,62 @@ def fp32x2_to_bf16x2(a: Float32, b: Float32, *, loc=None, ip=None) -> Uint32:
@dsl_user_op
def bf16x2_to_fp32x2(data, *, loc=None, ip=None) -> tuple[Float32, Float32]:
if isinstance(data, Uint32):
out = llvm.inline_asm(
llvm.StructType.get_literal([T.f32(), T.f32()]),
[data.ir_value(loc=loc, ip=ip)],
"shl.b32 $0, $2, 16;\n\tand.b32 $1, $2, 0xFFFF0000;",
"=f,=f,r",
has_side_effects=False,
is_align_stack=False,
loc=loc,
ip=ip,
)
return (
Float32(llvm.extractvalue(T.f32(), out, [0], loc=loc, ip=ip)),
Float32(llvm.extractvalue(T.f32(), out, [1], loc=loc, ip=ip)),
)
elif isinstance(data, (cute.Tensor, cute.TensorSSA)):
# NOTE: the output is always 1D
size = cute.size(data.shape)
out = cute.make_rmem_tensor(size * 2, Float32)
for i in range(size):
out[i * 2], out[i * 2 + 1] = bf16x2_to_fp32x2(data[i])
return out
else:
raise ValueError(f"Unsupported type {type(data)}")
def _bf16x2_to_fp32(data: Uint32, *, loc=None, ip=None) -> tuple[Float32, Float32]:
out = llvm.inline_asm(
llvm.StructType.get_literal([T.f32(), T.f32()]),
[data.ir_value(loc=loc, ip=ip)],
"shl.b32 $0, $2, 16;\n\tand.b32 $1, $2, 0xFFFF0000;\n",
"=f,=f,r",
has_side_effects=False,
is_align_stack=False,
)
return (
Float32(llvm.extractvalue(T.f32(), out, [0], loc=loc, ip=ip)),
Float32(llvm.extractvalue(T.f32(), out, [1], loc=loc, ip=ip)),
)
@dsl_user_op
def fp8x4_to_bf16x4(x: Uint32, *, loc=None, ip=None) -> cute.TensorSSA:
def _bf16x2_abs(a: Uint32, *, loc=None, ip=None) -> Uint32:
out = llvm.inline_asm(
T.i32(),
[a.ir_value(loc=loc, ip=ip)],
"abs.bf16x2 $0, $1;",
"=r,r",
has_side_effects=False,
is_align_stack=False,
)
return Uint32(out)
@dsl_user_op
def _bf16x2_max(a: Uint32, b: Uint32, *, loc=None, ip=None) -> Uint32:
out = llvm.inline_asm(
T.i32(),
[a.ir_value(loc=loc, ip=ip), b.ir_value(loc=loc, ip=ip)],
"max.bf16x2 $0, $1, $2;",
"=r,r,r",
has_side_effects=False,
is_align_stack=False,
)
return Uint32(out)
@dsl_user_op
def _bf16x2_mul(a: Uint32, b: Uint32, *, loc=None, ip=None) -> Uint32:
out = llvm.inline_asm(
T.i32(),
[a.ir_value(loc=loc, ip=ip), b.ir_value(loc=loc, ip=ip)],
"mul.rn.bf16x2 $0, $1, $2;",
"=r,r,r",
has_side_effects=False,
is_align_stack=False,
)
return Uint32(out)
@dsl_user_op
def _fp8x4_to_bf16x4(x: Uint32, *, loc=None, ip=None) -> cute.TensorSSA:
# there is only fp8->fp16 conversion, hence we need to go
# round trip through fp16.
out = llvm.inline_asm(
@@ -85,7 +118,7 @@ def fp8x4_to_bf16x4(x: Uint32, *, loc=None, ip=None) -> cute.TensorSSA:
@dsl_user_op
def fp32x4_to_fp8x4(
def _fp32x4_to_fp8x4(
a0: Float32,
a1: Float32,
a2: Float32,
@@ -118,9 +151,9 @@ def fp32x4_to_fp8x4(
@dsl_user_op
def fp32x8_to_fp4x8(
def _fp32x8_to_fp4x8(
vals: cute.Tensor,
offset: Constexpr[int],
offset: cutlass.Constexpr[int],
*,
loc=None,
ip=None,
@@ -11,7 +11,10 @@ from cutlass import BFloat16, Int32, Uint8, Uint32
from cutlass.cute.nvgpu import cpasync
from quack.compile_utils import make_fake_tensor
from vllm.cute_utils import _bf16x2_mul, cvt
from vllm.models.deepseek_v4.nvidia.ops.cutedsl_utils import (
_bf16x2_mul,
_fp8x4_to_bf16x4,
)
def dequantize_and_gather_k_cache_cutedsl(
@@ -265,8 +268,8 @@ class DequantGatherKCacheKernel:
dequant0 = cute.make_rmem_tensor(4, Uint32)
dequant1 = cute.make_rmem_tensor(4, Uint32)
for j in cutlass.range_constexpr(2):
tmp0 = cvt.fp8x4_to_bf16x4(data0[j])
tmp1 = cvt.fp8x4_to_bf16x4(data1[j])
tmp0 = _fp8x4_to_bf16x4(data0[j])
tmp1 = _fp8x4_to_bf16x4(data1[j])
# BF16 multiply is safe because the scales are exact powers of 2.
dequant0[j * 2] = _bf16x2_mul(tmp0[0], scale0_bf16x2)
@@ -9,11 +9,14 @@ from cuda.bindings.driver import CUstream
from cutlass import BFloat16, Float32, Int64, Uint8, Uint32, const_expr
from quack.compile_utils import make_fake_tensor
from vllm.cute_utils import (
from vllm.models.deepseek_v4.nvidia.ops.cutedsl_utils import (
_bf16x2_abs,
_bf16x2_max,
cvt,
recast_val,
_bf16x2_to_fp32,
_fp32x2_to_bf16x2,
_fp32x4_to_fp8x4,
_fp32x8_to_fp4x8,
_recast_val,
)
from vllm.vllm_flash_attn.cute import utils as cute_utils
@@ -222,8 +225,8 @@ class IndexerQRopeQuantKernel:
cute.copy(cp_u32x4, cute.recast_tensor(sin_src, Uint32), sin_bf16x2)
for i in cutlass.range_constexpr(4):
cos0, cos1 = cvt.bf16x2_to_fp32x2(cos_bf16x2[i])
sin0, sin1 = cvt.bf16x2_to_fp32x2(sin_bf16x2[i])
cos0, cos1 = _bf16x2_to_fp32(cos_bf16x2[i])
sin0, sin1 = _bf16x2_to_fp32(sin_bf16x2[i])
cos_vals[i * 2] = cos0
cos_vals[i * 2 + 1] = cos1
sin_vals[i * 2] = sin0
@@ -231,11 +234,11 @@ class IndexerQRopeQuantKernel:
for i in cutlass.range_constexpr(self.coarsen):
for j in cutlass.range_constexpr(8):
q0, q1 = cvt.bf16x2_to_fp32x2(q_bf16x2[i, j])
q0, q1 = _bf16x2_to_fp32(q_bf16x2[i, j])
rot0 = q0 * cos_vals[j] - q1 * sin_vals[j]
rot1 = q0 * sin_vals[j] + q1 * cos_vals[j]
# convert back to BF16 to match numerics
q_bf16x2[i, j] = cvt.fp32x2_to_bf16x2(rot0, rot1)
q_bf16x2[i, j] = _fp32x2_to_bf16x2(rot0, rot1)
return (
q_bf16x2,
@@ -324,7 +327,7 @@ class IndexerQMxFp4Kernel(IndexerQRopeQuantKernel):
_bf16x2_max,
width=MXFP4_BLOCK_SIZE // 16,
)
amax_pair = cvt.bf16x2_to_fp32x2(amax_bf16x2)
amax_pair = _bf16x2_to_fp32(amax_bf16x2)
amax = cute_utils.fmax(amax_pair[0], amax_pair[1])
if in_bounds:
@@ -333,7 +336,7 @@ class IndexerQMxFp4Kernel(IndexerQRopeQuantKernel):
# increments the exponent whenever fp4_scale is not exactly a power of 2
eps = cutlass.const_expr(float.fromhex("0x6p-126"))
fp4_scale = cute_utils.fmax(amax, eps) * Float32(1.0 / 6.0)
bits = recast_val(fp4_scale, Uint32)
bits = _recast_val(fp4_scale, Uint32)
ue8m0 = cute_utils.shr_u32(
bits + Uint32(0x7FFFFF), Uint32(23)
) & Uint32(0xFF)
@@ -346,18 +349,18 @@ class IndexerQMxFp4Kernel(IndexerQRopeQuantKernel):
# If scale = 2^A and ue8m0 = A + 127, then inverse scale has exponent
# -A + 127 = 254 - ue8m0.
inv_scale_bits = (Uint32(254) - ue8m0) << Uint32(23)
inv_fp4_scale = recast_val(inv_scale_bits, Float32)
inv_fp4_scale = _recast_val(inv_scale_bits, Float32)
vals = cute.make_rmem_tensor(16, Float32)
for j in cutlass.range_constexpr(8):
q0, q1 = cvt.bf16x2_to_fp32x2(q_bf16x2[i, j])
q0, q1 = _bf16x2_to_fp32(q_bf16x2[i, j])
vals[j * 2] = q0 * inv_fp4_scale
vals[j * 2 + 1] = q1 * inv_fp4_scale
# pack to FP4
packed = cute.make_rmem_tensor((2,), Uint32)
packed[0] = cvt.fp32x8_to_fp4x8(vals, 0)
packed[1] = cvt.fp32x8_to_fp4x8(vals, 8)
packed[0] = _fp32x8_to_fp4x8(vals, 0)
packed[1] = _fp32x8_to_fp4x8(vals, 8)
dst = q_fp4_tile[i, None]
cp_u32x2 = cute.make_copy_atom(cp_op, Uint32, num_bits_per_copy=64)
@@ -516,24 +519,24 @@ class IndexerQFp8Kernel(IndexerQRopeQuantKernel):
_bf16x2_max,
width=self.subwarp_size,
)
amax_pair = cvt.bf16x2_to_fp32x2(amax_bf16x2)
amax_pair = _bf16x2_to_fp32(amax_bf16x2)
amax = cute_utils.fmax(amax_pair[0], amax_pair[1])
# scale = max(amax, eps) / fp8_max, then rounded UP to the next
# power of two. Adding the mantissa mask before shifting out the
# mantissa bumps the exponent whenever s isn't a pure pow2.
fp32_scale = cute_utils.fmax(amax, Float32(1e-4)) * Float32(1.0 / 448.0)
bits = recast_val(fp32_scale, Uint32)
bits = _recast_val(fp32_scale, Uint32)
scale_exp = cute_utils.shr_u32(
bits + Uint32(0x7FFFFF), Uint32(23)
) & Uint32(0xFF)
# rounded scale = 2^(scale_exp - 127); bit pattern is scale_exp << 23
fp8_scale_bits = scale_exp << Uint32(23)
fp8_scale = recast_val(fp8_scale_bits, Float32)
fp8_scale = _recast_val(fp8_scale_bits, Float32)
# inverse = 2^-(scale_exp - 127); bit pattern is (254 - scale_exp) << 23
inv_scale_bits = (Uint32(254) - scale_exp) << Uint32(23)
inv_fp8_scale = recast_val(inv_scale_bits, Float32)
inv_fp8_scale = _recast_val(inv_scale_bits, Float32)
# Weight fold: weights_out = weights * q_scale * scale_combined.
# All threads in the subwarp share the same fp8_scale after the
@@ -550,9 +553,9 @@ class IndexerQFp8Kernel(IndexerQRopeQuantKernel):
# (one cp.async-shaped 128-bit store per row).
packed = cute.make_rmem_tensor((4,), Uint32)
for j in cutlass.range_constexpr(4):
q0, q1 = cvt.bf16x2_to_fp32x2(q_bf16x2[i, j * 2])
q2, q3 = cvt.bf16x2_to_fp32x2(q_bf16x2[i, j * 2 + 1])
packed[j] = cvt.fp32x4_to_fp8x4(
q0, q1 = _bf16x2_to_fp32(q_bf16x2[i, j * 2])
q2, q3 = _bf16x2_to_fp32(q_bf16x2[i, j * 2 + 1])
packed[j] = _fp32x4_to_fp8x4(
q0 * inv_fp8_scale,
q1 * inv_fp8_scale,
q2 * inv_fp8_scale,
@@ -1,173 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Triton input-staging kernel for DeepSeek V4 MegaMoE.
Quantizes hidden states to fp8 with E8M0 group scales and repacks the
routing top-k tensors into the int64/float32 layout that the DeepGEMM
MegaMoE kernels consume.
"""
import torch
from vllm.triton_utils import tl, triton
@triton.jit
def _prepare_megamoe_inputs_kernel(
hidden_states,
x_fp8,
x_sf,
topk_ids,
topk_weights,
topk_idx_out,
topk_weights_out,
hidden_stride_m: tl.constexpr,
hidden_stride_k: tl.constexpr,
x_stride_m: tl.constexpr,
x_stride_k: tl.constexpr,
x_sf_stride_m: tl.constexpr,
x_sf_stride_k: tl.constexpr,
topk_ids_stride_m: tl.constexpr,
topk_ids_stride_k: tl.constexpr,
topk_weights_stride_m: tl.constexpr,
topk_weights_stride_k: tl.constexpr,
topk_idx_stride_m: tl.constexpr,
topk_idx_stride_k: tl.constexpr,
topk_weights_out_stride_m: tl.constexpr,
topk_weights_out_stride_k: tl.constexpr,
hidden_size: tl.constexpr,
top_k: tl.constexpr,
BLOCK_K: tl.constexpr,
GROUP_K: tl.constexpr,
BLOCK_TOPK: tl.constexpr,
) -> None:
token_id = tl.program_id(0)
k_block_id = tl.program_id(1)
k_offsets = k_block_id * BLOCK_K + tl.arange(0, BLOCK_K)
k_mask = k_offsets < hidden_size
hidden = tl.load(
hidden_states + token_id * hidden_stride_m + k_offsets * hidden_stride_k,
mask=k_mask,
other=0.0,
).to(tl.float32)
num_groups: tl.constexpr = BLOCK_K // GROUP_K
hidden_groups = tl.reshape(tl.abs(hidden), [num_groups, GROUP_K])
amax = tl.max(hidden_groups, axis=1)
amax = tl.maximum(amax, 1.0e-4)
scale = amax / 448.0
scale_bits = scale.to(tl.uint32, bitcast=True)
scale_exp = ((scale_bits >> 23) & 0xFF) + ((scale_bits & 0x7FFFFF) != 0).to(
tl.uint32
)
scale_exp = tl.minimum(tl.maximum(scale_exp, 1), 254)
rounded_scale = (scale_exp << 23).to(tl.float32, bitcast=True)
hidden_groups = tl.reshape(hidden, [num_groups, GROUP_K])
scaled = hidden_groups * (1.0 / rounded_scale)[:, None]
scaled = tl.reshape(scaled, [BLOCK_K])
fp8 = scaled.to(tl.float8e4nv)
tl.store(
x_fp8 + token_id * x_stride_m + k_offsets * x_stride_k,
fp8,
mask=k_mask,
)
scale_offsets = tl.arange(0, num_groups)
packed_scale = tl.sum(scale_exp << (scale_offsets * 8), axis=0).to(tl.int32)
tl.store(
x_sf + token_id * x_sf_stride_m + k_block_id * x_sf_stride_k,
packed_scale,
)
if k_block_id == 0:
topk_offsets = tl.arange(0, BLOCK_TOPK)
topk_mask = topk_offsets < top_k
ids = tl.load(
topk_ids + token_id * topk_ids_stride_m + topk_offsets * topk_ids_stride_k,
mask=topk_mask,
other=0,
).to(tl.int64)
tl.store(
topk_idx_out
+ token_id * topk_idx_stride_m
+ topk_offsets * topk_idx_stride_k,
ids,
mask=topk_mask,
)
weights = tl.load(
topk_weights
+ token_id * topk_weights_stride_m
+ topk_offsets * topk_weights_stride_k,
mask=topk_mask,
other=0.0,
)
tl.store(
topk_weights_out
+ token_id * topk_weights_out_stride_m
+ topk_offsets * topk_weights_out_stride_k,
weights,
mask=topk_mask,
)
def prepare_megamoe_inputs(
hidden_states: torch.Tensor,
topk_weights: torch.Tensor,
topk_ids: torch.Tensor,
x_fp8: torch.Tensor,
x_sf: torch.Tensor,
topk_idx_out: torch.Tensor,
topk_weights_out: torch.Tensor,
) -> None:
num_tokens, hidden_size = hidden_states.shape
if num_tokens == 0:
return
if hidden_size % 128 != 0:
raise ValueError(
"DeepSeek V4 MegaMoE input staging requires hidden_size to be "
"a multiple of 128."
)
top_k = topk_ids.shape[1]
if topk_weights.shape != topk_ids.shape:
raise ValueError(
"DeepSeek V4 MegaMoE input staging requires topk_weights and "
"topk_ids to have the same shape."
)
block_k = 128
grid = (num_tokens, triton.cdiv(hidden_size, block_k))
block_topk = triton.next_power_of_2(top_k)
_prepare_megamoe_inputs_kernel[grid](
hidden_states,
x_fp8,
x_sf,
topk_ids,
topk_weights,
topk_idx_out,
topk_weights_out,
hidden_states.stride(0),
hidden_states.stride(1),
x_fp8.stride(0),
x_fp8.stride(1),
x_sf.stride(0),
x_sf.stride(1),
topk_ids.stride(0),
topk_ids.stride(1),
topk_weights.stride(0),
topk_weights.stride(1),
topk_idx_out.stride(0),
topk_idx_out.stride(1),
topk_weights_out.stride(0),
topk_weights_out.stride(1),
hidden_size,
top_k,
BLOCK_K=block_k,
GROUP_K=32,
BLOCK_TOPK=block_topk,
num_warps=4,
)
-1
View File
@@ -11,7 +11,6 @@ PoolingTask = Literal[
"token_embed",
"token_classify",
"plugin",
"embed&token_classify",
]
POOLING_TASKS: tuple[PoolingTask, ...] = get_args(PoolingTask)
+14 -37
View File
@@ -3,7 +3,6 @@
"""Backend for GatedDeltaNet attention."""
from dataclasses import dataclass
from typing import Literal
import torch
@@ -91,12 +90,6 @@ class GDNAttentionMetadataBuilder(AttentionMetadataBuilder[GDNAttentionMetadata]
self.compilation_config = vllm_config.compilation_config
self.speculative_config = vllm_config.speculative_config
self.kv_cache_spec = kv_cache_spec
from vllm.model_executor.layers.mamba.gdn.qwen_gdn_linear_attn import (
_resolve_gdn_prefill_backend,
)
self.gdn_prefill_backend: Literal["triton", "flashinfer", "cutedsl"]
_, self.gdn_prefill_backend = _resolve_gdn_prefill_backend(vllm_config)
if self.speculative_config:
assert self.speculative_config.num_speculative_tokens is not None
@@ -323,38 +316,22 @@ class GDNAttentionMetadataBuilder(AttentionMetadataBuilder[GDNAttentionMetadata]
chunk_indices: torch.Tensor | None = None
chunk_offsets: torch.Tensor | None = None
if num_prefills > 0:
# Only prefill batches use FLA chunk ops.
# Pre-compute on CPU and async-copy to GPU to avoid
# GPU→CPU sync (.tolist()) in prepare_chunk_indices.
from vllm.model_executor.layers.fla.ops.index import (
prepare_chunk_indices,
prepare_chunk_offsets,
)
from vllm.model_executor.layers.fla.ops.utils import FLA_CHUNK_SIZE
if self.gdn_prefill_backend == "cutedsl":
from vllm.model_executor.layers.mamba.ops.gdn_chunk_cutedsl import (
prepare_metadata_cutedsl,
)
assert non_spec_query_start_loc is not None
assert non_spec_query_start_loc_cpu is not None
total_tokens = int(non_spec_query_start_loc_cpu[-1].item())
chunk_indices, chunk_offsets = prepare_metadata_cutedsl(
non_spec_query_start_loc,
total_tokens,
FLA_CHUNK_SIZE,
)
else:
gpu_device = query_start_loc.device
# Only prefill batches use FLA chunk ops.
# Pre-compute on CPU and async-copy to GPU to avoid
# GPU→CPU sync (.tolist()) in prepare_chunk_indices.
from vllm.model_executor.layers.fla.ops.index import (
prepare_chunk_indices,
prepare_chunk_offsets,
)
assert non_spec_query_start_loc_cpu is not None
chunk_indices = prepare_chunk_indices(
non_spec_query_start_loc_cpu, FLA_CHUNK_SIZE
).to(device=gpu_device, non_blocking=True)
chunk_offsets = prepare_chunk_offsets(
non_spec_query_start_loc_cpu, FLA_CHUNK_SIZE
).to(device=gpu_device, non_blocking=True)
gpu_device = query_start_loc.device
chunk_indices = prepare_chunk_indices(
non_spec_query_start_loc_cpu, FLA_CHUNK_SIZE
).to(device=gpu_device, non_blocking=True)
chunk_offsets = prepare_chunk_offsets(
non_spec_query_start_loc_cpu, FLA_CHUNK_SIZE
).to(device=gpu_device, non_blocking=True)
if num_prefills > 0:
has_initial_state = context_lens_tensor > 0
+1 -2
View File
@@ -291,8 +291,7 @@ class TopKTopPSampler(nn.Module):
)
# The custom XPU sampler kernel consumes RNG values internally, so advance
# the default generator's offset to keep future draws deterministic.
# pytorch: offset must be multiple of 4
offset = (offset + logits.numel() + 3) // 4 * 4
offset += logits.numel()
state.view(torch.int64)[1] = offset
generator.set_state(state)
return random_sampled, logits_to_return