h3-blackwell-runtime/research/cute_nvfp4_ring/patches/0001-cute-qkv-ring-module.patch

289 lines
10 KiB
Diff
Raw Permalink Normal View History

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)