forked from Karylab-cklius/vllm
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
4ac64ec057 |
@@ -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"
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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
|
||||
-206
@@ -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):
|
||||
|
||||
@@ -1540,7 +1540,6 @@ class ModelConfig:
|
||||
return "token_classify"
|
||||
|
||||
priority: list[PoolingTask] = [
|
||||
"embed&token_classify",
|
||||
"embed",
|
||||
"classify",
|
||||
"token_embed",
|
||||
|
||||
@@ -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)
|
||||
@@ -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,
|
||||
|
||||
@@ -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
@@ -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:
|
||||
|
||||
@@ -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))
|
||||
@@ -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
|
||||
|
||||
@@ -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"))
|
||||
),
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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.
|
||||
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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]]):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
@@ -11,7 +11,6 @@ PoolingTask = Literal[
|
||||
"token_embed",
|
||||
"token_classify",
|
||||
"plugin",
|
||||
"embed&token_classify",
|
||||
]
|
||||
POOLING_TASKS: tuple[PoolingTask, ...] = get_args(PoolingTask)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user