Skip to content
Merged
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
18 changes: 2 additions & 16 deletions tests/pytorch/test_onnx_export.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -756,16 +756,10 @@ def test_export_layernorm_mlp(
(torch.float16, True, "padding"), # calls ScaledMaskedSoftmax
(torch.float16, False, "padding"), # calls ScaledSoftmax
])
@pytest.mark.parametrize("attention_softmax_in_fp32",
[True, False])
@pytest.mark.parametrize("apply_query_key_layer_scaling",
[True, False])
def test_export_core_attention(
precision: torch.dtype,
use_mask: bool,
attn_mask_type: str,
attention_softmax_in_fp32: bool,
apply_query_key_layer_scaling: bool,
):
# Set dimensions (these are arbitrary).
kv_channels = 64
Expand All@@ -784,11 +778,9 @@ def test_export_core_attention(
input_names.append("attention_mask")
inp = (query_layer, key_layer, value_layer, attention_mask)

sm_prec_str = "_sm-fp32" if attention_softmax_in_fp32 else "_sm-fp16"
qk_scaling_str = "_qk-scaling" if apply_query_key_layer_scaling else ""
mask_str = get_attn_mask_str(use_mask, attn_mask_type)
high_prec_str = dtype2str(precision)
fname = f"te.core_attention{mask_str}{qk_scaling_str}{sm_prec_str}{high_prec_str}.onnx"
fname = f"te.core_attention{mask_str}{high_prec_str}.onnx"

if attn_mask_type is None:
attn_mask_type = 'causal'
Expand All@@ -798,8 +790,6 @@ def test_export_core_attention(
kv_channels=kv_channels,
attention_dropout=0.5,
attn_mask_type=attn_mask_type,
attention_softmax_in_fp32=attention_softmax_in_fp32,
apply_query_key_layer_scaling=apply_query_key_layer_scaling,
).to(device='cuda')
do_export(model,
inp,
Expand DownExpand Up@@ -911,7 +901,6 @@ def test_export_multihead_attention(
])
@pytest.mark.parametrize("precision", [torch.float32, torch.float16])
@pytest.mark.parametrize("fuse_qkv_params", [False, True])
@pytest.mark.parametrize("apply_query_key_layer_scaling", [True, False])
@pytest.mark.parametrize("zero_centered_gamma", [False, True])
def test_export_transformer_layer(
use_fp8: bool,
Expand All@@ -920,7 +909,6 @@ def test_export_transformer_layer(
output_layernorm: bool,
precision: torch.dtype,
fuse_qkv_params: bool,
apply_query_key_layer_scaling: bool,
zero_centered_gamma: bool
):
# Skip FP8 tests on non-hopper devices
Expand All@@ -946,10 +934,9 @@ def test_export_transformer_layer(

fp8_str = "_fp8" if use_fp8 else ""
fuse_qkv_params_str = "_fused-qkv" if fuse_qkv_params else ""
qk_scaling_str = "_qk-scaling" if apply_query_key_layer_scaling else ""
high_prec_str = dtype2str(precision)
attn_mask_str = get_attn_mask_str(use_mask, attn_mask_type)
fname = f"te.transformer_layer{fp8_str}{attn_mask_str}{fuse_qkv_params_str}{qk_scaling_str}{high_prec_str}.onnx"
fname = f"te.transformer_layer{fp8_str}{attn_mask_str}{fuse_qkv_params_str}{high_prec_str}.onnx"

model = te.TransformerLayer(
hidden_size,
Expand All@@ -959,7 +946,6 @@ def test_export_transformer_layer(
output_layernorm=output_layernorm,
params_dtype=precision,
fuse_qkv_params=fuse_qkv_params,
apply_query_key_layer_scaling=apply_query_key_layer_scaling,
zero_centered_gamma=zero_centered_gamma).to(device='cuda')
do_export(model, inp, fname, use_fp8)
if not use_fp8:
Expand Down
36 changes: 19 additions & 17 deletions transformer_engine/pytorch/softmax.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -4,7 +4,7 @@

"""Fused scaled masked softmax functions"""
import os
from typing import Callable, Tuple, Union
from typing import Callable, Tuple, Union, Optional
import torch
from torch import nn
import torch._C._onnx as _C_onnx
Expand DownExpand Up@@ -198,15 +198,13 @@ class FusedScaleMaskSoftmax(nn.Module):
attn_mask_type: attention mask type (pad or causal)
mask_func: mask function to be applied.
softmax_in_fp32: if true, softmax in performed at fp32 precision.
scale: scaling factor used in input tensor scaling.
"""

def __init__(
self,
attn_mask_type: str,
mask_func: Callable,
softmax_in_fp32: bool,
scale: float,
softmax_in_fp32: bool = True,
Comment thread
ptrendx marked this conversation as resolved.
) -> None:
super().__init__()
self.attn_mask_type = attn_mask_type
Expand All@@ -215,23 +213,27 @@ def __init__(
)
self.mask_func = mask_func
self.softmax_in_fp32 = softmax_in_fp32
self.scale = scale

assert (
self.scale is None or softmax_in_fp32
), "softmax should be in fp32 when scaled"

def forward(self, inp: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
def forward(
self,
inp: torch.Tensor,
mask: torch.Tensor,
scale: Optional[float] = None,
) -> torch.Tensor:
"""FusedScaleMaskSoftmax fprop"""
# [b, np, sq, sk]
assert inp.dim() == 4
self.input_in_fp16 = inp.dtype == torch.float16
self.input_in_bf16 = inp.dtype == torch.bfloat16
self.input_in_float16 = self.input_in_fp16 or self.input_in_bf16

assert (
scale is None or self.softmax_in_fp32
), "softmax should be in fp32 when scaled"

if self.is_kernel_available(*inp.size()):
return self.forward_fused_softmax(inp, mask)
return self.forward_torch_softmax(inp, mask)
return self.forward_fused_softmax(inp, mask, scale)
return self.forward_torch_softmax(inp, mask, scale)

def is_kernel_available(self, b: int, np: int, sq: int, sk: int) -> bool:
"""Check FusedScaleMaskSoftmax kernel availability based on size"""
Expand All@@ -256,11 +258,11 @@ def is_kernel_available(self, b: int, np: int, sq: int, sk: int) -> bool:
return False

def forward_fused_softmax(
self, inp: torch.Tensor, mask: torch.Tensor
self, inp: torch.Tensor, mask: torch.Tensor, scale: Optional[float] = None
) -> torch.Tensor:
"""Fused masked softmax kernel"""
b, np, sq, sk = inp.size()
scale = self.scale if self.scale is not None else 1.0
scale = 1.0 if scale is None else scale

if self.attn_mask_type == "causal":
assert sq == sk, "causal mask is only for self attention"
Expand All@@ -275,14 +277,14 @@ def forward_fused_softmax(
return ScaledSoftmax.apply(inp, scale)

def forward_torch_softmax(
self, inp: torch.Tensor, mask: torch.Tensor
self, inp: torch.Tensor, mask: torch.Tensor, scale: Optional[float] = None
) -> torch.Tensor:
"""Framework softmax"""
if self.input_in_float16 and self.softmax_in_fp32:
inp = inp.float()

if self.scale is not None:
inp = inp * self.scale
if scale is not None:
inp = inp * scale

if self.attn_mask_type == "causal":
mask = _get_default_causal_mask(inp.size()[2])
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { // Add copy buttons to all
 blocks
(function() {
function addCopyButtons() {
document.querySelectorAll('pre code').forEach(function(codeBlock) {
if (codeBlock.parentElement.hasAttribute('data-copy-added')) return;
codeBlock.parentElement.setAttribute('data-copy-added', 'true');
var btn = document.createElement('button');
btn.textContent = 'Copy';
btn.style.cssText = 'position:absolute;top:4px;right:4px;padding:2px 8px;font-size:11px;background:#4ecdc4;border:none;border-radius:4px;color:#1a1a2e;cursor:pointer;opacity:0.7;transition:opacity 0.2s;';
btn.onmouseover = function() { this.style.opacity = '1'; };
btn.onmouseout = function() { this.style.opacity = '0.7'; };
btn.onclick = function() {
navigator.clipboard.writeText(codeBlock.textContent).then(function() {
btn.textContent = 'Copied!';
setTimeout(function() { btn.textContent = 'Copy'; }, 1500);
});
};
codeBlock.parentElement.style.position = 'relative';
codeBlock.parentElement.appendChild(btn);
});
}
addCopyButtons();
// Re-run on dynamic content
var observer = new MutationObserver(addCopyButtons);
observer.observe(document.body, { childList: true, subtree: true });
})();
}
} catch(__e) { console.warn('[Userscript:Add Copy Buttons to Code Blocks]', __e); }
})();
(function(){
try {
var __m = "github.com";
var __re = new RegExp('^' + "github\\.com" + '
deprecate qk layer scaling and fp32 softmax args by ksivaman · Pull Request #90 · NVIDIA/TransformerEngine · GitHub
Skip to content
Merged
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
18 changes: 2 additions & 16 deletions tests/pytorch/test_onnx_export.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -756,16 +756,10 @@ def test_export_layernorm_mlp(
(torch.float16, True, "padding"), # calls ScaledMaskedSoftmax
(torch.float16, False, "padding"), # calls ScaledSoftmax
])
@pytest.mark.parametrize("attention_softmax_in_fp32",
[True, False])
@pytest.mark.parametrize("apply_query_key_layer_scaling",
[True, False])
def test_export_core_attention(
precision: torch.dtype,
use_mask: bool,
attn_mask_type: str,
attention_softmax_in_fp32: bool,
apply_query_key_layer_scaling: bool,
):
# Set dimensions (these are arbitrary).
kv_channels = 64
Expand All@@ -784,11 +778,9 @@ def test_export_core_attention(
input_names.append("attention_mask")
inp = (query_layer, key_layer, value_layer, attention_mask)

sm_prec_str = "_sm-fp32" if attention_softmax_in_fp32 else "_sm-fp16"
qk_scaling_str = "_qk-scaling" if apply_query_key_layer_scaling else ""
mask_str = get_attn_mask_str(use_mask, attn_mask_type)
high_prec_str = dtype2str(precision)
fname = f"te.core_attention{mask_str}{qk_scaling_str}{sm_prec_str}{high_prec_str}.onnx"
fname = f"te.core_attention{mask_str}{high_prec_str}.onnx"

if attn_mask_type is None:
attn_mask_type = 'causal'
Expand All@@ -798,8 +790,6 @@ def test_export_core_attention(
kv_channels=kv_channels,
attention_dropout=0.5,
attn_mask_type=attn_mask_type,
attention_softmax_in_fp32=attention_softmax_in_fp32,
apply_query_key_layer_scaling=apply_query_key_layer_scaling,
).to(device='cuda')
do_export(model,
inp,
Expand DownExpand Up@@ -911,7 +901,6 @@ def test_export_multihead_attention(
])
@pytest.mark.parametrize("precision", [torch.float32, torch.float16])
@pytest.mark.parametrize("fuse_qkv_params", [False, True])
@pytest.mark.parametrize("apply_query_key_layer_scaling", [True, False])
@pytest.mark.parametrize("zero_centered_gamma", [False, True])
def test_export_transformer_layer(
use_fp8: bool,
Expand All@@ -920,7 +909,6 @@ def test_export_transformer_layer(
output_layernorm: bool,
precision: torch.dtype,
fuse_qkv_params: bool,
apply_query_key_layer_scaling: bool,
zero_centered_gamma: bool
):
# Skip FP8 tests on non-hopper devices
Expand All@@ -946,10 +934,9 @@ def test_export_transformer_layer(

fp8_str = "_fp8" if use_fp8 else ""
fuse_qkv_params_str = "_fused-qkv" if fuse_qkv_params else ""
qk_scaling_str = "_qk-scaling" if apply_query_key_layer_scaling else ""
high_prec_str = dtype2str(precision)
attn_mask_str = get_attn_mask_str(use_mask, attn_mask_type)
fname = f"te.transformer_layer{fp8_str}{attn_mask_str}{fuse_qkv_params_str}{qk_scaling_str}{high_prec_str}.onnx"
fname = f"te.transformer_layer{fp8_str}{attn_mask_str}{fuse_qkv_params_str}{high_prec_str}.onnx"

model = te.TransformerLayer(
hidden_size,
Expand All@@ -959,7 +946,6 @@ def test_export_transformer_layer(
output_layernorm=output_layernorm,
params_dtype=precision,
fuse_qkv_params=fuse_qkv_params,
apply_query_key_layer_scaling=apply_query_key_layer_scaling,
zero_centered_gamma=zero_centered_gamma).to(device='cuda')
do_export(model, inp, fname, use_fp8)
if not use_fp8:
Expand Down
36 changes: 19 additions & 17 deletions transformer_engine/pytorch/softmax.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -4,7 +4,7 @@

"""Fused scaled masked softmax functions"""
import os
from typing import Callable, Tuple, Union
from typing import Callable, Tuple, Union, Optional
import torch
from torch import nn
import torch._C._onnx as _C_onnx
Expand DownExpand Up@@ -198,15 +198,13 @@ class FusedScaleMaskSoftmax(nn.Module):
attn_mask_type: attention mask type (pad or causal)
mask_func: mask function to be applied.
softmax_in_fp32: if true, softmax in performed at fp32 precision.
scale: scaling factor used in input tensor scaling.
"""

def __init__(
self,
attn_mask_type: str,
mask_func: Callable,
softmax_in_fp32: bool,
scale: float,
softmax_in_fp32: bool = True,
Comment thread
ptrendx marked this conversation as resolved.
) -> None:
super().__init__()
self.attn_mask_type = attn_mask_type
Expand All@@ -215,23 +213,27 @@ def __init__(
)
self.mask_func = mask_func
self.softmax_in_fp32 = softmax_in_fp32
self.scale = scale

assert (
self.scale is None or softmax_in_fp32
), "softmax should be in fp32 when scaled"

def forward(self, inp: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
def forward(
self,
inp: torch.Tensor,
mask: torch.Tensor,
scale: Optional[float] = None,
) -> torch.Tensor:
"""FusedScaleMaskSoftmax fprop"""
# [b, np, sq, sk]
assert inp.dim() == 4
self.input_in_fp16 = inp.dtype == torch.float16
self.input_in_bf16 = inp.dtype == torch.bfloat16
self.input_in_float16 = self.input_in_fp16 or self.input_in_bf16

assert (
scale is None or self.softmax_in_fp32
), "softmax should be in fp32 when scaled"

if self.is_kernel_available(*inp.size()):
return self.forward_fused_softmax(inp, mask)
return self.forward_torch_softmax(inp, mask)
return self.forward_fused_softmax(inp, mask, scale)
return self.forward_torch_softmax(inp, mask, scale)

def is_kernel_available(self, b: int, np: int, sq: int, sk: int) -> bool:
"""Check FusedScaleMaskSoftmax kernel availability based on size"""
Expand All@@ -256,11 +258,11 @@ def is_kernel_available(self, b: int, np: int, sq: int, sk: int) -> bool:
return False

def forward_fused_softmax(
self, inp: torch.Tensor, mask: torch.Tensor
self, inp: torch.Tensor, mask: torch.Tensor, scale: Optional[float] = None
) -> torch.Tensor:
"""Fused masked softmax kernel"""
b, np, sq, sk = inp.size()
scale = self.scale if self.scale is not None else 1.0
scale = 1.0 if scale is None else scale

if self.attn_mask_type == "causal":
assert sq == sk, "causal mask is only for self attention"
Expand All@@ -275,14 +277,14 @@ def forward_fused_softmax(
return ScaledSoftmax.apply(inp, scale)

def forward_torch_softmax(
self, inp: torch.Tensor, mask: torch.Tensor
self, inp: torch.Tensor, mask: torch.Tensor, scale: Optional[float] = None
) -> torch.Tensor:
"""Framework softmax"""
if self.input_in_float16 and self.softmax_in_fp32:
inp = inp.float()

if self.scale is not None:
inp = inp * self.scale
if scale is not None:
inp = inp * scale

if self.attn_mask_type == "causal":
mask = _get_default_causal_mask(inp.size()[2])
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { // Force GitHub README to respect dark mode (function() { var style = document.createElement('style'); style.textContent = ' .markdown-body { color-scheme: dark light; } .markdown-body pre { background: #161b22 !important; } .markdown-body code { background: rgba(110, 118, 129, 0.4) !important; } .markdown-body table th, .markdown-body table td { border-color: #30363d !important; } .markdown-body img { background: #0d1117; } .markdown-body blockquote { border-left-color: #8b949e; } .markdown-body hr { border-color: #30363d; } '; document.head.appendChild(style); })(); } } catch(__e) { console.warn('[Userscript:GitHub Dark Mode README Fix]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + ' deprecate qk layer scaling and fp32 softmax args by ksivaman · Pull Request #90 · NVIDIA/TransformerEngine · GitHub
Skip to content
Merged
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
18 changes: 2 additions & 16 deletions tests/pytorch/test_onnx_export.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -756,16 +756,10 @@ def test_export_layernorm_mlp(
(torch.float16, True, "padding"), # calls ScaledMaskedSoftmax
(torch.float16, False, "padding"), # calls ScaledSoftmax
])
@pytest.mark.parametrize("attention_softmax_in_fp32",
[True, False])
@pytest.mark.parametrize("apply_query_key_layer_scaling",
[True, False])
def test_export_core_attention(
precision: torch.dtype,
use_mask: bool,
attn_mask_type: str,
attention_softmax_in_fp32: bool,
apply_query_key_layer_scaling: bool,
):
# Set dimensions (these are arbitrary).
kv_channels = 64
Expand All@@ -784,11 +778,9 @@ def test_export_core_attention(
input_names.append("attention_mask")
inp = (query_layer, key_layer, value_layer, attention_mask)

sm_prec_str = "_sm-fp32" if attention_softmax_in_fp32 else "_sm-fp16"
qk_scaling_str = "_qk-scaling" if apply_query_key_layer_scaling else ""
mask_str = get_attn_mask_str(use_mask, attn_mask_type)
high_prec_str = dtype2str(precision)
fname = f"te.core_attention{mask_str}{qk_scaling_str}{sm_prec_str}{high_prec_str}.onnx"
fname = f"te.core_attention{mask_str}{high_prec_str}.onnx"

if attn_mask_type is None:
attn_mask_type = 'causal'
Expand All@@ -798,8 +790,6 @@ def test_export_core_attention(
kv_channels=kv_channels,
attention_dropout=0.5,
attn_mask_type=attn_mask_type,
attention_softmax_in_fp32=attention_softmax_in_fp32,
apply_query_key_layer_scaling=apply_query_key_layer_scaling,
).to(device='cuda')
do_export(model,
inp,
Expand DownExpand Up@@ -911,7 +901,6 @@ def test_export_multihead_attention(
])
@pytest.mark.parametrize("precision", [torch.float32, torch.float16])
@pytest.mark.parametrize("fuse_qkv_params", [False, True])
@pytest.mark.parametrize("apply_query_key_layer_scaling", [True, False])
@pytest.mark.parametrize("zero_centered_gamma", [False, True])
def test_export_transformer_layer(
use_fp8: bool,
Expand All@@ -920,7 +909,6 @@ def test_export_transformer_layer(
output_layernorm: bool,
precision: torch.dtype,
fuse_qkv_params: bool,
apply_query_key_layer_scaling: bool,
zero_centered_gamma: bool
):
# Skip FP8 tests on non-hopper devices
Expand All@@ -946,10 +934,9 @@ def test_export_transformer_layer(

fp8_str = "_fp8" if use_fp8 else ""
fuse_qkv_params_str = "_fused-qkv" if fuse_qkv_params else ""
qk_scaling_str = "_qk-scaling" if apply_query_key_layer_scaling else ""
high_prec_str = dtype2str(precision)
attn_mask_str = get_attn_mask_str(use_mask, attn_mask_type)
fname = f"te.transformer_layer{fp8_str}{attn_mask_str}{fuse_qkv_params_str}{qk_scaling_str}{high_prec_str}.onnx"
fname = f"te.transformer_layer{fp8_str}{attn_mask_str}{fuse_qkv_params_str}{high_prec_str}.onnx"

model = te.TransformerLayer(
hidden_size,
Expand All@@ -959,7 +946,6 @@ def test_export_transformer_layer(
output_layernorm=output_layernorm,
params_dtype=precision,
fuse_qkv_params=fuse_qkv_params,
apply_query_key_layer_scaling=apply_query_key_layer_scaling,
zero_centered_gamma=zero_centered_gamma).to(device='cuda')
do_export(model, inp, fname, use_fp8)
if not use_fp8:
Expand Down
36 changes: 19 additions & 17 deletions transformer_engine/pytorch/softmax.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -4,7 +4,7 @@

"""Fused scaled masked softmax functions"""
import os
from typing import Callable, Tuple, Union
from typing import Callable, Tuple, Union, Optional
import torch
from torch import nn
import torch._C._onnx as _C_onnx
Expand DownExpand Up@@ -198,15 +198,13 @@ class FusedScaleMaskSoftmax(nn.Module):
attn_mask_type: attention mask type (pad or causal)
mask_func: mask function to be applied.
softmax_in_fp32: if true, softmax in performed at fp32 precision.
scale: scaling factor used in input tensor scaling.
"""

def __init__(
self,
attn_mask_type: str,
mask_func: Callable,
softmax_in_fp32: bool,
scale: float,
softmax_in_fp32: bool = True,
Comment thread
ptrendx marked this conversation as resolved.
) -> None:
super().__init__()
self.attn_mask_type = attn_mask_type
Expand All@@ -215,23 +213,27 @@ def __init__(
)
self.mask_func = mask_func
self.softmax_in_fp32 = softmax_in_fp32
self.scale = scale

assert (
self.scale is None or softmax_in_fp32
), "softmax should be in fp32 when scaled"

def forward(self, inp: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
def forward(
self,
inp: torch.Tensor,
mask: torch.Tensor,
scale: Optional[float] = None,
) -> torch.Tensor:
"""FusedScaleMaskSoftmax fprop"""
# [b, np, sq, sk]
assert inp.dim() == 4
self.input_in_fp16 = inp.dtype == torch.float16
self.input_in_bf16 = inp.dtype == torch.bfloat16
self.input_in_float16 = self.input_in_fp16 or self.input_in_bf16

assert (
scale is None or self.softmax_in_fp32
), "softmax should be in fp32 when scaled"

if self.is_kernel_available(*inp.size()):
return self.forward_fused_softmax(inp, mask)
return self.forward_torch_softmax(inp, mask)
return self.forward_fused_softmax(inp, mask, scale)
return self.forward_torch_softmax(inp, mask, scale)

def is_kernel_available(self, b: int, np: int, sq: int, sk: int) -> bool:
"""Check FusedScaleMaskSoftmax kernel availability based on size"""
Expand All@@ -256,11 +258,11 @@ def is_kernel_available(self, b: int, np: int, sq: int, sk: int) -> bool:
return False

def forward_fused_softmax(
self, inp: torch.Tensor, mask: torch.Tensor
self, inp: torch.Tensor, mask: torch.Tensor, scale: Optional[float] = None
) -> torch.Tensor:
"""Fused masked softmax kernel"""
b, np, sq, sk = inp.size()
scale = self.scale if self.scale is not None else 1.0
scale = 1.0 if scale is None else scale

if self.attn_mask_type == "causal":
assert sq == sk, "causal mask is only for self attention"
Expand All@@ -275,14 +277,14 @@ def forward_fused_softmax(
return ScaledSoftmax.apply(inp, scale)

def forward_torch_softmax(
self, inp: torch.Tensor, mask: torch.Tensor
self, inp: torch.Tensor, mask: torch.Tensor, scale: Optional[float] = None
) -> torch.Tensor:
"""Framework softmax"""
if self.input_in_float16 and self.softmax_in_fp32:
inp = inp.float()

if self.scale is not None:
inp = inp * self.scale
if scale is not None:
inp = inp * scale

if self.attn_mask_type == "causal":
mask = _get_default_causal_mask(inp.size()[2])
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { // Highlight search terms from Google/DuckDuckGo/Bing referrer (function() { var ref = document.referrer; var terms = []; if (ref.includes('google.com') || ref.includes('duckduckgo.com') || ref.includes('bing.com')) { var url = new URL(ref); var q = url.searchParams.get('q') || url.searchParams.get('p'); if (q) { terms = q.split(/\s+/).filter(function(t) { return t.length > 2; }); } } if (terms.length === 0) return; var style = document.createElement('style'); style.textContent = '.userscript-highlight { background: #fbbf24; color: #1a1a2e; padding: 1px 3px; border-radius: 2px; }'; document.head.appendChild(style); function highlight(node) { if (node.nodeType === 3) { // text node var text = node.textContent; var found = false; terms.forEach(function(term) { var regex = new RegExp('(' + term.replace(/[.*+?^${}()|[\]\\]/g, '\\') + ')', 'gi'); if (regex.test(text)) { found = true; var frag = document.createDocumentFragment(); var parts = text.split(regex); parts.forEach(function(part, i) { if (i % 2 === 0) { frag.appendChild(document.createTextNode(part)); } else { var span = document.createElement('span'); span.className = 'userscript-highlight'; span.textContent = part; frag.appendChild(span); } }); node.parentNode.replaceChild(frag, node); } }); } else if (node.nodeType === 1 && node.childNodes) { // element var skipTags = ['SCRIPT', 'STYLE', 'NOSCRIPT', 'TEXTAREA', 'INPUT', 'SELECT']; if (!skipTags.includes(node.tagName)) { Array.from(node.childNodes).forEach(highlight); } } } highlight(document.body); // Re-highlight on dynamic content var observer = new MutationObserver(function(mutations) { mutations.forEach(function(m) { m.addedNodes.forEach(function(node) { if (node.nodeType === 1 || node.nodeType === 3) highlight(node); }); }); }); observer.observe(document.body, { childList: true, subtree: true }); })(); } } catch(__e) { console.warn('[Userscript:Highlight Search Terms]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + ' deprecate qk layer scaling and fp32 softmax args by ksivaman · Pull Request #90 · NVIDIA/TransformerEngine · GitHub
Skip to content
Merged
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
18 changes: 2 additions & 16 deletions tests/pytorch/test_onnx_export.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -756,16 +756,10 @@ def test_export_layernorm_mlp(
(torch.float16, True, "padding"), # calls ScaledMaskedSoftmax
(torch.float16, False, "padding"), # calls ScaledSoftmax
])
@pytest.mark.parametrize("attention_softmax_in_fp32",
[True, False])
@pytest.mark.parametrize("apply_query_key_layer_scaling",
[True, False])
def test_export_core_attention(
precision: torch.dtype,
use_mask: bool,
attn_mask_type: str,
attention_softmax_in_fp32: bool,
apply_query_key_layer_scaling: bool,
):
# Set dimensions (these are arbitrary).
kv_channels = 64
Expand All@@ -784,11 +778,9 @@ def test_export_core_attention(
input_names.append("attention_mask")
inp = (query_layer, key_layer, value_layer, attention_mask)

sm_prec_str = "_sm-fp32" if attention_softmax_in_fp32 else "_sm-fp16"
qk_scaling_str = "_qk-scaling" if apply_query_key_layer_scaling else ""
mask_str = get_attn_mask_str(use_mask, attn_mask_type)
high_prec_str = dtype2str(precision)
fname = f"te.core_attention{mask_str}{qk_scaling_str}{sm_prec_str}{high_prec_str}.onnx"
fname = f"te.core_attention{mask_str}{high_prec_str}.onnx"

if attn_mask_type is None:
attn_mask_type = 'causal'
Expand All@@ -798,8 +790,6 @@ def test_export_core_attention(
kv_channels=kv_channels,
attention_dropout=0.5,
attn_mask_type=attn_mask_type,
attention_softmax_in_fp32=attention_softmax_in_fp32,
apply_query_key_layer_scaling=apply_query_key_layer_scaling,
).to(device='cuda')
do_export(model,
inp,
Expand DownExpand Up@@ -911,7 +901,6 @@ def test_export_multihead_attention(
])
@pytest.mark.parametrize("precision", [torch.float32, torch.float16])
@pytest.mark.parametrize("fuse_qkv_params", [False, True])
@pytest.mark.parametrize("apply_query_key_layer_scaling", [True, False])
@pytest.mark.parametrize("zero_centered_gamma", [False, True])
def test_export_transformer_layer(
use_fp8: bool,
Expand All@@ -920,7 +909,6 @@ def test_export_transformer_layer(
output_layernorm: bool,
precision: torch.dtype,
fuse_qkv_params: bool,
apply_query_key_layer_scaling: bool,
zero_centered_gamma: bool
):
# Skip FP8 tests on non-hopper devices
Expand All@@ -946,10 +934,9 @@ def test_export_transformer_layer(

fp8_str = "_fp8" if use_fp8 else ""
fuse_qkv_params_str = "_fused-qkv" if fuse_qkv_params else ""
qk_scaling_str = "_qk-scaling" if apply_query_key_layer_scaling else ""
high_prec_str = dtype2str(precision)
attn_mask_str = get_attn_mask_str(use_mask, attn_mask_type)
fname = f"te.transformer_layer{fp8_str}{attn_mask_str}{fuse_qkv_params_str}{qk_scaling_str}{high_prec_str}.onnx"
fname = f"te.transformer_layer{fp8_str}{attn_mask_str}{fuse_qkv_params_str}{high_prec_str}.onnx"

model = te.TransformerLayer(
hidden_size,
Expand All@@ -959,7 +946,6 @@ def test_export_transformer_layer(
output_layernorm=output_layernorm,
params_dtype=precision,
fuse_qkv_params=fuse_qkv_params,
apply_query_key_layer_scaling=apply_query_key_layer_scaling,
zero_centered_gamma=zero_centered_gamma).to(device='cuda')
do_export(model, inp, fname, use_fp8)
if not use_fp8:
Expand Down
36 changes: 19 additions & 17 deletions transformer_engine/pytorch/softmax.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -4,7 +4,7 @@

"""Fused scaled masked softmax functions"""
import os
from typing import Callable, Tuple, Union
from typing import Callable, Tuple, Union, Optional
import torch
from torch import nn
import torch._C._onnx as _C_onnx
Expand DownExpand Up@@ -198,15 +198,13 @@ class FusedScaleMaskSoftmax(nn.Module):
attn_mask_type: attention mask type (pad or causal)
mask_func: mask function to be applied.
softmax_in_fp32: if true, softmax in performed at fp32 precision.
scale: scaling factor used in input tensor scaling.
"""

def __init__(
self,
attn_mask_type: str,
mask_func: Callable,
softmax_in_fp32: bool,
scale: float,
softmax_in_fp32: bool = True,
Comment thread
ptrendx marked this conversation as resolved.
) -> None:
super().__init__()
self.attn_mask_type = attn_mask_type
Expand All@@ -215,23 +213,27 @@ def __init__(
)
self.mask_func = mask_func
self.softmax_in_fp32 = softmax_in_fp32
self.scale = scale

assert (
self.scale is None or softmax_in_fp32
), "softmax should be in fp32 when scaled"

def forward(self, inp: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
def forward(
self,
inp: torch.Tensor,
mask: torch.Tensor,
scale: Optional[float] = None,
) -> torch.Tensor:
"""FusedScaleMaskSoftmax fprop"""
# [b, np, sq, sk]
assert inp.dim() == 4
self.input_in_fp16 = inp.dtype == torch.float16
self.input_in_bf16 = inp.dtype == torch.bfloat16
self.input_in_float16 = self.input_in_fp16 or self.input_in_bf16

assert (
scale is None or self.softmax_in_fp32
), "softmax should be in fp32 when scaled"

if self.is_kernel_available(*inp.size()):
return self.forward_fused_softmax(inp, mask)
return self.forward_torch_softmax(inp, mask)
return self.forward_fused_softmax(inp, mask, scale)
return self.forward_torch_softmax(inp, mask, scale)

def is_kernel_available(self, b: int, np: int, sq: int, sk: int) -> bool:
"""Check FusedScaleMaskSoftmax kernel availability based on size"""
Expand All@@ -256,11 +258,11 @@ def is_kernel_available(self, b: int, np: int, sq: int, sk: int) -> bool:
return False

def forward_fused_softmax(
self, inp: torch.Tensor, mask: torch.Tensor
self, inp: torch.Tensor, mask: torch.Tensor, scale: Optional[float] = None
) -> torch.Tensor:
"""Fused masked softmax kernel"""
b, np, sq, sk = inp.size()
scale = self.scale if self.scale is not None else 1.0
scale = 1.0 if scale is None else scale

if self.attn_mask_type == "causal":
assert sq == sk, "causal mask is only for self attention"
Expand All@@ -275,14 +277,14 @@ def forward_fused_softmax(
return ScaledSoftmax.apply(inp, scale)

def forward_torch_softmax(
self, inp: torch.Tensor, mask: torch.Tensor
self, inp: torch.Tensor, mask: torch.Tensor, scale: Optional[float] = None
) -> torch.Tensor:
"""Framework softmax"""
if self.input_in_float16 and self.softmax_in_fp32:
inp = inp.float()

if self.scale is not None:
inp = inp * self.scale
if scale is not None:
inp = inp * scale

if self.attn_mask_type == "causal":
mask = _get_default_causal_mask(inp.size()[2])
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { // Strip utm_, fbclid, gclid, etc. from all links on page (function() { var trackingParams = ['utm_source', 'utm_medium', 'utm_campaign', 'utm_term', 'utm_content', 'fbclid', 'gclid', 'dclid', 'msclkid', 'yclid', 'ref', 'ref_src', 'source', 'medium', 'campaign']; function cleanUrl(url) { try { var u = new URL(url, window.location.origin); var changed = false; trackingParams.forEach(function(p) { if (u.searchParams.has(p)) { u.searchParams.delete(p); changed = true; } }); return changed ? u.toString() : url; } catch (e) { return url; } } function cleanLinks() { document.querySelectorAll('a[href]').forEach(function(a) { var clean = cleanUrl(a.href); if (clean !== a.href) a.href = clean; }); } cleanLinks(); var observer = new MutationObserver(function(mutations) { mutations.forEach(function(m) { m.addedNodes.forEach(function(node) { if (node.nodeType === 1) { if (node.tagName === 'A') cleanLinks(); node.querySelectorAll('a[href]').forEach(function(a) { var clean = cleanUrl(a.href); if (clean !== a.href) a.href = clean; }); } }); }); }); observer.observe(document.body, { childList: true, subtree: true }); })(); } } catch(__e) { console.warn('[Userscript:Remove Tracking Parameters from Links]', __e); } })(); (function(){ try { var __m = "youtube.com"; var __re = new RegExp('^' + "youtube\\.com" + ' deprecate qk layer scaling and fp32 softmax args by ksivaman · Pull Request #90 · NVIDIA/TransformerEngine · GitHub
Skip to content
Merged
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
18 changes: 2 additions & 16 deletions tests/pytorch/test_onnx_export.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -756,16 +756,10 @@ def test_export_layernorm_mlp(
(torch.float16, True, "padding"), # calls ScaledMaskedSoftmax
(torch.float16, False, "padding"), # calls ScaledSoftmax
])
@pytest.mark.parametrize("attention_softmax_in_fp32",
[True, False])
@pytest.mark.parametrize("apply_query_key_layer_scaling",
[True, False])
def test_export_core_attention(
precision: torch.dtype,
use_mask: bool,
attn_mask_type: str,
attention_softmax_in_fp32: bool,
apply_query_key_layer_scaling: bool,
):
# Set dimensions (these are arbitrary).
kv_channels = 64
Expand All@@ -784,11 +778,9 @@ def test_export_core_attention(
input_names.append("attention_mask")
inp = (query_layer, key_layer, value_layer, attention_mask)

sm_prec_str = "_sm-fp32" if attention_softmax_in_fp32 else "_sm-fp16"
qk_scaling_str = "_qk-scaling" if apply_query_key_layer_scaling else ""
mask_str = get_attn_mask_str(use_mask, attn_mask_type)
high_prec_str = dtype2str(precision)
fname = f"te.core_attention{mask_str}{qk_scaling_str}{sm_prec_str}{high_prec_str}.onnx"
fname = f"te.core_attention{mask_str}{high_prec_str}.onnx"

if attn_mask_type is None:
attn_mask_type = 'causal'
Expand All@@ -798,8 +790,6 @@ def test_export_core_attention(
kv_channels=kv_channels,
attention_dropout=0.5,
attn_mask_type=attn_mask_type,
attention_softmax_in_fp32=attention_softmax_in_fp32,
apply_query_key_layer_scaling=apply_query_key_layer_scaling,
).to(device='cuda')
do_export(model,
inp,
Expand DownExpand Up@@ -911,7 +901,6 @@ def test_export_multihead_attention(
])
@pytest.mark.parametrize("precision", [torch.float32, torch.float16])
@pytest.mark.parametrize("fuse_qkv_params", [False, True])
@pytest.mark.parametrize("apply_query_key_layer_scaling", [True, False])
@pytest.mark.parametrize("zero_centered_gamma", [False, True])
def test_export_transformer_layer(
use_fp8: bool,
Expand All@@ -920,7 +909,6 @@ def test_export_transformer_layer(
output_layernorm: bool,
precision: torch.dtype,
fuse_qkv_params: bool,
apply_query_key_layer_scaling: bool,
zero_centered_gamma: bool
):
# Skip FP8 tests on non-hopper devices
Expand All@@ -946,10 +934,9 @@ def test_export_transformer_layer(

fp8_str = "_fp8" if use_fp8 else ""
fuse_qkv_params_str = "_fused-qkv" if fuse_qkv_params else ""
qk_scaling_str = "_qk-scaling" if apply_query_key_layer_scaling else ""
high_prec_str = dtype2str(precision)
attn_mask_str = get_attn_mask_str(use_mask, attn_mask_type)
fname = f"te.transformer_layer{fp8_str}{attn_mask_str}{fuse_qkv_params_str}{qk_scaling_str}{high_prec_str}.onnx"
fname = f"te.transformer_layer{fp8_str}{attn_mask_str}{fuse_qkv_params_str}{high_prec_str}.onnx"

model = te.TransformerLayer(
hidden_size,
Expand All@@ -959,7 +946,6 @@ def test_export_transformer_layer(
output_layernorm=output_layernorm,
params_dtype=precision,
fuse_qkv_params=fuse_qkv_params,
apply_query_key_layer_scaling=apply_query_key_layer_scaling,
zero_centered_gamma=zero_centered_gamma).to(device='cuda')
do_export(model, inp, fname, use_fp8)
if not use_fp8:
Expand Down
36 changes: 19 additions & 17 deletions transformer_engine/pytorch/softmax.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -4,7 +4,7 @@

"""Fused scaled masked softmax functions"""
import os
from typing import Callable, Tuple, Union
from typing import Callable, Tuple, Union, Optional
import torch
from torch import nn
import torch._C._onnx as _C_onnx
Expand DownExpand Up@@ -198,15 +198,13 @@ class FusedScaleMaskSoftmax(nn.Module):
attn_mask_type: attention mask type (pad or causal)
mask_func: mask function to be applied.
softmax_in_fp32: if true, softmax in performed at fp32 precision.
scale: scaling factor used in input tensor scaling.
"""

def __init__(
self,
attn_mask_type: str,
mask_func: Callable,
softmax_in_fp32: bool,
scale: float,
softmax_in_fp32: bool = True,
Comment thread
ptrendx marked this conversation as resolved.
) -> None:
super().__init__()
self.attn_mask_type = attn_mask_type
Expand All@@ -215,23 +213,27 @@ def __init__(
)
self.mask_func = mask_func
self.softmax_in_fp32 = softmax_in_fp32
self.scale = scale

assert (
self.scale is None or softmax_in_fp32
), "softmax should be in fp32 when scaled"

def forward(self, inp: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
def forward(
self,
inp: torch.Tensor,
mask: torch.Tensor,
scale: Optional[float] = None,
) -> torch.Tensor:
"""FusedScaleMaskSoftmax fprop"""
# [b, np, sq, sk]
assert inp.dim() == 4
self.input_in_fp16 = inp.dtype == torch.float16
self.input_in_bf16 = inp.dtype == torch.bfloat16
self.input_in_float16 = self.input_in_fp16 or self.input_in_bf16

assert (
scale is None or self.softmax_in_fp32
), "softmax should be in fp32 when scaled"

if self.is_kernel_available(*inp.size()):
return self.forward_fused_softmax(inp, mask)
return self.forward_torch_softmax(inp, mask)
return self.forward_fused_softmax(inp, mask, scale)
return self.forward_torch_softmax(inp, mask, scale)

def is_kernel_available(self, b: int, np: int, sq: int, sk: int) -> bool:
"""Check FusedScaleMaskSoftmax kernel availability based on size"""
Expand All@@ -256,11 +258,11 @@ def is_kernel_available(self, b: int, np: int, sq: int, sk: int) -> bool:
return False

def forward_fused_softmax(
self, inp: torch.Tensor, mask: torch.Tensor
self, inp: torch.Tensor, mask: torch.Tensor, scale: Optional[float] = None
) -> torch.Tensor:
"""Fused masked softmax kernel"""
b, np, sq, sk = inp.size()
scale = self.scale if self.scale is not None else 1.0
scale = 1.0 if scale is None else scale

if self.attn_mask_type == "causal":
assert sq == sk, "causal mask is only for self attention"
Expand All@@ -275,14 +277,14 @@ def forward_fused_softmax(
return ScaledSoftmax.apply(inp, scale)

def forward_torch_softmax(
self, inp: torch.Tensor, mask: torch.Tensor
self, inp: torch.Tensor, mask: torch.Tensor, scale: Optional[float] = None
) -> torch.Tensor:
"""Framework softmax"""
if self.input_in_float16 and self.softmax_in_fp32:
inp = inp.float()

if self.scale is not None:
inp = inp * self.scale
if scale is not None:
inp = inp * scale

if self.attn_mask_type == "causal":
mask = _get_default_causal_mask(inp.size()[2])
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { // Auto-enable theater mode on YouTube (function() { function tryTheater() { var btn = document.querySelector('button[aria-label="Theater mode"], ytd-player #player button[title="Theater mode"]'); if (btn && !btn.classList.contains('activated')) { btn.click(); } } // Try immediately tryTheater(); // Try after navigation (SPA) var lastUrl = location.href; setInterval(function() { if (location.href !== lastUrl) { lastUrl = location.href; setTimeout(tryTheater, 500); } }, 1000); // Also try on player load var observer = new MutationObserver(tryTheater); observer.observe(document.body, { childList: true, subtree: true }); })(); } } catch(__e) { console.warn('[Userscript:YouTube Theater Mode Default]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + ' deprecate qk layer scaling and fp32 softmax args by ksivaman · Pull Request #90 · NVIDIA/TransformerEngine · GitHub
Skip to content
Merged
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
18 changes: 2 additions & 16 deletions tests/pytorch/test_onnx_export.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -756,16 +756,10 @@ def test_export_layernorm_mlp(
(torch.float16, True, "padding"), # calls ScaledMaskedSoftmax
(torch.float16, False, "padding"), # calls ScaledSoftmax
])
@pytest.mark.parametrize("attention_softmax_in_fp32",
[True, False])
@pytest.mark.parametrize("apply_query_key_layer_scaling",
[True, False])
def test_export_core_attention(
precision: torch.dtype,
use_mask: bool,
attn_mask_type: str,
attention_softmax_in_fp32: bool,
apply_query_key_layer_scaling: bool,
):
# Set dimensions (these are arbitrary).
kv_channels = 64
Expand All@@ -784,11 +778,9 @@ def test_export_core_attention(
input_names.append("attention_mask")
inp = (query_layer, key_layer, value_layer, attention_mask)

sm_prec_str = "_sm-fp32" if attention_softmax_in_fp32 else "_sm-fp16"
qk_scaling_str = "_qk-scaling" if apply_query_key_layer_scaling else ""
mask_str = get_attn_mask_str(use_mask, attn_mask_type)
high_prec_str = dtype2str(precision)
fname = f"te.core_attention{mask_str}{qk_scaling_str}{sm_prec_str}{high_prec_str}.onnx"
fname = f"te.core_attention{mask_str}{high_prec_str}.onnx"

if attn_mask_type is None:
attn_mask_type = 'causal'
Expand All@@ -798,8 +790,6 @@ def test_export_core_attention(
kv_channels=kv_channels,
attention_dropout=0.5,
attn_mask_type=attn_mask_type,
attention_softmax_in_fp32=attention_softmax_in_fp32,
apply_query_key_layer_scaling=apply_query_key_layer_scaling,
).to(device='cuda')
do_export(model,
inp,
Expand DownExpand Up@@ -911,7 +901,6 @@ def test_export_multihead_attention(
])
@pytest.mark.parametrize("precision", [torch.float32, torch.float16])
@pytest.mark.parametrize("fuse_qkv_params", [False, True])
@pytest.mark.parametrize("apply_query_key_layer_scaling", [True, False])
@pytest.mark.parametrize("zero_centered_gamma", [False, True])
def test_export_transformer_layer(
use_fp8: bool,
Expand All@@ -920,7 +909,6 @@ def test_export_transformer_layer(
output_layernorm: bool,
precision: torch.dtype,
fuse_qkv_params: bool,
apply_query_key_layer_scaling: bool,
zero_centered_gamma: bool
):
# Skip FP8 tests on non-hopper devices
Expand All@@ -946,10 +934,9 @@ def test_export_transformer_layer(

fp8_str = "_fp8" if use_fp8 else ""
fuse_qkv_params_str = "_fused-qkv" if fuse_qkv_params else ""
qk_scaling_str = "_qk-scaling" if apply_query_key_layer_scaling else ""
high_prec_str = dtype2str(precision)
attn_mask_str = get_attn_mask_str(use_mask, attn_mask_type)
fname = f"te.transformer_layer{fp8_str}{attn_mask_str}{fuse_qkv_params_str}{qk_scaling_str}{high_prec_str}.onnx"
fname = f"te.transformer_layer{fp8_str}{attn_mask_str}{fuse_qkv_params_str}{high_prec_str}.onnx"

model = te.TransformerLayer(
hidden_size,
Expand All@@ -959,7 +946,6 @@ def test_export_transformer_layer(
output_layernorm=output_layernorm,
params_dtype=precision,
fuse_qkv_params=fuse_qkv_params,
apply_query_key_layer_scaling=apply_query_key_layer_scaling,
zero_centered_gamma=zero_centered_gamma).to(device='cuda')
do_export(model, inp, fname, use_fp8)
if not use_fp8:
Expand Down
36 changes: 19 additions & 17 deletions transformer_engine/pytorch/softmax.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -4,7 +4,7 @@

"""Fused scaled masked softmax functions"""
import os
from typing import Callable, Tuple, Union
from typing import Callable, Tuple, Union, Optional
import torch
from torch import nn
import torch._C._onnx as _C_onnx
Expand DownExpand Up@@ -198,15 +198,13 @@ class FusedScaleMaskSoftmax(nn.Module):
attn_mask_type: attention mask type (pad or causal)
mask_func: mask function to be applied.
softmax_in_fp32: if true, softmax in performed at fp32 precision.
scale: scaling factor used in input tensor scaling.
"""

def __init__(
self,
attn_mask_type: str,
mask_func: Callable,
softmax_in_fp32: bool,
scale: float,
softmax_in_fp32: bool = True,
Comment thread
ptrendx marked this conversation as resolved.
) -> None:
super().__init__()
self.attn_mask_type = attn_mask_type
Expand All@@ -215,23 +213,27 @@ def __init__(
)
self.mask_func = mask_func
self.softmax_in_fp32 = softmax_in_fp32
self.scale = scale

assert (
self.scale is None or softmax_in_fp32
), "softmax should be in fp32 when scaled"

def forward(self, inp: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
def forward(
self,
inp: torch.Tensor,
mask: torch.Tensor,
scale: Optional[float] = None,
) -> torch.Tensor:
"""FusedScaleMaskSoftmax fprop"""
# [b, np, sq, sk]
assert inp.dim() == 4
self.input_in_fp16 = inp.dtype == torch.float16
self.input_in_bf16 = inp.dtype == torch.bfloat16
self.input_in_float16 = self.input_in_fp16 or self.input_in_bf16

assert (
scale is None or self.softmax_in_fp32
), "softmax should be in fp32 when scaled"

if self.is_kernel_available(*inp.size()):
return self.forward_fused_softmax(inp, mask)
return self.forward_torch_softmax(inp, mask)
return self.forward_fused_softmax(inp, mask, scale)
return self.forward_torch_softmax(inp, mask, scale)

def is_kernel_available(self, b: int, np: int, sq: int, sk: int) -> bool:
"""Check FusedScaleMaskSoftmax kernel availability based on size"""
Expand All@@ -256,11 +258,11 @@ def is_kernel_available(self, b: int, np: int, sq: int, sk: int) -> bool:
return False

def forward_fused_softmax(
self, inp: torch.Tensor, mask: torch.Tensor
self, inp: torch.Tensor, mask: torch.Tensor, scale: Optional[float] = None
) -> torch.Tensor:
"""Fused masked softmax kernel"""
b, np, sq, sk = inp.size()
scale = self.scale if self.scale is not None else 1.0
scale = 1.0 if scale is None else scale

if self.attn_mask_type == "causal":
assert sq == sk, "causal mask is only for self attention"
Expand All@@ -275,14 +277,14 @@ def forward_fused_softmax(
return ScaledSoftmax.apply(inp, scale)

def forward_torch_softmax(
self, inp: torch.Tensor, mask: torch.Tensor
self, inp: torch.Tensor, mask: torch.Tensor, scale: Optional[float] = None
) -> torch.Tensor:
"""Framework softmax"""
if self.input_in_float16 and self.softmax_in_fp32:
inp = inp.float()

if self.scale is not None:
inp = inp * self.scale
if scale is not None:
inp = inp * scale

if self.attn_mask_type == "causal":
mask = _get_default_causal_mask(inp.size()[2])
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { // Remove or un-stick sticky/fixed headers that block content (function() { function unstick() { document.querySelectorAll('header, nav, [role="banner"], .header, .navbar, .sticky, .fixed-top, [style*="position: fixed"], [style*="position:sticky"]').forEach(function(el) { if (el.style.position === 'fixed' || el.style.position === 'sticky' || getComputedStyle(el).position === 'fixed' || getComputedStyle(el).position === 'sticky') { el.style.position = 'static'; el.style.top = 'auto'; el.style.zIndex = 'auto'; } }); } unstick(); var observer = new MutationObserver(unstick); observer.observe(document.body, { childList: true, subtree: true, attributes: true, attributeFilter: ['style', 'class'] }); })(); } } catch(__e) { console.warn('[Userscript:Kill Sticky Headers]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + ' deprecate qk layer scaling and fp32 softmax args by ksivaman · Pull Request #90 · NVIDIA/TransformerEngine · GitHub
Skip to content
Merged
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
18 changes: 2 additions & 16 deletions tests/pytorch/test_onnx_export.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -756,16 +756,10 @@ def test_export_layernorm_mlp(
(torch.float16, True, "padding"), # calls ScaledMaskedSoftmax
(torch.float16, False, "padding"), # calls ScaledSoftmax
])
@pytest.mark.parametrize("attention_softmax_in_fp32",
[True, False])
@pytest.mark.parametrize("apply_query_key_layer_scaling",
[True, False])
def test_export_core_attention(
precision: torch.dtype,
use_mask: bool,
attn_mask_type: str,
attention_softmax_in_fp32: bool,
apply_query_key_layer_scaling: bool,
):
# Set dimensions (these are arbitrary).
kv_channels = 64
Expand All@@ -784,11 +778,9 @@ def test_export_core_attention(
input_names.append("attention_mask")
inp = (query_layer, key_layer, value_layer, attention_mask)

sm_prec_str = "_sm-fp32" if attention_softmax_in_fp32 else "_sm-fp16"
qk_scaling_str = "_qk-scaling" if apply_query_key_layer_scaling else ""
mask_str = get_attn_mask_str(use_mask, attn_mask_type)
high_prec_str = dtype2str(precision)
fname = f"te.core_attention{mask_str}{qk_scaling_str}{sm_prec_str}{high_prec_str}.onnx"
fname = f"te.core_attention{mask_str}{high_prec_str}.onnx"

if attn_mask_type is None:
attn_mask_type = 'causal'
Expand All@@ -798,8 +790,6 @@ def test_export_core_attention(
kv_channels=kv_channels,
attention_dropout=0.5,
attn_mask_type=attn_mask_type,
attention_softmax_in_fp32=attention_softmax_in_fp32,
apply_query_key_layer_scaling=apply_query_key_layer_scaling,
).to(device='cuda')
do_export(model,
inp,
Expand DownExpand Up@@ -911,7 +901,6 @@ def test_export_multihead_attention(
])
@pytest.mark.parametrize("precision", [torch.float32, torch.float16])
@pytest.mark.parametrize("fuse_qkv_params", [False, True])
@pytest.mark.parametrize("apply_query_key_layer_scaling", [True, False])
@pytest.mark.parametrize("zero_centered_gamma", [False, True])
def test_export_transformer_layer(
use_fp8: bool,
Expand All@@ -920,7 +909,6 @@ def test_export_transformer_layer(
output_layernorm: bool,
precision: torch.dtype,
fuse_qkv_params: bool,
apply_query_key_layer_scaling: bool,
zero_centered_gamma: bool
):
# Skip FP8 tests on non-hopper devices
Expand All@@ -946,10 +934,9 @@ def test_export_transformer_layer(

fp8_str = "_fp8" if use_fp8 else ""
fuse_qkv_params_str = "_fused-qkv" if fuse_qkv_params else ""
qk_scaling_str = "_qk-scaling" if apply_query_key_layer_scaling else ""
high_prec_str = dtype2str(precision)
attn_mask_str = get_attn_mask_str(use_mask, attn_mask_type)
fname = f"te.transformer_layer{fp8_str}{attn_mask_str}{fuse_qkv_params_str}{qk_scaling_str}{high_prec_str}.onnx"
fname = f"te.transformer_layer{fp8_str}{attn_mask_str}{fuse_qkv_params_str}{high_prec_str}.onnx"

model = te.TransformerLayer(
hidden_size,
Expand All@@ -959,7 +946,6 @@ def test_export_transformer_layer(
output_layernorm=output_layernorm,
params_dtype=precision,
fuse_qkv_params=fuse_qkv_params,
apply_query_key_layer_scaling=apply_query_key_layer_scaling,
zero_centered_gamma=zero_centered_gamma).to(device='cuda')
do_export(model, inp, fname, use_fp8)
if not use_fp8:
Expand Down
36 changes: 19 additions & 17 deletions transformer_engine/pytorch/softmax.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -4,7 +4,7 @@

"""Fused scaled masked softmax functions"""
import os
from typing import Callable, Tuple, Union
from typing import Callable, Tuple, Union, Optional
import torch
from torch import nn
import torch._C._onnx as _C_onnx
Expand DownExpand Up@@ -198,15 +198,13 @@ class FusedScaleMaskSoftmax(nn.Module):
attn_mask_type: attention mask type (pad or causal)
mask_func: mask function to be applied.
softmax_in_fp32: if true, softmax in performed at fp32 precision.
scale: scaling factor used in input tensor scaling.
"""

def __init__(
self,
attn_mask_type: str,
mask_func: Callable,
softmax_in_fp32: bool,
scale: float,
softmax_in_fp32: bool = True,
Comment thread
ptrendx marked this conversation as resolved.
) -> None:
super().__init__()
self.attn_mask_type = attn_mask_type
Expand All@@ -215,23 +213,27 @@ def __init__(
)
self.mask_func = mask_func
self.softmax_in_fp32 = softmax_in_fp32
self.scale = scale

assert (
self.scale is None or softmax_in_fp32
), "softmax should be in fp32 when scaled"

def forward(self, inp: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
def forward(
self,
inp: torch.Tensor,
mask: torch.Tensor,
scale: Optional[float] = None,
) -> torch.Tensor:
"""FusedScaleMaskSoftmax fprop"""
# [b, np, sq, sk]
assert inp.dim() == 4
self.input_in_fp16 = inp.dtype == torch.float16
self.input_in_bf16 = inp.dtype == torch.bfloat16
self.input_in_float16 = self.input_in_fp16 or self.input_in_bf16

assert (
scale is None or self.softmax_in_fp32
), "softmax should be in fp32 when scaled"

if self.is_kernel_available(*inp.size()):
return self.forward_fused_softmax(inp, mask)
return self.forward_torch_softmax(inp, mask)
return self.forward_fused_softmax(inp, mask, scale)
return self.forward_torch_softmax(inp, mask, scale)

def is_kernel_available(self, b: int, np: int, sq: int, sk: int) -> bool:
"""Check FusedScaleMaskSoftmax kernel availability based on size"""
Expand All@@ -256,11 +258,11 @@ def is_kernel_available(self, b: int, np: int, sq: int, sk: int) -> bool:
return False

def forward_fused_softmax(
self, inp: torch.Tensor, mask: torch.Tensor
self, inp: torch.Tensor, mask: torch.Tensor, scale: Optional[float] = None
) -> torch.Tensor:
"""Fused masked softmax kernel"""
b, np, sq, sk = inp.size()
scale = self.scale if self.scale is not None else 1.0
scale = 1.0 if scale is None else scale

if self.attn_mask_type == "causal":
assert sq == sk, "causal mask is only for self attention"
Expand All@@ -275,14 +277,14 @@ def forward_fused_softmax(
return ScaledSoftmax.apply(inp, scale)

def forward_torch_softmax(
self, inp: torch.Tensor, mask: torch.Tensor
self, inp: torch.Tensor, mask: torch.Tensor, scale: Optional[float] = None
) -> torch.Tensor:
"""Framework softmax"""
if self.input_in_float16 and self.softmax_in_fp32:
inp = inp.float()

if self.scale is not None:
inp = inp * self.scale
if scale is not None:
inp = inp * scale

if self.attn_mask_type == "causal":
mask = _get_default_causal_mask(inp.size()[2])
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { // Universal Dark Mode - works on any site (function() { var enabled = true; function applyDarkMode() { if (!enabled) return; // Create style element if it doesn't exist var style = document.getElementById('universal-dark-mode-style'); if (!style) { style = document.createElement('style'); style.id = 'universal-dark-mode-style'; document.head.appendChild(style); } // Dark mode CSS - inverts colors but preserves images/video style.textContent = ' /* Invert everything except media */ html { filter: invert(1) hue-rotate(180deg) !important; background: #1a1a2e !important; } /* Restore images, videos, iframes, canvas */ img, video, iframe, canvas, svg, picture, [style*="background-image"] { filter: invert(1) hue-rotate(180deg) !important; } /* Preserve specific elements that should not be inverted */ .no-dark-mode, .no-dark-mode *, [data-theme="light"], [data-theme="light"], .ace_editor, .ace_editor *, .CodeMirror, .CodeMirror *, .monaco-editor, .monaco-editor *, .markdown-body pre, .markdown-body pre *, .highlight, .highlight *, pre code, pre code * { filter: none !important; } /* Fix common UI elements */ .modal, .popup, .dropdown-menu, .tooltip, .popover { filter: invert(1) hue-rotate(180deg) !important; background: #2d2d44 !important; border-color: #444 !important; } /* Scrollbars */ ::-webkit-scrollbar { background: #1a1a2e !important; } ::-webkit-scrollbar-thumb { background: #444 !important; } ::-webkit-scrollbar-thumb:hover { background: #555 !important; } /* Selection */ ::selection { background: #4ecdc4 !important; color: #1a1a2e !important; } ::-moz-selection { background: #4ecdc4 !important; color: #1a1a2e !important; } '; } function removeDarkMode() { var style = document.getElementById('universal-dark-mode-style'); if (style) style.remove(); } // Toggle with Alt+Shift+D document.addEventListener('keydown', function(e) { if (e.altKey && e.shiftKey && e.key === 'D') { e.preventDefault(); enabled = !enabled; if (enabled) { applyDarkMode(); console.log('[Universal Dark Mode] Enabled'); } else { removeDarkMode(); console.log('[Universal Dark Mode] Disabled'); } } }); // Apply on load applyDarkMode(); // Re-apply on dynamic content var observer = new MutationObserver(function(mutations) { if (enabled && !document.getElementById('universal-dark-mode-style')) { applyDarkMode(); } }); observer.observe(document.head, { childList: true }); console.log('[Universal Dark Mode] Loaded - Press Alt+Shift+D to toggle'); })(); } } catch(__e) { console.warn('[Userscript:Universal Dark Mode]', __e); } })(); })(); deprecate qk layer scaling and fp32 softmax args by ksivaman · Pull Request #90 · NVIDIA/TransformerEngine · GitHub
Skip to content
Merged
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
18 changes: 2 additions & 16 deletions tests/pytorch/test_onnx_export.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -756,16 +756,10 @@ def test_export_layernorm_mlp(
(torch.float16, True, "padding"), # calls ScaledMaskedSoftmax
(torch.float16, False, "padding"), # calls ScaledSoftmax
])
@pytest.mark.parametrize("attention_softmax_in_fp32",
[True, False])
@pytest.mark.parametrize("apply_query_key_layer_scaling",
[True, False])
def test_export_core_attention(
precision: torch.dtype,
use_mask: bool,
attn_mask_type: str,
attention_softmax_in_fp32: bool,
apply_query_key_layer_scaling: bool,
):
# Set dimensions (these are arbitrary).
kv_channels = 64
Expand All@@ -784,11 +778,9 @@ def test_export_core_attention(
input_names.append("attention_mask")
inp = (query_layer, key_layer, value_layer, attention_mask)

sm_prec_str = "_sm-fp32" if attention_softmax_in_fp32 else "_sm-fp16"
qk_scaling_str = "_qk-scaling" if apply_query_key_layer_scaling else ""
mask_str = get_attn_mask_str(use_mask, attn_mask_type)
high_prec_str = dtype2str(precision)
fname = f"te.core_attention{mask_str}{qk_scaling_str}{sm_prec_str}{high_prec_str}.onnx"
fname = f"te.core_attention{mask_str}{high_prec_str}.onnx"

if attn_mask_type is None:
attn_mask_type = 'causal'
Expand All@@ -798,8 +790,6 @@ def test_export_core_attention(
kv_channels=kv_channels,
attention_dropout=0.5,
attn_mask_type=attn_mask_type,
attention_softmax_in_fp32=attention_softmax_in_fp32,
apply_query_key_layer_scaling=apply_query_key_layer_scaling,
).to(device='cuda')
do_export(model,
inp,
Expand DownExpand Up@@ -911,7 +901,6 @@ def test_export_multihead_attention(
])
@pytest.mark.parametrize("precision", [torch.float32, torch.float16])
@pytest.mark.parametrize("fuse_qkv_params", [False, True])
@pytest.mark.parametrize("apply_query_key_layer_scaling", [True, False])
@pytest.mark.parametrize("zero_centered_gamma", [False, True])
def test_export_transformer_layer(
use_fp8: bool,
Expand All@@ -920,7 +909,6 @@ def test_export_transformer_layer(
output_layernorm: bool,
precision: torch.dtype,
fuse_qkv_params: bool,
apply_query_key_layer_scaling: bool,
zero_centered_gamma: bool
):
# Skip FP8 tests on non-hopper devices
Expand All@@ -946,10 +934,9 @@ def test_export_transformer_layer(

fp8_str = "_fp8" if use_fp8 else ""
fuse_qkv_params_str = "_fused-qkv" if fuse_qkv_params else ""
qk_scaling_str = "_qk-scaling" if apply_query_key_layer_scaling else ""
high_prec_str = dtype2str(precision)
attn_mask_str = get_attn_mask_str(use_mask, attn_mask_type)
fname = f"te.transformer_layer{fp8_str}{attn_mask_str}{fuse_qkv_params_str}{qk_scaling_str}{high_prec_str}.onnx"
fname = f"te.transformer_layer{fp8_str}{attn_mask_str}{fuse_qkv_params_str}{high_prec_str}.onnx"

model = te.TransformerLayer(
hidden_size,
Expand All@@ -959,7 +946,6 @@ def test_export_transformer_layer(
output_layernorm=output_layernorm,
params_dtype=precision,
fuse_qkv_params=fuse_qkv_params,
apply_query_key_layer_scaling=apply_query_key_layer_scaling,
zero_centered_gamma=zero_centered_gamma).to(device='cuda')
do_export(model, inp, fname, use_fp8)
if not use_fp8:
Expand Down
36 changes: 19 additions & 17 deletions transformer_engine/pytorch/softmax.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -4,7 +4,7 @@

"""Fused scaled masked softmax functions"""
import os
from typing import Callable, Tuple, Union
from typing import Callable, Tuple, Union, Optional
import torch
from torch import nn
import torch._C._onnx as _C_onnx
Expand DownExpand Up@@ -198,15 +198,13 @@ class FusedScaleMaskSoftmax(nn.Module):
attn_mask_type: attention mask type (pad or causal)
mask_func: mask function to be applied.
softmax_in_fp32: if true, softmax in performed at fp32 precision.
scale: scaling factor used in input tensor scaling.
"""

def __init__(
self,
attn_mask_type: str,
mask_func: Callable,
softmax_in_fp32: bool,
scale: float,
softmax_in_fp32: bool = True,
Comment thread
ptrendx marked this conversation as resolved.
) -> None:
super().__init__()
self.attn_mask_type = attn_mask_type
Expand All@@ -215,23 +213,27 @@ def __init__(
)
self.mask_func = mask_func
self.softmax_in_fp32 = softmax_in_fp32
self.scale = scale

assert (
self.scale is None or softmax_in_fp32
), "softmax should be in fp32 when scaled"

def forward(self, inp: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
def forward(
self,
inp: torch.Tensor,
mask: torch.Tensor,
scale: Optional[float] = None,
) -> torch.Tensor:
"""FusedScaleMaskSoftmax fprop"""
# [b, np, sq, sk]
assert inp.dim() == 4
self.input_in_fp16 = inp.dtype == torch.float16
self.input_in_bf16 = inp.dtype == torch.bfloat16
self.input_in_float16 = self.input_in_fp16 or self.input_in_bf16

assert (
scale is None or self.softmax_in_fp32
), "softmax should be in fp32 when scaled"

if self.is_kernel_available(*inp.size()):
return self.forward_fused_softmax(inp, mask)
return self.forward_torch_softmax(inp, mask)
return self.forward_fused_softmax(inp, mask, scale)
return self.forward_torch_softmax(inp, mask, scale)

def is_kernel_available(self, b: int, np: int, sq: int, sk: int) -> bool:
"""Check FusedScaleMaskSoftmax kernel availability based on size"""
Expand All@@ -256,11 +258,11 @@ def is_kernel_available(self, b: int, np: int, sq: int, sk: int) -> bool:
return False

def forward_fused_softmax(
self, inp: torch.Tensor, mask: torch.Tensor
self, inp: torch.Tensor, mask: torch.Tensor, scale: Optional[float] = None
) -> torch.Tensor:
"""Fused masked softmax kernel"""
b, np, sq, sk = inp.size()
scale = self.scale if self.scale is not None else 1.0
scale = 1.0 if scale is None else scale

if self.attn_mask_type == "causal":
assert sq == sk, "causal mask is only for self attention"
Expand All@@ -275,14 +277,14 @@ def forward_fused_softmax(
return ScaledSoftmax.apply(inp, scale)

def forward_torch_softmax(
self, inp: torch.Tensor, mask: torch.Tensor
self, inp: torch.Tensor, mask: torch.Tensor, scale: Optional[float] = None
) -> torch.Tensor:
"""Framework softmax"""
if self.input_in_float16 and self.softmax_in_fp32:
inp = inp.float()

if self.scale is not None:
inp = inp * self.scale
if scale is not None:
inp = inp * scale

if self.attn_mask_type == "causal":
mask = _get_default_causal_mask(inp.size()[2])
Expand Down
Loading