diff --git a/src/h3_blackwell_runtime/cute_qkv_ring.py b/src/h3_blackwell_runtime/cute_qkv_ring.py new file mode 100644 index 0000000..5108a4d --- /dev/null +++ b/src/h3_blackwell_runtime/cute_qkv_ring.py @@ -0,0 +1,282 @@ +"""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)