h3-blackwell-runtime/src/h3_blackwell_runtime/cute_qkv_ring.py

283 lines
9.8 KiB
Python
Raw Normal View History

"""Opt-in CuTe bounded-ring backend for full-width H3 QKV projections."""
from __future__ import annotations
import importlib.util
import os
import sys
import threading
import weakref
from dataclasses import dataclass
from pathlib import Path
import torch
from .nvfp4_quant import nvfp4_activation_scale, vortex_native_quantize_nvfp4_into
def _enabled(name: str) -> bool:
return os.getenv(name, "").lower() in {"1", "true", "yes", "on"}
def _output_tensor(storage: torch.Tensor):
from cutlass.cute.runtime import from_dlpack
tensor = from_dlpack(storage.unsqueeze(-1), assumed_align=16)
return tensor.mark_compact_shape_dynamic(
mode=1, stride_order=(2, 0, 1), divisibility=1,
)
def _scale_tensor(storage: torch.Tensor):
import cutlass
from cutlass.cute.runtime import from_dlpack
tensor = from_dlpack(storage.view(torch.uint8).unsqueeze(-1), assumed_align=16)
tensor.element_type = cutlass.Float8E4M3FN
return tensor.mark_layout_dynamic(leading_dim=1)
def _weight_tensor(storage: torch.Tensor):
import cutlass
import cutlass.torch as cutlass_torch
lookup = torch.tensor(
[0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0,
-0.0, -0.5, -1.0, -1.5, -2.0, -3.0, -4.0, -6.0],
device=storage.device,
dtype=torch.float32,
)
codes = torch.stack((storage >> 4, storage & 0x0F), dim=-1).reshape(
storage.shape[0], -1,
)
logical = lookup[codes.long()].unsqueeze(-1)
tensor, backing = cutlass_torch.cute_tensor_like(
logical,
cutlass.Float4E2M1FN,
is_dynamic_layout=True,
assumed_align=16,
)
return tensor, backing
def _load_kernel(path: Path):
if not path.is_file():
raise FileNotFoundError(f"CuTe QKV kernel not found: {path}")
sys.path.insert(0, str(path.parent))
spec = importlib.util.spec_from_file_location("h3_cute_qkv_ring_kernel", path)
if spec is None or spec.loader is None:
raise ImportError(f"Cannot load CuTe QKV kernel: {path}")
module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module)
return module
@dataclass
class _PreparedWeight:
b: object
backing: torch.Tensor
sfb: object
@dataclass
class _Workspace:
a: object
a_backing: torch.Tensor
qdata: torch.Tensor
block_scale: torch.Tensor
sfa: object
alpha: torch.Tensor
alpha_argument: object
outputs: dict[int, tuple[torch.Tensor, list[object]]]
class _QkvRingBackend:
def __init__(self) -> None:
self.capacity = int(os.getenv("H3_CUTE_QKV_RING_CAPACITY", "2048"))
if self.capacity <= 0 or self.capacity % 128:
raise ValueError("H3_CUTE_QKV_RING_CAPACITY must be a positive multiple of 128")
default_path = Path("/opt/h3-blackwell-runtime/tools/dense_blockscaled_gemm_persistent_cooperative_vortex_alpha.py")
self.kernel_path = Path(os.getenv("H3_CUTE_QKV_KERNEL", str(default_path)))
self.strict = _enabled("H3_CUTE_QKV_RING_STRICT")
self._lock = threading.Lock()
self._module = None
self._workspace: dict[int, _Workspace] = {}
self._weights: weakref.WeakKeyDictionary = weakref.WeakKeyDictionary()
self._compiled: dict[int, tuple[object, object]] = {}
self.disabled_reason: str | None = None
def _ineligible(self, message: str):
if self.strict:
raise RuntimeError(message)
return None
def _eligible(self, linear, x: torch.Tensor) -> str | None:
if linear.role != "h3_attn_qkv":
return "projection role is not H3 QKV"
if linear.in_features != 5376 or linear.out_features != 21504:
return "QKV projection is sharded or has an unsupported shape"
if linear.full_precision_matrix_mult or linear.pre_quant_scale is not None:
return "QKV projection uses an unsupported quantization policy"
if linear.bias is not None or linear.output_dtype != torch.bfloat16:
return "QKV bias/output dtype is unsupported"
if not x.is_cuda or x.dtype != torch.bfloat16 or x.dim() != 2 or not x.is_contiguous():
return "QKV activation must be contiguous 2D CUDA BF16"
if x.requires_grad or torch.is_grad_enabled():
return "QKV ring is inference-only"
major, _ = torch.cuda.get_device_capability(x.device)
if major != 12:
return "QKV ring is currently validated only on SM12x"
return None
def _prepare_workspace(self, device: torch.device) -> _Workspace:
import cutlass
import cutlass.torch as cutlass_torch
from cutlass.cute.runtime import from_dlpack
index = device.index or 0
workspace = self._workspace.get(index)
if workspace is not None:
return workspace
a, a_backing = cutlass_torch.cute_tensor_like(
torch.zeros(
self.capacity, 5376, 1, device=device, dtype=torch.float32,
),
cutlass.Float4E2M1FN,
is_dynamic_layout=True,
assumed_align=16,
)
qdata = a_backing.view(torch.uint8).flatten()[
: self.capacity * 5376 // 2
].reshape(self.capacity, 5376 // 2)
block_scale = torch.empty(
self.capacity, 5376 // 16, device=device, dtype=torch.float8_e4m3fn,
)
alpha = torch.empty(1, device=device, dtype=torch.float32)
workspace = _Workspace(
a=a,
a_backing=a_backing,
qdata=qdata,
block_scale=block_scale,
sfa=_scale_tensor(block_scale),
alpha=alpha,
alpha_argument=from_dlpack(alpha, assumed_align=4),
outputs={},
)
self._workspace[index] = workspace
return workspace
def _prepare_weight(self, linear) -> _PreparedWeight:
prepared = self._weights.get(linear)
if prepared is not None:
return prepared
b, backing = _weight_tensor(linear.weight)
prepared = _PreparedWeight(
b=b,
backing=backing,
sfb=_scale_tensor(linear.weight_scale),
)
self._weights[linear] = prepared
return prepared
def _output(self, workspace: _Workspace, rows: int):
padded_rows = ((rows + self.capacity - 1) // self.capacity) * self.capacity
cached = workspace.outputs.get(padded_rows)
if cached is not None:
return cached
output = torch.empty(
padded_rows,
21504,
device=workspace.qdata.device,
dtype=torch.bfloat16,
)
chunks = [
_output_tensor(output[start : start + self.capacity])
for start in range(0, padded_rows, self.capacity)
]
cached = (output, chunks)
workspace.outputs[padded_rows] = cached
return cached
def _compile(self, device: torch.device, workspace: _Workspace, weight: _PreparedWeight, c):
import cutlass
import cutlass.cute as cute
import cutlass.torch as cutlass_torch
index = device.index or 0
cached = self._compiled.get(index)
if cached is not None:
return cached
if self._module is None:
self._module = _load_kernel(self.kernel_path)
gemm = self._module.Sm120BlockScaledGemmKernel(
cutlass.Float32, 16, (128, 128, 128), (128, 128),
)
stream = cutlass_torch.default_stream()
max_active_clusters = cutlass.utils.HardwareInfo().get_max_active_clusters(1)
compiled = cute.compile(
gemm,
workspace.a,
weight.b,
workspace.sfa,
weight.sfb,
c,
workspace.alpha_argument,
max_active_clusters,
stream,
)
cached = (compiled, stream)
self._compiled[index] = cached
return cached
def __call__(self, linear, x: torch.Tensor):
reason = self._eligible(linear, x)
if reason is not None:
return self._ineligible(reason)
try:
with self._lock:
workspace = self._prepare_workspace(x.device)
weight = self._prepare_weight(linear)
output, c_chunks = self._output(workspace, x.shape[0])
compiled, stream = self._compile(
x.device, workspace, weight, c_chunks[0],
)
scale = nvfp4_activation_scale(x).float()
workspace.alpha.copy_(scale * linear.weight_scale_2.float())
for index, start in enumerate(range(0, x.shape[0], self.capacity)):
stop = min(start + self.capacity, x.shape[0])
vortex_native_quantize_nvfp4_into(
x[start:stop],
scale,
workspace.qdata,
workspace.block_scale,
hi_first=False,
)
compiled(
workspace.a,
weight.b,
workspace.sfa,
weight.sfb,
c_chunks[index],
workspace.alpha_argument,
stream,
)
return output[: x.shape[0], : linear.out_features]
except (ImportError, FileNotFoundError, RuntimeError) as error:
self.disabled_reason = str(error)
return self._ineligible(f"QKV ring initialization failed: {error}")
_BACKEND: _QkvRingBackend | None = None
def qkv_ring_linear(linear, x: torch.Tensor):
global _BACKEND
if _BACKEND is None:
try:
_BACKEND = _QkvRingBackend()
except (ImportError, ValueError) as error:
if _enabled("H3_CUTE_QKV_RING_STRICT"):
raise
return None
return _BACKEND(linear, x)