diff --git a/src/h3_blackwell_runtime/nvfp4_quant.py b/src/h3_blackwell_runtime/nvfp4_quant.py new file mode 100644 index 0000000..18b1ccb --- /dev/null +++ b/src/h3_blackwell_runtime/nvfp4_quant.py @@ -0,0 +1,88 @@ +"""Experimental NVFP4 activation quantization helpers for profiling.""" + +from __future__ import annotations + +import time + +import torch + + +def _sync() -> None: + if torch.cuda.is_available(): + torch.cuda.synchronize() + + +def _record_timing(timings: dict[str, list[float]] | None, name: str, fn): + if timings is None: + return fn() + _sync() + started = time.perf_counter() + value = fn() + _sync() + timings.setdefault(name, []).append(time.perf_counter() - started) + return value + + +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 + + amax = _record_timing(timings, "scale_absmax", lambda: torch.amax(tensor.abs())) + return _record_timing(timings, "scale_finalize", lambda: amax / (F8_E4M3_MAX * F4_E2M1_MAX)) + + +def vortex_quantize_nvfp4( + tensor: torch.Tensor, + *, + scale: torch.Tensor | float | None = None, + timings: dict[str, list[float]] | None = None, +): + """Create a TensorCoreNVFP4 QuantizedTensor through an explicit Vortex seam. + + This prototype still uses Comfy Kitchen's low-level ``quantize_nvfp4`` CUDA op + for the pack/block-scale step, but bypasses ``QuantizedTensor.from_float`` and + lets callers provide a precomputed/global scale. It is intentionally isolated so a + native Vortex quantizer can replace this implementation without touching the + linear call sites or benchmarks. + """ + if tensor.dim() != 2: + raise ValueError(f"NVFP4 activation quantization requires a 2D tensor, got {tensor.dim()}D") + if not tensor.is_contiguous(): + raise ValueError("vortex_quantize_nvfp4 requires contiguous input; fix the caller rather than hiding a copy here") + + import comfy_kitchen as ck + from comfy_kitchen.tensor import QuantizedTensor, TensorCoreNVFP4Layout + + orig_dtype = tensor.dtype + orig_shape = tuple(tensor.shape) + if scale is None: + scale = nvfp4_activation_scale(tensor, timings=timings) + scale = _record_timing( + timings, + "scale_to_device", + lambda: torch.as_tensor(scale, device=tensor.device, dtype=torch.float32), + ) + + qdata, block_scale = _record_timing( + timings, + "ck_quantize_nvfp4", + lambda: ck.quantize_nvfp4( + tensor, + scale, + pad_16x=TensorCoreNVFP4Layout.get_padded_shape(orig_shape) != orig_shape, + ), + ) + return _record_timing( + timings, + "params_wrap", + lambda: QuantizedTensor( + qdata, + "TensorCoreNVFP4Layout", + TensorCoreNVFP4Layout.Params( + scale=scale, + orig_dtype=orig_dtype, + orig_shape=orig_shape, + block_scale=block_scale, + ), + ), + ) diff --git a/tools/profile_nvfp4_linear.py b/tools/profile_nvfp4_linear.py index 230964e..5f7d005 100644 --- a/tools/profile_nvfp4_linear.py +++ b/tools/profile_nvfp4_linear.py @@ -19,6 +19,7 @@ from h3_blackwell_runtime.attention import AVAILABLE_BACKENDS, rms_norm, rms_rop from h3_blackwell_runtime.block import H3DiTBlock, gate_segments, modulate_segments from h3_blackwell_runtime.checkpoint import H3Checkpoint from h3_blackwell_runtime.nvfp4 import Nvfp4Linear +from h3_blackwell_runtime.nvfp4_quant import nvfp4_activation_scale, vortex_quantize_nvfp4 from h3_blackwell_runtime.packing import H3PromptPacker from h3_blackwell_runtime.rope import h3_rope_rotation from h3_blackwell_runtime.sampler import _audio_sigma, _model_sigma, beta_sigmas @@ -41,18 +42,47 @@ def timed(stats: dict[str, list[float]], name: str, fn): def summarize(values: list[float]) -> dict[str, float]: ordered = sorted(values) + def percentile(percent: float) -> float: + if len(ordered) == 1: + return ordered[0] + rank = (len(ordered) - 1) * percent + low = int(rank) + high = min(low + 1, len(ordered) - 1) + weight = rank - low + return ordered[low] * (1.0 - weight) + ordered[high] * weight + return { "count": len(values), "mean_s": sum(values) / len(values), + "p50_s": percentile(0.50), + "p90_s": percentile(0.90), + "p95_s": percentile(0.95), + "p99_s": percentile(0.99), "min_s": ordered[0], "max_s": ordered[-1], } -def profile_linear_stages(module: Nvfp4Linear, x: torch.Tensor, *, iterations: int) -> dict[str, dict[str, float]]: +def quantize_activation(flat_x: torch.Tensor, quantizer: str, scale: torch.Tensor | None, timings: dict[str, list[float]] | None = None): from comfy_kitchen.tensor import QuantizedTensor + if quantizer == "comfy": + return QuantizedTensor.from_float(flat_x, "TensorCoreNVFP4Layout") + if quantizer == "vortex_recalculate": + return vortex_quantize_nvfp4(flat_x, timings=timings) + if quantizer == "vortex_precomputed_scale": + if scale is None: + raise ValueError("vortex_precomputed_scale requires a precomputed scale") + return vortex_quantize_nvfp4(flat_x, scale=scale, timings=timings) + raise ValueError(f"Unsupported quantizer: {quantizer}") + + +def profile_linear_stages(module: Nvfp4Linear, x: torch.Tensor, *, iterations: int, quantizer: str) -> dict[str, dict[str, float]]: + stats: dict[str, list[float]] = {} + precomputed_scale = None + if quantizer == "vortex_precomputed_scale": + precomputed_scale = timed(stats, "precomputed_scale_calibration", lambda: nvfp4_activation_scale(x.reshape(-1, module.in_features).contiguous())) for _ in range(iterations): original_shape = x.shape[:-1] flat_x = timed(stats, "flatten_contiguous", lambda: x.reshape(-1, module.in_features).contiguous()) @@ -66,23 +96,38 @@ def profile_linear_stages(module: Nvfp4Linear, x: torch.Tensor, *, iterations: i weight = timed(stats, "weight_dequantize", lambda: packed_weight.dequantize().to(flat_x)) output = timed(stats, "linear", lambda: functional.linear(flat_x, weight, bias)) else: - packed_x = timed(stats, "activation_quantize", lambda: QuantizedTensor.from_float(flat_x, "TensorCoreNVFP4Layout")) + packed_x = timed(stats, "activation_quantize", lambda: quantize_activation(flat_x, quantizer, precomputed_scale, stats)) output = timed(stats, "linear", lambda: functional.linear(packed_x, packed_weight, bias)) timed(stats, "slice_reshape", lambda: output[:flat_x.shape[0], :module.out_features].reshape(*original_shape, module.out_features)) return {name: summarize(values) for name, values in stats.items()} -def profile_cuda_kernels(module: Nvfp4Linear, x: torch.Tensor, *, warmup: int, iterations: int, row_limit: int) -> list[dict]: +def run_linear_with_quantizer(module: Nvfp4Linear, x: torch.Tensor, quantizer: str, precomputed_scale: torch.Tensor | None = None) -> torch.Tensor: + original_shape = x.shape[:-1] + flat_x = x.reshape(-1, module.in_features).contiguous() + if module.pre_quant_scale is not None: + flat_x = flat_x * module.pre_quant_scale.to(flat_x) + packed_weight = module._packed_weight() + bias = module.bias.to(flat_x) if module.bias is not None else None + if module.full_precision_matrix_mult: + output = functional.linear(flat_x, packed_weight.dequantize().to(flat_x), bias) + else: + output = functional.linear(quantize_activation(flat_x, quantizer, precomputed_scale), packed_weight, bias) + return output[:flat_x.shape[0], :module.out_features].reshape(*original_shape, module.out_features) + + +def profile_cuda_kernels(module: Nvfp4Linear, x: torch.Tensor, *, warmup: int, iterations: int, row_limit: int, quantizer: str) -> list[dict]: with torch.inference_mode(): + precomputed_scale = nvfp4_activation_scale(x.reshape(-1, module.in_features).contiguous()) if quantizer == "vortex_precomputed_scale" else None for _ in range(warmup): - module(x) + run_linear_with_quantizer(module, x, quantizer, precomputed_scale) sync() activities = [ProfilerActivity.CPU] if torch.cuda.is_available(): activities.append(ProfilerActivity.CUDA) with profile(activities=activities, record_shapes=True) as prof: for _ in range(iterations): - module(x) + run_linear_with_quantizer(module, x, quantizer, precomputed_scale) sync() rows = [] for event in prof.key_averages(group_by_input_shape=True): @@ -200,6 +245,7 @@ def parse_args() -> argparse.Namespace: parser.add_argument("--block-index", type=int, default=24) parser.add_argument("--attention", choices=AVAILABLE_BACKENDS, default="sage2") parser.add_argument("--linears", nargs="+", choices=("attn_qkv_proj", "attn_out_proj", "mlp_fc1", "mlp_fc2"), default=("attn_qkv_proj", "attn_out_proj", "mlp_fc1", "mlp_fc2")) + parser.add_argument("--quantizers", nargs="+", choices=("comfy", "vortex_recalculate", "vortex_precomputed_scale"), default=("comfy", "vortex_recalculate", "vortex_precomputed_scale")) parser.add_argument("--warmup", type=int, default=2) parser.add_argument("--iterations", type=int, default=5) parser.add_argument("--profiler-iterations", type=int, default=2) @@ -216,18 +262,25 @@ def main() -> None: for name in args.linears: module = module_for_name(block, name) x = inputs[name] - for _ in range(args.warmup): - module(x) - stage_timings = profile_linear_stages(module, x, iterations=args.iterations) - kernels = profile_cuda_kernels(module, x, warmup=args.warmup, iterations=args.profiler_iterations, row_limit=args.profiler_row_limit) - results.append( - { - "name": name, - "module": module_info(module, x), - "stage_timings": stage_timings, - "profiler_top_events": kernels, - } - ) + reference = run_linear_with_quantizer(module, x, "comfy") + for quantizer in args.quantizers: + precomputed_scale = nvfp4_activation_scale(x.reshape(-1, module.in_features).contiguous()) if quantizer == "vortex_precomputed_scale" else None + for _ in range(args.warmup): + run_linear_with_quantizer(module, x, quantizer, precomputed_scale) + stage_timings = profile_linear_stages(module, x, iterations=args.iterations, quantizer=quantizer) + candidate = run_linear_with_quantizer(module, x, quantizer, precomputed_scale) + diff = (candidate.float() - reference.float()).abs() + kernels = profile_cuda_kernels(module, x, warmup=args.warmup, iterations=args.profiler_iterations, row_limit=args.profiler_row_limit, quantizer=quantizer) + results.append( + { + "name": name, + "quantizer": quantizer, + "module": module_info(module, x), + "stage_timings": stage_timings, + "reference_diff": {"max": diff.max().item(), "mean": diff.mean().item()}, + "profiler_top_events": kernels, + } + ) output = { "model_path": args.model_path, @@ -236,6 +289,7 @@ def main() -> None: "warmup": args.warmup, "iterations": args.iterations, "profiler_iterations": args.profiler_iterations, + "quantizers": args.quantizers, "metadata": metadata, "results": results, }