class QutlassNvFP4LinearMethod(CompressedTensorsLinearTransformMethod):
def create_weights(
self,
layer,
input_size_per_partition,
output_partition_sizes,
input_size,
output_size,
params_dtype,
**extra_weight_attrs,
):
# initializes fp4 qparams
assert isinstance(layer.scheme, (CompressedTensorsW4A4Fp4,))
ret = super().create_weights(
layer,
input_size_per_partition,
output_partition_sizes,
input_size,
output_size,
params_dtype,
**extra_weight_attrs,
)
assert self.input_transform is not None
assert len(self.input_transform.weight.partitions) >= 1
return ret
@staticmethod
def _get_flashinfer_gemm_backend(kernel: NvFp4LinearKernel) -> str:
"""
Given a kernel, find the string that is needed to be passed into
`flashinfer_scaled_fp4_mm`, using
vllm.model_executor.kernels.linear._LINEAR_BACKEND_KERNEL_MAP as source of truth
"""
kernel_type = type(kernel)
for key, kernels in _LINEAR_BACKEND_KERNEL_MAP.items():
if not key.startswith("flashinfer_") or kernel_type not in kernels:
continue
backend = key.removeprefix("flashinfer_")
# flashinfer GEMM backend uses "cute-dsl", not "cutedsl"
return backend.replace("cutedsl", "cute-dsl")
raise ValueError(
f"QutlassNvFP4 transform requires a FlashInfer kernel, "
f"got {kernel_type.__name__}"
)
def process_weights_after_loading(self, layer):
super().process_weights_after_loading(layer)
assert self.input_transform is not None
layer.hadamard_matrix = self.input_transform.weight.partitions[0].data
# fusedQuantizeNv stores raw absmax as block scales (sf = absmax),
# while CT weights use sf = absmax * SFScaleVal / 6.0. The GEMM
# computes alpha * sum(fp4_a * sf_a * fp4_w * sf_w), so alpha must
# compensate: alpha = weight_global_scale / 6.0
layer.fused_alpha = Parameter(
layer.weight_global_scale / NVFP4_MAX, requires_grad=False
)
layer.fused_global_scale = Parameter(
torch.tensor(
[NVFP4_MAX],
dtype=torch.float32,
device=layer.weight_global_scale.device,
),
requires_grad=False,
)
layer.flashinfer_gemm_backend = self._get_flashinfer_gemm_backend(
layer.scheme.kernel
)
def apply(
self,
layer: torch.nn.Module,
x: torch.Tensor,
bias: torch.Tensor | None = None,
) -> torch.Tensor:
assert bias is None
output_size = layer.output_size_per_partition
output_shape = [*x.shape[:-1], output_size]
x_flat = x.contiguous().flatten(end_dim=-2)
x_fp4, x_scales = fusedQuantizeNv(
x_flat, layer.hadamard_matrix, layer.fused_global_scale
)
x_scales_blocked = to_blocked(x_scales, backend="triton").view(x_scales.shape)
out = flashinfer_scaled_fp4_mm(
x_fp4,
layer.weight,
x_scales_blocked,
layer.weight_scale,
layer.fused_alpha,
x.dtype,
backend=layer.flashinfer_gemm_backend,
)
out = slice_nvfp4_output(out, output_size)
if self.output_transform is not None:
for part_id, (start, length) in enumerate(self.partition_ranges):
out[:, start : start + length] = self.output_transform(
out[:, start : start + length].clone(), part_id=part_id
)
return out.view(*output_shape)