diff --git a/CHANGELOG.md b/CHANGELOG.md new file mode 100644 index 0000000..4a448ad --- /dev/null +++ b/CHANGELOG.md @@ -0,0 +1,34 @@ +# Changelog + + + + +## linghe 0.3.0 + +- use faster and more accurate implementation for `embedding` backward +- add multiple embedding_lookup implementations +- use 2048 block number for `batch_count_zero` kernel + + +## linghe 0.2.9 + +- support stride in grad tensor of `embedding` kernel +- return grad for dummy tensor in `embedding` kernel +- support bf16 in batch mul/clip/norm kernels + + +## linghe 0.2.8 + +- fix racing condition bug in softmax_cross_entropy kernel +- use tl.rsqrt instead of 1/tl.sqrt in all kernels +- add the parameter `tp_group` to `softmax_cross_entropy` + + +## linghe 0.2.7 + +- add the parameter `ignore_index` to `softmax_cross_entropy` +- support parallel `softmax_cross_entropy` +- add dtype and numel assertion in multiple batch kernels + +- Known issues: + - performance of `softmax_cross_entropy` degrades when vocab size is not multiple of 16 diff --git a/README.md b/README.md index cdd7d39..86ad03a 100644 --- a/README.md +++ b/README.md @@ -25,7 +25,7 @@ ## Introduction --- -Our repo, linghe, is designed for LLM training, especially for MoE training with FP8 quantizaiton. It provides 2 main categories of kernels: +Our repo, linghe, is designed for LLM training, especially for MoE training with FP8 quantizaiton. It provides 3 main categories of kernels: - **Fused quantization kernels**: fuse quantization with previous layer, e.g., RMS norm and Silu. - **Memory-efficiency kernels**: fuse multiple IO-itensive operations, e.g., ROPE with qk-norm. @@ -66,4 +66,14 @@ Examples can be found in tests. ## Api Reference --- -Please refer to [API](https://inclusionai.github.io/linghe/) \ No newline at end of file +Please refer to [API](https://inclusionai.github.io/linghe/) + +## Citations + +[TBD] +``` +@misc{zhao2025linghe, +title={Linghe: Enabling Efficient Trillion-Scale LLM Training via Optimized Kernels}, +author={Yao Zhao and Chen Liang and Jingyu Hu and Zixuan Cheng and Longfei Li} +} +``` diff --git a/benchmark/bench_gemm.py b/benchmark/bench_gemm.py index 862e7d9..cb624a0 100644 --- a/benchmark/bench_gemm.py +++ b/benchmark/bench_gemm.py @@ -5,13 +5,11 @@ import torch - from linghe.tools.benchmark import benchmark_func from linghe.tools.util import fp16_forward from linghe.utils.add import triton_inplace_add - def triton_accum_weight(x, w, out, x_scale, w_scale): output = torch._scaled_mm( x, @@ -19,7 +17,7 @@ def triton_accum_weight(x, w, out, x_scale, w_scale): scale_a=x_scale, scale_b=w_scale, out_dtype=torch.bfloat16, - use_fast_accum=True + use_fast_accum=True, ) triton_inplace_add(out, output) return out @@ -32,16 +30,15 @@ def torch_accum_weight(x, w, out, x_scale, w_scale): scale_a=x_scale, scale_b=w_scale, out_dtype=torch.bfloat16, - use_fast_accum=True + use_fast_accum=True, ) out.add_(output) return out -def test_cublas_blockwise_gemm(M=4096, N=4096, K=4096): - +def bench_cublas_channelwise_gemm(M=4096, N=4096, K=4096): dtype = torch.bfloat16 - device = 'cuda:0' + device = "cuda:0" n_repeat = 100 x = torch.randn(M, K, dtype=dtype, device=device) @@ -57,74 +54,134 @@ def test_cublas_blockwise_gemm(M=4096, N=4096, K=4096): out = torch.zeros((M, N), dtype=torch.float32, device=device) o = torch.empty((M, N), dtype=dtype, device=device) - ref_time = benchmark_func(fp16_forward, x, w.t(), n_repeat=n_repeat, - ref_flops=ref_flops, name=f'M:{M}') - benchmark_func(torch_accum_weight, x_q, w_q.t(), out, xrs, wcs.view(1, -1), - n_repeat=n_repeat, ref_flops=ref_flops, ref_time=ref_time, - name=f'M:{M}') - benchmark_func(triton_accum_weight, x_q, w_q.t(), out, xrs, wcs.view(1, -1), - n_repeat=n_repeat, ref_flops=ref_flops, ref_time=ref_time, - name=f'M:{M}') - + ref_time = benchmark_func( + fp16_forward, x, w.t(), n_repeat=n_repeat, ref_flops=ref_flops, name=f"M:{M}" + ) + benchmark_func( + torch_accum_weight, + x_q, + w_q.t(), + out, + xrs, + wcs.view(1, -1), + n_repeat=n_repeat, + ref_flops=ref_flops, + ref_time=ref_time, + name=f"M:{M}", + ) + benchmark_func( + triton_accum_weight, + x_q, + w_q.t(), + out, + xrs, + wcs.view(1, -1), + n_repeat=n_repeat, + ref_flops=ref_flops, + ref_time=ref_time, + name=f"M:{M}", + ) -def test_te_blockwise_gemm(M=4096, N=4096, K=4096): +def bench_te_blockwise_gemm(M=4096, N=4096, K=4096, round_scale=False): # layout == 'TN': # forward, y=x@w - # x_q = B._rowwise_data - # x_scale = B._rowwise_scale_inv - # w_q = A._rowwise_data - # w_scale = A._rowwise_scale_inv - - # import transformer_engine_torch as tex - import transformer_engine as te + from linghe.quant.block import triton_block_quant, triton_blockwise_quant import transformer_engine_torch as tex - from transformer_engine.pytorch.tensor.float8_blockwise_tensor import Float8BlockwiseQTensor + from transformer_engine.pytorch.tensor.float8_blockwise_tensor import ( + Float8BlockwiseQTensor, + Float8BlockQuantizer, + ) from transformer_engine.pytorch.module.base import get_workspace from transformer_engine.pytorch.constants import TE_DType - row_data = torch.randn((M,K), device='cuda:0').to(torch.float8_e4m3fn) - row_scales = torch.randn((K//128,M), device='cuda:0') - x = Float8BlockwiseQTensor(shape=(M,K), - dtype=torch.bfloat16, - fp8_dtype=TE_DType[torch.float8_e4m3fn], - rowwise_data=row_data, - rowwise_scale_inv=row_scales, - columnwise_data=None, - columnwise_scale_inv=None, - quantizer=None, - requires_grad=False, - is_2D_scaled=False - ) - - row_data = torch.randn((N,K), device='cuda:0').to(torch.float8_e4m3fn) - row_scales = torch.randn((K//128,N//128), device='cuda:0') - w = Float8BlockwiseQTensor(shape=(N,K), - dtype=torch.bfloat16, - fp8_dtype=TE_DType[torch.float8_e4m3fn], - rowwise_data=row_data, - rowwise_scale_inv=row_scales, - columnwise_data=None, - columnwise_scale_inv=None, - quantizer=None, - requires_grad=False, - is_2D_scaled=True - ) - A = w - transa = True - B = x - transb = False - out = None - quantization_params = None - out_dtype = TE_DType[torch.bfloat16] - bias = None - bias_dtype = TE_DType[torch.bfloat16] - gelu = False - gelu_in = None - grad = False - workspace = get_workspace() - workspace_size = workspace.shape[0] - accumulate = False - use_split_accumulator = True - args = ( + + quantizer = Float8BlockQuantizer( + TE_DType[torch.float8_e4m3fn], + rowwise=True, + columnwise=True, + amax_epsilon=0, + force_pow_2_scales=round_scale, + block_scaling_dim=1, + ) + dtype = torch.bfloat16 + device = "cuda:0" + x = torch.randn((M, K), device=device, dtype=dtype) ** 3 * 1e-10 + x[-(M // 2) :] = 0 + x[:, -(K // 2) :] = 0 + + weight_quantizer = Float8BlockQuantizer( + TE_DType[torch.float8_e4m3fn], + rowwise=True, + columnwise=True, + amax_epsilon=0, + force_pow_2_scales=round_scale, + block_scaling_dim=2, + ) + w = torch.randn((N, K), device=device, dtype=dtype) + + for manual in [False, True]: + if manual: + x_q, x_s, xt_q, xt_s = triton_blockwise_quant(x, round_scale=round_scale) + qx = Float8BlockwiseQTensor( + shape=(M, K), + dtype=torch.bfloat16, + fp8_dtype=TE_DType[torch.float8_e4m3fn], + rowwise_data=x_q, + rowwise_scale_inv=x_s, + columnwise_data=xt_q, + columnwise_scale_inv=xt_s, + quantizer=quantizer, + requires_grad=False, + is_2D_scaled=False, + ) + w_q, w_s = triton_block_quant(w, round_scale=round_scale) + wt_q, wt_s = ( + w_q.transpose(0, 1).contiguous(), + w_s.transpose(0, 1).contiguous(), + ) + qw = Float8BlockwiseQTensor( + shape=(N, K), + dtype=torch.bfloat16, + fp8_dtype=TE_DType[torch.float8_e4m3fn], + rowwise_data=w_q, + rowwise_scale_inv=w_s, + columnwise_data=wt_q, + columnwise_scale_inv=wt_s, + quantizer=weight_quantizer, + requires_grad=False, + is_2D_scaled=True, + ) + else: + qx = quantizer.make_empty( + (M, K), dtype=torch.bfloat16, device=device, requires_grad=False + ) + qx = quantizer.update_quantized(x, qx) + + qw = weight_quantizer.make_empty( + (N, K), dtype=torch.bfloat16, device=device, requires_grad=False + ) + qw = weight_quantizer.update_quantized(w, qw) + + # print(f'{qx._rowwise_data.shape=} {qx._rowwise_scale_inv.shape=} {qx._columnwise_data.shape=} {qx._columnwise_scale_inv.shape=}') + # print(f'{qw._rowwise_data.shape=} {qw._rowwise_scale_inv.shape=} {qw._columnwise_data.shape=} {qw._columnwise_scale_inv.shape=}') + + A = qw + transa = True + B = qx + transb = False + # out = torch.randn( (M, N), device='cuda:0', dtype=torch.bfloat16) + out = None + quantization_params = None + out_dtype = TE_DType[torch.bfloat16] + bias = None + bias_dtype = TE_DType[torch.bfloat16] + gelu = False + gelu_in = None + grad = False + workspace = get_workspace() + workspace_size = workspace.shape[0] + accumulate = False + use_split_accumulator = True + args = ( A, transa, # transa B, @@ -142,15 +199,119 @@ def test_te_blockwise_gemm(M=4096, N=4096, K=4096): accumulate, use_split_accumulator, ) + # kwargs = { + # "comm_overlap": None, + # "comm_type": None, + # "extra_output": None, + # "bulk_overlap": False, + # "alpha": 1.0, + # "beta": 0.0, + # } + out, bias_grad, gelu_input, extra_output = tex.generic_gemm(*args) + + ref_out = x @ w.t() + + rel_err = (out - ref_out).abs().sum().item() / ref_out.abs().sum().item() + print( + f"rel:{rel_err:.6f} ref:{ref_out.abs().mean().item():.3f} out:{out.abs().mean().item():.3f}" + ) + + ref_flops = M * N * K * 2 + ref_bytes = M * K + N * K + M * N * 2 + benchmark_func( + tex.generic_gemm, + *args, + n_repeat=100, + ref_flops=ref_flops, + ref_bytes=ref_bytes, + ) + + +def bench_te_mxfp8_gemm(M=4096, N=4096, K=4096): + if torch.cuda.get_device_properties(0).major < 10: + return + + # import transformer_engine_torch as tex + from linghe.quant.mxfp8 import triton_mxfp8_quant + import transformer_engine_torch as tex + from transformer_engine.pytorch.tensor.mxfp8_tensor import MXFP8Tensor + from transformer_engine.pytorch.module.base import get_workspace + from transformer_engine.pytorch.constants import TE_DType + + x = torch.randn((M, K), device="cuda:0", dtype=torch.bfloat16) + x_q, x_scale, xt_q, xt_scale = triton_mxfp8_quant(x) + + B = MXFP8Tensor( + shape=(M, K), + dtype=torch.bfloat16, + rowwise_data=x_q, + rowwise_scale_inv=x_scale, + columnwise_data=None, + columnwise_scale_inv=None, + fp8_dtype=TE_DType[torch.float8_e4m3fn], + quantizer=None, + ) + + w = torch.randn((N, K), device="cuda:0", dtype=torch.bfloat16) + w_q, w_scale, wt_q, wt_scale = triton_mxfp8_quant(w) + + A = MXFP8Tensor( + shape=(N, K), + dtype=torch.bfloat16, + rowwise_data=w_q, + rowwise_scale_inv=w_scale, + columnwise_data=None, + columnwise_scale_inv=None, + fp8_dtype=TE_DType[torch.float8_e4m3fn], + quantizer=None, + ) + transa = True + transb = False + out = None + quantization_params = None + out_dtype = TE_DType[torch.bfloat16] + bias = None + bias_dtype = TE_DType[torch.bfloat16] + gelu = False + gelu_in = None + grad = False + workspace = get_workspace() + workspace_size = workspace.shape[0] + accumulate = False + use_split_accumulator = True + args = ( + A, + transa, # transa + B, + transb, # transb + out, + quantization_params, + out_dtype, + bias, + bias_dtype, + gelu, + gelu_in, + grad, # grad + workspace, + workspace_size, + accumulate, + use_split_accumulator, + ) out, bias_grad, gelu_input, extra_output = tex.generic_gemm(*args) - ref_flops = M * N * K * 2 - ref_bytes = M * K + N * K + M * N *2 - benchmark_func(tex.generic_gemm, *args, - n_repeat=100, ref_flops=ref_flops, ref_bytes=ref_bytes) + out_ref = x @ w.t() + error = (out - out_ref).abs().sum() / out_ref.abs().sum() + print(error) + ref_flops = M * N * K * 2 + ref_bytes = M * K + N * K + M * N * 2 + benchmark_func( + tex.generic_gemm, *args, n_repeat=100, ref_flops=ref_flops, ref_bytes=ref_bytes + ) -if __name__ == '__main__': - test_cublas_blockwise_gemm(M=4096, N=4096, K=4096) - test_te_blockwise_gemm(M=4096, N=4096, K=4096) +if __name__ == "__main__": + # bench_cublas_channelwise_gemm(M=4096, N=4096, K=4096) + bench_te_blockwise_gemm(M=128, N=128, K=128) + bench_te_blockwise_gemm(M=4096, N=4096, K=4096) + # bench_te_mxfp8_gemm(M=4096, N=4096, K=4096) diff --git a/benchmark/bench_grad_norm.py b/benchmark/bench_grad_norm.py new file mode 100644 index 0000000..2b47cc4 --- /dev/null +++ b/benchmark/bench_grad_norm.py @@ -0,0 +1,48 @@ +import random + +import torch +from transformer_engine.pytorch.optimizers import ( + multi_tensor_applier, + multi_tensor_l2norm, +) + +from linghe.tools.benchmark import benchmark_func +from linghe.tools.check import output_check +from linghe.utils.reduce import triton_batch_norm + + +def bench_batch_norm(M=4096, N=2048, k=32): + xs = [ + torch.randn(random.randint(M // 10, M), N, dtype=torch.float32, device="cuda:0") + for i in range(k) + ] + + dummy_overflow_buf = torch.tensor([0], dtype=torch.int, device="cuda") + grad_norm_ref, _ = multi_tensor_applier( + multi_tensor_l2norm, + dummy_overflow_buf, + [xs], + False, # no per-parameter norm + ) + + grad_norm = triton_batch_norm(xs, ord=2, norm=True) + output_check(grad_norm_ref[0], grad_norm, "l2_norm") + + ref_bytes = sum([x.numel() for x in xs]) * 4 + n_repeat = 100 + ref_time = benchmark_func( + multi_tensor_applier, + multi_tensor_l2norm, + dummy_overflow_buf, + [xs], + False, + n_repeat=n_repeat, + ref_bytes=ref_bytes, + ) + benchmark_func( + triton_batch_norm, xs, n_repeat=n_repeat, ref_bytes=ref_bytes, ref_time=ref_time + ) + + +if __name__ == "__main__": + bench_batch_norm(M=1024, N=2048, k=512) diff --git a/benchmark/bench_la.py b/benchmark/bench_la.py new file mode 100644 index 0000000..d991d01 --- /dev/null +++ b/benchmark/bench_la.py @@ -0,0 +1,45 @@ +import torch +from fla.ops.lightning_attn import chunk_lightning_attn + +from linghe.tools.benchmark import benchmark_func + + +def bench_la(B=1, S=4096, H=32, D=128): + query = torch.randn( + B, S, H, D, device="cuda", dtype=torch.bfloat16, requires_grad=True + ) + key = torch.randn( + B, S, H, D, device="cuda", dtype=torch.bfloat16, requires_grad=True + ) + value = torch.randn( + B, S, H, D, device="cuda", dtype=torch.bfloat16, requires_grad=True + ) + grad = torch.randn(B, S, H, D, device="cuda", dtype=torch.bfloat16) + # decay_scales = 2**(-0.5 * torch.arange(1, H+1, dtype=torch.float32, device='cuda')) + + core_attn_out, _ = chunk_lightning_attn( + query, + key, + value, + layer_idx=1, # not used, starts from 0. + num_layers=20, # not used. + initial_state=None, + output_final_state=True, + cu_seqlens=None, # for varlen training + head_first=False, + ) + + benchmark_func( + chunk_lightning_attn, + query, + key, + value, + layer_idx=1, + num_layers=20, + output_final_state=True, + ) + benchmark_func(core_attn_out.backward, grad, retain_graph=True) + + +if __name__ == "__main__": + bench_la(B=2, S=4096, H=64, D=128) diff --git a/benchmark/bench_loss.py b/benchmark/bench_loss.py index 09240ef..2d45ac1 100644 --- a/benchmark/bench_loss.py +++ b/benchmark/bench_loss.py @@ -7,102 +7,124 @@ import torch import torch.distributed as dist -from megatron.core.fusions.fused_cross_entropy import \ - fused_vocab_parallel_cross_entropy +from megatron.core.fusions.fused_cross_entropy import fused_vocab_parallel_cross_entropy from transformer_engine.pytorch.cross_entropy import parallel_cross_entropy +from linghe.facade.loss import softmax_cross_entropy from linghe.tools.benchmark import benchmark_func -from linghe.tools.util import output_check -from linghe.utils.loss import (triton_softmax_cross_entropy_backward, - triton_softmax_cross_entropy_forward) - - -def torch_cross_entropy(logits, targets): - float_logits = logits.float() - losses = torch.nn.functional.cross_entropy( - float_logits.view(-1, float_logits.size()[-1]), - targets.view(-1), - reduction='none') - loss = losses.mean() - loss.backward() - return losses, logits.grad +from linghe.tools.check import output_check +from linghe.utils.loss import ( + triton_softmax_cross_entropy_backward, + triton_softmax_cross_entropy_forward, +) -def te_cross_entropy_forward_backward(logits, targets): - losses = parallel_cross_entropy(logits[None], - targets[None]) - loss = losses.mean() - loss.backward() +def fused_cross_entropy_forward_backward(logits, targets, input_grad, pg): + losses = fused_vocab_parallel_cross_entropy(logits[None], targets[None], pg)[0] + losses.backward(input_grad) return losses, logits.grad -def fused_cross_entropy_forward_backward(logits, targets, pg): - losses = fused_vocab_parallel_cross_entropy(logits[None], - targets[None], - pg)[0] - loss = losses.mean() - loss.backward() +def te_cross_entropy_forward_backward(logits, targets, input_grad): + losses = parallel_cross_entropy(logits[None], targets[None]) + losses.backward(input_grad[None]) return losses, logits.grad -def triton_cross_entropy_forward_backward(logits, targets, input_grad): - losses, sum_exp, max_logits = triton_softmax_cross_entropy_forward(logits, - targets) - output_grad = triton_softmax_cross_entropy_backward(logits, targets, - sum_exp, max_logits, - input_grad) - return losses, output_grad +def triton_cross_entropy_forward_backward(logits, targets, input_grad, inplace=True): + # losses, sum_exp, max_logits = triton_softmax_cross_entropy_forward(logits, + # targets) + # output_grad = triton_softmax_cross_entropy_backward(logits, targets, + # sum_exp, max_logits, + # input_grad, + # inplace=inplace) + # return losses, output_grad + losses = softmax_cross_entropy(logits, targets, inplace=inplace) + losses.backward(input_grad) + return losses, logits.grad def bench_triton_softmax_cross_entropy(M=4096, N=157184): - device = 'cuda:0' - logits = torch.randn((M, N), dtype=torch.bfloat16, device=device, - requires_grad=True) ** 1 + device = "cuda:0" + logits = torch.randn((M, N), dtype=torch.bfloat16, device=device) logits = logits.detach().clone().requires_grad_() - # targets = (torch.rand((M,), dtype=torch.float32, device=device) * N).to( - # torch.int64) - targets = torch.topk(logits, 4)[1][:, 3].contiguous() - input_grad = 1 / M * torch.ones((M,), dtype=torch.bfloat16, device=device) + targets = (torch.rand((M,), dtype=torch.float32, device=device) * N).to(torch.int64) + # targets = torch.topk(logits, 4)[1][:, 3].contiguous() + input_grad = 1 / M * torch.randn((M,), dtype=torch.float32, device=device) sum_exp = torch.rand((M,), dtype=torch.float32, device=device) max_logits = torch.rand((M,), dtype=torch.float32, device=device) - pg = dist.new_group(ranks=[0], backend='nccl') - torch_losses, torch_grad = torch_cross_entropy( - logits.detach().clone().requires_grad_(), targets) + pg = dist.new_group(ranks=[0], backend="nccl") fused_losses, fused_grad = fused_cross_entropy_forward_backward( - logits.detach().clone().requires_grad_(), targets, pg) - triton_losses, sum_exp_, max_logit_ = triton_softmax_cross_entropy_forward( - logits, targets) - output_check(torch_losses, triton_losses) + logits.detach().clone().requires_grad_(), targets, input_grad, pg + ) + te_losses, te_grad = te_cross_entropy_forward_backward( + logits.detach().clone().requires_grad_(), targets, input_grad + ) + triton_losses, triton_grad = triton_cross_entropy_forward_backward( + logits.detach().clone().requires_grad_(), targets, input_grad, inplace=False + ) output_check(fused_losses, triton_losses) - output_check(torch_losses, fused_losses) - - ref_time = benchmark_func(torch_cross_entropy, - logits.detach().clone().requires_grad_(), targets, - ref_bytes=M * N * 4) - benchmark_func(triton_softmax_cross_entropy_forward, logits, targets, - ref_bytes=M * N * 2, ref_time=ref_time) - benchmark_func(triton_softmax_cross_entropy_backward, logits, targets, - sum_exp, max_logits, input_grad, ref_bytes=M * N * 4, - ref_time=ref_time) - - benchmark_func(te_cross_entropy_forward_backward, - logits.detach().clone().requires_grad_(), targets, - ref_bytes=M * N * 4, - ref_time=ref_time) - benchmark_func(te_cross_entropy_forward_backward, - logits.detach().clone().requires_grad_(), targets, - ref_bytes=M * N * 4, - ref_time=ref_time) - benchmark_func(triton_cross_entropy_forward_backward, logits, targets, - input_grad, ref_bytes=M * N * 4, - ref_time=ref_time) - - -if __name__ == '__main__': + output_check(fused_grad, triton_grad) + + ref_time = benchmark_func( + fused_cross_entropy_forward_backward, + logits.detach().clone().requires_grad_(), + targets, + input_grad, + pg, + n_repeat=1, + ref_bytes=M * N * 6, + ) + benchmark_func( + te_cross_entropy_forward_backward, + logits.detach().clone().requires_grad_(), + targets, + input_grad, + n_repeat=1, + ref_bytes=M * N * 6, + ref_time=ref_time, + ) + benchmark_func( + triton_cross_entropy_forward_backward, + logits.detach().clone().requires_grad_(), + targets, + input_grad, + inplace=False, + n_repeat=1, + ref_bytes=M * N * 6, + ref_time=ref_time, + ) + benchmark_func( + triton_softmax_cross_entropy_forward, + logits, + targets, + ref_bytes=M * N * 2, + ref_time=ref_time, + ) + benchmark_func( + triton_softmax_cross_entropy_backward, + logits, + targets, + sum_exp, + max_logits, + input_grad, + ref_bytes=M * N * 4, + ref_time=ref_time, + ) + + +if __name__ == "__main__": + # torchrun bench_loss.py init_method = "env://" - dist.init_process_group(backend='nccl', init_method=init_method, - world_size=1, rank=0, - timeout=timedelta(seconds=30)) + dist.init_process_group( + backend="nccl", + init_method=init_method, + world_size=1, + rank=0, + timeout=timedelta(seconds=30), + ) + bench_triton_softmax_cross_entropy(M=4096, N=157184) bench_triton_softmax_cross_entropy(M=8192, N=157184) + bench_triton_softmax_cross_entropy(M=8192, N=128) diff --git a/benchmark/bench_mla.py b/benchmark/bench_mla.py new file mode 100644 index 0000000..b1d5e71 --- /dev/null +++ b/benchmark/bench_mla.py @@ -0,0 +1,605 @@ +import logging +import os +import sys +import pathlib +from typing import Any, Dict, Tuple, Union + +import torch + +from transformer_engine.pytorch import DotProductAttention + +import transformer_engine_torch as tex +from linghe.tools.benchmark import benchmark_func + + +class ModelConfig: + def __init__( + self, + batch_size: int, + max_seqlen_q: int, + num_heads: int, + head_dim_qk: int, + max_seqlen_kv: int = None, + num_gqa_groups: int = None, + head_dim_v: int = None, + softmax_type: str = "vanilla", + dropout_p: float = 0.0, + attn_mask_type: str = "no_mask", + attn_bias_type: str = "no_bias", + alibi_type: str = "none", + bias_shape: str = "1hss", + window_size: Tuple[int, int] = (-1, -1), + context_parallel: bool = False, + cp_comm_type: str = "p2p", + return_max_logit=False, + total_requests: int = None, + max_ctx_len: int = None, + num_layers: int = 1, + eps: float = 1e-5, + ): + self.batch_size = batch_size + self.max_seqlen_q = max_seqlen_q + self.max_seqlen_kv = max_seqlen_q if max_seqlen_kv is None else max_seqlen_kv + self.num_heads = num_heads + self.num_gqa_groups = num_heads if num_gqa_groups is None else num_gqa_groups + self.head_dim_qk = head_dim_qk + self.head_dim_v = head_dim_qk if head_dim_v is None else head_dim_v + if self.head_dim_qk == self.head_dim_v: + self.kv_channels = self.head_dim_qk + else: + self.kv_channels = (self.head_dim_qk, self.head_dim_v) + self.hidden_size = self.num_heads * self.head_dim_qk + self.hidden_size_kv = self.num_gqa_groups * self.head_dim_v + self.softmax_type = softmax_type + self.dropout_p = dropout_p + self.attn_mask_type = attn_mask_type + self.attn_bias_type = attn_bias_type + self.alibi_type = alibi_type + self.attn_type = ( + "self" if (self.max_seqlen_q == self.max_seqlen_kv) else "cross" + ) + self.bias_shape = bias_shape + self.window_size = window_size + self.context_parallel = context_parallel + self.cp_comm_type = cp_comm_type + self.return_max_logit = return_max_logit + self.total_requests = total_requests + self.max_ctx_len = max_ctx_len + self.num_layers = num_layers + self.eps = eps + + +def _run_dot_product_attention( + dtype: torch.dtype, + config, + backend: str, + ckpt_attn: bool, + qkv_layout: str, + workspace_opt: bool, + pad_between_seqs: bool, + is_training: bool, +) -> Tuple[torch.Tensor, Tuple[torch.Tensor, torch.Tensor, torch.Tensor]]: + """Run DotProductAttention module with one forward pass and one backward pass""" + # Set RNG and environment varables + os.environ["NVTE_FLASH_ATTN"] = "0" + os.environ["NVTE_FUSED_ATTN"] = "0" + if backend == "FlashAttention": + os.environ["NVTE_FLASH_ATTN"] = "1" + if backend == "FusedAttention": + os.environ["NVTE_FUSED_ATTN"] = "1" + os.environ["NVTE_FUSED_ATTN_FORCE_WORKSPACE_OPT"] = ( + "1" if workspace_opt else "0" + ) + + # Create seqlens + qkv_format = "".join([i for i in qkv_layout.split("_")[0] if i.isalpha()]) + if ("padding" in config.attn_mask_type or qkv_format == "thd") and False: + if config.attn_type == "self": + seqlens_q = torch.randint( + 1, + config.max_seqlen_q, + [config.batch_size], + dtype=torch.int32, + device="cuda", + ) + seqlens_kv = seqlens_q + if config.attn_type == "cross": + if config.max_seqlen_q > 1: + seqlens_q = torch.randint( + 1, + config.max_seqlen_q, + [config.batch_size], + dtype=torch.int32, + device="cuda", + ) + else: + seqlens_q = torch.ones( + [config.batch_size], dtype=torch.int32, device="cuda" + ) + seqlens_kv = torch.randint( + 1, + config.max_seqlen_kv, + [config.batch_size], + dtype=torch.int32, + device="cuda", + ) + else: + seqlens_q = torch.full( + [config.batch_size], config.max_seqlen_q, dtype=torch.int32, device="cuda" + ) + seqlens_kv = torch.full( + [config.batch_size], config.max_seqlen_kv, dtype=torch.int32, device="cuda" + ) + cu_seqlens_q = torch.zeros(config.batch_size + 1, dtype=torch.int32, device="cuda") + cu_seqlens_kv = torch.zeros(config.batch_size + 1, dtype=torch.int32, device="cuda") + cu_seqlens_q[1:] = torch.cumsum(seqlens_q, dim=0) + cu_seqlens_kv[1:] = torch.cumsum(seqlens_kv, dim=0) + + seqlens_q_after_pad = seqlens_q.clone() + seqlens_kv_after_pad = seqlens_kv.clone() + cu_seqlens_q_after_pad = cu_seqlens_q.clone() + cu_seqlens_kv_after_pad = cu_seqlens_kv.clone() + pad_len = [0] * config.batch_size + if pad_between_seqs: + max_pad_len = 3 + pad_len = torch.randint( + 0, max_pad_len + 1, [config.batch_size], device="cuda" + ) # 3 + seqlens_q_after_pad = seqlens_q + pad_len + seqlens_kv_after_pad = seqlens_kv + pad_len + cu_seqlens_q_after_pad[1:] = torch.cumsum(seqlens_q_after_pad, dim=0) + cu_seqlens_kv_after_pad[1:] = torch.cumsum(seqlens_kv_after_pad, dim=0) + + # Create attention mask if padding + attention_mask = None + if "padding" in config.attn_mask_type: + if config.attn_type == "self": + attention_mask_q = torch.Tensor([]).to(dtype=torch.bool) + for i in range(config.batch_size): + attention_mask_q = torch.cat( + [ + attention_mask_q, + torch.Tensor( + [False] * seqlens_q[i] + + [True] * (config.max_seqlen_q - seqlens_q[i]) + ) + .to(dtype=torch.bool) + .unsqueeze(0) + .unsqueeze(0) + .unsqueeze(0), + ], + dim=0, + ) + attention_mask = attention_mask_q.to(device="cuda") + if config.attn_type == "cross": + attention_mask_q = torch.Tensor([]).to(dtype=torch.bool) + attention_mask_kv = torch.Tensor([]).to(dtype=torch.bool) + for i in range(config.batch_size): + attention_mask_q = torch.cat( + [ + attention_mask_q, + torch.Tensor( + [False] * seqlens_q[i] + + [True] * (config.max_seqlen_q - seqlens_q[i]) + ) + .to(dtype=torch.bool) + .unsqueeze(0) + .unsqueeze(0) + .unsqueeze(0), + ], + dim=0, + ) + attention_mask_kv = torch.cat( + [ + attention_mask_kv, + torch.Tensor( + [False] * seqlens_kv[i] + + [True] * (config.max_seqlen_kv - seqlens_kv[i]) + ) + .to(dtype=torch.bool) + .unsqueeze(0) + .unsqueeze(0) + .unsqueeze(0), + ], + dim=0, + ) + attention_mask = ( + attention_mask_q.to(device="cuda"), + attention_mask_kv.to(device="cuda"), + ) + + alibi_slopes = None + + # Create input tensors + dim_to_num = { + "b": config.batch_size, + "sq": config.max_seqlen_q, + "skv": config.max_seqlen_kv, + "h": config.num_heads, + "hg": config.num_gqa_groups, + "dqk": config.head_dim_qk, + "dv": config.head_dim_v, + "t": cu_seqlens_q_after_pad[-1], + "tg": cu_seqlens_kv_after_pad[-1], + "3": 3, + "2": 2, + "1": 1, + } + inp = [] + inp_orig = [] + for i, layout in enumerate(qkv_layout.split("_")): + layout = "_".join(layout) + if i == 0: + layout = layout.replace("s", "sq") + else: + layout = layout.replace("s", "skv") + layout = layout.replace("h", "hg") + layout = layout.replace("t", "tg") + if i == 2: + layout = layout.replace("d", "dv") + else: + layout = layout.replace("d", "dqk") + tensor_shape = [dim_to_num[j] for j in layout.split("_")] + tensor = 0.1 * torch.randn(tensor_shape, dtype=dtype, device="cuda") + # tensor: with padding tokens + # tensor_orig: without padding tokens + tensor_orig = tensor + if qkv_format == "thd" and pad_between_seqs: + tensor_orig = torch.Tensor([]).to(device="cuda", dtype=dtype) + if layout in ["t_h_dqk", "t_3_h_dqk", "t_h_3_dqk"]: + for i in range(1, config.batch_size + 1): + valid_range = ( + cu_seqlens_q_after_pad[i - 1], + cu_seqlens_q_after_pad[i] - pad_len[i - 1], + ) + pad_range = ( + cu_seqlens_q_after_pad[i] - pad_len[i - 1], + cu_seqlens_q_after_pad[i], + ) + tensor[pad_range[0] : pad_range[1]] = 0.0 + tensor_orig = torch.cat( + [tensor_orig, tensor[valid_range[0] : valid_range[1]]], dim=0 + ) + if layout in ["tg_hg_dqk", "tg_2_hg_dqk", "tg_hg_2_dqk", "tg_hg_dv"]: + for i in range(1, config.batch_size + 1): + valid_range = ( + cu_seqlens_kv_after_pad[i - 1], + cu_seqlens_kv_after_pad[i] - pad_len[i - 1], + ) + pad_range = ( + cu_seqlens_kv_after_pad[i] - pad_len[i - 1], + cu_seqlens_kv_after_pad[i], + ) + tensor[pad_range[0] : pad_range[1]] = 0.0 + tensor_orig = torch.cat( + [tensor_orig, tensor[valid_range[0] : valid_range[1]]], dim=0 + ) + tensor_count = 1 + split_dim = 0 + for dim, l in enumerate(layout.split("_")): + if l.isdigit(): + tensor_count = int(l) + split_dim = dim + break + tensors = torch.split(tensor, 1, dim=split_dim) if split_dim != 0 else [tensor] + tensors_orig = ( + torch.split(tensor_orig, 1, dim=split_dim) + if split_dim != 0 + else [tensor_orig] + ) + for j in range(tensor_count): + if split_dim != 0: + inp.append(tensors[j].squeeze(split_dim)) + inp_orig.append(tensors_orig[j].squeeze(split_dim)) + else: + inp.append(tensors[j]) + inp_orig.append(tensors_orig[j]) + for i in range(3): + inp[i].requires_grad = True + inp_orig[i].requires_grad = True + + # Create output gradient + qkv_format_kv = "_".join(qkv_format) + qkv_format_kv = qkv_format_kv.replace("s", "sq") + qkv_format_kv = qkv_format_kv.replace("d", "dv") + out_grad_shape = [dim_to_num[i] for i in qkv_format_kv.split("_")] + out_grad_shape_new = [*out_grad_shape[:-2], out_grad_shape[-2] * out_grad_shape[-1]] + out_grad = 0.001 * torch.randint( + 0, 200, out_grad_shape_new, dtype=dtype, device="cuda" + ) + out_grad_orig = out_grad + if qkv_format == "thd" and pad_between_seqs: + out_grad_orig = torch.Tensor([]).to(device="cuda", dtype=dtype) + if qkv_format_kv == "t_h_dv": + for i in range(1, config.batch_size + 1): + valid_range = ( + cu_seqlens_q_after_pad[i - 1], + cu_seqlens_q_after_pad[i] - pad_len[i - 1], + ) + pad_range = ( + cu_seqlens_q_after_pad[i] - pad_len[i - 1], + cu_seqlens_q_after_pad[i], + ) + out_grad[pad_range[0] : pad_range[1]] = 0.0 + out_grad_orig = torch.cat( + [out_grad_orig, out_grad[valid_range[0] : valid_range[1]]], dim=0 + ) + + # Create bias + if config.attn_bias_type in ["no_bias"]: + bias = None + + # # Create RNG + # _DUMMY_CUDA_RNG_STATE_TRACKER = CudaRNGStatesTracker() + # _DUMMY_CUDA_RNG_STATE_TRACKER.add("model-parallel-rng", seed) + + # def get_dummy_cuda_rng_tracker() -> CudaRNGStatesTracker: + # """Get cuda rng tracker.""" + # return _DUMMY_CUDA_RNG_STATE_TRACKER + + # Set up model + block = DotProductAttention( + config.num_heads, + (config.head_dim_qk, config.head_dim_v), + num_gqa_groups=config.num_gqa_groups, + attention_dropout=config.dropout_p, + qkv_format=qkv_format, + attn_mask_type=config.attn_mask_type, + sequence_parallel=False, + tp_size=1, + get_rng_state_tracker=None, + tp_group=None, + layer_number=1, + attention_type=config.attn_type, + softmax_type=config.softmax_type, + return_max_logit=config.return_max_logit, + ).to(dtype=dtype, device="cuda") + if not is_training: + block = block.eval() + if is_training and config.softmax_type != "vanilla": + block.softmax_offset.requires_grad = True + + # Run a forward and backward pass + if backend in ["FlashAttention", "UnfusedDotProductAttention"]: + q = inp_orig[0] + k = inp_orig[1] + v = inp_orig[2] + d_out = out_grad_orig + if backend == "FusedAttention": + q = inp[0] + k = inp[1] + v = inp[2] + d_out = out_grad + out = block( + q, + k, + v, + window_size=config.window_size, + attention_mask=attention_mask, + qkv_format=qkv_format, + max_seqlen_q=config.max_seqlen_q, + max_seqlen_kv=config.max_seqlen_kv, + cu_seqlens_q=cu_seqlens_q, + cu_seqlens_kv=cu_seqlens_kv, + cu_seqlens_q_padded=( + cu_seqlens_q_after_pad if backend == "FusedAttention" else None + ), + cu_seqlens_kv_padded=( + cu_seqlens_kv_after_pad if backend == "FusedAttention" else None + ), + attn_mask_type=config.attn_mask_type, + checkpoint_core_attention=ckpt_attn, + core_attention_bias_type=config.attn_bias_type, + core_attention_bias=bias, + alibi_slopes=alibi_slopes, + fast_zero_fill=True, + ) + max_logit = None + if config.return_max_logit: + out, max_logit = out + if is_training: + out.backward(d_out) + + d_softmax_offset = None + if is_training and config.softmax_type != "vanilla": + d_softmax_offset = block.softmax_offset.grad + + if backend in ["FlashAttention", "UnfusedDotProductAttention"]: + if is_training: + return out, max_logit, (q.grad, k.grad, v.grad, d_softmax_offset) + else: + return out, max_logit, (None, None, None, d_softmax_offset) + if backend == "FusedAttention": + if qkv_format == "thd" and pad_between_seqs: + out_orig = torch.Tensor([]).to(device="cuda", dtype=dtype) + if is_training: + q_grad_orig = torch.Tensor([]).to(device="cuda", dtype=dtype) + k_grad_orig = torch.Tensor([]).to(device="cuda", dtype=dtype) + v_grad_orig = torch.Tensor([]).to(device="cuda", dtype=dtype) + for i in range(1, config.batch_size + 1): + valid_range_q = ( + cu_seqlens_q_after_pad[i - 1], + cu_seqlens_q_after_pad[i] - pad_len[i - 1], + ) + valid_range_kv = ( + cu_seqlens_kv_after_pad[i - 1], + cu_seqlens_kv_after_pad[i] - pad_len[i - 1], + ) + out_orig = torch.cat( + [out_orig, out[valid_range_q[0] : valid_range_q[1]]], dim=0 + ) + if is_training: + q_grad_orig = torch.cat( + [q_grad_orig, q.grad[valid_range_q[0] : valid_range_q[1]]], + dim=0, + ) + k_grad_orig = torch.cat( + [k_grad_orig, k.grad[valid_range_kv[0] : valid_range_kv[1]]], + dim=0, + ) + v_grad_orig = torch.cat( + [v_grad_orig, v.grad[valid_range_kv[0] : valid_range_kv[1]]], + dim=0, + ) + if is_training: + return ( + out_orig, + max_logit, + (q_grad_orig, k_grad_orig, v_grad_orig, d_softmax_offset), + ) + else: + return out_orig, max_logit, (None, None, None, d_softmax_offset) + else: + if is_training: + return out, max_logit, (q.grad, k.grad, v.grad, d_softmax_offset) + else: + return out, max_logit, (None, None, None, d_softmax_offset) + + +def fused_attn(block, q, k, v, cu_seqlens_q, cu_seqlens_kv, mask): + out = block( + q, + k, + v, + window_size=(-1, 0), + attention_mask=mask, + qkv_format="thd", + max_seqlen_q=q.size(0), + max_seqlen_kv=q.size(0), + cu_seqlens_q=cu_seqlens_q, + cu_seqlens_kv=cu_seqlens_kv, + cu_seqlens_q_padded=cu_seqlens_q, + cu_seqlens_kv_padded=cu_seqlens_kv, + attn_mask_type="padding_causal", + checkpoint_core_attention=False, + core_attention_bias_type="no_bias", + core_attention_bias=None, + alibi_slopes=None, + fast_zero_fill=True, + ) + return out + + +def test_fused_attn(B=1, S=8192, H=64): + dtype = torch.bfloat16 + config = ModelConfig( + B, + S, + H, + 192, + max_seqlen_kv=S, + head_dim_v=128, + attn_mask_type="padding_causal", + window_size=(-1, 0), + ) + backend = "FusedAttention" + ckpt_attn = False + qkv_layout = "thd_thd_thd" + workspace_opt = True + pad_between_seqs = True + is_training = True + + _run_dot_product_attention( + dtype, + config, + backend, + ckpt_attn, + qkv_layout, + workspace_opt, + pad_between_seqs, + is_training, + ) + benchmark_func( + _run_dot_product_attention, + dtype, + config, + backend, + ckpt_attn, + qkv_layout, + workspace_opt, + pad_between_seqs, + is_training, + n_profile=0, + trace_dir=None, + ) + + +def bench_fused_attn(B=1, S=8192, H=64): + dtype = torch.bfloat16 + config = ModelConfig( + B, + S, + H, + 192, + max_seqlen_kv=S, + head_dim_v=128, + attn_mask_type="padding_causal", + window_size=(-1, 0), + ) + backend = "FusedAttention" + ckpt_attn = False + # qkv_layout = 'sbhd_sbhd_sbhd' + # qkv_layout = 'bshd_bshd_bshd' + qkv_layout = "thd_thd_thd" + workspace_opt = True + pad_between_seqs = True + is_training = True + + os.environ["NVTE_FLASH_ATTN"] = "0" + os.environ["NVTE_FUSED_ATTN"] = "0" + if backend == "FlashAttention": + os.environ["NVTE_FLASH_ATTN"] = "1" + if backend == "FusedAttention": + os.environ["NVTE_FUSED_ATTN"] = "1" + os.environ["NVTE_FUSED_ATTN_FORCE_WORKSPACE_OPT"] = ( + "1" if workspace_opt else "0" + ) + + block = DotProductAttention( + H, + (192, 128), + num_gqa_groups=64, + attention_dropout=0.0, + qkv_format="thd", + attn_mask_type="padding_causal", + sequence_parallel=False, + tp_size=1, + get_rng_state_tracker=None, + tp_group=None, + layer_number=1, + attention_type="self", + softmax_type="vanilla", + return_max_logit=False, + ).to(dtype=dtype, device="cuda") + if not is_training: + block = block.eval() + seqlens_q = torch.full( + [config.batch_size], config.max_seqlen_q, dtype=torch.int32, device="cuda" + ) + seqlens_kv = torch.full( + [config.batch_size], config.max_seqlen_kv, dtype=torch.int32, device="cuda" + ) + + cu_seqlens_q = torch.zeros(B + 1, dtype=torch.int32, device="cuda") + cu_seqlens_kv = torch.zeros(B + 1, dtype=torch.int32, device="cuda") + cu_seqlens_q[1:] = torch.cumsum(seqlens_q, dim=0) + cu_seqlens_kv[1:] = torch.cumsum(seqlens_kv, dim=0) + q = torch.randn((S, H, 192), device="cuda", dtype=dtype, requires_grad=True) + k = torch.randn((S, H, 192), device="cuda", dtype=dtype, requires_grad=True) + v = torch.randn((S, H, 128), device="cuda", dtype=dtype, requires_grad=True) + g = torch.randn((S, H * 128), device="cuda", dtype=dtype) + + mask = torch.zeros((1, 1, 1, S), device="cuda", dtype=torch.bool) + out = fused_attn(block, q, k, v, cu_seqlens_q, cu_seqlens_kv, mask) + out.backward(g, retain_graph=True) + + benchmark_func( + fused_attn, block, q, k, v, cu_seqlens_q, cu_seqlens_kv, mask, n_profile=1 + ) + + benchmark_func(out.backward, g, retain_graph=True, n_profile=1) + + +if __name__ == "__main__": + test_fused_attn(B=1, S=8192, H=64) + bench_fused_attn(B=1, S=8192, H=64) diff --git a/benchmark/bench_mla_rope.py b/benchmark/bench_mla_rope.py new file mode 100644 index 0000000..71ef061 --- /dev/null +++ b/benchmark/bench_mla_rope.py @@ -0,0 +1,355 @@ +import torch +from megatron.core.fusions.fused_mla_yarn_rope_apply import ( + fused_apply_mla_rope_for_kv, + fused_apply_mla_rope_for_q, +) + +from linghe.facade.rope import mla_rope +from linghe.tools.benchmark import benchmark_func +from linghe.tools.check import output_check + + +def rope_freqs(length, dim, rope_theta=10000.0): + inv_freq = 1.0 / ( + rope_theta ** (torch.arange(0, dim, 2, device="cuda:0").float() / dim) + ) + t = torch.arange(length, device="cuda:0", dtype=torch.int64).float() + freqs = torch.outer(t, inv_freq) + return freqs + + +def bench_mla_rope(B=2, L=4096, H=32, rope_theta=10000.0, transpose=True): + dtype = torch.bfloat16 + device = "cuda:0" + q = torch.randn(L, B, H, 192, dtype=dtype, device=device).requires_grad_() + kv = torch.randn(L, B, H, 256, dtype=dtype, device=device).requires_grad_() + k_pos_emb = ( + torch.randn(L, B, 64 + 512, dtype=dtype, device=device)[:, :, :64] + .view(L, B, 1, 64) + .requires_grad_() + ) + freqs = rope_freqs(L, 64, rope_theta=rope_theta) + freqs = torch.cat([freqs, freqs], -1) + freqs = freqs[:, None, None] + if transpose: + q_grad = torch.randn(B, L, H, 192, dtype=dtype, device=device) + k_grad = torch.randn(B, L, H, 192, dtype=dtype, device=device) + v_grad = torch.randn(B, L, H, 128, dtype=dtype, device=device) + else: + q_grad = torch.randn(L, B, H, 192, dtype=dtype, device=device) + k_grad = torch.randn(L, B, H, 192, dtype=dtype, device=device) + v_grad = torch.randn(L, B, H, 128, dtype=dtype, device=device) + + mscale = 1.0 + + rotary_pos_cos = freqs.cos() + rotary_pos_sin = freqs.sin() + q_ref = q.detach().clone().requires_grad_() + kv_ref = kv.detach().clone().requires_grad_() + k_pos_emb_ref = k_pos_emb.detach().clone().requires_grad_() + query_ref = fused_apply_mla_rope_for_q( + q_ref, + rotary_pos_cos, + rotary_pos_sin, + 128, + 64, + cu_seqlens_q=None, + cp_rank=0, + cp_size=1, + ) + key_ref, value_ref = fused_apply_mla_rope_for_kv( + kv_ref, + k_pos_emb_ref, + rotary_pos_cos, + rotary_pos_sin, + 64, + 128, + 128, + cu_seqlens_kv=None, + cp_rank=0, + cp_size=1, + ) + if transpose: + query_ref = query_ref.transpose(0, 1) + key_ref = key_ref.transpose(0, 1) + value_ref = value_ref.transpose(0, 1) + + query_ref.backward(gradient=q_grad.detach().clone(), retain_graph=True) + key_ref.backward(gradient=k_grad.detach().clone(), retain_graph=True) + value_ref.backward(gradient=v_grad.detach().clone(), retain_graph=True) + dq_ref = q_ref.grad + dkv_ref = kv_ref.grad + dp_ref = k_pos_emb_ref.grad + + qo, ko, vo = mla_rope( + q, + kv, + k_pos_emb, + freqs, + cu_seqlens_q=None, + cu_seqlens_kv=None, + mscale=mscale, + cp_size=1, + cp_rank=0, + transpose=transpose, + ) + qo.backward(gradient=q_grad, retain_graph=True) + ko.backward(gradient=k_grad, retain_graph=True) + vo.backward(gradient=v_grad, retain_graph=True) + dq = q.grad + dkv = kv.grad + dp = k_pos_emb.grad + + output_check(query_ref, qo, name="q") + output_check(key_ref, ko, name="k") + output_check(value_ref, vo, name="v") + + output_check(dq_ref, dq, name="dq") + output_check(dkv_ref, dkv, name="dkv") + output_check(dp_ref, dp, name="dp", atol=0.1, rtol=0.02) + + lbh = L * B * H + benchmark_func( + fused_apply_mla_rope_for_q, + q, + rotary_pos_cos, + rotary_pos_sin, + 128, + 64, + cu_seqlens_q=None, + cp_rank=0, + cp_size=1, + ref_bytes=lbh * (64 * 2 + 256 * 2 + 64 * 2 + 192 * 2 + 128 * 2), + n_profile=0, + ) + benchmark_func( + fused_apply_mla_rope_for_kv, + kv, + k_pos_emb, + rotary_pos_cos, + rotary_pos_sin, + 64, + 128, + 128, + cu_seqlens_kv=None, + cp_rank=0, + cp_size=1, + ref_bytes=lbh * (64 * 2 + 256 * 2 + 64 * 2 + 192 * 2 + 128 * 2), + n_profile=0, + ) + benchmark_func( + mla_rope, + q, + kv, + k_pos_emb, + freqs, + ref_bytes=lbh * (64 * 2 + 256 * 2 + 64 * 2 + 192 * 2 + 128 * 2), + n_profile=0, + ) + + +def bench_varlen_mla_rope( + lengths=[2048, 2048], H=32, rope_theta=10000.0, cp_size=1, cp_rank=0, stride=True +): + dtype = torch.bfloat16 + device = "cuda:0" + q = torch.randn( + sum(lengths) // cp_size, H, 192, dtype=dtype, device=device + ).requires_grad_() + kv = torch.randn( + sum(lengths) // cp_size, H, 256, dtype=dtype, device=device + ).requires_grad_() + if stride: + k_pos_emb = torch.randn( + sum(lengths) // cp_size, 576, dtype=dtype, device=device + ) + k_pos_emb = ( + k_pos_emb[:, 512:].view(sum(lengths) // cp_size, 1, 64).requires_grad_() + ) + else: + k_pos_emb = torch.randn( + sum(lengths) // cp_size, 1, 64, dtype=dtype, device=device + ).requires_grad_() + cu_seqlens_q = torch.cumsum( + torch.tensor([0] + lengths, device=device, dtype=torch.int32), 0 + ).to(torch.int32) + cu_seqlens_kv = cu_seqlens_q + + freqs = rope_freqs((max(lengths) - 1) // 32 * 32 + 32, 64, rope_theta=rope_theta) + freqs = torch.cat([freqs, freqs], -1)[:, None, None] + + q_grad = torch.randn(sum(lengths) // cp_size, H, 192, dtype=dtype, device=device) + k_grad = torch.randn(sum(lengths) // cp_size, H, 192, dtype=dtype, device=device) + v_grad = torch.randn(sum(lengths) // cp_size, H, 128, dtype=dtype, device=device) + + mscale = 1.0 + + rotary_pos_cos = freqs.cos() + rotary_pos_sin = freqs.sin() + q_ref = q.detach().clone().requires_grad_() + kv_ref = kv.detach().clone().requires_grad_() + k_pos_emb_ref = k_pos_emb.detach().clone().requires_grad_() + query_ref = fused_apply_mla_rope_for_q( + q_ref, + rotary_pos_cos, + rotary_pos_sin, + 128, + 64, + cu_seqlens_q, + cp_rank, + cp_size, + ) + key_ref, value_ref = fused_apply_mla_rope_for_kv( + kv_ref, + k_pos_emb_ref, + rotary_pos_cos, + rotary_pos_sin, + 64, + 128, + 128, + cu_seqlens_kv, + cp_rank, + cp_size, + ) + + query_ref.backward(gradient=q_grad.clone().detach(), retain_graph=True) + key_ref.backward(gradient=k_grad.clone().detach(), retain_graph=True) + value_ref.backward(gradient=v_grad.clone().detach(), retain_graph=True) + dq_ref = q_ref.grad + dkv_ref = kv_ref.grad + dp_ref = k_pos_emb_ref.grad + + qo, ko, vo = mla_rope( + q, + kv, + k_pos_emb, + freqs, + cu_seqlens_q=cu_seqlens_q, + cu_seqlens_kv=cu_seqlens_kv, + mscale=mscale, + cp_size=cp_size, + cp_rank=cp_rank, + transpose=False, + ) + + qo.backward(gradient=q_grad.clone().detach(), retain_graph=True) + ko.backward(gradient=k_grad, retain_graph=True) + vo.backward(gradient=v_grad, retain_graph=True) + dq = q.grad + dkv = kv.grad + dp = k_pos_emb.grad + + output_check(query_ref, qo, name="q") + output_check(key_ref, ko, name="k") + output_check(value_ref, vo, name="v") + + output_check(dq_ref, dq, name="dq") + output_check(dkv_ref, dkv, name="dkv") + output_check(dp_ref, dp, name="dp", atol=0.1, rtol=0.02) + + lbh = sum(lengths) // cp_size * H + benchmark_func( + fused_apply_mla_rope_for_q, + q, + rotary_pos_cos, + rotary_pos_sin, + 128, + 64, + cu_seqlens_q, + cp_size=cp_size, + cp_rank=cp_rank, + ref_bytes=lbh * (64 * 2 + 256 * 2 + 64 * 2 + 192 * 2 + 128 * 2), + n_profile=0, + ) + benchmark_func( + fused_apply_mla_rope_for_kv, + kv, + k_pos_emb, + rotary_pos_cos, + rotary_pos_sin, + 64, + 128, + 128, + cu_seqlens_kv, + cp_size=cp_size, + cp_rank=cp_rank, + ref_bytes=lbh * (64 * 2 + 256 * 2 + 64 * 2 + 192 * 2 + 128 * 2), + n_profile=0, + ) + benchmark_func( + mla_rope, + q, + kv, + k_pos_emb, + freqs, + mscale=mscale, + cu_seqlens_q=cu_seqlens_q, + cu_seqlens_kv=cu_seqlens_kv, + cp_size=cp_size, + cp_rank=cp_rank, + ref_bytes=lbh * (64 * 2 + 256 * 2 + 64 * 2 + 192 * 2 + 128 * 2), + n_profile=0, + ) + + benchmark_func( + query_ref.backward, + q_grad, + retain_graph=True, + ref_bytes=lbh * (64 * 2 + 256 * 2 + 64 * 2 + 192 * 2 + 128 * 2), + n_profile=0, + ) + benchmark_func( + key_ref.backward, + k_grad, + retain_graph=True, + ref_bytes=lbh * (64 * 2 + 256 * 2 + 64 * 2 + 192 * 2 + 128 * 2), + n_profile=0, + ) + benchmark_func( + qo.backward, + q_grad, + retain_graph=True, + ref_bytes=lbh * (64 * 2 + 256 * 2 + 64 * 2 + 192 * 2 + 128 * 2), + n_profile=0, + ) + + +if __name__ == "__main__": + # bench_mla_rope(L=4096, B=2, H=32, transpose=True) + # bench_mla_rope(L=4096, B=2, H=16, transpose=True) + # bench_mla_rope(L=4096, B=2, H=64, transpose=True) + # bench_mla_rope(L=4096, B=2, H=16, transpose=True) + # bench_mla_rope(L=4096, B=1, H=16, transpose=True) + # bench_mla_rope(L=4096, B=1, H=16, transpose=False) + bench_varlen_mla_rope( + lengths=[444, 503, 434, 433, 472, 483, 557, 770], + H=16, + rope_theta=10000.0, + cp_size=1, + cp_rank=0, + stride=True, + ) + bench_varlen_mla_rope( + lengths=[444, 503, 434, 433, 472, 483, 557, 770], + H=32, + rope_theta=10000.0, + cp_size=2, + cp_rank=0, + stride=False, + ) + bench_varlen_mla_rope( + lengths=[444, 503, 434, 433, 472, 483, 557, 770], + H=32, + rope_theta=10000.0, + cp_size=2, + cp_rank=1, + stride=False, + ) + bench_varlen_mla_rope( + lengths=[444, 503, 434, 433, 472, 483, 557, 770], + H=32, + rope_theta=10000.0, + cp_size=4, + cp_rank=3, + stride=False, + ) diff --git a/benchmark/bench_norm.py b/benchmark/bench_norm.py new file mode 100644 index 0000000..ee1c91a --- /dev/null +++ b/benchmark/bench_norm.py @@ -0,0 +1,125 @@ +import torch +import transformer_engine as te +from transformer_engine.pytorch.constants import TE_DType +from transformer_engine.pytorch.tensor.float8_blockwise_tensor import ( + Float8BlockwiseQTensor, + Float8BlockQuantizer, +) +from transformer_engine.pytorch.tensor.mxfp8_tensor import MXFP8Tensor, MXFP8Quantizer + +from linghe.facade.norm import rms_norm, block_rms_norm, mxfp8_rms_norm +from linghe.tools.benchmark import benchmark_func + + +def bench_rmsnorm(bs=1, M=4096, N=4096): + # M, N, K = 8192, 4096, 13312 + # M, N, K = 4096, 4096, 6144 + # M, N, K = 4096, 4096, 4096 + # M, N, K = 4096, 8192, 4096 + + dtype = torch.bfloat16 + device = "cuda:0" + n_repeat = 100 + + x = torch.randn(bs, M, N, dtype=dtype, requires_grad=True, device=device) + weight = torch.randn(N, dtype=dtype, requires_grad=True, device=device) + dy = torch.randn(bs, M, N, dtype=dtype, device=device) + + rmsnorm_torch = torch.nn.RMSNorm( + normalized_shape=N, eps=1e-6, dtype=torch.bfloat16, device="cuda" + ) + + rmsnorm_torch = torch.compile(rmsnorm_torch) + + rmsnorm_te = te.pytorch.RMSNorm(normalized_shape=N, eps=1e-6) + + def torch_forward_backward(x_torch_back, dy): + y_torch_back = rmsnorm_torch(x_torch_back) + y_torch_back.backward(gradient=dy) + return x_torch_back.grad, rmsnorm_torch.weight.grad + + def te_forward_backward(x_te_back, dy): + y_te_back = rmsnorm_te(x_te_back) + y_te_back.backward(gradient=dy) + return x_te_back.grad, rmsnorm_te.weight.grad + + def triton_forward_backward(x_triton_back, g_triton_back, dy): + y_triton_back = rms_norm(x_triton_back, g_triton_back) + y_triton_back.backward(gradient=dy) + return x_triton_back.grad, g_triton_back.grad + + ref_time = benchmark_func( + rmsnorm_torch, x, n_repeat=n_repeat, name="rms_torch", ref_bytes=M * N * 4 + ) + benchmark_func( + rmsnorm_te, + x, + n_repeat=n_repeat, + ref_bytes=M * N * 4, + name="rms_te", + ref_time=ref_time, + ) + benchmark_func( + rms_norm, + x, + weight, + n_repeat=n_repeat, + ref_bytes=M * N * 4, + name="rms_triton", + ref_time=ref_time, + ) + + quantizer = Float8BlockQuantizer( + TE_DType[torch.float8_e4m3fn], + rowwise=True, + columnwise=True, + amax_epsilon=0, + force_pow_2_scales=True, + block_scaling_dim=1, + ) + y = block_rms_norm( + x, weight, None, quantizer, Float8BlockwiseQTensor, is_recomputing=None + ) + y[0].backward(dy) + benchmark_func( + block_rms_norm, + x, + weight, + None, + quantizer, + Float8BlockwiseQTensor, + is_recomputing=None, + n_repeat=n_repeat, + ref_bytes=M * N * 4, + name="rms_triton", + ref_time=ref_time, + ) + + quantizer = MXFP8Quantizer(fp8_dtype=TE_DType[torch.float8_e4m3fn]) + y = mxfp8_rms_norm(x, weight, None, quantizer, MXFP8Tensor, is_recomputing=None) + y[0].backward(dy) + benchmark_func( + mxfp8_rms_norm, + x, + weight, + None, + quantizer, + MXFP8Tensor, + is_recomputing=None, + n_repeat=n_repeat, + ref_bytes=M * N * 4, + name="rms_triton", + ref_time=ref_time, + ) + + ref_time = benchmark_func(torch_forward_backward, x, dy, n_repeat=n_repeat) + + benchmark_func(te_forward_backward, x, dy, n_repeat=n_repeat, ref_time=ref_time) + + benchmark_func( + triton_forward_backward, x, weight, dy, n_repeat=n_repeat, ref_time=ref_time + ) + + +if __name__ == "__main__": + bench_rmsnorm(1, 4096, 4096) diff --git a/benchmark/bench_permutation.py b/benchmark/bench_permutation.py index b435096..03b146b 100644 --- a/benchmark/bench_permutation.py +++ b/benchmark/bench_permutation.py @@ -1,15 +1,23 @@ import torch -import random import transformer_engine.pytorch.triton.permutation as triton_permutation - +from transformer_engine.pytorch.constants import TE_DType +from transformer_engine.pytorch.module.fp8_padding import Fp8Padding +from transformer_engine.pytorch.tensor.float8_blockwise_tensor import ( + Float8BlockQuantizer, +) from linghe.tools.benchmark import benchmark_func -from linghe.utils.gather import triton_permute_with_mask_map, triton_make_row_id_map -from linghe.utils.scatter import (triton_scatter_add, - triton_unpermute_with_mask_map, - ) from linghe.tools.util import torch_make_indices -import transformer_engine_torch as tex +from linghe.utils.gather import ( + triton_permute_with_mask_map, + triton_make_row_id_map, + triton_batch_block_pad_permute_with_indices, + triton_make_row_id_map_and_index, +) +from linghe.utils.scatter import ( + triton_scatter_add, + triton_unpermute_with_mask_map, +) def torch_index_select(y, indices): @@ -29,43 +37,152 @@ def torch_fp16_scatter_add(x, outputs, indices, weights): return outputs +def split_permute_pad_quantize( + x, probs, mask_map, fp8_padding, out_tokens, token_count_per_expert_list +): + M, N = x.shape + n_experts = mask_map.size(1) + row_id_map = triton_permutation.make_row_id_map(mask_map, M, n_experts) + output, permuted_scale, permuted_probs = triton_permutation.permute_with_mask_map( + x, row_id_map, probs, None, M, n_experts, out_tokens, N, 1 + ) + output, _ = fp8_padding(output, token_count_per_expert_list) + permuted_probs, _ = fp8_padding( + permuted_probs.view(-1, 1), token_count_per_expert_list + ) + + quantizer = Float8BlockQuantizer( + TE_DType[torch.float8_e4m3fn], + rowwise=True, + columnwise=True, + amax_epsilon=0, + force_pow_2_scales=True, + block_scaling_dim=1, + ) + + qx = quantizer.make_empty( + output.shape, dtype=x.dtype, device=x.device, requires_grad=False + ) + qx = quantizer.update_quantized(output, qx) + + return qx, permuted_probs + + +def fused_permute_pad_quantize( + x, probs, mask_map, token_count_per_expert, token_count_per_expert_list +): + num_out_tokens = sum([(x + 15) // 16 * 16 for x in token_count_per_expert_list]) + row_id_map, pad_indices = triton_make_row_id_map_and_index( + mask_map, num_out_tokens, multiple_of=16 + ) + x_q, x_s, xt_q, xt_s, p = triton_batch_block_pad_permute_with_indices( + x, + token_count_per_expert, + pad_indices, + token_count_per_expert_list, + probs=probs, + round_scale=True, + ) + return x_q, x_s, xt_q, xt_s, p + + def bench_triton_permute_with_mask_map(M=4096, N=4096, n_experts=256, topk=8): - device = 'cuda:0' + device = "cuda:0" dtype = torch.bfloat16 x = torch.randn(M, N, dtype=dtype, device=device) - x_q = x.to(torch.float8_e4m3fn) scales = torch.randn(M, dtype=dtype, device=device) logits = torch.randn((M, n_experts), dtype=torch.float32, device=device) probs, mask_map, token_count_per_expert, indices, row_id_map = torch_make_indices( - logits, topk=topk, bias=0.0) + logits, topk=topk, bias=0.0 + ) out_tokens = sum(token_count_per_expert.tolist()) mega_row_id_map = triton_permutation.make_row_id_map(mask_map, M, n_experts) n_repeat = 100 - ref_time = benchmark_func(torch_fp16_index_select, x, scales, indices, - n_repeat=n_repeat) - benchmark_func(triton_permute_with_mask_map, x, scales, probs, row_id_map, - out_tokens, n_repeat=n_repeat, ref_time=ref_time) + ref_time = benchmark_func( + torch_fp16_index_select, x, scales, indices, n_repeat=n_repeat + ) + benchmark_func( + triton_permute_with_mask_map, + x, + scales, + probs, + row_id_map, + out_tokens, + n_repeat=n_repeat, + ref_time=ref_time, + ) + + scales_m = torch.randn((M, 1), dtype=dtype, device=device) + + benchmark_func( + triton_permutation.permute_with_mask_map, + x, + mega_row_id_map, + probs, + scales_m, + M, + n_experts, + out_tokens, + N, + 1, + n_repeat=n_repeat, + ref_time=ref_time, + ) + + +def bench_permute_pad_quantization(M=4096, N=4096, n_experts=32, topk=2): + device = "cuda:0" + dtype = torch.bfloat16 + x = torch.randn(M, N, dtype=dtype, device=device) + fp8_padding = Fp8Padding(32, 16) - scales_m = torch.randn((M,1), dtype=dtype, device=device) + logits = torch.randn((M, n_experts), dtype=torch.float32, device=device) + probs, mask_map, token_count_per_expert, indices, row_id_map = torch_make_indices( + logits, topk=topk, bias=0.0 + ) + token_count_per_expert_list = token_count_per_expert.tolist() + out_tokens = sum(token_count_per_expert_list) - benchmark_func(triton_permutation.permute_with_mask_map, x, - mega_row_id_map, probs, scales_m, M, - n_experts, out_tokens, N, 1, n_repeat=n_repeat, - ref_time=ref_time) + split_permute_pad_quantize( + x, probs, mask_map, fp8_padding, out_tokens, token_count_per_expert_list + ) + fused_permute_pad_quantize( + x, probs, mask_map, token_count_per_expert, token_count_per_expert_list + ) + + ref_time = benchmark_func( + split_permute_pad_quantize, + x, + probs, + mask_map, + fp8_padding, + out_tokens, + token_count_per_expert_list, + ) + benchmark_func( + fused_permute_pad_quantize, + x, + probs, + mask_map, + token_count_per_expert, + token_count_per_expert_list, + ref_time=ref_time, + ) def bench_triton_unpermute_with_mask_map(M=4098, N=4096, n_experts=32, topk=2): dtype = torch.bfloat16 - device = 'cuda:0' + device = "cuda:0" weights = torch.randn(M * topk, dtype=dtype, device=device) logits = torch.randn((M, n_experts), dtype=torch.float32, device=device) probs, mask_map, token_count_per_expert, indices, row_id_map = torch_make_indices( - logits, topk=topk, bias=0.0) + logits, topk=topk, bias=0.0 + ) token_count_per_expert_list = token_count_per_expert.tolist() out_tokens = sum(token_count_per_expert_list) @@ -73,24 +190,41 @@ def bench_triton_unpermute_with_mask_map(M=4098, N=4096, n_experts=32, topk=2): x = torch.randn(out_tokens, N, dtype=dtype, device=device) outputs = torch.zeros((M, N), dtype=dtype, device=device) - + mega_row_id_map = triton_permutation.make_row_id_map(mask_map, M, n_experts) n_repeat = 100 - ref_time = benchmark_func(triton_scatter_add, x, outputs, indices, - n_repeat=n_repeat) - benchmark_func(triton_unpermute_with_mask_map, x, row_id_map, - probs, n_repeat=n_repeat, ref_time=ref_time) - benchmark_func(triton_permutation.unpermute_with_mask_map, x, mega_row_id_map, - probs, None , M, n_experts, N) - - ref_time = benchmark_func(triton_permutation.make_row_id_map, mask_map, - M, n_experts, n_repeat=n_repeat) - benchmark_func(triton_make_row_id_map, mask_map, n_repeat=n_repeat, - ref_time=ref_time) - - -if __name__ == '__main__': - bench_triton_permute_with_mask_map(M=8192, N=2048, n_experts=32, topk=2) - bench_triton_unpermute_with_mask_map(M=2048*32, N=2048, n_experts=32, topk=2) - + ref_time = benchmark_func( + triton_scatter_add, x, outputs, indices, n_repeat=n_repeat + ) + benchmark_func( + triton_unpermute_with_mask_map, + x, + row_id_map, + probs, + n_repeat=n_repeat, + ref_time=ref_time, + ) + benchmark_func( + triton_permutation.unpermute_with_mask_map, + x, + mega_row_id_map, + probs, + None, + M, + n_experts, + N, + ) + + ref_time = benchmark_func( + triton_permutation.make_row_id_map, mask_map, M, n_experts, n_repeat=n_repeat + ) + benchmark_func( + triton_make_row_id_map, mask_map, n_repeat=n_repeat, ref_time=ref_time + ) + + +if __name__ == "__main__": + bench_triton_permute_with_mask_map(M=8192 * 4, N=2048, n_experts=32, topk=2) + bench_triton_unpermute_with_mask_map(M=8192 * 4, N=2048, n_experts=32, topk=2) + bench_permute_pad_quantization(M=8192 * 4, N=4096, n_experts=32, topk=2) diff --git a/benchmark/bench_quantization.py b/benchmark/bench_quantization.py new file mode 100644 index 0000000..79f75fd --- /dev/null +++ b/benchmark/bench_quantization.py @@ -0,0 +1,115 @@ +import random + +import torch +import transformer_engine_torch as tex +from transformer_engine.pytorch.constants import TE_DType +from transformer_engine.pytorch.tensor.float8_blockwise_tensor import ( + Float8BlockQuantizer, +) +from transformer_engine.pytorch.tensor.mxfp8_tensor import MXFP8Quantizer + +from linghe.quant.block import triton_block_quant, triton_blockwise_quant +from linghe.quant.mxfp8 import triton_batch_mxfp8_quant +from linghe.tools.benchmark import benchmark_func +from linghe.tools.check import output_check + + +def bench_blockwise_quantization(M=8192, N=4096, round_scale=True): + quantizer = Float8BlockQuantizer( + TE_DType[torch.float8_e4m3fn], + rowwise=True, + columnwise=True, + amax_epsilon=0, + force_pow_2_scales=round_scale, + block_scaling_dim=1, + ) + dtype = torch.bfloat16 + device = "cuda:0" + x = torch.randn((M, N), device=device, dtype=dtype) + x[:, -2:] = 0.0 + + qx = quantizer.make_empty((M, N), dtype=dtype, device=device, requires_grad=False) + qx = quantizer.update_quantized(x, qx) + xq_ref = qx._rowwise_data.view(torch.float8_e4m3fn) + xs_ref = qx._rowwise_scale_inv + xt_q_ref = qx._columnwise_data.view(torch.float8_e4m3fn) + xt_s_ref = qx._columnwise_scale_inv + xq, xs, xt_q, xt_s = triton_blockwise_quant(x, round_scale=round_scale) + + output_check(xq_ref, xq, "x.data") + output_check(xs_ref, xs, "x.scale") + output_check(xt_q_ref, xt_q, "xt.data") + output_check(xt_s_ref, xt_s, "xt.scale") + + +def bench_block_quantization(M=8192, N=4096, round_scale=True): + dtype = torch.bfloat16 + device = "cuda:0" + weight_quantizer = Float8BlockQuantizer( + TE_DType[torch.float8_e4m3fn], + rowwise=True, + columnwise=True, + amax_epsilon=0, + force_pow_2_scales=round_scale, + block_scaling_dim=2, + ) + w = torch.randn((N, N), device=device, dtype=dtype) + qw = weight_quantizer.make_empty( + (N, N), dtype=dtype, device=device, requires_grad=False + ) + qw = weight_quantizer.update_quantized(w, qw) + wq_ref = qw._rowwise_data.view(torch.float8_e4m3fn) + ws_ref = qw._rowwise_scale_inv + wq, ws = triton_block_quant(w, round_scale=round_scale) + + output_check(wq_ref, wq, "w.data") + output_check(ws_ref, ws, "w.scale") + + +def bench_batch_mxfp8_quant(M=4096, N=4096, n_experts=32, bench=False): + dtype = torch.bfloat16 + device = "cuda:0" + + splits = [max(random.randint(M - 256, M + 256), 0) for x in range(n_experts)] + splits = [(x + 32) // 32 * 32 for x in splits] + print(sum(splits)) + token_count_per_expert = torch.tensor(splits, device=device) + quantizers = [ + MXFP8Quantizer(fp8_dtype=tex.DType.kFloat8E4M3) for _ in range(len(splits)) + ] + + x = torch.randn((sum(splits), N), dtype=dtype, device=device) + + inputmats = tex.split_quantize(x, splits, quantizers) + + x_q, x_scale, xt_q, xt_scale = triton_batch_mxfp8_quant( + x, token_count_per_expert, splits, output_mode=2 + ) + + # output_check(x_q_ref, x_q, 'x_q') + # output_check(x_scale_ref, x_scale, 'x_scale') + # output_check(xt_q_ref, xt_q, 'xt_q') + # output_check(xt_scale_ref, xt_scale, 'xt_scale') + + if bench: + ref_bytes = M * N * n_experts * 4 + benchmark_func(tex.split_quantize, x, splits, quantizers) + benchmark_func( + triton_batch_mxfp8_quant, + x, + token_count_per_expert, + splits, + output_mode=2, + ref_bytes=ref_bytes, + ) + + +if __name__ == "__main__": + bench_blockwise_quantization(M=8192, N=4096, round_scale=True) + # bench_blockwise_quantization(M=8192, N=4096, round_scale=False) + # bench_blockwise_quantization(M=16, N=4096, round_scale=True) + # bench_blockwise_quantization(M=16, N=4096, round_scale=False) + # bench_block_quantization(M=128, N=4096, round_scale=True) + # bench_block_quantization(M=128, N=4096, round_scale=False) + # bench_batch_mxfp8_quant(M=4096, N=2048, n_experts=32, bench=False) + # bench_batch_mxfp8_quant(M=4096, N=2048, n_experts=32, bench=True) diff --git a/benchmark/bench_rmsnorm.py b/benchmark/bench_rmsnorm.py deleted file mode 100644 index 1af6f11..0000000 --- a/benchmark/bench_rmsnorm.py +++ /dev/null @@ -1,64 +0,0 @@ -import torch -import transformer_engine as te - -from linghe.facade.norm import RMSNormFunction -from linghe.tools.benchmark import benchmark_func - - -def bench_rmsnorm(M=4096, N=4096): - # M, N, K = 8192, 4096, 13312 - # M, N, K = 4096, 4096, 6144 - # M, N, K = 4096, 4096, 4096 - # M, N, K = 4096, 8192, 4096 - - dtype = torch.bfloat16 - device = 'cuda:0' - n_repeat = 100 - - x = torch.randn(M, N, dtype=dtype, requires_grad=True, device=device) - weight = torch.randn(N, dtype=dtype, requires_grad=True, device=device) - dy = torch.randn(M, N, dtype=dtype, device=device) - - rmsnorm_torch = torch.nn.RMSNorm( - normalized_shape=N, - eps=1e-6, - dtype=torch.bfloat16, - device='cuda' - ) - - rmsnorm_torch = torch.compile(rmsnorm_torch) - - te_norm = te.pytorch.RMSNorm(normalized_shape=N, eps=1e-6) - - def torch_forward_backward(x_torch_back, dy): - y_torch_back = rmsnorm_torch(x_torch_back) - y_torch_back.backward(gradient=dy) - return x_torch_back.grad, rmsnorm_torch.weight.grad - - def te_forward_backward(x_te_back, dy): - y_te_back = te_norm(x_te_back) - y_te_back.backward(gradient=dy) - return x_te_back.grad, te_norm.weight.grad - - def triton_forward_backward(x_triton_back, g_triton_back, dy): - y_triton_back = RMSNormFunction.apply(x_triton_back, g_triton_back) - y_triton_back.backward(gradient=dy) - return x_triton_back.grad, g_triton_back.grad - - ref_time = benchmark_func(rmsnorm_torch, x, n_repeat=n_repeat, - name="rms_torch", ref_bytes=M * N * 4) - benchmark_func(te_norm, x, n_repeat=n_repeat, ref_bytes=M * N * 4, - name="rms_te", ref_time=ref_time) - benchmark_func(RMSNormFunction.apply, x, weight, n_repeat=n_repeat, - ref_bytes=M * N * 4, name="rms_triton", ref_time=ref_time) - - ref_time = benchmark_func(torch_forward_backward, x, dy, n_repeat=n_repeat) - - benchmark_func(te_forward_backward, x, dy, n_repeat=n_repeat, - ref_time = ref_time) - - benchmark_func(triton_forward_backward, x, weight, dy, n_repeat=n_repeat, - ref_time=ref_time) - -if __name__ == '__main__': - bench_rmsnorm(4096, 4096) diff --git a/benchmark/bench_topk.py b/benchmark/bench_topk.py new file mode 100644 index 0000000..4204653 --- /dev/null +++ b/benchmark/bench_topk.py @@ -0,0 +1,104 @@ +# -*- coding: utf-8 -*- +""" +Copyright (c) Ant Financial Service Group and its affiliates. +""" + +import torch +from megatron.core.transformer.moe.moe_utils import topk_softmax_with_capacity + +from linghe.facade.topk import fused_topk, group_topk_score +from linghe.tools.benchmark import benchmark_func + + +def bench_topk(M=4096, N=256, k=8): + device = "cuda:0" + logits = torch.randn((M, N), dtype=torch.float32, device=device) + logits = logits.detach().clone().requires_grad_() + + values_ref, indices_ref = torch.topk(logits, k) + # print(f'{indices_ref[0]}') + values, indices = fused_topk(logits, k) + # print(f'{indices[0]}') + + ref_time = benchmark_func(torch.topk, logits, k, ref_bytes=M * N * 4) + benchmark_func(fused_topk, logits, k, ref_time=ref_time, ref_bytes=M * N * 4) + + +def bench_group_topk_score(M=4096, N=256, k=8): + device = "cuda:0" + # logits = torch.randn((M, N), dtype=torch.float32, device=device) + logits = torch.zeros((M, N), dtype=torch.float32, device=device) + 1.0 + logits = logits.detach().clone().requires_grad_() + + input_grad = 1 / M * torch.randn((M, N), dtype=torch.float32, device=device) + num_groups = 32 + group_topk = 4 + scaling_factor = 2.5 + deterministic_mode = True + score_function = "sigmoid" + expert_bias = torch.randn((N,), dtype=torch.float32, device=device) + moe_router_fusion = True + + probs_ref, routing_map_ref, tokens_per_expert_ref = topk_softmax_with_capacity( + logits, + k, + None, + None, + None, + False, + num_groups, + group_topk, + scaling_factor, + deterministic_mode, + score_function, + expert_bias, + moe_router_fusion, + ) + # print((-routing_map_ref[0].float()).argsort(0)) + + probs, routing_map, tokens_per_expert = group_topk_score( + logits, + k, + expert_bias, + num_groups=num_groups, + group_topk=group_topk, + scaling_factor=scaling_factor, + score_function=score_function, + ) + # print((-routing_map[0].float()).argsort(0)) + + ref_time = benchmark_func( + topk_softmax_with_capacity, + logits, + k, + None, + None, + None, + False, + num_groups, + group_topk, + scaling_factor, + deterministic_mode, + score_function, + expert_bias, + moe_router_fusion, + ref_bytes=M * N * 4, + ) + + benchmark_func( + group_topk_score, + logits, + k, + expert_bias, + num_groups=num_groups, + group_topk=group_topk, + scaling_factor=scaling_factor, + score_function=score_function, + ref_bytes=M * N * 4, + ref_time=ref_time, + ) + + +if __name__ == "__main__": + bench_topk(M=8192, N=256, k=8) + bench_group_topk_score(M=8192, N=256, k=8) diff --git a/docs/linghe.html b/docs/linghe.html index 27dbb27..849f07e 100644 --- a/docs/linghe.html +++ b/docs/linghe.html @@ -24,9 +24,12 @@
Copyright (c) Ant Financial Service Group and its affiliates.
+Copyright (c) Ant Financial Service Group and its affiliates.
+kernels should be run with torch above 2.9.0
+Copyright (c) Ant Financial Service Group and its affiliates.
+inplace update embedding weight gradient
+ +++None
+
inplace update embedding weight gradient
+ +++None
+
Copyright (c) Ant Financial Service Group and its affiliates.
+Copyright (c) Ant Financial Service Group and its affiliates.
+tensor-parallel fc2 in the shared expert, use split-k implementation +y = all_reduce(x @ fc2)
+ +++c: all-reduced output
+
Wait for a global memory barrier to reach the expected state.
+ +This function implements a spin-wait loop that continuously checks a memory location +until it reaches the expected value, providing synchronization across GPU threads.
+ +Synchronizes blocks with matching block_id across participating devices.
+ +Note: the function itself is not a system level barrier/fence. It is a +building block for expressing different synchronization patterns.
+ +Pattern 0: Ensures that all writes to symm_mem buffers from previous +kernels across all devices are visible to the current kernel:
+ +symm_mem_sync(..., hasPreviousMemAccess=False, hasSubsequentMemAccess=True)
+
+
+Pattern 1: Ensures that all writes to symm_mem buffers from the current +block are visible to all remote blocks with matching blockIdx:
+ +symm_mem_sync(..., hasPreviousMemAccess=True, hasSubsequentMemAccess=True)
+
+
+Pattern 2: Ensures that symm_mem buffers read by the current kernel are safe +for writing by subsequent kernels across all devices.
+ +symm_mem_sync(..., hasPreviousMemAccess=True, hasSubsequentMemAccess=False)
+
+
+++This barrier operates through atomic operations on a zero-filled signal + pad, which resets to a zero-filled state after each successful + synchronization. This design eliminates the need for incrementing a + flag from host.
+
Copyright (c) Ant Financial Service Group and its affiliates.
+Copyright (c) Ant Financial Service Group and its affiliates.
+Copyright (c) Ant Financial Service Group and its affiliates.
+Copyright (c) Ant Financial Service Group and its affiliates.
+embedding lookup
+ +++lookup output
+
embedding lookup
+ +grad_name tensor++lookup output
+
embedding lookup
+ +++lookup output
+
Copyright (c) Ant Financial Service Group and its affiliates.
+return group_rms_norm(transpose(attn_output, [0,1]), weight) * sigmoid(gate)
+ +++output with shape [length, bs, dim]
+
softmax cross entropy
+ +++z loss
+