[Diffusion] Tune QK head LayerNorm for SM103 (#34503)
This commit is contained in:
@@ -341,9 +341,13 @@ def _is_bf16_cuda(t: torch.Tensor) -> bool:
|
|||||||
return t.is_cuda and t.dtype is torch.bfloat16
|
return t.is_cuda and t.dtype is torch.bfloat16
|
||||||
|
|
||||||
|
|
||||||
def _is_sm120_or_newer() -> bool:
|
def _qk_head_launch_config() -> tuple[int, int]:
|
||||||
arch = get_jit_cuda_arch()
|
arch = get_jit_cuda_arch()
|
||||||
return arch.major * 10 + arch.minor >= 120
|
if arch.major == 10 and arch.minor == 3:
|
||||||
|
return 16, 1
|
||||||
|
if arch.major * 10 + arch.minor >= 120:
|
||||||
|
return 8, 4
|
||||||
|
return 64, 2
|
||||||
|
|
||||||
|
|
||||||
def _mod_row_stride(t: torch.Tensor, batch: int, hidden: int) -> int | None:
|
def _mod_row_stride(t: torch.Tensor, batch: int, hidden: int) -> int | None:
|
||||||
@@ -460,12 +464,10 @@ def fused_qk_head_layernorm(
|
|||||||
launch, bit-exact vs the eager aten kernel."""
|
launch, bit-exact vs the eager aten kernel."""
|
||||||
head_dim = q.shape[-1]
|
head_dim = q.shape[-1]
|
||||||
n_rows = q.numel() // head_dim
|
n_rows = q.numel() // head_dim
|
||||||
# SM120 has a smaller register file per SM than H200. Grouping 64 exact
|
# Architecture sweeps at the production GLM shape select 16 rows / 1 warp
|
||||||
# Welford rows in one program depresses occupancy on RTX 5090; an exhaustive
|
# on B300 (SM103) and 8 rows / 4 warps on RTX 5090 (SM120). Preserve the
|
||||||
# production-shape sweep selects 8 rows / 4 warps there (about 10% faster).
|
# independently tuned H100/H200 launch on all other architectures.
|
||||||
# Preserve the independently tuned SM90 launch byte-for-byte.
|
rows, num_warps = _qk_head_launch_config()
|
||||||
is_sm120 = _is_sm120_or_newer()
|
|
||||||
rows = 8 if is_sm120 else 64
|
|
||||||
q_out = torch.empty_like(q)
|
q_out = torch.empty_like(q)
|
||||||
k_out = torch.empty_like(k)
|
k_out = torch.empty_like(k)
|
||||||
with torch.cuda.device(q.device):
|
with torch.cuda.device(q.device):
|
||||||
@@ -481,6 +483,6 @@ def fused_qk_head_layernorm(
|
|||||||
ROWS=rows,
|
ROWS=rows,
|
||||||
# H200-tuned: 62us at (1, 4360, 32, 128) vs the 301us of the two
|
# H200-tuned: 62us at (1, 4360, 32, 128) vs the 301us of the two
|
||||||
# aten launches (one 128-thread block per head_dim-element row).
|
# aten launches (one 128-thread block per head_dim-element row).
|
||||||
num_warps=4 if is_sm120 else 2,
|
num_warps=num_warps,
|
||||||
)
|
)
|
||||||
return q_out, k_out
|
return q_out, k_out
|
||||||
|
|||||||
Reference in New Issue
Block a user