diff --git a/src/h3_blackwell_runtime/csrc/nvfp4_scale.cpp b/src/h3_blackwell_runtime/csrc/nvfp4_scale.cpp index d5a8371..f67f1e2 100644 --- a/src/h3_blackwell_runtime/csrc/nvfp4_scale.cpp +++ b/src/h3_blackwell_runtime/csrc/nvfp4_scale.cpp @@ -1,6 +1,7 @@ #include torch::Tensor nvfp4_activation_scale_cuda(torch::Tensor input, double divisor); +torch::Tensor nvfp4_activation_scale_into_cuda(torch::Tensor input, double divisor, torch::Tensor partials, torch::Tensor output, int64_t blocks, int64_t threads); torch::Tensor nvfp4_activation_scale(torch::Tensor input, double divisor) { TORCH_CHECK(input.is_cuda(), "nvfp4_activation_scale expects a CUDA tensor"); @@ -9,6 +10,23 @@ torch::Tensor nvfp4_activation_scale(torch::Tensor input, double divisor) { return nvfp4_activation_scale_cuda(input, divisor); } +torch::Tensor nvfp4_activation_scale_into(torch::Tensor input, double divisor, torch::Tensor partials, torch::Tensor output, int64_t blocks, int64_t threads) { + TORCH_CHECK(input.is_cuda(), "nvfp4_activation_scale_into expects a CUDA tensor"); + TORCH_CHECK(input.is_contiguous(), "nvfp4_activation_scale_into expects contiguous input"); + TORCH_CHECK(input.dim() == 2, "nvfp4_activation_scale_into expects a 2D tensor"); + TORCH_CHECK(partials.is_cuda() && output.is_cuda(), "nvfp4_activation_scale_into workspace must be CUDA tensors"); + TORCH_CHECK(partials.device() == input.device() && output.device() == input.device(), "nvfp4_activation_scale_into workspace must be on the input device"); + TORCH_CHECK(partials.is_contiguous() && output.is_contiguous(), "nvfp4_activation_scale_into workspace must be contiguous"); + TORCH_CHECK(partials.scalar_type() == torch::kFloat32 && output.scalar_type() == torch::kFloat32, "nvfp4_activation_scale_into workspace must be float32"); + TORCH_CHECK(blocks > 0, "nvfp4_activation_scale_into blocks must be positive"); + TORCH_CHECK(threads > 0 && threads <= 1024, "nvfp4_activation_scale_into threads must be between 1 and 1024"); + TORCH_CHECK((threads & (threads - 1)) == 0, "nvfp4_activation_scale_into threads must be a power of two"); + TORCH_CHECK(partials.numel() >= blocks, "nvfp4_activation_scale_into partial workspace is too small"); + TORCH_CHECK(output.numel() >= 1, "nvfp4_activation_scale_into output workspace is too small"); + return nvfp4_activation_scale_into_cuda(input, divisor, partials, output, blocks, threads); +} + PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { m.def("nvfp4_activation_scale", &nvfp4_activation_scale, "Vortex NVFP4 activation scale"); + m.def("nvfp4_activation_scale_into", &nvfp4_activation_scale_into, "Vortex NVFP4 activation scale with caller workspace"); } diff --git a/src/h3_blackwell_runtime/csrc/nvfp4_scale.cu b/src/h3_blackwell_runtime/csrc/nvfp4_scale.cu index 9783f07..77184bd 100644 --- a/src/h3_blackwell_runtime/csrc/nvfp4_scale.cu +++ b/src/h3_blackwell_runtime/csrc/nvfp4_scale.cu @@ -3,12 +3,46 @@ #include #include +#include #include namespace { constexpr int kThreads = 256; +__inline__ __device__ float warp_reduce_max(float value) { + for (int offset = 16; offset > 0; offset >>= 1) { + value = fmaxf(value, __shfl_down_sync(0xffffffff, value, offset)); + } + return value; +} + +__inline__ __device__ float block_reduce_max(float value) { + __shared__ float warp_values[32]; + const int lane = threadIdx.x & 31; + const int warp = threadIdx.x >> 5; + value = warp_reduce_max(value); + if (lane == 0) { + warp_values[warp] = value; + } + __syncthreads(); + value = threadIdx.x < ((blockDim.x + 31) >> 5) ? warp_values[lane] : 0.0f; + if (warp == 0) { + value = warp_reduce_max(value); + } + return value; +} + +__inline__ __device__ uint32_t max_bf16_abs_bits(uint32_t current, uint32_t packed) { + const uint32_t lo = packed & 0x7fffu; + const uint32_t hi = (packed >> 16) & 0x7fffu; + return max(current, max(lo, hi)); +} + +__inline__ __device__ float bf16_abs_bits_to_float(uint32_t bits) { + return __uint_as_float(bits << 16); +} + template __global__ void partial_absmax_kernel(const scalar_t* __restrict__ input, float* __restrict__ partials, int64_t numel) { __shared__ float shared[kThreads]; @@ -54,6 +88,40 @@ __global__ void final_scale_kernel(const float* __restrict__ partials, float* __ } } +__global__ void partial_absmax_bf16_vec_kernel(const uint4* __restrict__ input, const uint16_t* __restrict__ scalar_input, float* __restrict__ partials, int64_t vector_count, int64_t numel) { + const int64_t stride = static_cast(blockDim.x) * gridDim.x; + int64_t index = static_cast(blockIdx.x) * blockDim.x + threadIdx.x; + uint32_t local_bits = 0; + while (index < vector_count) { + const uint4 values = input[index]; + local_bits = max_bf16_abs_bits(local_bits, values.x); + local_bits = max_bf16_abs_bits(local_bits, values.y); + local_bits = max_bf16_abs_bits(local_bits, values.z); + local_bits = max_bf16_abs_bits(local_bits, values.w); + index += stride; + } + const int64_t tail_start = vector_count * 8; + for (int64_t tail = tail_start + static_cast(blockIdx.x) * blockDim.x + threadIdx.x; tail < numel; tail += stride) { + local_bits = max(local_bits, static_cast(scalar_input[tail] & 0x7fffu)); + } + const float local_max = bf16_abs_bits_to_float(local_bits); + const float block_max = block_reduce_max(local_max); + if (threadIdx.x == 0) { + partials[blockIdx.x] = block_max; + } +} + +__global__ void final_scale_warp_kernel(const float* __restrict__ partials, float* __restrict__ output, int64_t count, float divisor) { + float local_max = 0.0f; + for (int64_t index = threadIdx.x; index < count; index += blockDim.x) { + local_max = fmaxf(local_max, partials[index]); + } + const float block_max = block_reduce_max(local_max); + if (threadIdx.x == 0) { + output[0] = block_max / divisor; + } +} + } // namespace torch::Tensor nvfp4_activation_scale_cuda(torch::Tensor input, double divisor) { @@ -76,3 +144,30 @@ torch::Tensor nvfp4_activation_scale_cuda(torch::Tensor input, double divisor) { C10_CUDA_KERNEL_LAUNCH_CHECK(); return output; } + +torch::Tensor nvfp4_activation_scale_into_cuda(torch::Tensor input, double divisor, torch::Tensor partials, torch::Tensor output, int64_t blocks, int64_t threads) { + c10::cuda::CUDAGuard device_guard(input.device()); + const auto numel = input.numel(); + TORCH_CHECK(numel > 0, "nvfp4_activation_scale_into input must be non-empty"); + TORCH_CHECK(divisor > 0.0, "nvfp4_activation_scale_into divisor must be positive"); + TORCH_CHECK(input.scalar_type() == at::ScalarType::BFloat16, "nvfp4_activation_scale_into v1 currently expects bfloat16 input"); + TORCH_CHECK(threads == 128 || threads == 256 || threads == 512, "nvfp4_activation_scale_into threads must be 128, 256, or 512"); + TORCH_CHECK(blocks > 0, "nvfp4_activation_scale_into blocks must be positive"); + TORCH_CHECK(reinterpret_cast(input.data_ptr()) % alignof(uint4) == 0, "nvfp4_activation_scale_into expects 16-byte aligned input"); + + const int launch_blocks = static_cast(blocks); + const int launch_threads = static_cast(threads); + const int64_t vector_count = numel / 8; + auto stream = at::cuda::getCurrentCUDAStream(); + + partial_absmax_bf16_vec_kernel<<>>( + reinterpret_cast(input.data_ptr()), + reinterpret_cast(input.data_ptr()), + partials.data_ptr(), + vector_count, + numel); + C10_CUDA_KERNEL_LAUNCH_CHECK(); + final_scale_warp_kernel<<<1, launch_threads, 0, stream>>>(partials.data_ptr(), output.data_ptr(), blocks, static_cast(divisor)); + C10_CUDA_KERNEL_LAUNCH_CHECK(); + return output; +} diff --git a/src/h3_blackwell_runtime/nvfp4_quant.py b/src/h3_blackwell_runtime/nvfp4_quant.py index b4e0662..bd23ed6 100644 --- a/src/h3_blackwell_runtime/nvfp4_quant.py +++ b/src/h3_blackwell_runtime/nvfp4_quant.py @@ -10,6 +10,9 @@ from pathlib import Path import torch +_VORTEX_SCALE_WORKSPACES: dict[tuple[int, int], tuple[torch.Tensor, torch.Tensor]] = {} + + def _sync() -> None: if torch.cuda.is_available(): torch.cuda.synchronize() @@ -40,6 +43,36 @@ def _vortex_scale_extension(): ) +def _env_int(name: str, default: int) -> int: + value = os.getenv(name) + if value is None or value == "": + return default + return int(value) + + +def _vortex_scale_geometry(numel: int) -> tuple[int, int]: + threads = _env_int("H3_NVFP4_SCALE_THREADS", 256) + if threads not in {128, 256, 512}: + raise ValueError(f"H3_NVFP4_SCALE_THREADS must be 128, 256, or 512, got {threads}") + default_blocks = min((numel + threads - 1) // threads, 4096) + blocks = _env_int("H3_NVFP4_SCALE_BLOCKS", default_blocks) + if blocks <= 0: + raise ValueError(f"H3_NVFP4_SCALE_BLOCKS must be positive, got {blocks}") + return blocks, threads + + +def _vortex_scale_workspace(tensor: torch.Tensor, blocks: int) -> tuple[torch.Tensor, torch.Tensor]: + key = (tensor.device.index or 0, blocks) + workspace = _VORTEX_SCALE_WORKSPACES.get(key) + if workspace is None or workspace[0].device != tensor.device: + workspace = ( + torch.empty((blocks,), device=tensor.device, dtype=torch.float32), + torch.empty((), device=tensor.device, dtype=torch.float32), + ) + _VORTEX_SCALE_WORKSPACES[key] = workspace + return workspace + + def nvfp4_activation_scale(tensor: torch.Tensor, *, timings: dict[str, list[float]] | None = None) -> torch.Tensor: """Compute Comfy Kitchen's current per-tensor NVFP4 activation scale.""" from comfy_kitchen.float_utils import F4_E2M1_MAX, F8_E4M3_MAX @@ -55,7 +88,17 @@ def nvfp4_activation_scale(tensor: torch.Tensor, *, timings: dict[str, list[floa raise ValueError("Vortex NVFP4 scale backend requires contiguous input") try: extension = _vortex_scale_extension() - scale = _record_timing(timings, "vortex_absmax_scale", lambda: extension.nvfp4_activation_scale(tensor, divisor)) + version = os.getenv("H3_NVFP4_SCALE_VERSION", "1").lower() + if version in {"1", "v1"} and tensor.dtype == torch.bfloat16 and tensor.data_ptr() % 16 == 0: + blocks, threads = _vortex_scale_geometry(tensor.numel()) + partials, output = _vortex_scale_workspace(tensor, blocks) + scale = _record_timing( + timings, + "vortex_absmax_scale", + lambda: extension.nvfp4_activation_scale_into(tensor, divisor, partials, output, blocks, threads), + ) + else: + scale = _record_timing(timings, "vortex_absmax_scale", lambda: extension.nvfp4_activation_scale(tensor, divisor)) return _record_timing(timings, "scale_compat_cast", lambda: scale.to(tensor.dtype) if tensor.dtype != torch.float32 else scale) except Exception: if os.getenv("H3_NVFP4_SCALE_STRICT", "").lower() in {"1", "true", "yes", "on"}: