289 lines
10 KiB
Diff
289 lines
10 KiB
Diff
|
|
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)
|