def fused_int8_dequant_bias(acc: torch.Tensor, x_scale: torch.Tensor,
weight_scale: torch.Tensor, bias: torch.Tensor | None,
dtype: torch.dtype) -> torch.Tensor:
"""Avoid full-size FP32 scaling intermediates; retain both rounding steps."""
rows, cols = acc.shape
if not acc.is_cuda or acc.dtype != torch.int32 or not acc.is_contiguous():
raise ValueError("INT8 VAE epilogue requires a contiguous CUDA INT32 matrix")
if x_scale.shape != (rows, 1) or weight_scale.shape != (cols, 1):
raise ValueError("INT8 VAE epilogue requires per-row and per-output-channel scales")
if any(t.device != acc.device for t in (x_scale, weight_scale)):
raise ValueError("INT8 VAE epilogue scales must be on the accumulator device")
if bias is not None and (bias.device != acc.device or bias.shape != (cols,) or not bias.is_contiguous()):
raise ValueError("INT8 VAE epilogue bias must be contiguous on the accumulator device")
if dtype not in (torch.float32, torch.float16, torch.bfloat16):
raise ValueError("INT8 VAE epilogue supports FP32, FP16 and BF16 outputs")
out = torch.empty((rows, cols), device=acc.device, dtype=dtype)
_dequant_bias[(triton.cdiv(rows * cols, 1024),)](
acc, x_scale, weight_scale, bias if bias is not None else acc, out,
rows, cols, x_scale.stride(0), weight_scale.stride(0), bias is not None,
BLOCK=1024, num_warps=4, enable_fp_fusion=False)
return out