Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions transformer_engine/pytorch/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,6 +29,7 @@
from transformer_engine.pytorch.module import destroy_ub
from transformer_engine.pytorch.module import UserBufferQuantizationMode
from transformer_engine.pytorch.attention import DotProductAttention
from transformer_engine.pytorch.attention import FusedMLAQUpProjRopeQuant
from transformer_engine.pytorch.attention import MultiheadAttention
from transformer_engine.pytorch.attention import InferenceParams
from transformer_engine.pytorch.attention import RotaryPositionEmbedding
Expand Down
2 changes: 2 additions & 0 deletions transformer_engine/pytorch/attention/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,12 +5,14 @@
"""Python interface for attention"""

from .dot_product_attention import DotProductAttention
from .fused_mla_q_uproj import FusedMLAQUpProjRopeQuant
from .multi_head_attention import MultiheadAttention
from .inference import InferenceParams
from .rope import RotaryPositionEmbedding

__all__ = [
"DotProductAttention",
"FusedMLAQUpProjRopeQuant",
"MultiheadAttention",
"InferenceParams",
"RotaryPositionEmbedding",
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -1363,6 +1363,7 @@ def forward(
deterministic,
softmax_offset,
fp8_output,
bf16_backward,
layer_number,
return_max_logit,
packed_qkv=None,
Expand Down Expand Up @@ -1427,6 +1428,9 @@ def forward(
# fp8_dtype = tex.DType.kFloat8E4M3
if is_input_fp8:
q_fp8, k_fp8, v_fp8 = q, k, v

if fp8_recipe.mxfp8():
qkv_scale_inv_format = "bhsd" # Same as what combine_and_quantize would give
else:
q_fp8, k_fp8, v_fp8, qkv_layout, qkv_scale_inv_format = combine_and_quantize(
qkv_layout,
Expand Down Expand Up @@ -1602,6 +1606,8 @@ def forward(

ctx.is_input_fp8 = is_input_fp8
ctx.is_output_fp8 = is_output_fp8
# Return dQ/dK/dV in bf16 even if is_input_fp8
ctx.bf16_backward = bf16_backward

tensors_to_save, tensor_objects = prepare_for_saving(
*fp8_tensors,
Expand Down Expand Up @@ -1860,7 +1866,8 @@ def backward(ctx, d_out, *_args):
# dq, dk, dv: torch.Tensor; dtype = torch.float16 or torch.bfloat16
dq, dk, dv = dq_, dk_, dv_
is_quantized_tensor = isinstance(dq_, QuantizedTensorStorage)
if is_quantized_tensor and not ctx.is_input_fp8:

if is_quantized_tensor and (not ctx.is_input_fp8 or ctx.bf16_backward):
# return in F16
dq, dk, dv = combine_and_dequantize(
ctx.dqkv_layout,
Expand All @@ -1869,7 +1876,7 @@ def backward(ctx, d_out, *_args):
dv_,
src_nominal_dtype=dq_.dtype,
)
if not is_quantized_tensor and ctx.is_input_fp8:
if not is_quantized_tensor and ctx.is_input_fp8 and not ctx.bf16_backward:
# return in FP8
dq, dk, dv, _, _ = combine_and_quantize(
ctx.dqkv_layout, dq_, dk_, dv_, ctx.dQKV_quantizer
Expand Down Expand Up @@ -1968,6 +1975,7 @@ def backward(ctx, d_out, *_args):
None,
None, # packed_qkv
None, # packed_kv
None,
)


Expand Down Expand Up @@ -2064,6 +2072,7 @@ def forward(
score_mod_bprop_tensors: Optional[Dict[str, torch.Tensor]] = None,
packed_qkv: Optional[torch.Tensor] = None,
packed_kv: Optional[torch.Tensor] = None,
bf16_backward: bool = False,
) -> torch.Tensor:
"""fused attention fprop"""
assert (
Expand Down Expand Up @@ -2280,6 +2289,7 @@ def forward(
self.deterministic,
softmax_offset,
fp8_output,
bf16_backward,
self.layer_number,
self.return_max_logit,
packed_qkv,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,7 @@
Float8BlockScalingRecipeState,
)
from transformer_engine.pytorch.tensor.storage.float8_tensor_storage import Float8TensorStorage
from transformer_engine.pytorch.tensor.storage.mxfp8_tensor_storage import MXFP8TensorStorage
from transformer_engine.pytorch.module.base import TransformerEngineBaseModule
from transformer_engine.pytorch.export import is_in_onnx_export_mode
from transformer_engine.pytorch.constants import AttnMaskTypes, AttnTypes, dist_group_type, DType
Expand Down Expand Up @@ -1115,6 +1116,7 @@ def forward(
inference_params: Optional[InferenceParams] = None,
pad_between_seqs: Optional[bool] = None,
fp8_output: Optional[bool] = False,
bf16_backward: Optional[bool] = False,
num_splits: Optional[int] = 1,
score_mod: Optional[Callable] = None,
score_mod_bprop: Optional[Callable] = None,
Expand Down Expand Up @@ -1585,6 +1587,25 @@ def forward(
qkv_format=qkv_format,
inference_params=inference_params,
)
elif all(
isinstance(x, MXFP8TensorStorage) for x in [query_layer, key_layer, value_layer]
):
# Pre-quantized MXFP8 q/k/v: the wrapper has no real storage, so run
# layout detection on the underlying rowwise data (mirrors the Float8 path).
(
qkv_layout,
query_layer._rowwise_data,
key_layer._rowwise_data,
value_layer._rowwise_data,
q_format,
kv_format,
) = dpa_utils.get_qkv_layout(
query_layer._rowwise_data,
key_layer._rowwise_data,
value_layer._rowwise_data,
qkv_format=qkv_format,
inference_params=inference_params,
)
else:
(
qkv_layout,
Expand Down Expand Up @@ -1958,6 +1979,7 @@ def forward(
fp8_output=fp8_output,
packed_qkv=qkv_layer,
packed_kv=kv_layer,
bf16_backward=bf16_backward,
)
return self.fused_attention(
query_layer,
Expand Down Expand Up @@ -1995,6 +2017,7 @@ def forward(
score_mod_bprop_tensors=score_mod_bprop_tensors,
packed_qkv=qkv_layer,
packed_kv=kv_layer,
bf16_backward=bf16_backward,
)

if use_unfused_attention:
Expand Down
156 changes: 156 additions & 0 deletions transformer_engine/pytorch/attention/dot_product_attention/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -2923,6 +2923,162 @@ def _build_outputs(scale_list, alignment):
return result, "bhsd"


def mxfp8_quantize_only(tensor_quantizer_pairs, src_format):
"""Phase 1 of mxfp8_quantize_fast_path: quantize only, no BHSD transpose or GEMM swizzle.

Returns MXFP8Tensors with data and scale_invs reshaped to src_format layout.
Call mxfp8_transpose_swizzle to complete the BHSD permute + swizzle when ready
(e.g. after pre-quantized tensors from fused kernels are also available).

Parameters
----------
tensor_quantizer_pairs : list of (torch.Tensor, MXFP8Quantizer)
Same contract as mxfp8_quantize_fast_path.
src_format : str
``"bshd"`` or ``"sbhd"``.

Returns
-------
fp8_tensors : list of MXFP8Tensor
Data and scale_invs in src_format layout; NOT yet BHSD-permuted or swizzled.
"""
if not tensor_quantizer_pairs:
return []
assert src_format in (
"bshd",
"sbhd",
), f"mxfp8_quantize_only only supports bshd/sbhd, got {src_format!r}."
_s_dim = {"bshd": 1, "sbhd": 0}
_d_dim = {"bshd": 3, "sbhd": 3}

fp8_tensors = []
for tensor, quantizer in tensor_quantizer_pairs:
original_shape = tensor.shape
rs_shape = list(original_shape)
rs_shape[_d_dim[src_format]] //= MXFP8_BLOCK_SCALING_SIZE
cs_shape = list(original_shape)
cs_shape[_s_dim[src_format]] //= MXFP8_BLOCK_SCALING_SIZE
if src_format == "bshd":
t2d = tensor.view(*tensor.shape[:2], -1)
else:
t2d = tensor.view(tensor.shape[0], -1)
orig_optimize = quantizer.optimize_for_gemm
quantizer.optimize_for_gemm = False
fp8_2d = quantizer(t2d)
quantizer.optimize_for_gemm = orig_optimize
# Re-wrap with the original 4D SBHD shape so that shape[-1] equals the per-head
# dimension (matching Q's wrapper shape) and fused_attn_bwd produces 4D dkv that
# matches key/value's expected gradient shape in _KFQuantizeKVForAttn.backward.
fp8_t = MXFP8Tensor(
shape=original_shape,
dtype=tensor.dtype,
rowwise_data=(
fp8_2d._rowwise_data.view(original_shape)
if fp8_2d._rowwise_data is not None
else None
),
rowwise_scale_inv=(
fp8_2d._rowwise_scale_inv.view(rs_shape)
if fp8_2d._rowwise_scale_inv is not None
else None
),
columnwise_data=(
fp8_2d._columnwise_data.view(original_shape)
if fp8_2d._columnwise_data is not None
else None
),
columnwise_scale_inv=(
fp8_2d._columnwise_scale_inv.view(cs_shape)
if fp8_2d._columnwise_scale_inv is not None
else None
),
quantizer=quantizer,
requires_grad=False,
fp8_dtype=fp8_2d._fp8_dtype,
with_gemm_swizzled_scales=False,
)
fp8_tensors.append(fp8_t)
return fp8_tensors


def mxfp8_transpose_swizzle(fp8_tensors, src_format):
"""Phase 2 of mxfp8_quantize_fast_path: batched BHSD-transpose + GEMM-swizzle.

For tensors whose data is already quantized (e.g. from a fused GEMM+quant kernel
or from mxfp8_quantize_only), permutes each tensor's scale_invs from src_format to
BHSD and applies the GEMM swizzle in-place. Complements mxfp8_quantize_only to
allow pre-quantized tensors (like a fused-kernel Q) to be processed in the same
batched operation as freshly quantized K/V.

Parameters
----------
fp8_tensors : list of MXFP8Tensor
Tensors with _rowwise_scale_inv / _columnwise_scale_inv in src_format layout.
Modified in-place: scale_invs are replaced with BHSD-permuted, swizzled versions.
src_format : str
``"bshd"`` or ``"sbhd"``.
"""
if not fp8_tensors:
return

assert src_format in (
"bshd",
"sbhd",
), f"mxfp8_transpose_swizzle only supports bshd/sbhd, got {src_format!r}."

rs_list = [t._rowwise_scale_inv for t in fp8_tensors]
cs_list = [t._columnwise_scale_inv for t in fp8_tensors]

def _align_up(x, a):
return ((x + a - 1) // a) * a

def _bhsd_shape(src_4d, d_pad):
if src_format == "sbhd":
S, B, H, _ = src_4d.shape
else:
B, S, H, _ = src_4d.shape
return (B, H, S, d_pad)

def _build_outputs(scale_list, alignment):
entries = []
total = 0
for s in scale_list:
if s is None:
entries.append(None)
continue
d_pad = _align_up(s.shape[-1], alignment)
shape = _bhsd_shape(s, d_pad)
numel = 1
for dim in shape:
numel *= dim
entries.append((total, numel, shape))
total += numel
if total == 0:
return [None] * len(scale_list)
device = next(s for s in scale_list if s is not None).device
buf = torch.empty(total, dtype=torch.uint8, device=device)
return [buf[e[0] : e[0] + e[1]].view(e[2]) if e is not None else None for e in entries]

rs_outs = _build_outputs(rs_list, 4)
cs_outs = _build_outputs(cs_list, 128)

rs_permuted = tex.multi_tensor_transpose_to_bhsd(
rs_list, original_format=src_format, outputs=rs_outs
)
cs_permuted = tex.multi_tensor_transpose_to_bhsd(
cs_list, original_format=src_format, outputs=cs_outs
)

for t, rp, cp in zip(fp8_tensors, rs_permuted, cs_permuted):
t._rowwise_scale_inv = rp.view(-1, rp.shape[-1]) if rp is not None else None
t._columnwise_scale_inv = cp.view(-1, cp.shape[-1]) if cp is not None else None

tex.multi_tensor_swizzle_scales_for_gemm_unchecked_(fp8_tensors, True, False)
tex.multi_tensor_swizzle_scales_for_gemm_unchecked_(fp8_tensors, False, True)
for t in fp8_tensors:
t._with_gemm_swizzled_scales = True


def combine_and_quantize(
qkv_layout,
q,
Expand Down
Loading
Loading