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 @@

Submodules

diff --git a/docs/linghe/attn.html b/docs/linghe/attn.html new file mode 100644 index 0000000..b5522d2 --- /dev/null +++ b/docs/linghe/attn.html @@ -0,0 +1,238 @@ + + + + + + + linghe.attn API documentation + + + + + + + + + +
+
+

+linghe.attn

+ + + + + +
+
+ + \ No newline at end of file diff --git a/docs/linghe/attn/la.html b/docs/linghe/attn/la.html new file mode 100644 index 0000000..d6ab560 --- /dev/null +++ b/docs/linghe/attn/la.html @@ -0,0 +1,239 @@ + + + + + + + linghe.attn.la API documentation + + + + + + + + + +
+
+

+linghe.attn.la

+ +

Copyright (c) Ant Financial Service Group and its affiliates.

+
+ + + + +
+
+ + \ No newline at end of file diff --git a/docs/linghe/attn/mla.html b/docs/linghe/attn/mla.html new file mode 100644 index 0000000..8b60593 --- /dev/null +++ b/docs/linghe/attn/mla.html @@ -0,0 +1,239 @@ + + + + + + + linghe.attn.mla API documentation + + + + + + + + + +
+
+

+linghe.attn.mla

+ +

Copyright (c) Ant Financial Service Group and its affiliates.

+
+ + + + +
+
+ + \ No newline at end of file diff --git a/docs/linghe/experimental.html b/docs/linghe/experimental.html new file mode 100644 index 0000000..b173b8e --- /dev/null +++ b/docs/linghe/experimental.html @@ -0,0 +1,246 @@ + + + + + + + linghe.experimental API documentation + + + + + + + + + +
+
+

+linghe.experimental

+ +

kernels should be run with torch above 2.9.0

+
+ + + + +
+
+ + \ No newline at end of file diff --git a/docs/linghe/experimental/demb.html b/docs/linghe/experimental/demb.html new file mode 100644 index 0000000..89affba --- /dev/null +++ b/docs/linghe/experimental/demb.html @@ -0,0 +1,303 @@ + + + + + + + linghe.experimental.demb API documentation + + + + + + + + + +
+
+

+linghe.experimental.demb

+ +

Copyright (c) Ant Financial Service Group and its affiliates.

+
+ + + + +
+
+
+ + def + triton_tp_embedding_lookup_backward(grad_output, x, g_ptr, vocab_size, hdl, group, dtype=torch.bfloat16): + + +
+ + +

inplace update embedding weight gradient

+ +
Arguments:
+ +
    +
  • y: gradient of output
  • +
  • x: input ids Tensor
  • +
  • g_ptr: data_ptr of embedding weight gradient
  • +
+ +
Returns:
+ +
+

None

+
+
+ + +
+
+
+ + def + triton_sp_embedding_lookup_backward( grad_output, input_ids, g_ptr, vocab_size, hdl, group, dtype=torch.bfloat16, gathered_input_ids=None): + + +
+ + +

inplace update embedding weight gradient

+ +
Arguments:
+ +
    +
  • y: gradient of output
  • +
  • x: input ids Tensor
  • +
  • g_ptr: data_ptr of embedding weight gradient
  • +
+ +
Returns:
+ +
+

None

+
+
+ + +
+
+ + \ No newline at end of file diff --git a/docs/linghe/experimental/dla.html b/docs/linghe/experimental/dla.html new file mode 100644 index 0000000..6bcbd3d --- /dev/null +++ b/docs/linghe/experimental/dla.html @@ -0,0 +1,239 @@ + + + + + + + linghe.experimental.dla API documentation + + + + + + + + + +
+
+

+linghe.experimental.dla

+ +

Copyright (c) Ant Financial Service Group and its affiliates.

+
+ + + + +
+
+ + \ No newline at end of file diff --git a/docs/linghe/experimental/dmm.html b/docs/linghe/experimental/dmm.html new file mode 100644 index 0000000..1a3f426 --- /dev/null +++ b/docs/linghe/experimental/dmm.html @@ -0,0 +1,271 @@ + + + + + + + linghe.experimental.dmm API documentation + + + + + + + + + +
+
+

+linghe.experimental.dmm

+ +

Copyright (c) Ant Financial Service Group and its affiliates.

+
+ + + + +
+
+
+ + def + triton_split_tp_gemm( x: torch.Tensor, w: torch.Tensor, hdl, group: torch.distributed.distributed_c10d.ProcessGroup): + + +
+ + +

tensor-parallel fc2 in the shared expert, use split-k implementation +y = all_reduce(x @ fc2)

+ +
Arguments:
+ +
    +
  • a: left matrix with bf16 precision
  • +
  • b: right matrix with bf16 precision
  • +
+ +
Returns:
+ +
+

c: all-reduced output

+
+
+ + +
+
+ + \ No newline at end of file diff --git a/docs/linghe/experimental/gmem_barrier_arrive_wait.html b/docs/linghe/experimental/gmem_barrier_arrive_wait.html new file mode 100644 index 0000000..9af94f9 --- /dev/null +++ b/docs/linghe/experimental/gmem_barrier_arrive_wait.html @@ -0,0 +1,272 @@ + + + + + + + linghe.experimental.gmem_barrier_arrive_wait API documentation + + + + + + + + + +
+
+

+linghe.experimental.gmem_barrier_arrive_wait

+ + + + + + +
+
+
+
@triton.jit
+ + def + wait_gmem_barrier( addr, expect: int = 1, update: int = 0, sem: int = 'acquire', scope: int = 'gpu', op: int = 'ld', skip_sync: int = False): + + +
+ + +

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.

+ +
Arguments:
+ +
    +
  • addr: Memory address of the barrier to wait on (Must be a scalar)
  • +
  • expect: Expected value to wait for (default: 1)
  • +
  • update: Update the barrier with once acquired (default: 0)
  • +
  • sem: Memory semantics for the atomic operation (default: "acquire")
  • +
  • scope: Scope of the atomic operation. Options: "gpu", "sys" (default: "gpu")
  • +
  • op: Atomic operation type (default: "ld", currently only supported option)
  • +
+
+ + +
+
+ + \ No newline at end of file diff --git a/docs/linghe/experimental/symm_mem_barrier.html b/docs/linghe/experimental/symm_mem_barrier.html new file mode 100644 index 0000000..102af58 --- /dev/null +++ b/docs/linghe/experimental/symm_mem_barrier.html @@ -0,0 +1,288 @@ + + + + + + + linghe.experimental.symm_mem_barrier API documentation + + + + + + + + + +
+
+

+linghe.experimental.symm_mem_barrier

+ + + + + + +
+
+
+
@triton.jit
+ + def + symm_mem_sync( signal_pad_ptrs, block_id, rank: int, world_size: int, hasPreviousMemAccess: int = False, hasSubsequentMemAccess: int = False): + + +
+ + +

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)
+
+ +
CUDA graph friendliness:
+ +
+

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.

+
+
+ + +
+
+ + \ No newline at end of file diff --git a/docs/linghe/experimental/test_demb.html b/docs/linghe/experimental/test_demb.html new file mode 100644 index 0000000..53ce88f --- /dev/null +++ b/docs/linghe/experimental/test_demb.html @@ -0,0 +1,239 @@ + + + + + + + linghe.experimental.test_demb API documentation + + + + + + + + + +
+
+

+linghe.experimental.test_demb

+ +

Copyright (c) Ant Financial Service Group and its affiliates.

+
+ + + + +
+
+ + \ No newline at end of file diff --git a/docs/linghe/experimental/test_dla.html b/docs/linghe/experimental/test_dla.html new file mode 100644 index 0000000..2490444 --- /dev/null +++ b/docs/linghe/experimental/test_dla.html @@ -0,0 +1,239 @@ + + + + + + + linghe.experimental.test_dla API documentation + + + + + + + + + +
+
+

+linghe.experimental.test_dla

+ +

Copyright (c) Ant Financial Service Group and its affiliates.

+
+ + + + +
+
+ + \ No newline at end of file diff --git a/docs/linghe/experimental/test_dmm.html b/docs/linghe/experimental/test_dmm.html new file mode 100644 index 0000000..c897db6 --- /dev/null +++ b/docs/linghe/experimental/test_dmm.html @@ -0,0 +1,239 @@ + + + + + + + linghe.experimental.test_dmm API documentation + + + + + + + + + +
+
+

+linghe.experimental.test_dmm

+ +

Copyright (c) Ant Financial Service Group and its affiliates.

+
+ + + + +
+
+ + \ No newline at end of file diff --git a/docs/linghe/facade.html b/docs/linghe/facade.html index ad2c300..3d4cbdc 100644 --- a/docs/linghe/facade.html +++ b/docs/linghe/facade.html @@ -30,12 +30,18 @@

Submodules

diff --git a/docs/linghe/facade/emb.html b/docs/linghe/facade/emb.html new file mode 100644 index 0000000..fcbbade --- /dev/null +++ b/docs/linghe/facade/emb.html @@ -0,0 +1,336 @@ + + + + + + + linghe.facade.emb API documentation + + + + + + + + + +
+
+

+linghe.facade.emb

+ +

Copyright (c) Ant Financial Service Group and its affiliates.

+
+ + + + +
+
+
+ + def + deprecated_fused_accumulation_embedding_lookup(x: torch.Tensor, w_ptr, g_ptr, dim, dtype, grad_dtype): + + +
+ + +

embedding lookup

+ +
Arguments:
+ +
    +
  • x: input ids
  • +
  • w_ptr:
  • +
  • g_ptr:
  • +
  • dim:
  • +
  • dtype:
  • +
  • grad_dtype:
  • +
+ +
Returns:
+ +
+

lookup output

+
+
+ + +
+
+
+ + def + fused_accumulation_embedding_lookup( x: torch.Tensor, w: torch.nn.parameter.Parameter, grad_name: str = 'grad'): + + +
+ + +

embedding lookup

+ +
Arguments:
+ +
    +
  • x: input ids
  • +
  • w: embedding weight, should contain a grad_name tensor
  • +
+ +
Returns:
+ +
+

lookup output

+
+
+ + +
+
+
+ + def + embedding_lookup(x: torch.Tensor, w: torch.nn.parameter.Parameter): + + +
+ + +

embedding lookup

+ +
Arguments:
+ +
    +
  • x: input ids
  • +
  • w: embedding weight
  • +
+ +
Returns:
+ +
+

lookup output

+
+
+ + +
+
+ + \ No newline at end of file diff --git a/docs/linghe/facade/gate.html b/docs/linghe/facade/gate.html new file mode 100644 index 0000000..aa2ce3c --- /dev/null +++ b/docs/linghe/facade/gate.html @@ -0,0 +1,273 @@ + + + + + + + linghe.facade.gate API documentation + + + + + + + + + +
+
+

+linghe.facade.gate

+ +

Copyright (c) Ant Financial Service Group and its affiliates.

+
+ + + + +
+
+
+ + def + group_rms_norm_gate( attn_output: torch.Tensor, gate: torch.Tensor, weight: torch.Tensor, eps: float = 1e-06, group_size: int = 4): + + +
+ + +

return group_rms_norm(transpose(attn_output, [0,1]), weight) * sigmoid(gate)

+ +
Arguments:
+ +
    +
  • attn_output: output of core attn, shape [bs, length, n_heads, head_dim]
  • +
  • gate: gate tensor for attention output, shape [length, bs, dim]
  • +
  • weight: weight of RMS norm, shape [dim]
  • +
  • eps: epsilon for RMS
  • +
  • group_size: group size of group RMS norm
  • +
+ +
Returns:
+ +
+

output with shape [length, bs, dim]

+
+
+ + +
+
+ + \ No newline at end of file diff --git a/docs/linghe/facade/loss.html b/docs/linghe/facade/loss.html index 8b6d236..024c99f 100644 --- a/docs/linghe/facade/loss.html +++ b/docs/linghe/facade/loss.html @@ -33,6 +33,9 @@

API Documentation

  • softmax_cross_entropy
  • +
  • + moe_z_loss +
  • @@ -60,7 +63,7 @@

    def - softmax_cross_entropy(logits: torch.Tensor, labels: torch.Tensor, inplace: bool = False): + softmax_cross_entropy( logits: torch.Tensor, labels: torch.Tensor, ignore_index: int = -100, inplace: bool = False, tp_group=None):
    @@ -84,6 +87,34 @@

    Returns:
    + +
    +
    + + def + moe_z_loss(logits: torch.Tensor, coef: float = 0.001): + + +
    + + +

    softmax cross entropy

    + +
    Arguments:
    + + + +
    Returns:
    + +
    +

    z loss

    +
    +
    + +
    + \ No newline at end of file diff --git a/docs/linghe/facade/norm.html b/docs/linghe/facade/norm.html index 279f1c6..2b3e878 100644 --- a/docs/linghe/facade/norm.html +++ b/docs/linghe/facade/norm.html @@ -34,7 +34,16 @@

    API Documentation

    rms_norm
  • - group_rms_norm_gate + BlockRMSNorm + +
  • @@ -88,36 +97,146 @@
    Returns:
    -
    -
    +
    +
    + class + BlockRMSNorm(torch.autograd.function.Function): + + +
    + + +

    Base class to create custom autograd.Function.

    + +

    To create a custom autograd.Function, subclass this class and implement +the forward() and backward() static methods. Then, to use your custom +op in the forward pass, call the class method apply. Do not call +forward() directly.

    + +

    To ensure correctness and best performance, make sure you are calling the +correct methods on ctx and validating your backward function using +torch.autograd.gradcheck().

    + +

    See :ref:extending-autograd for more details on how to use this class.

    + +

    Examples::

    + +
    >>> # xdoctest: +REQUIRES(env:TORCH_DOCTEST_AUTOGRAD)
    +>>> class Exp(Function):
    +>>>     @staticmethod
    +>>>     def forward(ctx, i):
    +>>>         result = i.exp()
    +>>>         ctx.save_for_backward(result)
    +>>>         return result
    +>>>
    +>>>     @staticmethod
    +>>>     def backward(ctx, grad_output):
    +>>>         result, = ctx.saved_tensors
    +>>>         return grad_output * result
    +>>>
    +>>> # Use it by calling the apply method:
    +>>> # xdoctest: +SKIP
    +>>> output = Exp.apply(input)
    +
    +
    + + +
    +
    +
    @staticmethod
    + def - group_rms_norm_gate( attn_output: torch.Tensor, gate: torch.Tensor, weight: torch.Tensor, eps: float = 1e-06, group_size: int = 4): + forward(ctx, input, weight, rms, eps, quantizer, cls, is_recomputing):
    - + -

    return group_rms_norm(transpose(attn_output, [0,1]), weight) * sigmoid(gate)

    +

    Define the forward of the custom autograd Function.

    -
    Arguments:
    +

    This function is to be overridden by all subclasses. +There are two ways to define forward:

    + +

    Usage 1 (Combined forward and ctx)::

    + +
    @staticmethod
    +def forward(ctx: Any, *args: Any, **kwargs: Any) -> Any:
    +    pass
    +
      -
    • attn_output: output of core attn, shape [bs, length, n_heads, head_dim]
    • -
    • gate: gate tensor for attention output, shape [length, bs, dim]
    • -
    • weight: weight of RMS norm, shape [dim]
    • -
    • eps: epsilon for RMS
    • -
    • group_size: group size of group RMS norm
    • +
    • It must accept a context ctx as the first argument, followed by any +number of arguments (tensors or other types).
    • +
    • See :ref:combining-forward-context for more details
    -
    Returns:
    +

    Usage 2 (Separate forward and ctx)::

    -
    -

    output with shape [length, bs, dim]

    -
    +
    @staticmethod
    +def forward(*args: Any, **kwargs: Any) -> Any:
    +    pass
    +
    +@staticmethod
    +def setup_context(ctx: Any, inputs: Tuple[Any, ...], output: Any) -> None:
    +    pass
    +
    + +
      +
    • The forward no longer accepts a ctx argument.
    • +
    • Instead, you must also override the torch.autograd.Function.setup_context() +staticmethod to handle setting up the ctx object. +output is the output of the forward, inputs are a Tuple of inputs +to the forward.
    • +
    • See :ref:extending-autograd for more details
    • +
    + +

    The context can be used to store arbitrary data that can be then +retrieved during the backward pass. Tensors should not be stored +directly on ctx (though this is not currently enforced for +backward compatibility). Instead, tensors should be saved either with +ctx.save_for_backward() if they are intended to be used in +backward (equivalently, vjp) or ctx.save_for_forward() +if they are intended to be used for in jvp.

    +
    + + +
    +
    +
    +
    @staticmethod
    + + def + backward(ctx, grad_output, grad_rms): + + +
    + + +

    Define a formula for differentiating the operation with backward mode automatic differentiation.

    + +

    This function is to be overridden by all subclasses. +(Defining this function is equivalent to defining the vjp function.)

    + +

    It must accept a context ctx as the first argument, followed by +as many outputs as the forward() returned (None will be passed in +for non tensor outputs of the forward function), +and it should return as many tensors, as there were inputs to +forward(). Each argument is the gradient w.r.t the given output, +and each returned value should be the gradient w.r.t. the +corresponding input. If an input is not a Tensor or is a Tensor not +requiring grads, you can just pass None as a gradient for that input.

    + +

    The context can be used to retrieve tensors saved during the forward +pass. It also has an attribute ctx.needs_input_grad as a tuple +of booleans representing whether each input needs gradient. E.g., +backward() will have ctx.needs_input_grad[0] = True if the +first input to forward() needs gradient computed w.r.t. the +output.

    +
    + \ No newline at end of file diff --git a/docs/linghe/facade/rope.html b/docs/linghe/facade/rope.html index b9189f9..0d5b122 100644 --- a/docs/linghe/facade/rope.html +++ b/docs/linghe/facade/rope.html @@ -33,6 +33,9 @@

    API Documentation

  • qk_norm_half_rope
  • +
  • + mla_rope +
  • @@ -60,7 +63,7 @@

    def - qk_norm_half_rope( qkv: torch.Tensor, q_norm_weight: torch.Tensor, k_norm_weight: torch.Tensor, freqs: torch.Tensor, H: int = 32, h: int = 4, eps: float = 1e-06): + qk_norm_half_rope( qkv: torch.Tensor, q_norm_weight: torch.Tensor, k_norm_weight: torch.Tensor, freqs: torch.Tensor, cu_seqlens_q: Optional[torch.Tensor] = None, cu_seqlens_kv: Optional[torch.Tensor] = None, H: int = 32, h: int = 4, eps: float = 1e-06, cp_rank=0, cp_size=1, mscale=1.0, silu=False, reuse=False):
    @@ -71,22 +74,70 @@

    Arguments:
    + +
    Returns:
    + +
    +
      +
    • qo: shape [B, S, H, head_dim] or [T, H, head_dim]
    • +
    • ko: shape [B, S, h, head_dim] or [T, h, head_dim]
    • +
    • vo: shape [B, S, h, head_dim] or [T, h, head_dim]
    • +
    +
    +
    + + +
    +
    +
    + + def + mla_rope( q: torch.Tensor, kv: torch.Tensor, k_pos_emb: torch.Tensor, freqs: torch.Tensor, cu_seqlens_q: Optional[torch.Tensor] = None, cu_seqlens_kv: Optional[torch.Tensor] = None, mscale: float = 1.0, transpose: bool = False, cp_size: int = 1, cp_rank: int = 0, reuse: bool = False): + + +
    + + +

    inplace apply rope to tail 64 dims, split kv and apply rope to k_pos_emb and copy to k

    + +
    Arguments:
    + +
    Returns:
      -
    • qo: shape [B, S, H, head_dim]
    • -
    • ko: shape [B, S, h, head_dim]
    • -
    • vo: shape [B, S, h, head_dim]
    • +
    • qo: shape [S, B, H, 192] or [N, H, 192]
    • +
    • ko: shape [S, B, H, 192] or [N, H, 192]
    • +
    • vo: shape [S, B, H, 128] or [N, H, 128]
    diff --git a/docs/linghe/facade/silu.html b/docs/linghe/facade/silu.html new file mode 100644 index 0000000..0f9dcda --- /dev/null +++ b/docs/linghe/facade/silu.html @@ -0,0 +1,543 @@ + + + + + + + linghe.facade.silu API documentation + + + + + + + + + +
    +
    +

    +linghe.facade.silu

    + + + + + +
    +
    +
    + + class + BlockSiluFunction(torch.autograd.function.Function): + + +
    + + +

    Base class to create custom autograd.Function.

    + +

    To create a custom autograd.Function, subclass this class and implement +the forward() and backward() static methods. Then, to use your custom +op in the forward pass, call the class method apply. Do not call +forward() directly.

    + +

    To ensure correctness and best performance, make sure you are calling the +correct methods on ctx and validating your backward function using +torch.autograd.gradcheck().

    + +

    See :ref:extending-autograd for more details on how to use this class.

    + +

    Examples::

    + +
    >>> # xdoctest: +REQUIRES(env:TORCH_DOCTEST_AUTOGRAD)
    +>>> class Exp(Function):
    +>>>     @staticmethod
    +>>>     def forward(ctx, i):
    +>>>         result = i.exp()
    +>>>         ctx.save_for_backward(result)
    +>>>         return result
    +>>>
    +>>>     @staticmethod
    +>>>     def backward(ctx, grad_output):
    +>>>         result, = ctx.saved_tensors
    +>>>         return grad_output * result
    +>>>
    +>>> # Use it by calling the apply method:
    +>>> # xdoctest: +SKIP
    +>>> output = Exp.apply(input)
    +
    +
    + + +
    +
    +
    @staticmethod
    + + def + forward(ctx, input, quantizer, grad_quantizer, cls): + + +
    + + +

    Define the forward of the custom autograd Function.

    + +

    This function is to be overridden by all subclasses. +There are two ways to define forward:

    + +

    Usage 1 (Combined forward and ctx)::

    + +
    @staticmethod
    +def forward(ctx: Any, *args: Any, **kwargs: Any) -> Any:
    +    pass
    +
    + +
      +
    • It must accept a context ctx as the first argument, followed by any +number of arguments (tensors or other types).
    • +
    • See :ref:combining-forward-context for more details
    • +
    + +

    Usage 2 (Separate forward and ctx)::

    + +
    @staticmethod
    +def forward(*args: Any, **kwargs: Any) -> Any:
    +    pass
    +
    +@staticmethod
    +def setup_context(ctx: Any, inputs: Tuple[Any, ...], output: Any) -> None:
    +    pass
    +
    + +
      +
    • The forward no longer accepts a ctx argument.
    • +
    • Instead, you must also override the torch.autograd.Function.setup_context() +staticmethod to handle setting up the ctx object. +output is the output of the forward, inputs are a Tuple of inputs +to the forward.
    • +
    • See :ref:extending-autograd for more details
    • +
    + +

    The context can be used to store arbitrary data that can be then +retrieved during the backward pass. Tensors should not be stored +directly on ctx (though this is not currently enforced for +backward compatibility). Instead, tensors should be saved either with +ctx.save_for_backward() if they are intended to be used in +backward (equivalently, vjp) or ctx.save_for_forward() +if they are intended to be used for in jvp.

    +
    + + +
    +
    +
    +
    @staticmethod
    + + def + backward(ctx, grad_output): + + +
    + + +

    Define a formula for differentiating the operation with backward mode automatic differentiation.

    + +

    This function is to be overridden by all subclasses. +(Defining this function is equivalent to defining the vjp function.)

    + +

    It must accept a context ctx as the first argument, followed by +as many outputs as the forward() returned (None will be passed in +for non tensor outputs of the forward function), +and it should return as many tensors, as there were inputs to +forward(). Each argument is the gradient w.r.t the given output, +and each returned value should be the gradient w.r.t. the +corresponding input. If an input is not a Tensor or is a Tensor not +requiring grads, you can just pass None as a gradient for that input.

    + +

    The context can be used to retrieve tensors saved during the forward +pass. It also has an attribute ctx.needs_input_grad as a tuple +of booleans representing whether each input needs gradient. E.g., +backward() will have ctx.needs_input_grad[0] = True if the +first input to forward() needs gradient computed w.r.t. the +output.

    +
    + + +
    +
    +
    +
    + + class + BlockBatchWeightedSiluFunction(torch.autograd.function.Function): + + +
    + + +

    Base class to create custom autograd.Function.

    + +

    To create a custom autograd.Function, subclass this class and implement +the forward() and backward() static methods. Then, to use your custom +op in the forward pass, call the class method apply. Do not call +forward() directly.

    + +

    To ensure correctness and best performance, make sure you are calling the +correct methods on ctx and validating your backward function using +torch.autograd.gradcheck().

    + +

    See :ref:extending-autograd for more details on how to use this class.

    + +

    Examples::

    + +
    >>> # xdoctest: +REQUIRES(env:TORCH_DOCTEST_AUTOGRAD)
    +>>> class Exp(Function):
    +>>>     @staticmethod
    +>>>     def forward(ctx, i):
    +>>>         result = i.exp()
    +>>>         ctx.save_for_backward(result)
    +>>>         return result
    +>>>
    +>>>     @staticmethod
    +>>>     def backward(ctx, grad_output):
    +>>>         result, = ctx.saved_tensors
    +>>>         return grad_output * result
    +>>>
    +>>> # Use it by calling the apply method:
    +>>> # xdoctest: +SKIP
    +>>> output = Exp.apply(input)
    +
    +
    + + +
    +
    +
    @staticmethod
    + + def + forward( ctx, input, weights, counts, splits, quantizers, grad_quantizers, cls, is_recomputing): + + +
    + + +

    Define the forward of the custom autograd Function.

    + +

    This function is to be overridden by all subclasses. +There are two ways to define forward:

    + +

    Usage 1 (Combined forward and ctx)::

    + +
    @staticmethod
    +def forward(ctx: Any, *args: Any, **kwargs: Any) -> Any:
    +    pass
    +
    + +
      +
    • It must accept a context ctx as the first argument, followed by any +number of arguments (tensors or other types).
    • +
    • See :ref:combining-forward-context for more details
    • +
    + +

    Usage 2 (Separate forward and ctx)::

    + +
    @staticmethod
    +def forward(*args: Any, **kwargs: Any) -> Any:
    +    pass
    +
    +@staticmethod
    +def setup_context(ctx: Any, inputs: Tuple[Any, ...], output: Any) -> None:
    +    pass
    +
    + +
      +
    • The forward no longer accepts a ctx argument.
    • +
    • Instead, you must also override the torch.autograd.Function.setup_context() +staticmethod to handle setting up the ctx object. +output is the output of the forward, inputs are a Tuple of inputs +to the forward.
    • +
    • See :ref:extending-autograd for more details
    • +
    + +

    The context can be used to store arbitrary data that can be then +retrieved during the backward pass. Tensors should not be stored +directly on ctx (though this is not currently enforced for +backward compatibility). Instead, tensors should be saved either with +ctx.save_for_backward() if they are intended to be used in +backward (equivalently, vjp) or ctx.save_for_forward() +if they are intended to be used for in jvp.

    +
    + + +
    +
    +
    +
    @staticmethod
    + + def + backward(ctx, grad_output): + + +
    + + +

    Define a formula for differentiating the operation with backward mode automatic differentiation.

    + +

    This function is to be overridden by all subclasses. +(Defining this function is equivalent to defining the vjp function.)

    + +

    It must accept a context ctx as the first argument, followed by +as many outputs as the forward() returned (None will be passed in +for non tensor outputs of the forward function), +and it should return as many tensors, as there were inputs to +forward(). Each argument is the gradient w.r.t the given output, +and each returned value should be the gradient w.r.t. the +corresponding input. If an input is not a Tensor or is a Tensor not +requiring grads, you can just pass None as a gradient for that input.

    + +

    The context can be used to retrieve tensors saved during the forward +pass. It also has an attribute ctx.needs_input_grad as a tuple +of booleans representing whether each input needs gradient. E.g., +backward() will have ctx.needs_input_grad[0] = True if the +first input to forward() needs gradient computed w.r.t. the +output.

    +
    + + +
    +
    +
    + + \ No newline at end of file diff --git a/docs/linghe/facade/topk.html b/docs/linghe/facade/topk.html new file mode 100644 index 0000000..0ba01fd --- /dev/null +++ b/docs/linghe/facade/topk.html @@ -0,0 +1,310 @@ + + + + + + + linghe.facade.topk API documentation + + + + + + + + + +
    +
    +

    +linghe.facade.topk

    + +

    Copyright (c) Ant Financial Service Group and its affiliates.

    +
    + + + + +
    +
    +
    + + def + fused_topk(x, k, dim=-1): + + +
    + + +

    topk

    + +
    Arguments:
    + +
      +
    • x: input tensor
    • +
    • k: topk
    • +
    • dim: dimension to apply topk, only support -1 currently
    • +
    + +
    Returns:
    + +
    +

    values: topk values + indices: topk indices

    +
    +
    + + +
    +
    +
    + + def + group_topk_score( x, topk, expert_bias=None, num_groups=32, group_topk=4, scaling_factor=1.0, score_function='sigmoid'): + + +
    + + +

    group topk with softmax/sigmoid function

    + +
    Arguments:
    + +
      +
    • x: input logit tensor
    • +
    • topk: topk
    • +
    • expert_bias: expert bias
    • +
    • num_groups: number of groups
    • +
    • group_topk: group to apply topk
    • +
    • scaling_factor: scaling factor
    • +
    • score_function: scaling function
    • +
    + +
    Returns:
    + +
    +

    probs: topk probs + routing_map: topk binary map + counts: token count per expert

    +
    +
    + + +
    +
    + + \ No newline at end of file diff --git a/docs/linghe/facade/transpose.html b/docs/linghe/facade/transpose.html index 6198ec3..2c986e6 100644 --- a/docs/linghe/facade/transpose.html +++ b/docs/linghe/facade/transpose.html @@ -31,7 +31,7 @@

    API Documentation

    @@ -56,22 +56,24 @@

    -
    +
    def - transpose_dim01(x): + transpose(x, inner=True):
    - + -

    transpose a tensor with the first two dims, x.ndims should not greater than 4

    +

    transpose a tensor, x.ndims should not greater than 4

    Arguments:
    • x: input tensor
    • +
    • inner: if True, transpose the first two dimensions +if False, transpose the last two dimensions
    Returns:
    diff --git a/docs/linghe/gemm/fp32_gemm.html b/docs/linghe/gemm/fp32_gemm.html index 358a970..02341ed 100644 --- a/docs/linghe/gemm/fp32_gemm.html +++ b/docs/linghe/gemm/fp32_gemm.html @@ -40,10 +40,13 @@

    API Documentation

    triton_fp32_gemm_for_update
  • - triton_scaled_fp32_gemm + triton_split_fp32_gemm
  • - triton_scaled_fp32_gemm_for_update + triton_split_fp32_gemm_for_backward +
  • +
  • + triton_split_fp32_gemm_for_update
  • @@ -72,7 +75,7 @@

    def - triton_fp32_gemm(a: torch.Tensor, b: torch.Tensor): + triton_fp32_gemm(x: torch.Tensor, w: torch.Tensor):
    @@ -102,7 +105,7 @@

    Returns:
    def - triton_fp32_gemm_for_backward(a: torch.Tensor, b: torch.Tensor): + triton_fp32_gemm_for_backward(y: torch.Tensor, w: torch.Tensor):
    @@ -130,7 +133,7 @@
    Returns:
    def - triton_fp32_gemm_for_update(a: torch.Tensor, b: torch.Tensor): + triton_fp32_gemm_for_update(y: torch.Tensor, x: torch.Tensor):
    @@ -141,8 +144,8 @@
    Returns:
    Arguments:
      -
    • a: gradient of output, fp32
    • -
    • b: input activation, bf16/fp16
    • +
    • y: gradient of output, fp32
    • +
    • x: input activation, bf16/fp16
    Returns:
    @@ -154,70 +157,87 @@
    Returns:
    -
    +
    def - triton_scaled_fp32_gemm(a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor): + triton_split_fp32_gemm(x: torch.Tensor, w: torch.Tensor):
    - + -

    c = (ascale[:,None])b -this kernel is used to fuse RMSNorm and quantization in MoE layer -native implementation: - y = rms_norm(x), - y_q = quantization(y), - router_logits = y@w -we can not fuse rms_norm and quantization -as we still need bf16 y for moe router gemm -fused implementation: - y_q, rms = quantization(rms_norm(x)) - router_logits = (x/rms)@y -so we need a scaled fp32 gemm kernel

    +

    return fp32 gemm result with fp16/bf16 inputs, + it's mainly used for MoE router GEMM + and DO NOT suitable for large size GEMM

    Arguments:
      -
    • a: activation tensor
    • -
    • b: weight tensor
    • -
    • scale: scale for activation tensor, 1/rms
    • +
    • a: left matrix with fp16/bf16 precision
    • +
    • b: right matrix with fp16/bf16 precision
    Returns:
    -

    output tensor

    +

    c: output with fp32 precision

    -
    +
    def - triton_scaled_fp32_gemm_for_update(a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor): + triton_split_fp32_gemm_for_backward(y: torch.Tensor, w: torch.Tensor):
    - + -

    see triton_scaled_fp32_gemm

    +

    mix precision gemm for backward, a@b.float()

    Arguments:
      -
    • a: y
    • -
    • b: activation before RMS norm
    • -
    • scale: 1/rms
    • +
    • a: input gradient, fp32
    • +
    • b: gemm weight, bf16/fp16
    Returns:
    -

    dw

    +

    c: gradient of activation

    +
    +
    + + +
    +
    +
    + + def + triton_split_fp32_gemm_for_update(y: torch.Tensor, x: torch.Tensor): + + +
    + + +

    mix precision gemm for updaing weight

    + +
    Arguments:
    + +
      +
    • y: gradient of output, fp32
    • +
    • x: input activation, bf16/fp16
    • +
    + +
    Returns:
    + +
    +

    c: gradient of weight

    diff --git a/docs/linghe/quant/block.html b/docs/linghe/quant/block.html index 957cbb2..016bd56 100644 --- a/docs/linghe/quant/block.html +++ b/docs/linghe/quant/block.html @@ -33,6 +33,12 @@

    API Documentation

  • triton_block_quant
  • +
  • + triton_blockwise_quant +
  • +
  • + triton_batch_blockwise_quant +
  • @@ -66,7 +72,7 @@

    -

    blockwise quantize x

    +

    blockwise quantize x, used for blockwise recipe for weight in megatron

    Arguments:
    @@ -87,6 +93,76 @@
    Returns:
    +

    +
    +
    + + def + triton_blockwise_quant(x, round_scale=False, output_mode=2): + + +
    + + +

    blockwise quantization, used in blockwise recipt in megatron

    + +
    Arguments:
    + +
      +
    • x: input tensor
    • +
    • round_scale: whether round scale to power of 2
    • +
    • output_mode: one of {0, 1, 2} +0: only output non-transposed quantized tensor +1: only output transposed quantized tensor +2: output both
    • +
    + +
    Returns:
    + +
    +

    x_q: + x_scale: + xt_q: + xt_scale:

    +
    +
    + + +
    +
    +
    + + def + triton_batch_blockwise_quant(xs, token_count_per_expert, splits, round_scale=False): + + +
    + + +

    select and quant, used in megatron 0.12 flex moe

    + +
    Arguments:
    + +
      +
    • xs: [bs, dim]
    • +
    • token_count_per_expert: [n_experts]
    • +
    • splits: python int list of token_count_per_expert
    • +
    • round_scale: whether round scale to power of 2
    • +
    + +
    Returns:
    + +
    +
      +
    • x_q:
    • +
    • x_scale:
    • +
    • xt_q:
    • +
    • xt_scale:
    • +
    +
    +
    + +
    + \ No newline at end of file diff --git a/docs/linghe/tools/benchmark.html b/docs/linghe/tools/benchmark.html new file mode 100644 index 0000000..ec7edd4 --- /dev/null +++ b/docs/linghe/tools/benchmark.html @@ -0,0 +1,239 @@ + + + + + + + linghe.tools.benchmark API documentation + + + + + + + + + +
    +
    +

    +linghe.tools.benchmark

    + +

    Copyright (c) Ant Financial Service Group and its affiliates.

    +
    + + + + +
    +
    + + \ No newline at end of file diff --git a/docs/linghe/tools/check.html b/docs/linghe/tools/check.html new file mode 100644 index 0000000..5df8cd4 --- /dev/null +++ b/docs/linghe/tools/check.html @@ -0,0 +1,239 @@ + + + + + + + linghe.tools.check API documentation + + + + + + + + + +
    +
    +

    +linghe.tools.check

    + +

    Copyright (c) Ant Financial Service Group and its affiliates.

    +
    + + + + +
    +
    + + \ No newline at end of file diff --git a/docs/linghe/tools/util.html b/docs/linghe/tools/util.html new file mode 100644 index 0000000..f4cdeb4 --- /dev/null +++ b/docs/linghe/tools/util.html @@ -0,0 +1,239 @@ + + + + + + + linghe.tools.util API documentation + + + + + + + + + +
    +
    +

    +linghe.tools.util

    + +

    Copyright (c) Ant Financial Service Group and its affiliates.

    +
    + + + + +
    +
    + + \ No newline at end of file diff --git a/docs/linghe/utils.html b/docs/linghe/utils.html index afd4a60..0219c76 100644 --- a/docs/linghe/utils.html +++ b/docs/linghe/utils.html @@ -30,16 +30,20 @@

    Submodules

    diff --git a/docs/linghe/utils/emb.html b/docs/linghe/utils/emb.html new file mode 100644 index 0000000..ccd5cc6 --- /dev/null +++ b/docs/linghe/utils/emb.html @@ -0,0 +1,366 @@ + + + + + + + linghe.utils.emb API documentation + + + + + + + + + +
    +
    +

    +linghe.utils.emb

    + +

    Copyright (c) Ant Financial Service Group and its affiliates.

    +
    + + + + +
    +
    +
    + + def + triton_embedding_forward(x, w_ptr, dim=4096, dtype=torch.bfloat16): + + +
    + + +

    inplace add y to x

    + +
    Arguments:
    + +
      +
    • x: input ids Tensor
    • +
    • w_ptr: data_ptr of embedding weight
    • +
    + +
    Returns:
    + +
    +

    embedding output

    +
    +
    + + +
    +
    +
    + + def + triton_atomic_embedding_backward(y, x, g_ptr, dtype=torch.bfloat16): + + +
    + + +

    inplace update embedding weight gradient

    + +
    Arguments:
    + +
      +
    • y: gradient of output
    • +
    • x: input ids Tensor
    • +
    • g_ptr: data_ptr of embedding weight gradient
    • +
    + +
    Returns:
    + +
    +

    None

    +
    +
    + + +
    +
    +
    + + def + triton_sync_embedding_backward(grad_output, x, g_ptr, dtype=torch.bfloat16): + + +
    + + +

    inplace update embedding weight gradient

    + +
    Arguments:
    + +
      +
    • y: gradient of output
    • +
    • x: input ids Tensor
    • +
    • g_ptr: data_ptr of embedding weight gradient
    • +
    + +
    Returns:
    + +
    +

    None

    +
    +
    + + +
    +
    +
    + + def + triton_embedding_backward(grad_output, x, g_ptr, dtype=torch.bfloat16): + + +
    + + +

    inplace update embedding weight gradient

    + +
    Arguments:
    + +
      +
    • y: gradient of output
    • +
    • x: input ids Tensor
    • +
    • g_ptr: data_ptr of embedding weight gradient
    • +
    + +
    Returns:
    + +
    +

    None

    +
    +
    + + +
    +
    + + \ No newline at end of file diff --git a/docs/linghe/utils/gate.html b/docs/linghe/utils/gate.html new file mode 100644 index 0000000..cd3c1b4 --- /dev/null +++ b/docs/linghe/utils/gate.html @@ -0,0 +1,272 @@ + + + + + + + linghe.utils.gate API documentation + + + + + + + + + +
    +
    +

    +linghe.utils.gate

    + + + + + +
    +
    +
    + + def + triton_group_rms_norm_gate_forward( x: torch.Tensor, gate: torch.Tensor, weight: torch.Tensor, eps=1e-06, group_size=4, transpose=True): + + +
    + + +

    norm and gate in linear attention

    + +
    Arguments:
    + +
      +
    • x: output of attn, [bs, length, n_heads, head_dim]
    • +
    • gate: gate tensor, [length, bs, dim] if transpose=True else [bs, length, dim]
    • +
    • weight: rms norm weight, [dim]
    • +
    • eps: epsilon of rms norm
    • +
    • group_size: group size of group rms norm
    • +
    • transpose: whether gate tensor has been transposed and output will be transposed
    • +
    + +
    Returns:
    + +
    +

    output tensor, [length, bs, dim] if transpose=True else [bs, length, dim]

    +
    +
    + + +
    +
    + + \ No newline at end of file diff --git a/docs/linghe/utils/gather.html b/docs/linghe/utils/gather.html index a7444c8..41a2779 100644 --- a/docs/linghe/utils/gather.html +++ b/docs/linghe/utils/gather.html @@ -34,7 +34,7 @@

    API Documentation

    triton_make_row_id_map
  • - triton_make_row_id_map_and_indices + triton_make_row_id_map_and_index
  • triton_index_select @@ -54,6 +54,9 @@

    API Documentation

  • triton_smooth_permute_with_mask_map
  • +
  • + triton_batch_block_pad_permute_with_indices +
  • @@ -105,15 +108,15 @@
    Returns:
    -
    +
    def - triton_make_row_id_map_and_indices(routing_map: torch.Tensor, num_out_tokens: int, multiple_of: int = 1): + triton_make_row_id_map_and_index(routing_map: torch.Tensor, num_out_tokens: int, multiple_of: int = 1):
    - +

    similar with triton_make_row_id_map, but output an indices tensor as well

    @@ -181,8 +184,8 @@
    Arguments:
    • inp: [num_tokens, hidden_size], rowwise quantized tensor
    • -
    • scale: [num_tokens], quantization scale
    • -
    • probs: router prob, used as weight
    • +
    • scale: optional, [num_tokens], quantization scale
    • +
    • probs: optional, router prob, used as weight
    • row_id_map: [n_experts, num_tokens] index >= 0: row index of output tensor index == -1: ignore @@ -345,6 +348,42 @@
      Returns:
    +
    +
    +
    + + def + triton_batch_block_pad_permute_with_indices( xs, token_count_per_expert, indices, splits, probs=None, round_scale=False): + + +
    + + +

    select and quant, used in megatron 0.12 flex moe

    + +
    Arguments:
    + +
      +
    • xs: [bs, dim]
    • +
    • token_count_per_expert: [n_experts]
    • +
    • indices: [n_experts*topk]
    • +
    • splits: python int list of token_count_per_expert
    • +
    • probs: route weights, [bs, n_experts]
    • +
    • round_scale: whether round scale to power of 2
    • +
    + +
    Returns:
    + +
    +

    x_q: + x_scale: + xt_q: + xt_scale: + prob_output:

    +
    +
    + +
    + \ No newline at end of file diff --git a/docs/linghe/utils/norm.html b/docs/linghe/utils/norm.html index a496817..d8a3b92 100644 --- a/docs/linghe/utils/norm.html +++ b/docs/linghe/utils/norm.html @@ -37,7 +37,7 @@

    API Documentation

    triton_rms_norm_and_block_quant_forward
  • - triton_group_rms_norm_gate_forward + triton_rms_norm_fp32_gemm_block_quant_forward
  • @@ -55,7 +55,9 @@

    API Documentation

    linghe.utils.norm

    - +

    Copyright (c) Ant Financial Service Group and its affiliates.

    +
    + @@ -64,7 +66,7 @@

    def - triton_rms_norm_forward(x, weight, eps=1e-06, out=None): + triton_rms_norm_forward(x, weight, eps=1e-06, out=None, rms=None):
    @@ -78,12 +80,15 @@

    Arguments:
  • x: input tensor
  • weight: weight of rms norm
  • eps: epsilon of rms norm
  • +
  • rms: use x*rms to calculate output if rms is not None, +it will accelerate recompute of rms norm
  • Returns:
    -

    out: output tensor

    +

    out: output tensor + rms: 1/rms of input tensor

    @@ -133,33 +138,44 @@
    Returns:
    -
    +
    def - triton_group_rms_norm_gate_forward( x: torch.Tensor, gate: torch.Tensor, weight: torch.Tensor, eps=1e-06, group_size=4, transpose=True): + triton_rms_norm_fp32_gemm_block_quant_forward( x: torch.Tensor, norm_weight: torch.Tensor, route_weight: torch.Tensor, rms: Optional[torch.Tensor] = None, eps: float = 1e-06, output_mode: int = 0, round_scale=False):
    - + -

    norm and gate in linear attention

    +

    y = rms_norm(x) +logits = y@w_route +x_q, x_s, xt_q, xt_s = quantization(y)

    Arguments:
      -
    • x: output of attn, [bs, length, n_heads, head_dim]
    • -
    • gate: gate tensor, [length, bs, dim] if transpose=True else [bs, length, dim]
    • -
    • weight: rms norm weight, [dim]
    • +
    • x: input tensor
    • +
    • norm weight: weight tensor of rms norm
    • +
    • route_weight: moe router weight
    • eps: epsilon of rms norm
    • -
    • group_size: group size of group rms norm
    • -
    • transpose: whether gate is transposed and output will be transposed
    • +
    • output_mode: 0 or 1 +0: only output non-transpose quantizatino tensor +1: only output transposed quantizatino tensor
    Returns:
    -

    output tensor, [length, bs, dim] if transpose=True else [bs, length, dim]

    +
      +
    • y: rms normed tensor
    • +
    • rms: 1/rms
    • +
    • logits: router logit
    • +
    • x_q:
    • +
    • x_s:
    • +
    • xt_q:
    • +
    • xt_s:
    • +
    diff --git a/docs/linghe/utils/rearange.html b/docs/linghe/utils/rearange.html index ec9cc8e..7a5772c 100644 --- a/docs/linghe/utils/rearange.html +++ b/docs/linghe/utils/rearange.html @@ -31,7 +31,7 @@

    API Documentation

    @@ -56,18 +56,18 @@

    -
    +
    def - triton_split_and_cat(x, counts, indices, scales=None): + triton_sort_chunks_by_index(x, counts, indices, scales=None):
    - +

    split x to multiple tensors and cat with indices, -it is used for permutation in moe

    +it is used for permutation in moe with all2all communication

    Arguments:
    diff --git a/docs/linghe/utils/reduce.html b/docs/linghe/utils/reduce.html index 0aeb0ae..9fedb01 100644 --- a/docs/linghe/utils/reduce.html +++ b/docs/linghe/utils/reduce.html @@ -37,7 +37,10 @@

    API Documentation

    triton_batch_count_zero
  • - triton_batch_sum_with_ord + triton_norm +
  • +
  • + triton_batch_norm
  • @@ -120,29 +123,63 @@
    Returns:
    -
    +
    +
    + + def + triton_norm(x, ord=2, norm=True, scalar=True): + + +
    + + +

    calculate norm.

    + +
    Arguments:
    + +
      +
    • x: input tensor.
    • +
    • ord: the order of tensor. -1 means 'inf' ord.
    • +
    • norm: only used with ord in (1, 2) +True: (sum(sum(abs(x)ord) x for x in xs))(1/ord) +False: sum(sum(abs(x)**ord) x for x in xs))
    • +
    + +
    Returns:
    + +
    +

    a scalar if scalar=True else a single-value fp32 tensor

    +
    +
    + + +
    +
    def - triton_batch_sum_with_ord(xs, ord=2): + triton_batch_norm(xs, ord=2, norm=True, scalar=True, high_precision=True):
    - + -

    return sum(abs(x)**ord).

    +

    treat multiple tensors as a single tensor and calculate norm.

    Arguments:
    • xs: Tensor lists.
    • -
    • ord: the order of tensor.
    • +
    • ord: the order of tensor. -1 means 'inf' ord.
    • +
    • norm: only used with ord in (1, 2) +True: (sum(sum(abs(x)ord) x for x in xs))(1/ord) +False: sum(sum(abs(x)**ord) x for x in xs))
    Returns:
    -

    a single-value fp32 tensor

    +

    a scalar if scalar=True else a single-value fp32 tensor

    diff --git a/docs/linghe/utils/rope.html b/docs/linghe/utils/rope.html index c531737..7054a59 100644 --- a/docs/linghe/utils/rope.html +++ b/docs/linghe/utils/rope.html @@ -39,6 +39,15 @@

    API Documentation

  • triton_qk_norm_and_half_rope_backward
  • +
  • + triton_varlen_qk_norm_and_half_rope_forward +
  • +
  • + triton_varlen_qk_norm_and_half_rope_backward +
  • +
  • + triton_mla_rope_forward +
  • @@ -72,7 +81,7 @@

    -

    apply norm to qk, then apply half rope to qk

    +

    apply half rope to qk

    Arguments:
    @@ -98,7 +107,7 @@
    Returns:
    def - triton_qk_norm_and_half_rope_forward( qkv, q_norm_weight, k_norm_weight, freqs, H=32, h=4, eps=1e-06, interleaved=True, transposed=False): + triton_qk_norm_and_half_rope_forward( qkv, q_norm_weight, k_norm_weight, freqs, H=32, h=4, eps=1e-06, interleaved=True, transposed=True, silu=False):
    @@ -122,8 +131,8 @@
    Arguments:
    non-interleaved: [q...qk...kv...v]
  • transposed: whether qkv is tranposed transposed: [S, B, dim] -non-transposed: [B, S, dim] -only support transpose format currently
  • +non-transposed: [B, S, dim] +
  • silu: apply silu on qkv before qk norm and rope
  • Returns:
    @@ -143,7 +152,7 @@
    Returns:
    def - triton_qk_norm_and_half_rope_backward( gq, gk, gv, qkv, q_norm_weight, k_norm_weight, freqs, eps=1e-06, interleaved=True, transposed=False): + triton_qk_norm_and_half_rope_backward( gq, gk, gv, qkv, q_norm_weight, k_norm_weight, freqs, eps=1e-06, interleaved=True, transposed=True, silu=False):
    @@ -158,12 +167,101 @@
    Arguments:
  • gk: gradient of ko, [len, bs, q_head, head_dim]
  • gv: gradient of vo, [len, bs, q_head, head_dim]
  • qkv: input qkv
  • -
  • q_norm_weight:
  • -
  • k_norm_weight:
  • -
  • freqs:
  • -
  • eps:
  • -
  • interleaved:
  • -
  • transposed:
  • +
  • q_norm_weight: rms norm weight for query
  • +
  • k_norm_weight: rms norm weight for key
  • +
  • freqs: Freqs tensor based on half dim.
  • +
  • eps: epsilon value for L2 normalization.
  • +
  • interleaved: whether head of qkv is interleaved, +interleaved: [q...qkvq...qkv] +non-interleaved: [q...qk...kv...v]
  • +
  • transposed: whether qkv is tranposed +transposed: [S, B, dim] +non-transposed: [B, S, dim]
  • +
  • silu: whether silu is applied to qkv
  • + + +
    Returns:
    + +
    +
      +
    • dqkv: gradient of qkv
    • +
    • dqw: gradient of q_norm_weight
    • +
    • dkw: gradient of k_norm_weight
    • +
    +
    +
    + + +
    +
    +
    + + def + triton_varlen_qk_norm_and_half_rope_forward( qkv, q_norm_weight, k_norm_weight, freqs, cu_seqlens_q, cu_seqlens_kv, H=32, h=4, eps=1e-06, interleaved=True, silu=False, cp_rank=0, cp_size=1, mscale=1.0, reuse=False): + + +
    + + +

    split qkv to q/k/v, apply qk norm and half rope to q/k, + transpose q/k/v to flash-attention layout

    + +
    Arguments:
    + +
      +
    • qkv: QKV tensor with size of [S, B, dim], heads are interleaved
    • +
    • q_norm_weight: rms norm weight for query
    • +
    • k_norm_weight: rms norm weight for key
    • +
    • freqs: Freqs tensor based on half dim.
    • +
    • H: Number of attention heads.
    • +
    • h: Number of key/value heads.
    • +
    • eps: epsilon value for L2 normalization.
    • +
    • interleaved: whether head of qkv is interleaved, +interleaved: [q...qkvq...qkv] +non-interleaved: [q...qk...kv...v]
    • +
    • silu: apply silu on qkv before qk norm and rope
    • +
    + +
    Returns:
    + +
    +
      +
    • qo: shape [B, S, H, head_dim]
    • +
    • ko: shape [B, S, h, head_dim]
    • +
    • vo: shape [B, S, h, head_dim]
    • +
    +
    +
    + + +
    +
    +
    + + def + triton_varlen_qk_norm_and_half_rope_backward( gq, gk, gv, qkv, q_norm_weight, k_norm_weight, freqs, cu_seqlens_q, cu_seqlens_kv, eps=1e-06, interleaved=True, silu=False, cp_rank=0, cp_size=1, mscale=1.0, reuse=False): + + +
    + + +

    backward kernel of triton_qk_norm_and_half_rope_forward

    + +
    Arguments:
    + +
      +
    • gq: gradient of qo, [len, bs, q_head, head_dim]
    • +
    • gk: gradient of ko, [len, bs, q_head, head_dim]
    • +
    • gv: gradient of vo, [len, bs, q_head, head_dim]
    • +
    • qkv: input qkv
    • +
    • q_norm_weight: rms norm weight for query
    • +
    • k_norm_weight: rms norm weight for key
    • +
    • freqs: Freqs tensor based on half dim.
    • +
    • eps: epsilon value for L2 normalization.
    • +
    • interleaved: whether head of qkv is interleaved, +interleaved: [q...qkvq...qkv] +non-interleaved: [q...qk...kv...v]
    • +
    • silu: whether silu is applied to qkv
    Returns:
    @@ -178,6 +276,49 @@
    Returns:
    +
    +
    +
    + + def + triton_mla_rope_forward( q, kv, k_pos_emb, freqs, mscale=1.0, transpose=False, cu_seqlens_q=None, cu_seqlens_kv=None, cp_rank=0, cp_size=1, reuse=False): + + +
    + + +

    apply MLA-type rope to qkv

    + +
    Arguments:
    + +
      +
    • q: query tensor, [len, bs, n_heads, 192]
    • +
    • kv: key-value tensor, [len, bs, n_heads, 256]
    • +
    • k_pos_emb: k pos emb, [len, bs, 1, 64]
    • +
    • freqs: rope freqs, [len, 64]
    • +
    • mscale: mscale for rope
    • +
    • transpose: whether transpose the output to [bs, len, n_heads, dim] layout
    • +
    • cu_seqlens_q: accummulated query length
    • +
    • cu_seqlens_kv: accummulated kv length
    • +
    • cp_rank: rank of context parallel
    • +
    • cp_size: size of context parallel
    • +
    + +
    Returns:
    + +
    +
      +
    • qo: inplace updated query, [len, bs, n_heads, 192] if not transpose + else [bs, len, n_heads, 192]
    • +
    • ko: key output, [len, bs, n_heads, 192] if not transpose + else [bs, len, n_heads, 192]
    • +
    • vo: value output, [len, bs, n_heads, 128] if not transpose + else [bs, len, n_heads, 128]
    • +
    +
    +
    + +
    + \ No newline at end of file diff --git a/docs/linghe/utils/transpose.html b/docs/linghe/utils/transpose.html index df81d84..50fc051 100644 --- a/docs/linghe/utils/transpose.html +++ b/docs/linghe/utils/transpose.html @@ -69,7 +69,7 @@

    def - triton_transpose( x: torch.Tensor, dim0: Optional[int] = None, dim1: Optional[int] = None): + triton_transpose(x: torch.Tensor, inner=True):
    @@ -81,8 +81,7 @@

    Arguments:
    • x: input tensor
    • -
    • dim0: dim 0
    • -
    • dim1: dim 1
    • +
    • inner: inner dim if True, outer dim if False
    Returns:
    diff --git a/docs/linghe/utils/unary.html b/docs/linghe/utils/unary.html new file mode 100644 index 0000000..fda4c62 --- /dev/null +++ b/docs/linghe/utils/unary.html @@ -0,0 +1,271 @@ + + + + + + + linghe.utils.unary API documentation + + + + + + + + + +
    +
    +

    +linghe.utils.unary

    + +

    Copyright (c) Ant Financial Service Group and its affiliates.

    +
    + + + + +
    +
    +
    + + def + triton_batch_clip(xs, clip_value=100.0): + + +
    + + +

    return [clip(x, -clip_value, clip_value) for x in xs], +used to clip gradient.

    + +
    Arguments:
    + +
      +
    • xs: Tensor lists.
    • +
    • clip_value: a python float scale
    • +
    + +
    Returns:
    + +
    +

    updated xs

    +
    +
    + + +
    +
    + + \ No newline at end of file diff --git a/docs/search.js b/docs/search.js index f4e8fab..43db75a 100644 --- a/docs/search.js +++ b/docs/search.js @@ -1,6 +1,6 @@ window.pdocSearch = (function(){ /** elasticlunr - http://weixsong.github.io * Copyright (C) 2017 Oliver Nightingale * Copyright (C) 2017 Wei Song * MIT Licensed */!function(){function e(e){if(null===e||"object"!=typeof e)return e;var t=e.constructor();for(var n in e)e.hasOwnProperty(n)&&(t[n]=e[n]);return t}var t=function(e){var n=new t.Index;return n.pipeline.add(t.trimmer,t.stopWordFilter,t.stemmer),e&&e.call(n,n),n};t.version="0.9.5",lunr=t,t.utils={},t.utils.warn=function(e){return function(t){e.console&&console.warn&&console.warn(t)}}(this),t.utils.toString=function(e){return void 0===e||null===e?"":e.toString()},t.EventEmitter=function(){this.events={}},t.EventEmitter.prototype.addListener=function(){var e=Array.prototype.slice.call(arguments),t=e.pop(),n=e;if("function"!=typeof t)throw new TypeError("last argument must be a function");n.forEach(function(e){this.hasHandler(e)||(this.events[e]=[]),this.events[e].push(t)},this)},t.EventEmitter.prototype.removeListener=function(e,t){if(this.hasHandler(e)){var n=this.events[e].indexOf(t);-1!==n&&(this.events[e].splice(n,1),0==this.events[e].length&&delete this.events[e])}},t.EventEmitter.prototype.emit=function(e){if(this.hasHandler(e)){var t=Array.prototype.slice.call(arguments,1);this.events[e].forEach(function(e){e.apply(void 0,t)},this)}},t.EventEmitter.prototype.hasHandler=function(e){return e in this.events},t.tokenizer=function(e){if(!arguments.length||null===e||void 0===e)return[];if(Array.isArray(e)){var n=e.filter(function(e){return null===e||void 0===e?!1:!0});n=n.map(function(e){return t.utils.toString(e).toLowerCase()});var i=[];return n.forEach(function(e){var n=e.split(t.tokenizer.seperator);i=i.concat(n)},this),i}return e.toString().trim().toLowerCase().split(t.tokenizer.seperator)},t.tokenizer.defaultSeperator=/[\s\-]+/,t.tokenizer.seperator=t.tokenizer.defaultSeperator,t.tokenizer.setSeperator=function(e){null!==e&&void 0!==e&&"object"==typeof e&&(t.tokenizer.seperator=e)},t.tokenizer.resetSeperator=function(){t.tokenizer.seperator=t.tokenizer.defaultSeperator},t.tokenizer.getSeperator=function(){return t.tokenizer.seperator},t.Pipeline=function(){this._queue=[]},t.Pipeline.registeredFunctions={},t.Pipeline.registerFunction=function(e,n){n in t.Pipeline.registeredFunctions&&t.utils.warn("Overwriting existing registered function: "+n),e.label=n,t.Pipeline.registeredFunctions[n]=e},t.Pipeline.getRegisteredFunction=function(e){return e in t.Pipeline.registeredFunctions!=!0?null:t.Pipeline.registeredFunctions[e]},t.Pipeline.warnIfFunctionNotRegistered=function(e){var n=e.label&&e.label in this.registeredFunctions;n||t.utils.warn("Function is not registered with pipeline. This may cause problems when serialising the index.\n",e)},t.Pipeline.load=function(e){var n=new t.Pipeline;return e.forEach(function(e){var i=t.Pipeline.getRegisteredFunction(e);if(!i)throw new Error("Cannot load un-registered function: "+e);n.add(i)}),n},t.Pipeline.prototype.add=function(){var e=Array.prototype.slice.call(arguments);e.forEach(function(e){t.Pipeline.warnIfFunctionNotRegistered(e),this._queue.push(e)},this)},t.Pipeline.prototype.after=function(e,n){t.Pipeline.warnIfFunctionNotRegistered(n);var i=this._queue.indexOf(e);if(-1===i)throw new Error("Cannot find existingFn");this._queue.splice(i+1,0,n)},t.Pipeline.prototype.before=function(e,n){t.Pipeline.warnIfFunctionNotRegistered(n);var i=this._queue.indexOf(e);if(-1===i)throw new Error("Cannot find existingFn");this._queue.splice(i,0,n)},t.Pipeline.prototype.remove=function(e){var t=this._queue.indexOf(e);-1!==t&&this._queue.splice(t,1)},t.Pipeline.prototype.run=function(e){for(var t=[],n=e.length,i=this._queue.length,o=0;n>o;o++){for(var r=e[o],s=0;i>s&&(r=this._queue[s](r,o,e),void 0!==r&&null!==r);s++);void 0!==r&&null!==r&&t.push(r)}return t},t.Pipeline.prototype.reset=function(){this._queue=[]},t.Pipeline.prototype.get=function(){return this._queue},t.Pipeline.prototype.toJSON=function(){return this._queue.map(function(e){return t.Pipeline.warnIfFunctionNotRegistered(e),e.label})},t.Index=function(){this._fields=[],this._ref="id",this.pipeline=new t.Pipeline,this.documentStore=new t.DocumentStore,this.index={},this.eventEmitter=new t.EventEmitter,this._idfCache={},this.on("add","remove","update",function(){this._idfCache={}}.bind(this))},t.Index.prototype.on=function(){var e=Array.prototype.slice.call(arguments);return this.eventEmitter.addListener.apply(this.eventEmitter,e)},t.Index.prototype.off=function(e,t){return this.eventEmitter.removeListener(e,t)},t.Index.load=function(e){e.version!==t.version&&t.utils.warn("version mismatch: current "+t.version+" importing "+e.version);var n=new this;n._fields=e.fields,n._ref=e.ref,n.documentStore=t.DocumentStore.load(e.documentStore),n.pipeline=t.Pipeline.load(e.pipeline),n.index={};for(var i in e.index)n.index[i]=t.InvertedIndex.load(e.index[i]);return n},t.Index.prototype.addField=function(e){return this._fields.push(e),this.index[e]=new t.InvertedIndex,this},t.Index.prototype.setRef=function(e){return this._ref=e,this},t.Index.prototype.saveDocument=function(e){return this.documentStore=new t.DocumentStore(e),this},t.Index.prototype.addDoc=function(e,n){if(e){var n=void 0===n?!0:n,i=e[this._ref];this.documentStore.addDoc(i,e),this._fields.forEach(function(n){var o=this.pipeline.run(t.tokenizer(e[n]));this.documentStore.addFieldLength(i,n,o.length);var r={};o.forEach(function(e){e in r?r[e]+=1:r[e]=1},this);for(var s in r){var u=r[s];u=Math.sqrt(u),this.index[n].addToken(s,{ref:i,tf:u})}},this),n&&this.eventEmitter.emit("add",e,this)}},t.Index.prototype.removeDocByRef=function(e){if(e&&this.documentStore.isDocStored()!==!1&&this.documentStore.hasDoc(e)){var t=this.documentStore.getDoc(e);this.removeDoc(t,!1)}},t.Index.prototype.removeDoc=function(e,n){if(e){var n=void 0===n?!0:n,i=e[this._ref];this.documentStore.hasDoc(i)&&(this.documentStore.removeDoc(i),this._fields.forEach(function(n){var o=this.pipeline.run(t.tokenizer(e[n]));o.forEach(function(e){this.index[n].removeToken(e,i)},this)},this),n&&this.eventEmitter.emit("remove",e,this))}},t.Index.prototype.updateDoc=function(e,t){var t=void 0===t?!0:t;this.removeDocByRef(e[this._ref],!1),this.addDoc(e,!1),t&&this.eventEmitter.emit("update",e,this)},t.Index.prototype.idf=function(e,t){var n="@"+t+"/"+e;if(Object.prototype.hasOwnProperty.call(this._idfCache,n))return this._idfCache[n];var i=this.index[t].getDocFreq(e),o=1+Math.log(this.documentStore.length/(i+1));return this._idfCache[n]=o,o},t.Index.prototype.getFields=function(){return this._fields.slice()},t.Index.prototype.search=function(e,n){if(!e)return[];e="string"==typeof e?{any:e}:JSON.parse(JSON.stringify(e));var i=null;null!=n&&(i=JSON.stringify(n));for(var o=new t.Configuration(i,this.getFields()).get(),r={},s=Object.keys(e),u=0;u0&&t.push(e);for(var i in n)"docs"!==i&&"df"!==i&&this.expandToken(e+i,t,n[i]);return t},t.InvertedIndex.prototype.toJSON=function(){return{root:this.root}},t.Configuration=function(e,n){var e=e||"";if(void 0==n||null==n)throw new Error("fields should not be null");this.config={};var i;try{i=JSON.parse(e),this.buildUserConfig(i,n)}catch(o){t.utils.warn("user configuration parse failed, will use default configuration"),this.buildDefaultConfig(n)}},t.Configuration.prototype.buildDefaultConfig=function(e){this.reset(),e.forEach(function(e){this.config[e]={boost:1,bool:"OR",expand:!1}},this)},t.Configuration.prototype.buildUserConfig=function(e,n){var i="OR",o=!1;if(this.reset(),"bool"in e&&(i=e.bool||i),"expand"in e&&(o=e.expand||o),"fields"in e)for(var r in e.fields)if(n.indexOf(r)>-1){var s=e.fields[r],u=o;void 0!=s.expand&&(u=s.expand),this.config[r]={boost:s.boost||0===s.boost?s.boost:1,bool:s.bool||i,expand:u}}else t.utils.warn("field name in user configuration not found in index instance fields");else this.addAllFields2UserConfig(i,o,n)},t.Configuration.prototype.addAllFields2UserConfig=function(e,t,n){n.forEach(function(n){this.config[n]={boost:1,bool:e,expand:t}},this)},t.Configuration.prototype.get=function(){return this.config},t.Configuration.prototype.reset=function(){this.config={}},lunr.SortedSet=function(){this.length=0,this.elements=[]},lunr.SortedSet.load=function(e){var t=new this;return t.elements=e,t.length=e.length,t},lunr.SortedSet.prototype.add=function(){var e,t;for(e=0;e1;){if(r===e)return o;e>r&&(t=o),r>e&&(n=o),i=n-t,o=t+Math.floor(i/2),r=this.elements[o]}return r===e?o:-1},lunr.SortedSet.prototype.locationFor=function(e){for(var t=0,n=this.elements.length,i=n-t,o=t+Math.floor(i/2),r=this.elements[o];i>1;)e>r&&(t=o),r>e&&(n=o),i=n-t,o=t+Math.floor(i/2),r=this.elements[o];return r>e?o:e>r?o+1:void 0},lunr.SortedSet.prototype.intersect=function(e){for(var t=new lunr.SortedSet,n=0,i=0,o=this.length,r=e.length,s=this.elements,u=e.elements;;){if(n>o-1||i>r-1)break;s[n]!==u[i]?s[n]u[i]&&i++:(t.add(s[n]),n++,i++)}return t},lunr.SortedSet.prototype.clone=function(){var e=new lunr.SortedSet;return e.elements=this.toArray(),e.length=e.elements.length,e},lunr.SortedSet.prototype.union=function(e){var t,n,i;this.length>=e.length?(t=this,n=e):(t=e,n=this),i=t.clone();for(var o=0,r=n.toArray();o

    \n"}, "linghe.facade": {"fullname": "linghe.facade", "modulename": "linghe.facade", "kind": "module", "doc": "

    \n"}, "linghe.facade.add": {"fullname": "linghe.facade.add", "modulename": "linghe.facade.add", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, "linghe.facade.add.inplace_add": {"fullname": "linghe.facade.add.inplace_add", "modulename": "linghe.facade.add", "qualname": "inplace_add", "kind": "function", "doc": "

    inplace add y to x with mix precise

    \n\n
    Arguments:
    \n\n
      \n
    • x: to be updated
    • \n
    • y: add to x
    • \n
    \n\n
    Returns:
    \n\n
    \n

    updated x tensor

    \n
    \n", "signature": "(x: torch.Tensor, y: torch.Tensor):", "funcdef": "def"}, "linghe.facade.fp32_gemm": {"fullname": "linghe.facade.fp32_gemm", "modulename": "linghe.facade.fp32_gemm", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, "linghe.facade.fp32_gemm.fp32_gemm": {"fullname": "linghe.facade.fp32_gemm.fp32_gemm", "modulename": "linghe.facade.fp32_gemm", "qualname": "fp32_gemm", "kind": "function", "doc": "

    gemm with bf16/fp16 inputs and float32 output,\ncurrently used in MoE router gemm.

    \n\n
    Arguments:
    \n\n
      \n
    • input: bf16/fp16 activation tensor
    • \n
    • weight: bf16/fp16 weight tensor
    • \n
    \n\n
    Returns:
    \n\n
    \n

    output of gemm

    \n
    \n", "signature": "(input: torch.Tensor, weight: torch.Tensor):", "funcdef": "def"}, "linghe.facade.hadamard_quant_linear": {"fullname": "linghe.facade.hadamard_quant_linear", "modulename": "linghe.facade.hadamard_quant_linear", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, "linghe.facade.hadamard_quant_linear.HadamardQuantLinear": {"fullname": "linghe.facade.hadamard_quant_linear.HadamardQuantLinear", "modulename": "linghe.facade.hadamard_quant_linear", "qualname": "HadamardQuantLinear", "kind": "class", "doc": "

    a naive implementation of hadamard transformation and quantization

    \n", "bases": "torch.nn.modules.module.Module"}, "linghe.facade.hadamard_quant_linear.HadamardQuantLinear.__init__": {"fullname": "linghe.facade.hadamard_quant_linear.HadamardQuantLinear.__init__", "modulename": "linghe.facade.hadamard_quant_linear", "qualname": "HadamardQuantLinear.__init__", "kind": "function", "doc": "
    Arguments:
    \n\n
      \n
    • in_features: in feature number
    • \n
    • out_features: out feature number
    • \n
    • bias: whether use bias
    • \n
    • device: weight device
    • \n
    • dtype: weight dtype
    • \n
    \n", "signature": "(\tin_features: int,\tout_features: int,\tbias: bool = True,\tdevice=None,\tdtype=None)"}, "linghe.facade.loss": {"fullname": "linghe.facade.loss", "modulename": "linghe.facade.loss", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, "linghe.facade.loss.softmax_cross_entropy": {"fullname": "linghe.facade.loss.softmax_cross_entropy", "modulename": "linghe.facade.loss", "qualname": "softmax_cross_entropy", "kind": "function", "doc": "

    softmax cross entropy

    \n\n
    Arguments:
    \n\n
      \n
    • logits: logits tensor, shape [...,dim]
    • \n
    • labels: labels tensor, shape [...]
    • \n
    • inplace: update gradient in the logits tensor if True
    • \n
    \n\n
    Returns:
    \n\n
    \n

    a tensor of per token loss

    \n
    \n", "signature": "(logits: torch.Tensor, labels: torch.Tensor, inplace: bool = False):", "funcdef": "def"}, "linghe.facade.norm": {"fullname": "linghe.facade.norm", "modulename": "linghe.facade.norm", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, "linghe.facade.norm.rms_norm": {"fullname": "linghe.facade.norm.rms_norm", "modulename": "linghe.facade.norm", "qualname": "rms_norm", "kind": "function", "doc": "

    rms norm of x with weight

    \n\n
    Arguments:
    \n\n
      \n
    • x: activation tensor
    • \n
    • weight: weight tensor
    • \n
    • eps: epsilon for RMS
    • \n
    \n\n
    Returns:
    \n\n
    \n

    rms output

    \n
    \n", "signature": "(x: torch.Tensor, weight: torch.Tensor, eps: float = 1e-06):", "funcdef": "def"}, "linghe.facade.norm.group_rms_norm_gate": {"fullname": "linghe.facade.norm.group_rms_norm_gate", "modulename": "linghe.facade.norm", "qualname": "group_rms_norm_gate", "kind": "function", "doc": "

    return group_rms_norm(transpose(attn_output, [0,1]), weight) * sigmoid(gate)

    \n\n
    Arguments:
    \n\n
      \n
    • attn_output: output of core attn, shape [bs, length, n_heads, head_dim]
    • \n
    • gate: gate tensor for attention output, shape [length, bs, dim]
    • \n
    • weight: weight of RMS norm, shape [dim]
    • \n
    • eps: epsilon for RMS
    • \n
    • group_size: group size of group RMS norm
    • \n
    \n\n
    Returns:
    \n\n
    \n

    output with shape [length, bs, dim]

    \n
    \n", "signature": "(\tattn_output: torch.Tensor,\tgate: torch.Tensor,\tweight: torch.Tensor,\teps: float = 1e-06,\tgroup_size: int = 4):", "funcdef": "def"}, "linghe.facade.rope": {"fullname": "linghe.facade.rope", "modulename": "linghe.facade.rope", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, "linghe.facade.rope.qk_norm_half_rope": {"fullname": "linghe.facade.rope.qk_norm_half_rope", "modulename": "linghe.facade.rope", "qualname": "qk_norm_half_rope", "kind": "function", "doc": "

    split qkv to q/k/v, apply qk norm and half rope to q/k, transpose q/k/v to flash-attention layout

    \n\n
    Arguments:
    \n\n
      \n
    • qkv: QKV tensor with size of [S, B, dim], heads are interleaved
    • \n
    • q_norm_weight: rms norm weight for query
    • \n
    • k_norm_weight: rms norm weight for key
    • \n
    • freqs: Freqs tensor based on half dim.
    • \n
    • H: Number of attention heads.
    • \n
    • h: Number of key/value heads.
    • \n
    • eps: epsilon value for L2 normalization.
    • \n
    \n\n
    Returns:
    \n\n
    \n
      \n
    • qo: shape [B, S, H, head_dim]
    • \n
    • ko: shape [B, S, h, head_dim]
    • \n
    • vo: shape [B, S, h, head_dim]
    • \n
    \n
    \n", "signature": "(\tqkv: torch.Tensor,\tq_norm_weight: torch.Tensor,\tk_norm_weight: torch.Tensor,\tfreqs: torch.Tensor,\tH: int = 32,\th: int = 4,\teps: float = 1e-06):", "funcdef": "def"}, "linghe.facade.smooth_quant_linear": {"fullname": "linghe.facade.smooth_quant_linear", "modulename": "linghe.facade.smooth_quant_linear", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear": {"fullname": "linghe.facade.smooth_quant_linear.SmoothQuantLinear", "modulename": "linghe.facade.smooth_quant_linear", "qualname": "SmoothQuantLinear", "kind": "class", "doc": "

    a naive implementation of smooth quantization linear

    \n", "bases": "torch.nn.modules.module.Module"}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear.__init__": {"fullname": "linghe.facade.smooth_quant_linear.SmoothQuantLinear.__init__", "modulename": "linghe.facade.smooth_quant_linear", "qualname": "SmoothQuantLinear.__init__", "kind": "function", "doc": "
    Arguments:
    \n\n
      \n
    • in_features: in feature number
    • \n
    • out_features: out feature number
    • \n
    • bias: whether use bias
    • \n
    • device: weight device
    • \n
    • dtype: weight dtype
    • \n
    \n", "signature": "(\tin_features: int,\tout_features: int,\tbias: bool = True,\tdevice=None,\tdtype=None)"}, "linghe.facade.transpose": {"fullname": "linghe.facade.transpose", "modulename": "linghe.facade.transpose", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, "linghe.facade.transpose.transpose_dim01": {"fullname": "linghe.facade.transpose.transpose_dim01", "modulename": "linghe.facade.transpose", "qualname": "transpose_dim01", "kind": "function", "doc": "

    transpose a tensor with the first two dims, x.ndims should not greater than 4

    \n\n
    Arguments:
    \n\n
      \n
    • x: input tensor
    • \n
    \n\n
    Returns:
    \n\n
    \n

    a transposed tensor

    \n
    \n", "signature": "(x):", "funcdef": "def"}, "linghe.gemm": {"fullname": "linghe.gemm", "modulename": "linghe.gemm", "kind": "module", "doc": "

    \n"}, "linghe.gemm.blockwise_fp8_gemm": {"fullname": "linghe.gemm.blockwise_fp8_gemm", "modulename": "linghe.gemm.blockwise_fp8_gemm", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, "linghe.gemm.channelwise_fp8_gemm": {"fullname": "linghe.gemm.channelwise_fp8_gemm", "modulename": "linghe.gemm.channelwise_fp8_gemm", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, "linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"fullname": "linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm", "modulename": "linghe.gemm.channelwise_fp8_gemm", "qualname": "triton_scaled_mm", "kind": "function", "doc": "

    similar to torch._scaled_mm, support accumulating gemm output to c\n and low precision output tensor

    \n\n
    Arguments:
    \n\n
      \n
    • a: left fp8 tensor
    • \n
    • b: right fp8 tensor, column-major
    • \n
    • a_scale: fp32 scale of a
    • \n
    • b_scale: fp32 scale of b
    • \n
    • out_dtype: output tensor dtype
    • \n
    • c: output tensor
    • \n
    • accum: accumulate output on c if True
    • \n
    \n\n
    Returns:
    \n\n
    \n

    c: output tensor

    \n
    \n", "signature": "(\ta: torch.Tensor,\tb: torch.Tensor,\ta_scale: torch.Tensor,\tb_scale: torch.Tensor,\tout_dtype=torch.float32,\tc=None,\taccum=True):", "funcdef": "def"}, "linghe.gemm.fp32_gemm": {"fullname": "linghe.gemm.fp32_gemm", "modulename": "linghe.gemm.fp32_gemm", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, "linghe.gemm.fp32_gemm.triton_fp32_gemm": {"fullname": "linghe.gemm.fp32_gemm.triton_fp32_gemm", "modulename": "linghe.gemm.fp32_gemm", "qualname": "triton_fp32_gemm", "kind": "function", "doc": "

    return fp32 gemm result with fp16/bf16 inputs,\n it's mainly used for MoE router GEMM\n and DO NOT suitable for large size GEMM

    \n\n
    Arguments:
    \n\n
      \n
    • a: left matrix with fp16/bf16 precision
    • \n
    • b: right matrix with fp16/bf16 precision
    • \n
    \n\n
    Returns:
    \n\n
    \n

    c: output with fp32 precision

    \n
    \n", "signature": "(a: torch.Tensor, b: torch.Tensor):", "funcdef": "def"}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_backward": {"fullname": "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_backward", "modulename": "linghe.gemm.fp32_gemm", "qualname": "triton_fp32_gemm_for_backward", "kind": "function", "doc": "

    mix precision gemm for backward, a@b.float()

    \n\n
    Arguments:
    \n\n
      \n
    • a: input gradient, fp32
    • \n
    • b: gemm weight, bf16/fp16
    • \n
    \n\n
    Returns:
    \n\n
    \n

    c: gradient of activation

    \n
    \n", "signature": "(a: torch.Tensor, b: torch.Tensor):", "funcdef": "def"}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_update": {"fullname": "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_update", "modulename": "linghe.gemm.fp32_gemm", "qualname": "triton_fp32_gemm_for_update", "kind": "function", "doc": "

    mix precision gemm for updaing weight

    \n\n
    Arguments:
    \n\n
      \n
    • a: gradient of output, fp32
    • \n
    • b: input activation, bf16/fp16
    • \n
    \n\n
    Returns:
    \n\n
    \n

    c: gradient of weight

    \n
    \n", "signature": "(a: torch.Tensor, b: torch.Tensor):", "funcdef": "def"}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"fullname": "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm", "modulename": "linghe.gemm.fp32_gemm", "qualname": "triton_scaled_fp32_gemm", "kind": "function", "doc": "

    c = (ascale[:,None])b\nthis kernel is used to fuse RMSNorm and quantization in MoE layer\nnative implementation:\n y = rms_norm(x),\n y_q = quantization(y),\n router_logits = y@w\nwe can not fuse rms_norm and quantization\nas we still need bf16 y for moe router gemm\nfused implementation:\n y_q, rms = quantization(rms_norm(x))\n router_logits = (x/rms)@y\nso we need a scaled fp32 gemm kernel

    \n\n
    Arguments:
    \n\n
      \n
    • a: activation tensor
    • \n
    • b: weight tensor
    • \n
    • scale: scale for activation tensor, 1/rms
    • \n
    \n\n
    Returns:
    \n\n
    \n

    output tensor

    \n
    \n", "signature": "(a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor):", "funcdef": "def"}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm_for_update": {"fullname": "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm_for_update", "modulename": "linghe.gemm.fp32_gemm", "qualname": "triton_scaled_fp32_gemm_for_update", "kind": "function", "doc": "

    see triton_scaled_fp32_gemm

    \n\n
    Arguments:
    \n\n
      \n
    • a: y
    • \n
    • b: activation before RMS norm
    • \n
    • scale: 1/rms
    • \n
    \n\n
    Returns:
    \n\n
    \n

    dw

    \n
    \n", "signature": "(a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor):", "funcdef": "def"}, "linghe.quant": {"fullname": "linghe.quant", "modulename": "linghe.quant", "kind": "module", "doc": "

    \n"}, "linghe.quant.block": {"fullname": "linghe.quant.block", "modulename": "linghe.quant.block", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, "linghe.quant.block.triton_block_quant": {"fullname": "linghe.quant.block.triton_block_quant", "modulename": "linghe.quant.block", "qualname": "triton_block_quant", "kind": "function", "doc": "

    blockwise quantize x

    \n\n
    Arguments:
    \n\n
      \n
    • x: input tensor
    • \n
    • block_size: block wise
    • \n
    • round_scale: whether round scale to power of 2
    • \n
    \n\n
    Returns:
    \n\n
    \n
      \n
    • y: quantized tensor, float8_e4m3fn
    • \n
    • s: quantization scale, float32
    • \n
    \n
    \n", "signature": "(x, block_size=128, round_scale=False):", "funcdef": "def"}, "linghe.quant.channel": {"fullname": "linghe.quant.channel", "modulename": "linghe.quant.channel", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, "linghe.quant.channel.triton_row_quant": {"fullname": "linghe.quant.channel.triton_row_quant", "modulename": "linghe.quant.channel", "qualname": "triton_row_quant", "kind": "function", "doc": "

    rowwise quantize x

    \n\n
    Arguments:
    \n\n
      \n
    • x: input x
    • \n
    • round_scale: whether round scale to power of 2
    • \n
    \n\n
    Returns:
    \n\n
    \n

    x_q: quantized tensor\n x_scale: quantization scale

    \n
    \n", "signature": "(x, round_scale=False):", "funcdef": "def"}, "linghe.quant.channel.triton_tokenwise_row_quant": {"fullname": "linghe.quant.channel.triton_tokenwise_row_quant", "modulename": "linghe.quant.channel", "qualname": "triton_tokenwise_row_quant", "kind": "function", "doc": "

    rowwise quantize x with power of 2 dim size

    \n\n
    Arguments:
    \n\n
      \n
    • x: input x
    • \n
    • round_scale: whether round scale to power of 2
    • \n
    \n\n
    Returns:
    \n\n
    \n

    out: quantized tensor\n scale: quantization scale

    \n
    \n", "signature": "(x, out=None, scale=None, round_scale=False):", "funcdef": "def"}, "linghe.quant.channel.triton_transpose_row_quant": {"fullname": "linghe.quant.channel.triton_transpose_row_quant", "modulename": "linghe.quant.channel", "qualname": "triton_transpose_row_quant", "kind": "function", "doc": "

    transpose x and row quantize x

    \n\n
    Arguments:
    \n\n
      \n
    • x: input x
    • \n
    • round_scale: whether round scale to power of 2
    • \n
    \n\n
    Returns:
    \n\n
    \n
      \n
    • x_q: quantized tensor
    • \n
    • x_scale: quantization scale
    • \n
    \n
    \n", "signature": "(x, round_scale=False):", "funcdef": "def"}, "linghe.quant.group": {"fullname": "linghe.quant.group", "modulename": "linghe.quant.group", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, "linghe.quant.group.triton_group_quant": {"fullname": "linghe.quant.group.triton_group_quant", "modulename": "linghe.quant.group", "qualname": "triton_group_quant", "kind": "function", "doc": "

    groupwise quantize x, group is in under rowwise format

    \n\n
    Arguments:
    \n\n
      \n
    • x: input tensor
    • \n
    • group_size: group wise
    • \n
    • round_scale: whether round scale to power of 2
    • \n
    \n\n
    Returns:
    \n\n
    \n
      \n
    • y: quantized tensor, float8_e4m3fn
    • \n
    • s: quantization scale, float32
    • \n
    \n
    \n", "signature": "(x, dtype=torch.float8_e4m3fn, group_size=128, round_scale=False):", "funcdef": "def"}, "linghe.quant.hadamard": {"fullname": "linghe.quant.hadamard", "modulename": "linghe.quant.hadamard", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, "linghe.quant.hadamard.triton_hadamard_quant": {"fullname": "linghe.quant.hadamard.triton_hadamard_quant", "modulename": "linghe.quant.hadamard", "qualname": "triton_hadamard_quant", "kind": "function", "doc": "

    apply hadamard transformation and then quantize transformed tensor

    \n\n
    Arguments:
    \n\n
      \n
    • x: input tensor
    • \n
    • hm: hamadard matrix
    • \n
    \n\n
    Returns:
    \n\n
    \n
      \n
    • x_q: rowwise quantized tensor of non-transposed x
    • \n
    • x_scale: rowwise quantization scale of non-transposed x
    • \n
    • xt_q: columnwise quantized tensor of transposed x
    • \n
    • xt_scale: columnwise quantization scale of transposed x
    • \n
    \n
    \n", "signature": "(x, hm):", "funcdef": "def"}, "linghe.quant.smooth": {"fullname": "linghe.quant.smooth", "modulename": "linghe.quant.smooth", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, "linghe.utils": {"fullname": "linghe.utils", "modulename": "linghe.utils", "kind": "module", "doc": "

    \n"}, "linghe.utils.add": {"fullname": "linghe.utils.add", "modulename": "linghe.utils.add", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, "linghe.utils.add.triton_inplace_add": {"fullname": "linghe.utils.add.triton_inplace_add", "modulename": "linghe.utils.add", "qualname": "triton_inplace_add", "kind": "function", "doc": "

    inplace add y to x

    \n\n
    Arguments:
    \n\n
      \n
    • x: Tensor
    • \n
    • y: Tensor
    • \n
    • accum: x += y if accum=True else x.copy_(y)
    • \n
    \n\n
    Returns:
    \n\n
    \n

    updated x

    \n
    \n", "signature": "(x: torch.Tensor, y: torch.Tensor, accum: bool = True):", "funcdef": "def"}, "linghe.utils.dot": {"fullname": "linghe.utils.dot", "modulename": "linghe.utils.dot", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, "linghe.utils.dot.triton_dot": {"fullname": "linghe.utils.dot.triton_dot", "modulename": "linghe.utils.dot", "qualname": "triton_dot", "kind": "function", "doc": "

    vector dot multiply, output = sum(x*y, 1),\nit is used to calculate gradient of router weight

    \n\n
    Arguments:
    \n\n
      \n
    • x:
    • \n
    • y:
    • \n
    \n\n
    Returns:
    \n\n
    \n

    output of sum(x*y, 1)

    \n
    \n", "signature": "(x, y):", "funcdef": "def"}, "linghe.utils.gather": {"fullname": "linghe.utils.gather", "modulename": "linghe.utils.gather", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, "linghe.utils.gather.triton_make_row_id_map": {"fullname": "linghe.utils.gather.triton_make_row_id_map", "modulename": "linghe.utils.gather", "qualname": "triton_make_row_id_map", "kind": "function", "doc": "

    make row id map, values in the tensor are the row indices

    \n\n
    Arguments:
    \n\n
      \n
    • routing_map: a tensor of 0/1 values, 1 indicates routed
    • \n
    • multiple_of: padding the tokens of each expert to multiple of this value
    • \n
    \n\n
    Returns:
    \n\n
    \n

    row id map with shape [n_tokens, n_experts]

    \n
    \n", "signature": "(routing_map: torch.Tensor, multiple_of: int = 1):", "funcdef": "def"}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"fullname": "linghe.utils.gather.triton_make_row_id_map_and_indices", "modulename": "linghe.utils.gather", "qualname": "triton_make_row_id_map_and_indices", "kind": "function", "doc": "

    similar with triton_make_row_id_map, but output an indices tensor as well

    \n\n
    Arguments:
    \n\n
      \n
    • routing_map: [n_tokens, n_experts]
    • \n
    • num_out_tokens: sum(round_up_to(n_tokens, multiple_of))
    • \n
    • multiple_of: padding the tokens of each expert to this value
    • \n
    \n\n
    Returns:
    \n\n
    \n

    row_in_map: [n_tokens, n_experts]\n row_indices: [num_out_tokens]

    \n
    \n", "signature": "(routing_map: torch.Tensor, num_out_tokens: int, multiple_of: int = 1):", "funcdef": "def"}, "linghe.utils.gather.triton_index_select": {"fullname": "linghe.utils.gather.triton_index_select", "modulename": "linghe.utils.gather", "qualname": "triton_index_select", "kind": "function", "doc": "

    index select for quantized tensor

    \n\n
    Arguments:
    \n\n
      \n
    • x: [bs, dim]
    • \n
    • indices: [K]
    • \n
    • scale: [bs]
    • \n
    \n\n
    Returns:
    \n\n
    \n

    out: output of selected x\n scale_out: scale of selected scale

    \n
    \n", "signature": "(x, indices, scale=None, out=None, scale_out=None):", "funcdef": "def"}, "linghe.utils.gather.triton_permute_with_mask_map": {"fullname": "linghe.utils.gather.triton_permute_with_mask_map", "modulename": "linghe.utils.gather", "qualname": "triton_permute_with_mask_map", "kind": "function", "doc": "

    gather quantized tensor with row id map

    \n\n
    Arguments:
    \n\n
      \n
    • inp: [num_tokens, hidden_size], rowwise quantized tensor
    • \n
    • scale: [num_tokens], quantization scale
    • \n
    • probs: router prob, used as weight
    • \n
    • row_id_map: [n_experts, num_tokens]\nindex >= 0: row index of output tensor\nindex == -1: ignore\nNote: index may not be contiguous
    • \n
    • num_out_tokens: output token count, including padding tokens
    • \n
    • contiguous: whether indices in row_id_map is contiguous,\nFalse means padded
    • \n
    • tokens_per_expert: [num_experts], token count per expert,\nnon-blocking cuda tensor
    • \n
    \n\n
    Returns:
    \n\n
    \n

    output: permuted quantized tensor\n permuted_scale: permuted quantization scale\n permuted_probs: permuted router prob

    \n
    \n", "signature": "(\tinp: torch.Tensor,\tscale: torch.Tensor,\tprobs: torch.Tensor,\trow_id_map: torch.Tensor,\tnum_out_tokens: int,\tcontiguous: bool = True,\ttokens_per_expert: Optional[torch.Tensor] = None):", "funcdef": "def"}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"fullname": "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices", "modulename": "linghe.utils.gather", "qualname": "triton_batch_transpose_smooth_permute_with_indices", "kind": "function", "doc": "

    used for smooth quantization backward in megatron 0.12,\nx is gathered, requantized, padded to multiple of 32 and tranposed

    \n\n
    Arguments:
    \n\n
      \n
    • x: dy, [bs, dim], it is smooth quantized
    • \n
    • scale: [bs], quantized scale
    • \n
    • org_smooth_scale: [dim]
    • \n
    • smooth_scales: [n_experts, dim]
    • \n
    • indices: [sum(tokens_per_experts)]
    • \n
    • token_count_per_expert: [n_experts], tensor of token count per expert
    • \n
    • splits: [n_experts], list of token_count_per_expert
    • \n
    • round_scale: round quantization scale to power of 2
    • \n
    \n\n
    Returns:
    \n\n
    \n

    x_q: [sum(roundup(tokens_per_experts)) * dim]\n x_scale: [sum(roundup(tokens_per_experts))]

    \n
    \n", "signature": "(\tx,\tscale,\torg_smooth_scale,\tsmooth_scales,\tindices,\ttoken_count_per_expert,\tsplits,\tx_q=None,\tx_scale=None,\tround_scale=False):", "funcdef": "def"}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"fullname": "linghe.utils.gather.triton_smooth_weighted_permute_with_indices", "modulename": "linghe.utils.gather", "qualname": "triton_smooth_weighted_permute_with_indices", "kind": "function", "doc": "

    select and smooth and quant, used in megatron 0.11 all2all moe

    \n\n
    Arguments:
    \n\n
      \n
    • grads: [bs, dim]
    • \n
    • tokens: [bs, dim]
    • \n
    • smooth_scales: [n_experts, dim]
    • \n
    • token_count_per_expert: [n_experts]
    • \n
    • indices: [n_experts*topk]
    • \n
    • reverse: whether scale is 1/scale
    • \n
    • round_scale: whether round scale to power of 2
    • \n
    \n\n
    Returns:
    \n\n
    \n

    x_q: [bs*topk, dim]\n x_scale: [bstopk]\n x_sum: [bstopk]

    \n
    \n", "signature": "(\tgrads,\ttokens,\tsmooth_scales,\ttoken_count_per_expert,\tindices,\tx_q=None,\tx_scale=None,\tx_sum=None,\treverse=False,\tround_scale=False):", "funcdef": "def"}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"fullname": "linghe.utils.gather.triton_smooth_permute_with_indices", "modulename": "linghe.utils.gather", "qualname": "triton_smooth_permute_with_indices", "kind": "function", "doc": "

    select and smooth and quant

    \n\n
    Arguments:
    \n\n
      \n
    • grad_data: [bs, dim]
    • \n
    • grad_scale: [bs]
    • \n
    • smooth_scales: [n_experts, dim]
    • \n
    • token_count_per_expert: [n_experts]
    • \n
    • indices: [n_experts*topk]
    • \n
    • x_q: [bs*topk, dim]
    • \n
    • x_scale: [bs*topk]
    • \n
    • reverse:
    • \n
    • round_scale:
    • \n
    \n\n

    Returns:

    \n", "signature": "(\tgrad_data,\tgrad_scale,\tsmooth_scales,\ttoken_count_per_expert,\tindices,\tx_q=None,\tx_scale=None,\treverse=False,\tround_scale=False):", "funcdef": "def"}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"fullname": "linghe.utils.gather.triton_smooth_permute_with_mask_map", "modulename": "linghe.utils.gather", "qualname": "triton_smooth_permute_with_mask_map", "kind": "function", "doc": "

    gather ( and optional dequant) and smooth quant

    \n\n
    Arguments:
    \n\n
      \n
    • inp: [num_tokens, hidden_size], rowwise quantized tensor
    • \n
    • row_id_map: [n_experts, num_tokens], indices
    • \n
    • scale: [num_tokens, hs], rowwise_scale_inv, optional
    • \n
    • num_tokens: [n_experts]
    • \n
    • num_experts:
    • \n
    • num_out_tokens:
    • \n
    • hidden_size:
    • \n
    • smooth_scales: [n_experts, hidden_size]
    • \n
    • reverse:
    • \n
    • round_scale:
    • \n
    \n\n
    Returns:
    \n\n
    \n
      \n
    • output: output tensor
    • \n
    • permuted_scale: permuted scale if scale is not None
    • \n
    \n
    \n", "signature": "(\tinp: torch.Tensor,\trow_id_map: torch.Tensor,\tscale: torch.Tensor,\tnum_tokens: int,\tnum_experts: int,\tnum_out_tokens: int,\thidden_size: int,\tsmooth_scales: torch.Tensor,\treverse=True,\tround_scale=False):", "funcdef": "def"}, "linghe.utils.loss": {"fullname": "linghe.utils.loss", "modulename": "linghe.utils.loss", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, "linghe.utils.loss.triton_softmax_cross_entropy_forward": {"fullname": "linghe.utils.loss.triton_softmax_cross_entropy_forward", "modulename": "linghe.utils.loss", "qualname": "triton_softmax_cross_entropy_forward", "kind": "function", "doc": "

    compute token-wise softmax cross entropy loss

    \n\n
    Arguments:
    \n\n
      \n
    • logits: logits tensor
    • \n
    • labels: labels tensor
    • \n
    \n\n
    Returns:
    \n\n
    \n

    loss of each token

    \n
    \n", "signature": "(logits, labels):", "funcdef": "def"}, "linghe.utils.loss.triton_softmax_cross_entropy_backward": {"fullname": "linghe.utils.loss.triton_softmax_cross_entropy_backward", "modulename": "linghe.utils.loss", "qualname": "triton_softmax_cross_entropy_backward", "kind": "function", "doc": "

    backward of softmax cross entropy loss

    \n\n
    Arguments:
    \n\n
      \n
    • logits: logit tensor, [bs, dim]
    • \n
    • labels: label tensor, [bs]
    • \n
    • sum_exp: [bs]
    • \n
    • max_logit: [bs]
    • \n
    • input_grad: gradient, [bs, dim]
    • \n
    \n\n
    Returns:
    \n\n
    \n

    output_grad: [bs, dim]

    \n
    \n", "signature": "(logits, labels, sum_exp, max_logit, input_grad, output_grad=None):", "funcdef": "def"}, "linghe.utils.norm": {"fullname": "linghe.utils.norm", "modulename": "linghe.utils.norm", "kind": "module", "doc": "

    \n"}, "linghe.utils.norm.triton_rms_norm_forward": {"fullname": "linghe.utils.norm.triton_rms_norm_forward", "modulename": "linghe.utils.norm", "qualname": "triton_rms_norm_forward", "kind": "function", "doc": "

    rms norm

    \n\n
    Arguments:
    \n\n
      \n
    • x: input tensor
    • \n
    • weight: weight of rms norm
    • \n
    • eps: epsilon of rms norm
    • \n
    \n\n
    Returns:
    \n\n
    \n

    out: output tensor

    \n
    \n", "signature": "(x, weight, eps=1e-06, out=None):", "funcdef": "def"}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"fullname": "linghe.utils.norm.triton_rms_norm_and_block_quant_forward", "modulename": "linghe.utils.norm", "qualname": "triton_rms_norm_and_block_quant_forward", "kind": "function", "doc": "

    Fused RMSNorm forward and block quantization.

    \n\n
    Arguments:
    \n\n
      \n
    • x: Input tensor, shape [M, N]
    • \n
    • weight: RMSNorm weight, shape [N]
    • \n
    • eps: epsilon value for L2 normalization.
    • \n
    • out: output of quantization data
    • \n
    • scale: output of quantization scale.
    • \n
    • rms: output of rms
    • \n
    • round_scale: Set whether to force power of 2 scales.
    • \n
    • output_mode: one of {0, 1, 2}.\n0: only output non-transpose tensor\n1: only output transposed tensor\n2: return both
    • \n
    \n\n
    Returns:
    \n\n
    \n
      \n
    • out: quantization data.
    • \n
    • scale: quantization scale.
    • \n
    • rms: Reciprocal of the root mean square of the\n input calculated over the last dimension.
    • \n
    • transpose_output: quantization data of transposed gradient.
    • \n
    • transpose_scale: quantization scale of transposed gradient.
    • \n
    \n
    \n", "signature": "(\tx: torch.Tensor,\tweight: torch.Tensor,\teps: float = 1e-06,\tout: Optional[torch.Tensor] = None,\tscale: Optional[torch.Tensor] = None,\trms: Optional[torch.Tensor] = None,\tround_scale: bool = False,\toutput_mode: int = 2):", "funcdef": "def"}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"fullname": "linghe.utils.norm.triton_group_rms_norm_gate_forward", "modulename": "linghe.utils.norm", "qualname": "triton_group_rms_norm_gate_forward", "kind": "function", "doc": "

    norm and gate in linear attention

    \n\n
    Arguments:
    \n\n
      \n
    • x: output of attn, [bs, length, n_heads, head_dim]
    • \n
    • gate: gate tensor, [length, bs, dim] if transpose=True else [bs, length, dim]
    • \n
    • weight: rms norm weight, [dim]
    • \n
    • eps: epsilon of rms norm
    • \n
    • group_size: group size of group rms norm
    • \n
    • transpose: whether gate is transposed and output will be transposed
    • \n
    \n\n
    Returns:
    \n\n
    \n

    output tensor, [length, bs, dim] if transpose=True else [bs, length, dim]

    \n
    \n", "signature": "(\tx: torch.Tensor,\tgate: torch.Tensor,\tweight: torch.Tensor,\teps=1e-06,\tgroup_size=4,\ttranspose=True):", "funcdef": "def"}, "linghe.utils.rearange": {"fullname": "linghe.utils.rearange", "modulename": "linghe.utils.rearange", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, "linghe.utils.rearange.triton_split_and_cat": {"fullname": "linghe.utils.rearange.triton_split_and_cat", "modulename": "linghe.utils.rearange", "qualname": "triton_split_and_cat", "kind": "function", "doc": "

    split x to multiple tensors and cat with indices,\nit is used for permutation in moe

    \n\n
    Arguments:
    \n\n
      \n
    • x: [bs, dim]
    • \n
    • counts: [n_split]
    • \n
    • indices: [n_split]
    • \n
    • scales: [bs]
    • \n
    \n\n
    Returns:
    \n\n
    \n
      \n
    • y: output tensor
    • \n
    • output_scales: output scales if scales is not None
    • \n
    \n
    \n", "signature": "(x, counts, indices, scales=None):", "funcdef": "def"}, "linghe.utils.reduce": {"fullname": "linghe.utils.reduce", "modulename": "linghe.utils.reduce", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, "linghe.utils.reduce.triton_abs_max": {"fullname": "linghe.utils.reduce.triton_abs_max", "modulename": "linghe.utils.reduce", "qualname": "triton_abs_max", "kind": "function", "doc": "

    columnwise abs max of x, it is used in smooth quantization

    \n\n
    Arguments:
    \n\n
      \n
    • x: input tensor, may be quantized tensor
    • \n
    • scale: quantization scale if x is quantized
    • \n
    • smooth_scale: optional smooth scale
    • \n
    • min_value: output = max(max(abs(x,0)), min_value)
    • \n
    • axis: reduce axis
    • \n
    \n\n
    Returns:
    \n\n
    \n

    max tensor

    \n
    \n", "signature": "(x, scale=None, smooth_scale=None, min_value=1e-30, axis=0):", "funcdef": "def"}, "linghe.utils.reduce.triton_batch_count_zero": {"fullname": "linghe.utils.reduce.triton_batch_count_zero", "modulename": "linghe.utils.reduce", "qualname": "triton_batch_count_zero", "kind": "function", "doc": "

    count zero in tensor list, it is used to monitor zeros in gradient tensor

    \n\n
    Arguments:
    \n\n
      \n
    • xs: input tensors
    • \n
    \n\n
    Returns:
    \n\n
    \n

    a single-value int64 tensor

    \n
    \n", "signature": "(xs):", "funcdef": "def"}, "linghe.utils.reduce.triton_batch_sum_with_ord": {"fullname": "linghe.utils.reduce.triton_batch_sum_with_ord", "modulename": "linghe.utils.reduce", "qualname": "triton_batch_sum_with_ord", "kind": "function", "doc": "

    return sum(abs(x)**ord).

    \n\n
    Arguments:
    \n\n
      \n
    • xs: Tensor lists.
    • \n
    • ord: the order of tensor.
    • \n
    \n\n
    Returns:
    \n\n
    \n

    a single-value fp32 tensor

    \n
    \n", "signature": "(xs, ord=2):", "funcdef": "def"}, "linghe.utils.rope": {"fullname": "linghe.utils.rope", "modulename": "linghe.utils.rope", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, "linghe.utils.rope.triton_half_rope_forward": {"fullname": "linghe.utils.rope.triton_half_rope_forward", "modulename": "linghe.utils.rope", "qualname": "triton_half_rope_forward", "kind": "function", "doc": "

    apply norm to qk, then apply half rope to qk

    \n\n
    Arguments:
    \n\n
      \n
    • q: query tensor, [len, bs, q_head, head_dim]
    • \n
    • k: key tensor, [len, bs, kv_head, head_dim]
    • \n
    • freqs: rope freqs
    • \n
    \n\n
    Returns:
    \n\n
    \n
      \n
    • qo: query output
    • \n
    • ko: key output
    • \n
    \n
    \n", "signature": "(q, k, freqs, transposed=True):", "funcdef": "def"}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"fullname": "linghe.utils.rope.triton_qk_norm_and_half_rope_forward", "modulename": "linghe.utils.rope", "qualname": "triton_qk_norm_and_half_rope_forward", "kind": "function", "doc": "

    split qkv to q/k/v, apply qk norm and half rope to q/k,\n transpose q/k/v to flash-attention layout

    \n\n
    Arguments:
    \n\n
      \n
    • qkv: QKV tensor with size of [S, B, dim], heads are interleaved
    • \n
    • q_norm_weight: rms norm weight for query
    • \n
    • k_norm_weight: rms norm weight for key
    • \n
    • freqs: Freqs tensor based on half dim.
    • \n
    • H: Number of attention heads.
    • \n
    • h: Number of key/value heads.
    • \n
    • eps: epsilon value for L2 normalization.
    • \n
    • interleaved: whether head of qkv is interleaved,\ninterleaved: [q...qkvq...qkv]\nnon-interleaved: [q...qk...kv...v]
    • \n
    • transposed: whether qkv is tranposed\ntransposed: [S, B, dim]\nnon-transposed: [B, S, dim]\nonly support transpose format currently
    • \n
    \n\n
    Returns:
    \n\n
    \n
      \n
    • qo: shape [B, S, H, head_dim]
    • \n
    • ko: shape [B, S, h, head_dim]
    • \n
    • vo: shape [B, S, h, head_dim]
    • \n
    \n
    \n", "signature": "(\tqkv,\tq_norm_weight,\tk_norm_weight,\tfreqs,\tH=32,\th=4,\teps=1e-06,\tinterleaved=True,\ttransposed=False):", "funcdef": "def"}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"fullname": "linghe.utils.rope.triton_qk_norm_and_half_rope_backward", "modulename": "linghe.utils.rope", "qualname": "triton_qk_norm_and_half_rope_backward", "kind": "function", "doc": "

    backward kernel of triton_qk_norm_and_half_rope_forward

    \n\n
    Arguments:
    \n\n
      \n
    • gq: gradient of qo, [len, bs, q_head, head_dim]
    • \n
    • gk: gradient of ko, [len, bs, q_head, head_dim]
    • \n
    • gv: gradient of vo, [len, bs, q_head, head_dim]
    • \n
    • qkv: input qkv
    • \n
    • q_norm_weight:
    • \n
    • k_norm_weight:
    • \n
    • freqs:
    • \n
    • eps:
    • \n
    • interleaved:
    • \n
    • transposed:
    • \n
    \n\n
    Returns:
    \n\n
    \n
      \n
    • dqkv: gradient of qkv
    • \n
    • dqw: gradient of q_norm_weight
    • \n
    • dkw: gradient of k_norm_weight
    • \n
    \n
    \n", "signature": "(\tgq,\tgk,\tgv,\tqkv,\tq_norm_weight,\tk_norm_weight,\tfreqs,\teps=1e-06,\tinterleaved=True,\ttransposed=False):", "funcdef": "def"}, "linghe.utils.scatter": {"fullname": "linghe.utils.scatter", "modulename": "linghe.utils.scatter", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, "linghe.utils.scatter.triton_aligned_scatter_add": {"fullname": "linghe.utils.scatter.triton_aligned_scatter_add", "modulename": "linghe.utils.scatter", "qualname": "triton_aligned_scatter_add", "kind": "function", "doc": "

    scatter_add for megatron 0.11

    \n\n
    Arguments:
    \n\n
      \n
    • x: input tensor
    • \n
    • outputs: output tensor
    • \n
    • indices: gather indices
    • \n
    • weights: rowwise weight, it is router prob in MoE router
    • \n
    \n\n
    Returns:
    \n\n
    \n

    output tensor

    \n
    \n", "signature": "(\tx: torch.Tensor,\toutputs: torch.Tensor,\tindices: torch.Tensor,\tweights: Optional[torch.Tensor] = None):", "funcdef": "def"}, "linghe.utils.scatter.triton_scatter_add": {"fullname": "linghe.utils.scatter.triton_scatter_add", "modulename": "linghe.utils.scatter", "qualname": "triton_scatter_add", "kind": "function", "doc": "

    naive version of scatter add, very slow

    \n\n
    Arguments:
    \n\n
      \n
    • x: input tensor
    • \n
    • outputs: output tensor
    • \n
    • indices: indices
    • \n
    \n\n
    Returns:
    \n\n
    \n

    output tensor

    \n
    \n", "signature": "(x, outputs, indices):", "funcdef": "def"}, "linghe.utils.scatter.triton_unpermute_with_mask_map": {"fullname": "linghe.utils.scatter.triton_unpermute_with_mask_map", "modulename": "linghe.utils.scatter", "qualname": "triton_unpermute_with_mask_map", "kind": "function", "doc": "

    scatter add with row id map

    \n\n
    Arguments:
    \n\n
      \n
    • grad: gradient tensor, [num_out_tokens, hidden_size]
    • \n
    • row_id_map: row id map, [n_experts, num_tokens]
    • \n
    • probs: [num_out_tokens]
    • \n
    \n\n
    Returns:
    \n\n
    \n
      \n
    • output: [num_tokens, hidden_size]
    • \n
    • restore_probs: [num_tokens, num_experts]
    • \n
    \n
    \n", "signature": "(grad: torch.Tensor, row_id_map: torch.Tensor, probs: torch.Tensor):", "funcdef": "def"}, "linghe.utils.silu": {"fullname": "linghe.utils.silu", "modulename": "linghe.utils.silu", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, "linghe.utils.silu.triton_weighted_silu_forward": {"fullname": "linghe.utils.silu.triton_weighted_silu_forward", "modulename": "linghe.utils.silu", "qualname": "triton_weighted_silu_forward", "kind": "function", "doc": "

    compute silu(x)*weight, used in bf16/fp16 training with MoE

    \n\n
    Arguments:
    \n\n
      \n
    • x: input tensor
    • \n
    • weight: tokenwise weight
    • \n
    \n\n
    Returns:
    \n\n
    \n

    out: output tensor

    \n
    \n", "signature": "(x, weight=None, out=None):", "funcdef": "def"}, "linghe.utils.silu.triton_weighted_silu_backward": {"fullname": "linghe.utils.silu.triton_weighted_silu_backward", "modulename": "linghe.utils.silu", "qualname": "triton_weighted_silu_backward", "kind": "function", "doc": "

    backward of triton_weighted_silu_forward

    \n\n
    Arguments:
    \n\n
      \n
    • g: gradient tensor
    • \n
    • x: input tensor
    • \n
    • weight: weight tensor
    • \n
    \n\n
    Returns:
    \n\n
    \n
      \n
    • dx: gradient of x
    • \n
    • dw: gradient of weight
    • \n
    \n
    \n", "signature": "(\tg: torch.Tensor,\tx: torch.Tensor,\tweight: Optional[torch.Tensor] = None):", "funcdef": "def"}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"fullname": "linghe.utils.silu.triton_silu_and_block_quant_forward", "modulename": "linghe.utils.silu", "qualname": "triton_silu_and_block_quant_forward", "kind": "function", "doc": "

    fused silu and blockwise quantization, used in shared expert

    \n\n
    Arguments:
    \n\n
      \n
    • x: input tensor
    • \n
    • round_scale: whether round scale to power of 2
    • \n
    • output_mode: one of {0, 1, 2}\n0: only output non-transposed quantized tensor\n1: only output transposed quantized tensor\n2: output both
    • \n
    \n\n
    Returns:
    \n\n
    \n
      \n
    • out: quantized tensor
    • \n
    • scale: quantization scale
    • \n
    • transpose_output: quantized tensor of transposed output
    • \n
    • transpose_scale: quantization scale of transposed output
    • \n
    \n
    \n", "signature": "(x, out=None, scale=None, round_scale=False, output_mode=2):", "funcdef": "def"}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"fullname": "linghe.utils.silu.triton_silu_and_block_quant_backward", "modulename": "linghe.utils.silu", "qualname": "triton_silu_and_block_quant_backward", "kind": "function", "doc": "

    backward of triton_silu_and_block_quant_forward

    \n\n
    Arguments:
    \n\n
      \n
    • g: gradient
    • \n
    • x: input tensor
    • \n
    • round_scale: whether round to power of 2
    • \n
    \n\n
    Returns:
    \n\n
    \n
      \n
    • dx: quantized non-transposed gradient
    • \n
    • dx_scale: scales of quantization non-transposed gradient
    • \n
    • transpose_dx: quantized transposed gradient
    • \n
    • transpose_dx_scale: scales of quantization transposed gradient
    • \n
    \n
    \n", "signature": "(g, x, round_scale=False):", "funcdef": "def"}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"fullname": "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward", "modulename": "linghe.utils.silu", "qualname": "triton_batch_weighted_silu_and_block_quant_forward", "kind": "function", "doc": "

    silu and blockwise quantize activation in routed experts

    \n\n
    Arguments:
    \n\n
      \n
    • x: activation tensor in routed experts
    • \n
    • weight: router prob tensor
    • \n
    • counts: cuda tensor of token count per expert
    • \n
    • splits: python int list of token count per expert
    • \n
    • round_scale: whether round scale to power of 2
    • \n
    • output_mode: one of {0, 1, 2}\n0: only output non-transposed quantized tensor\n1: only output transposed quantized tensor\n2: output both
    • \n
    \n\n
    Returns:
    \n\n
    \n
      \n
    • out: quantized tensor
    • \n
    • scale: quantization scale
    • \n
    • transpose_output: quantized tensor of transposed output
    • \n
    • transpose_scale: quantization scale of transposed output
    • \n
    \n
    \n", "signature": "(\tx,\tweight,\tcounts,\tsplits=None,\tout=None,\tscale=None,\tround_scale=False,\toutput_mode=2):", "funcdef": "def"}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"fullname": "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward", "modulename": "linghe.utils.silu", "qualname": "triton_batch_weighted_silu_and_block_quant_backward", "kind": "function", "doc": "

    backward of triton_batch_weighted_silu_and_block_quant_forward

    \n\n
    Arguments:
    \n\n
      \n
    • g: gradient
    • \n
    • x: input tensor
    • \n
    • weight: router prob tensor
    • \n
    • counts: cuda tensor of token count per expert
    • \n
    • splits: python int list of token count per expert
    • \n
    • round_scale: whether round scale to power of 2
    • \n
    \n\n
    Returns:
    \n\n
    \n
      \n
    • dx: quantized non-transposed gradient
    • \n
    • dx_scale: scales of quantization non-transposed gradient
    • \n
    • dw: gradient of weight
    • \n
    • transpose_dx: quantized transposed gradient
    • \n
    • transpose_dx_scale: scales of quantization transposed gradient
    • \n
    \n
    \n", "signature": "(g, x, weight, counts, splits=None, round_scale=False):", "funcdef": "def"}, "linghe.utils.transpose": {"fullname": "linghe.utils.transpose", "modulename": "linghe.utils.transpose", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, "linghe.utils.transpose.triton_transpose": {"fullname": "linghe.utils.transpose.triton_transpose", "modulename": "linghe.utils.transpose", "qualname": "triton_transpose", "kind": "function", "doc": "

    transpose x with dim0 and dim1

    \n\n
    Arguments:
    \n\n
      \n
    • x: input tensor
    • \n
    • dim0: dim 0
    • \n
    • dim1: dim 1
    • \n
    \n\n
    Returns:
    \n\n
    \n

    transposed tensor

    \n
    \n", "signature": "(\tx: torch.Tensor,\tdim0: Optional[int] = None,\tdim1: Optional[int] = None):", "funcdef": "def"}, "linghe.utils.transpose.triton_transpose_and_pad": {"fullname": "linghe.utils.transpose.triton_transpose_and_pad", "modulename": "linghe.utils.transpose", "qualname": "triton_transpose_and_pad", "kind": "function", "doc": "

    transpose x and padding the column size to be mutiplier of 32,\nit is used for calculated gradient of weight with torch._scaled__mm

    \n\n
    Arguments:
    \n\n
      \n
    • x: input tensor
    • \n
    • out:
    • \n
    • pad: whether need padding
    • \n
    \n\n
    Returns:
    \n\n
    \n

    out: output tensor

    \n
    \n", "signature": "(x, out=None, pad=True):", "funcdef": "def"}, "linghe.utils.transpose.triton_batch_transpose": {"fullname": "linghe.utils.transpose.triton_batch_transpose", "modulename": "linghe.utils.transpose", "qualname": "triton_batch_transpose", "kind": "function", "doc": "

    batch transpose x

    \n\n
    Arguments:
    \n\n
      \n
    • xs: input tensor list, [M, N]*expert
    • \n
    \n\n
    Returns:
    \n\n
    \n

    xts: output tensor list, [N,M]*expert

    \n
    \n", "signature": "(xs, xts=None):", "funcdef": "def"}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"fullname": "linghe.utils.transpose.triton_batch_transpose_and_pad", "modulename": "linghe.utils.transpose", "qualname": "triton_batch_transpose_and_pad", "kind": "function", "doc": "

    transpose and pad each tensor stored in x

    \n\n
    Arguments:
    \n\n
      \n
    • x: [sum(bs), N]
    • \n
    • count_list: a python list of token count
    • \n
    • pad: whether pad to mutiplier of 32,\npadding value should be filled with 0 if padded
    • \n
    \n\n
    Returns:
    \n\n
    \n

    x_t: output tensor

    \n
    \n", "signature": "(x, count_list, x_t=None, pad=True):", "funcdef": "def"}}, "docInfo": {"linghe": {"qualname": 0, "fullname": 1, "annotation": 0, "default_value": 0, "signature": 0, "bases": 0, "doc": 3}, "linghe.facade": {"qualname": 0, "fullname": 2, "annotation": 0, "default_value": 0, "signature": 0, "bases": 0, "doc": 3}, "linghe.facade.add": {"qualname": 0, "fullname": 3, "annotation": 0, "default_value": 0, "signature": 0, "bases": 0, "doc": 12}, "linghe.facade.add.inplace_add": {"qualname": 2, "fullname": 5, "annotation": 0, "default_value": 0, "signature": 36, "bases": 0, "doc": 45}, "linghe.facade.fp32_gemm": {"qualname": 0, "fullname": 4, "annotation": 0, "default_value": 0, "signature": 0, "bases": 0, "doc": 12}, "linghe.facade.fp32_gemm.fp32_gemm": {"qualname": 2, "fullname": 6, "annotation": 0, "default_value": 0, "signature": 36, "bases": 0, "doc": 51}, "linghe.facade.hadamard_quant_linear": {"qualname": 0, "fullname": 5, "annotation": 0, "default_value": 0, "signature": 0, "bases": 0, "doc": 12}, "linghe.facade.hadamard_quant_linear.HadamardQuantLinear": {"qualname": 1, "fullname": 6, "annotation": 0, "default_value": 0, "signature": 0, "bases": 5, "doc": 10}, "linghe.facade.hadamard_quant_linear.HadamardQuantLinear.__init__": {"qualname": 3, "fullname": 8, "annotation": 0, "default_value": 0, "signature": 68, "bases": 0, "doc": 47}, "linghe.facade.loss": {"qualname": 0, "fullname": 3, "annotation": 0, "default_value": 0, "signature": 0, "bases": 0, "doc": 12}, "linghe.facade.loss.softmax_cross_entropy": {"qualname": 3, "fullname": 6, "annotation": 0, "default_value": 0, "signature": 53, "bases": 0, "doc": 62}, "linghe.facade.norm": {"qualname": 0, "fullname": 3, "annotation": 0, "default_value": 0, "signature": 0, "bases": 0, "doc": 12}, "linghe.facade.norm.rms_norm": {"qualname": 2, "fullname": 5, "annotation": 0, "default_value": 0, "signature": 54, "bases": 0, "doc": 48}, "linghe.facade.norm.group_rms_norm_gate": {"qualname": 4, "fullname": 7, "annotation": 0, "default_value": 0, "signature": 93, "bases": 0, "doc": 100}, "linghe.facade.rope": {"qualname": 0, "fullname": 3, "annotation": 0, "default_value": 0, "signature": 0, "bases": 0, "doc": 12}, "linghe.facade.rope.qk_norm_half_rope": {"qualname": 4, "fullname": 7, "annotation": 0, "default_value": 0, "signature": 129, "bases": 0, "doc": 148}, "linghe.facade.smooth_quant_linear": {"qualname": 0, "fullname": 5, "annotation": 0, "default_value": 0, "signature": 0, "bases": 0, "doc": 12}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear": {"qualname": 1, "fullname": 6, "annotation": 0, "default_value": 0, "signature": 0, "bases": 5, "doc": 9}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear.__init__": {"qualname": 3, "fullname": 8, "annotation": 0, "default_value": 0, "signature": 68, "bases": 0, "doc": 47}, "linghe.facade.transpose": {"qualname": 0, "fullname": 3, "annotation": 0, "default_value": 0, "signature": 0, "bases": 0, "doc": 12}, "linghe.facade.transpose.transpose_dim01": {"qualname": 2, "fullname": 5, "annotation": 0, "default_value": 0, "signature": 11, "bases": 0, "doc": 43}, "linghe.gemm": {"qualname": 0, "fullname": 2, "annotation": 0, "default_value": 0, "signature": 0, "bases": 0, "doc": 3}, "linghe.gemm.blockwise_fp8_gemm": {"qualname": 0, "fullname": 5, "annotation": 0, "default_value": 0, "signature": 0, "bases": 0, "doc": 12}, "linghe.gemm.channelwise_fp8_gemm": {"qualname": 0, "fullname": 5, "annotation": 0, "default_value": 0, "signature": 0, "bases": 0, "doc": 12}, "linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"qualname": 3, "fullname": 8, "annotation": 0, "default_value": 0, "signature": 111, "bases": 0, "doc": 102}, "linghe.gemm.fp32_gemm": {"qualname": 0, "fullname": 4, "annotation": 0, "default_value": 0, "signature": 0, "bases": 0, "doc": 12}, "linghe.gemm.fp32_gemm.triton_fp32_gemm": {"qualname": 3, "fullname": 7, "annotation": 0, "default_value": 0, "signature": 36, "bases": 0, "doc": 66}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_backward": {"qualname": 5, "fullname": 9, "annotation": 0, "default_value": 0, "signature": 36, "bases": 0, "doc": 46}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_update": {"qualname": 5, "fullname": 9, "annotation": 0, "default_value": 0, "signature": 36, "bases": 0, "doc": 45}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"qualname": 4, "fullname": 8, "annotation": 0, "default_value": 0, "signature": 51, "bases": 0, "doc": 114}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm_for_update": {"qualname": 6, "fullname": 10, "annotation": 0, "default_value": 0, "signature": 51, "bases": 0, "doc": 45}, "linghe.quant": {"qualname": 0, "fullname": 2, "annotation": 0, "default_value": 0, "signature": 0, "bases": 0, "doc": 3}, "linghe.quant.block": {"qualname": 0, "fullname": 3, "annotation": 0, "default_value": 0, "signature": 0, "bases": 0, "doc": 12}, "linghe.quant.block.triton_block_quant": {"qualname": 3, "fullname": 6, "annotation": 0, "default_value": 0, "signature": 33, "bases": 0, "doc": 64}, "linghe.quant.channel": {"qualname": 0, "fullname": 3, "annotation": 0, "default_value": 0, "signature": 0, "bases": 0, "doc": 12}, "linghe.quant.channel.triton_row_quant": {"qualname": 3, "fullname": 6, "annotation": 0, "default_value": 0, "signature": 22, "bases": 0, "doc": 49}, "linghe.quant.channel.triton_tokenwise_row_quant": {"qualname": 4, "fullname": 7, "annotation": 0, "default_value": 0, "signature": 42, "bases": 0, "doc": 53}, "linghe.quant.channel.triton_transpose_row_quant": {"qualname": 4, "fullname": 7, "annotation": 0, "default_value": 0, "signature": 22, "bases": 0, "doc": 58}, "linghe.quant.group": {"qualname": 0, "fullname": 3, "annotation": 0, "default_value": 0, "signature": 0, "bases": 0, "doc": 12}, "linghe.quant.group.triton_group_quant": {"qualname": 3, "fullname": 6, "annotation": 0, "default_value": 0, "signature": 49, "bases": 0, "doc": 70}, "linghe.quant.hadamard": {"qualname": 0, "fullname": 3, "annotation": 0, "default_value": 0, "signature": 0, "bases": 0, "doc": 12}, "linghe.quant.hadamard.triton_hadamard_quant": {"qualname": 3, "fullname": 6, "annotation": 0, "default_value": 0, "signature": 16, "bases": 0, "doc": 84}, "linghe.quant.smooth": {"qualname": 0, "fullname": 3, "annotation": 0, "default_value": 0, "signature": 0, "bases": 0, "doc": 12}, "linghe.utils": {"qualname": 0, "fullname": 2, "annotation": 0, "default_value": 0, "signature": 0, "bases": 0, "doc": 3}, "linghe.utils.add": {"qualname": 0, "fullname": 3, "annotation": 0, "default_value": 0, "signature": 0, "bases": 0, "doc": 12}, "linghe.utils.add.triton_inplace_add": {"qualname": 3, "fullname": 6, "annotation": 0, "default_value": 0, "signature": 53, "bases": 0, "doc": 53}, "linghe.utils.dot": {"qualname": 0, "fullname": 3, "annotation": 0, "default_value": 0, "signature": 0, "bases": 0, "doc": 12}, "linghe.utils.dot.triton_dot": {"qualname": 2, "fullname": 5, "annotation": 0, "default_value": 0, "signature": 16, "bases": 0, "doc": 52}, "linghe.utils.gather": {"qualname": 0, "fullname": 3, "annotation": 0, "default_value": 0, "signature": 0, "bases": 0, "doc": 12}, "linghe.utils.gather.triton_make_row_id_map": {"qualname": 5, "fullname": 8, "annotation": 0, "default_value": 0, "signature": 40, "bases": 0, "doc": 70}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"qualname": 7, "fullname": 10, "annotation": 0, "default_value": 0, "signature": 52, "bases": 0, "doc": 85}, "linghe.utils.gather.triton_index_select": {"qualname": 3, "fullname": 6, "annotation": 0, "default_value": 0, "signature": 47, "bases": 0, "doc": 53}, "linghe.utils.gather.triton_permute_with_mask_map": {"qualname": 5, "fullname": 8, "annotation": 0, "default_value": 0, "signature": 134, "bases": 0, "doc": 144}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"qualname": 7, "fullname": 10, "annotation": 0, "default_value": 0, "signature": 90, "bases": 0, "doc": 144}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"qualname": 6, "fullname": 9, "annotation": 0, "default_value": 0, "signature": 99, "bases": 0, "doc": 104}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"qualname": 5, "fullname": 8, "annotation": 0, "default_value": 0, "signature": 89, "bases": 0, "doc": 86}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"qualname": 6, "fullname": 9, "annotation": 0, "default_value": 0, "signature": 145, "bases": 0, "doc": 132}, "linghe.utils.loss": {"qualname": 0, "fullname": 3, "annotation": 0, "default_value": 0, "signature": 0, "bases": 0, "doc": 12}, "linghe.utils.loss.triton_softmax_cross_entropy_forward": {"qualname": 5, "fullname": 8, "annotation": 0, "default_value": 0, "signature": 16, "bases": 0, "doc": 43}, "linghe.utils.loss.triton_softmax_cross_entropy_backward": {"qualname": 5, "fullname": 8, "annotation": 0, "default_value": 0, "signature": 45, "bases": 0, "doc": 68}, "linghe.utils.norm": {"qualname": 0, "fullname": 3, "annotation": 0, "default_value": 0, "signature": 0, "bases": 0, "doc": 3}, "linghe.utils.norm.triton_rms_norm_forward": {"qualname": 4, "fullname": 7, "annotation": 0, "default_value": 0, "signature": 37, "bases": 0, "doc": 48}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"qualname": 7, "fullname": 10, "annotation": 0, "default_value": 0, "signature": 182, "bases": 0, "doc": 174}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"qualname": 6, "fullname": 9, "annotation": 0, "default_value": 0, "signature": 89, "bases": 0, "doc": 111}, "linghe.utils.rearange": {"qualname": 0, "fullname": 3, "annotation": 0, "default_value": 0, "signature": 0, "bases": 0, "doc": 12}, "linghe.utils.rearange.triton_split_and_cat": {"qualname": 4, "fullname": 7, "annotation": 0, "default_value": 0, "signature": 31, "bases": 0, "doc": 79}, "linghe.utils.reduce": {"qualname": 0, "fullname": 3, "annotation": 0, "default_value": 0, "signature": 0, "bases": 0, "doc": 12}, "linghe.utils.reduce.triton_abs_max": {"qualname": 3, "fullname": 6, "annotation": 0, "default_value": 0, "signature": 54, "bases": 0, "doc": 84}, "linghe.utils.reduce.triton_batch_count_zero": {"qualname": 4, "fullname": 7, "annotation": 0, "default_value": 0, "signature": 11, "bases": 0, "doc": 44}, "linghe.utils.reduce.triton_batch_sum_with_ord": {"qualname": 5, "fullname": 8, "annotation": 0, "default_value": 0, "signature": 21, "bases": 0, "doc": 47}, "linghe.utils.rope": {"qualname": 0, "fullname": 3, "annotation": 0, "default_value": 0, "signature": 0, "bases": 0, "doc": 12}, "linghe.utils.rope.triton_half_rope_forward": {"qualname": 4, "fullname": 7, "annotation": 0, "default_value": 0, "signature": 31, "bases": 0, "doc": 73}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"qualname": 7, "fullname": 10, "annotation": 0, "default_value": 0, "signature": 90, "bases": 0, "doc": 192}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"qualname": 7, "fullname": 10, "annotation": 0, "default_value": 0, "signature": 86, "bases": 0, "doc": 141}, "linghe.utils.scatter": {"qualname": 0, "fullname": 3, "annotation": 0, "default_value": 0, "signature": 0, "bases": 0, "doc": 12}, "linghe.utils.scatter.triton_aligned_scatter_add": {"qualname": 4, "fullname": 7, "annotation": 0, "default_value": 0, "signature": 83, "bases": 0, "doc": 61}, "linghe.utils.scatter.triton_scatter_add": {"qualname": 3, "fullname": 6, "annotation": 0, "default_value": 0, "signature": 21, "bases": 0, "doc": 47}, "linghe.utils.scatter.triton_unpermute_with_mask_map": {"qualname": 5, "fullname": 8, "annotation": 0, "default_value": 0, "signature": 53, "bases": 0, "doc": 75}, "linghe.utils.silu": {"qualname": 0, "fullname": 3, "annotation": 0, "default_value": 0, "signature": 0, "bases": 0, "doc": 12}, "linghe.utils.silu.triton_weighted_silu_forward": {"qualname": 4, "fullname": 7, "annotation": 0, "default_value": 0, "signature": 31, "bases": 0, "doc": 45}, "linghe.utils.silu.triton_weighted_silu_backward": {"qualname": 4, "fullname": 7, "annotation": 0, "default_value": 0, "signature": 67, "bases": 0, "doc": 59}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"qualname": 6, "fullname": 9, "annotation": 0, "default_value": 0, "signature": 53, "bases": 0, "doc": 104}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"qualname": 6, "fullname": 9, "annotation": 0, "default_value": 0, "signature": 27, "bases": 0, "doc": 87}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"qualname": 8, "fullname": 11, "annotation": 0, "default_value": 0, "signature": 81, "bases": 0, "doc": 139}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"qualname": 8, "fullname": 11, "annotation": 0, "default_value": 0, "signature": 47, "bases": 0, "doc": 129}, "linghe.utils.transpose": {"qualname": 0, "fullname": 3, "annotation": 0, "default_value": 0, "signature": 0, "bases": 0, "doc": 12}, "linghe.utils.transpose.triton_transpose": {"qualname": 2, "fullname": 5, "annotation": 0, "default_value": 0, "signature": 70, "bases": 0, "doc": 47}, "linghe.utils.transpose.triton_transpose_and_pad": {"qualname": 4, "fullname": 7, "annotation": 0, "default_value": 0, "signature": 31, "bases": 0, "doc": 66}, "linghe.utils.transpose.triton_batch_transpose": {"qualname": 3, "fullname": 6, "annotation": 0, "default_value": 0, "signature": 21, "bases": 0, "doc": 37}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"qualname": 5, "fullname": 8, "annotation": 0, "default_value": 0, "signature": 38, "bases": 0, "doc": 70}}, "length": 90, "save": true}, "index": {"qualname": {"root": {"docs": {"linghe.facade.hadamard_quant_linear.HadamardQuantLinear.__init__": {"tf": 1}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear.__init__": {"tf": 1}}, "df": 2, "i": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "p": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "c": {"docs": {}, "df": 0, "e": {"docs": {"linghe.facade.add.inplace_add": {"tf": 1}, "linghe.utils.add.triton_inplace_add": {"tf": 1}}, "df": 2}}}}}, "i": {"docs": {}, "df": 0, "t": {"docs": {"linghe.facade.hadamard_quant_linear.HadamardQuantLinear.__init__": {"tf": 1}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear.__init__": {"tf": 1}}, "df": 2}}, "d": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "c": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "s": {"docs": {"linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1}}, "df": 4}}}}, "e": {"docs": {}, "df": 0, "x": {"docs": {"linghe.utils.gather.triton_index_select": {"tf": 1}}, "df": 1}}}}, "d": {"docs": {"linghe.utils.gather.triton_make_row_id_map": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}}, "df": 2}}, "a": {"docs": {}, "df": 0, "d": {"docs": {}, "df": 0, "d": {"docs": {"linghe.facade.add.inplace_add": {"tf": 1}, "linghe.utils.add.triton_inplace_add": {"tf": 1}, "linghe.utils.scatter.triton_aligned_scatter_add": {"tf": 1}, "linghe.utils.scatter.triton_scatter_add": {"tf": 1}}, "df": 4}}, "n": {"docs": {}, "df": 0, "d": {"docs": {"linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.rearange.triton_split_and_cat": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.transpose.triton_transpose_and_pad": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1}}, "df": 11}}, "b": {"docs": {}, "df": 0, "s": {"docs": {"linghe.utils.reduce.triton_abs_max": {"tf": 1}}, "df": 1}}, "l": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "g": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "d": {"docs": {"linghe.utils.scatter.triton_aligned_scatter_add": {"tf": 1}}, "df": 1}}}}}}}, "f": {"docs": {}, "df": 0, "p": {"3": {"2": {"docs": {"linghe.facade.fp32_gemm.fp32_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_backward": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_update": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm_for_update": {"tf": 1}}, "df": 6}, "docs": {}, "df": 0}, "docs": {}, "df": 0}, "o": {"docs": {}, "df": 0, "r": {"docs": {"linghe.gemm.fp32_gemm.triton_fp32_gemm_for_backward": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_update": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm_for_update": {"tf": 1}}, "df": 3, "w": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "d": {"docs": {"linghe.utils.loss.triton_softmax_cross_entropy_forward": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_forward": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1}, "linghe.utils.rope.triton_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}, "linghe.utils.silu.triton_weighted_silu_forward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}}, "df": 9}}}}}}}, "g": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "m": {"docs": {}, "df": 0, "m": {"docs": {"linghe.facade.fp32_gemm.fp32_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_backward": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_update": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm_for_update": {"tf": 1}}, "df": 6}}}, "r": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "p": {"docs": {"linghe.facade.norm.group_rms_norm_gate": {"tf": 1}, "linghe.quant.group.triton_group_quant": {"tf": 1}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1}}, "df": 3}}}}, "a": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "e": {"docs": {"linghe.facade.norm.group_rms_norm_gate": {"tf": 1}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1}}, "df": 2}}}}, "h": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "d": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "m": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "d": {"docs": {"linghe.quant.hadamard.triton_hadamard_quant": {"tf": 1}}, "df": 1, "q": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "r": {"docs": {"linghe.facade.hadamard_quant_linear.HadamardQuantLinear": {"tf": 1}, "linghe.facade.hadamard_quant_linear.HadamardQuantLinear.__init__": {"tf": 1}}, "df": 2}}}}}}}}}}}}}}}}}, "l": {"docs": {}, "df": 0, "f": {"docs": {"linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.utils.rope.triton_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}}, "df": 4}}}}, "s": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "f": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "m": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "x": {"docs": {"linghe.facade.loss.softmax_cross_entropy": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_forward": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_backward": {"tf": 1}}, "df": 3}}}}}}, "m": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "h": {"docs": {"linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}}, "df": 4, "q": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "r": {"docs": {"linghe.facade.smooth_quant_linear.SmoothQuantLinear": {"tf": 1}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear.__init__": {"tf": 1}}, "df": 2}}}}}}}}}}}}}}}}, "c": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "d": {"docs": {"linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm_for_update": {"tf": 1}}, "df": 3}}}, "t": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "r": {"docs": {"linghe.utils.scatter.triton_aligned_scatter_add": {"tf": 1}, "linghe.utils.scatter.triton_scatter_add": {"tf": 1}}, "df": 2}}}}}}, "e": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "c": {"docs": {}, "df": 0, "t": {"docs": {"linghe.utils.gather.triton_index_select": {"tf": 1}}, "df": 1}}}}}, "p": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "t": {"docs": {"linghe.utils.rearange.triton_split_and_cat": {"tf": 1}}, "df": 1}}}}, "u": {"docs": {}, "df": 0, "m": {"docs": {"linghe.utils.reduce.triton_batch_sum_with_ord": {"tf": 1}}, "df": 1}}, "i": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "u": {"docs": {"linghe.utils.silu.triton_weighted_silu_forward": {"tf": 1}, "linghe.utils.silu.triton_weighted_silu_backward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}}, "df": 6}}}}, "c": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "s": {"docs": {}, "df": 0, "s": {"docs": {"linghe.facade.loss.softmax_cross_entropy": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_forward": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_backward": {"tf": 1}}, "df": 3}}}}, "a": {"docs": {}, "df": 0, "t": {"docs": {"linghe.utils.rearange.triton_split_and_cat": {"tf": 1}}, "df": 1}}, "o": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "t": {"docs": {"linghe.utils.reduce.triton_batch_count_zero": {"tf": 1}}, "df": 1}}}}}, "e": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "p": {"docs": {}, "df": 0, "y": {"docs": {"linghe.facade.loss.softmax_cross_entropy": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_forward": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_backward": {"tf": 1}}, "df": 3}}}}}}}, "r": {"docs": {}, "df": 0, "m": {"docs": {}, "df": 0, "s": {"docs": {"linghe.facade.norm.rms_norm": {"tf": 1}, "linghe.facade.norm.group_rms_norm_gate": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_forward": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1}}, "df": 5}}, "o": {"docs": {}, "df": 0, "p": {"docs": {}, "df": 0, "e": {"docs": {"linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.utils.rope.triton_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}}, "df": 4}}, "w": {"docs": {"linghe.quant.channel.triton_row_quant": {"tf": 1}, "linghe.quant.channel.triton_tokenwise_row_quant": {"tf": 1}, "linghe.quant.channel.triton_transpose_row_quant": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}}, "df": 5}}}, "n": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "m": {"docs": {"linghe.facade.norm.rms_norm": {"tf": 1}, "linghe.facade.norm.group_rms_norm_gate": {"tf": 1}, "linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_forward": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}}, "df": 8}}}}, "q": {"docs": {}, "df": 0, "k": {"docs": {"linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}}, "df": 3}, "u": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "t": {"docs": {"linghe.quant.block.triton_block_quant": {"tf": 1}, "linghe.quant.channel.triton_row_quant": {"tf": 1}, "linghe.quant.channel.triton_tokenwise_row_quant": {"tf": 1}, "linghe.quant.channel.triton_transpose_row_quant": {"tf": 1}, "linghe.quant.group.triton_group_quant": {"tf": 1}, "linghe.quant.hadamard.triton_hadamard_quant": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}}, "df": 11}}}}}, "t": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "s": {"docs": {}, "df": 0, "p": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "s": {"docs": {}, "df": 0, "e": {"docs": {"linghe.facade.transpose.transpose_dim01": {"tf": 1}, "linghe.quant.channel.triton_transpose_row_quant": {"tf": 1}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.transpose.triton_transpose": {"tf": 1}, "linghe.utils.transpose.triton_transpose_and_pad": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1}}, "df": 7}}}}}}}, "i": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "n": {"docs": {"linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_backward": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_update": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm_for_update": {"tf": 1}, "linghe.quant.block.triton_block_quant": {"tf": 1}, "linghe.quant.channel.triton_row_quant": {"tf": 1}, "linghe.quant.channel.triton_tokenwise_row_quant": {"tf": 1}, "linghe.quant.channel.triton_transpose_row_quant": {"tf": 1}, "linghe.quant.group.triton_group_quant": {"tf": 1}, "linghe.quant.hadamard.triton_hadamard_quant": {"tf": 1}, "linghe.utils.add.triton_inplace_add": {"tf": 1}, "linghe.utils.dot.triton_dot": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}, "linghe.utils.gather.triton_index_select": {"tf": 1}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_forward": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_backward": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_forward": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1}, "linghe.utils.rearange.triton_split_and_cat": {"tf": 1}, "linghe.utils.reduce.triton_abs_max": {"tf": 1}, "linghe.utils.reduce.triton_batch_count_zero": {"tf": 1}, "linghe.utils.reduce.triton_batch_sum_with_ord": {"tf": 1}, "linghe.utils.rope.triton_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}, "linghe.utils.scatter.triton_aligned_scatter_add": {"tf": 1}, "linghe.utils.scatter.triton_scatter_add": {"tf": 1}, "linghe.utils.scatter.triton_unpermute_with_mask_map": {"tf": 1}, "linghe.utils.silu.triton_weighted_silu_forward": {"tf": 1}, "linghe.utils.silu.triton_weighted_silu_backward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.transpose.triton_transpose": {"tf": 1}, "linghe.utils.transpose.triton_transpose_and_pad": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1}}, "df": 47}}}}}, "o": {"docs": {}, "df": 0, "k": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "w": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "s": {"docs": {}, "df": 0, "e": {"docs": {"linghe.quant.channel.triton_tokenwise_row_quant": {"tf": 1}}, "df": 1}}}}}}}}}, "d": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "m": {"0": {"1": {"docs": {"linghe.facade.transpose.transpose_dim01": {"tf": 1}}, "df": 1}, "docs": {}, "df": 0}, "docs": {}, "df": 0}}, "o": {"docs": {}, "df": 0, "t": {"docs": {"linghe.utils.dot.triton_dot": {"tf": 1}}, "df": 1}}}, "m": {"docs": {}, "df": 0, "m": {"docs": {"linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1}}, "df": 1}, "a": {"docs": {}, "df": 0, "k": {"docs": {}, "df": 0, "e": {"docs": {"linghe.utils.gather.triton_make_row_id_map": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}}, "df": 2}}, "p": {"docs": {"linghe.utils.gather.triton_make_row_id_map": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}, "linghe.utils.scatter.triton_unpermute_with_mask_map": {"tf": 1}}, "df": 5}, "s": {"docs": {}, "df": 0, "k": {"docs": {"linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}, "linghe.utils.scatter.triton_unpermute_with_mask_map": {"tf": 1}}, "df": 3}}, "x": {"docs": {"linghe.utils.reduce.triton_abs_max": {"tf": 1}}, "df": 1}}}, "b": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "c": {"docs": {}, "df": 0, "k": {"docs": {}, "df": 0, "w": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "d": {"docs": {"linghe.gemm.fp32_gemm.triton_fp32_gemm_for_backward": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_backward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}, "linghe.utils.silu.triton_weighted_silu_backward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}}, "df": 6}}}}}}, "t": {"docs": {}, "df": 0, "c": {"docs": {}, "df": 0, "h": {"docs": {"linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.reduce.triton_batch_count_zero": {"tf": 1}, "linghe.utils.reduce.triton_batch_sum_with_ord": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1}}, "df": 7}}}}, "l": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "c": {"docs": {}, "df": 0, "k": {"docs": {"linghe.quant.block.triton_block_quant": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}}, "df": 6}}}}}, "u": {"docs": {}, "df": 0, "p": {"docs": {}, "df": 0, "d": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "e": {"docs": {"linghe.gemm.fp32_gemm.triton_fp32_gemm_for_update": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm_for_update": {"tf": 1}}, "df": 2}}}}}, "n": {"docs": {}, "df": 0, "p": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "m": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "e": {"docs": {"linghe.utils.scatter.triton_unpermute_with_mask_map": {"tf": 1}}, "df": 1}}}}}}}}}, "p": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "m": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "e": {"docs": {"linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}}, "df": 5}}}}}}, "a": {"docs": {}, "df": 0, "d": {"docs": {"linghe.utils.transpose.triton_transpose_and_pad": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1}}, "df": 2}}}, "w": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "h": {"docs": {"linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}, "linghe.utils.reduce.triton_batch_sum_with_ord": {"tf": 1}, "linghe.utils.scatter.triton_unpermute_with_mask_map": {"tf": 1}}, "df": 7}}}, "e": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "g": {"docs": {}, "df": 0, "h": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "d": {"docs": {"linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.silu.triton_weighted_silu_forward": {"tf": 1}, "linghe.utils.silu.triton_weighted_silu_backward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}}, "df": 5}}}}}}}}, "z": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "o": {"docs": {"linghe.utils.reduce.triton_batch_count_zero": {"tf": 1}}, "df": 1}}}}, "o": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "d": {"docs": {"linghe.utils.reduce.triton_batch_sum_with_ord": {"tf": 1}}, "df": 1}}}}}, "fullname": {"root": {"docs": {"linghe.facade.hadamard_quant_linear.HadamardQuantLinear.__init__": {"tf": 1}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear.__init__": {"tf": 1}}, "df": 2, "l": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "g": {"docs": {}, "df": 0, "h": {"docs": {}, "df": 0, "e": {"docs": {"linghe": {"tf": 1}, "linghe.facade": {"tf": 1}, "linghe.facade.add": {"tf": 1}, "linghe.facade.add.inplace_add": {"tf": 1}, "linghe.facade.fp32_gemm": {"tf": 1}, "linghe.facade.fp32_gemm.fp32_gemm": {"tf": 1}, "linghe.facade.hadamard_quant_linear": {"tf": 1}, "linghe.facade.hadamard_quant_linear.HadamardQuantLinear": {"tf": 1}, "linghe.facade.hadamard_quant_linear.HadamardQuantLinear.__init__": {"tf": 1}, "linghe.facade.loss": {"tf": 1}, "linghe.facade.loss.softmax_cross_entropy": {"tf": 1}, "linghe.facade.norm": {"tf": 1}, "linghe.facade.norm.rms_norm": {"tf": 1}, "linghe.facade.norm.group_rms_norm_gate": {"tf": 1}, "linghe.facade.rope": {"tf": 1}, "linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.facade.smooth_quant_linear": {"tf": 1}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear": {"tf": 1}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear.__init__": {"tf": 1}, "linghe.facade.transpose": {"tf": 1}, "linghe.facade.transpose.transpose_dim01": {"tf": 1}, "linghe.gemm": {"tf": 1}, "linghe.gemm.blockwise_fp8_gemm": {"tf": 1}, "linghe.gemm.channelwise_fp8_gemm": {"tf": 1}, "linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1}, "linghe.gemm.fp32_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_backward": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_update": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm_for_update": {"tf": 1}, "linghe.quant": {"tf": 1}, "linghe.quant.block": {"tf": 1}, "linghe.quant.block.triton_block_quant": {"tf": 1}, "linghe.quant.channel": {"tf": 1}, "linghe.quant.channel.triton_row_quant": {"tf": 1}, "linghe.quant.channel.triton_tokenwise_row_quant": {"tf": 1}, "linghe.quant.channel.triton_transpose_row_quant": {"tf": 1}, "linghe.quant.group": {"tf": 1}, "linghe.quant.group.triton_group_quant": {"tf": 1}, "linghe.quant.hadamard": {"tf": 1}, "linghe.quant.hadamard.triton_hadamard_quant": {"tf": 1}, "linghe.quant.smooth": {"tf": 1}, "linghe.utils": {"tf": 1}, "linghe.utils.add": {"tf": 1}, "linghe.utils.add.triton_inplace_add": {"tf": 1}, "linghe.utils.dot": {"tf": 1}, "linghe.utils.dot.triton_dot": {"tf": 1}, "linghe.utils.gather": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}, "linghe.utils.gather.triton_index_select": {"tf": 1}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}, "linghe.utils.loss": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_forward": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_backward": {"tf": 1}, "linghe.utils.norm": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_forward": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1}, "linghe.utils.rearange": {"tf": 1}, "linghe.utils.rearange.triton_split_and_cat": {"tf": 1}, "linghe.utils.reduce": {"tf": 1}, "linghe.utils.reduce.triton_abs_max": {"tf": 1}, "linghe.utils.reduce.triton_batch_count_zero": {"tf": 1}, "linghe.utils.reduce.triton_batch_sum_with_ord": {"tf": 1}, "linghe.utils.rope": {"tf": 1}, "linghe.utils.rope.triton_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}, "linghe.utils.scatter": {"tf": 1}, "linghe.utils.scatter.triton_aligned_scatter_add": {"tf": 1}, "linghe.utils.scatter.triton_scatter_add": {"tf": 1}, "linghe.utils.scatter.triton_unpermute_with_mask_map": {"tf": 1}, "linghe.utils.silu": {"tf": 1}, "linghe.utils.silu.triton_weighted_silu_forward": {"tf": 1}, "linghe.utils.silu.triton_weighted_silu_backward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.transpose": {"tf": 1}, "linghe.utils.transpose.triton_transpose": {"tf": 1}, "linghe.utils.transpose.triton_transpose_and_pad": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1}}, "df": 90}}}, "e": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "r": {"docs": {"linghe.facade.hadamard_quant_linear": {"tf": 1}, "linghe.facade.hadamard_quant_linear.HadamardQuantLinear": {"tf": 1}, "linghe.facade.hadamard_quant_linear.HadamardQuantLinear.__init__": {"tf": 1}, "linghe.facade.smooth_quant_linear": {"tf": 1}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear": {"tf": 1}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear.__init__": {"tf": 1}}, "df": 6}}}}}, "o": {"docs": {}, "df": 0, "s": {"docs": {}, "df": 0, "s": {"docs": {"linghe.facade.loss": {"tf": 1}, "linghe.facade.loss.softmax_cross_entropy": {"tf": 1}, "linghe.utils.loss": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_forward": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_backward": {"tf": 1}}, "df": 5}}}}, "f": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "c": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "d": {"docs": {}, "df": 0, "e": {"docs": {"linghe.facade": {"tf": 1}, "linghe.facade.add": {"tf": 1}, "linghe.facade.add.inplace_add": {"tf": 1}, "linghe.facade.fp32_gemm": {"tf": 1}, "linghe.facade.fp32_gemm.fp32_gemm": {"tf": 1}, "linghe.facade.hadamard_quant_linear": {"tf": 1}, "linghe.facade.hadamard_quant_linear.HadamardQuantLinear": {"tf": 1}, "linghe.facade.hadamard_quant_linear.HadamardQuantLinear.__init__": {"tf": 1}, "linghe.facade.loss": {"tf": 1}, "linghe.facade.loss.softmax_cross_entropy": {"tf": 1}, "linghe.facade.norm": {"tf": 1}, "linghe.facade.norm.rms_norm": {"tf": 1}, "linghe.facade.norm.group_rms_norm_gate": {"tf": 1}, "linghe.facade.rope": {"tf": 1}, "linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.facade.smooth_quant_linear": {"tf": 1}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear": {"tf": 1}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear.__init__": {"tf": 1}, "linghe.facade.transpose": {"tf": 1}, "linghe.facade.transpose.transpose_dim01": {"tf": 1}}, "df": 20}}}}}, "p": {"3": {"2": {"docs": {"linghe.facade.fp32_gemm": {"tf": 1}, "linghe.facade.fp32_gemm.fp32_gemm": {"tf": 1.4142135623730951}, "linghe.gemm.fp32_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm": {"tf": 1.4142135623730951}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_backward": {"tf": 1.4142135623730951}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_update": {"tf": 1.4142135623730951}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1.4142135623730951}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm_for_update": {"tf": 1.4142135623730951}}, "df": 8}, "docs": {}, "df": 0}, "8": {"docs": {"linghe.gemm.blockwise_fp8_gemm": {"tf": 1}, "linghe.gemm.channelwise_fp8_gemm": {"tf": 1}, "linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1}}, "df": 3}, "docs": {}, "df": 0}, "o": {"docs": {}, "df": 0, "r": {"docs": {"linghe.gemm.fp32_gemm.triton_fp32_gemm_for_backward": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_update": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm_for_update": {"tf": 1}}, "df": 3, "w": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "d": {"docs": {"linghe.utils.loss.triton_softmax_cross_entropy_forward": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_forward": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1}, "linghe.utils.rope.triton_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}, "linghe.utils.silu.triton_weighted_silu_forward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}}, "df": 9}}}}}}}, "a": {"docs": {}, "df": 0, "d": {"docs": {}, "df": 0, "d": {"docs": {"linghe.facade.add": {"tf": 1}, "linghe.facade.add.inplace_add": {"tf": 1.4142135623730951}, "linghe.utils.add": {"tf": 1}, "linghe.utils.add.triton_inplace_add": {"tf": 1.4142135623730951}, "linghe.utils.scatter.triton_aligned_scatter_add": {"tf": 1}, "linghe.utils.scatter.triton_scatter_add": {"tf": 1}}, "df": 6}}, "n": {"docs": {}, "df": 0, "d": {"docs": {"linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.rearange.triton_split_and_cat": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.transpose.triton_transpose_and_pad": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1}}, "df": 11}}, "b": {"docs": {}, "df": 0, "s": {"docs": {"linghe.utils.reduce.triton_abs_max": {"tf": 1}}, "df": 1}}, "l": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "g": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "d": {"docs": {"linghe.utils.scatter.triton_aligned_scatter_add": {"tf": 1}}, "df": 1}}}}}}}, "i": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "p": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "c": {"docs": {}, "df": 0, "e": {"docs": {"linghe.facade.add.inplace_add": {"tf": 1}, "linghe.utils.add.triton_inplace_add": {"tf": 1}}, "df": 2}}}}}, "i": {"docs": {}, "df": 0, "t": {"docs": {"linghe.facade.hadamard_quant_linear.HadamardQuantLinear.__init__": {"tf": 1}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear.__init__": {"tf": 1}}, "df": 2}}, "d": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "c": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "s": {"docs": {"linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1}}, "df": 4}}}}, "e": {"docs": {}, "df": 0, "x": {"docs": {"linghe.utils.gather.triton_index_select": {"tf": 1}}, "df": 1}}}}, "d": {"docs": {"linghe.utils.gather.triton_make_row_id_map": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}}, "df": 2}}, "g": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "m": {"docs": {}, "df": 0, "m": {"docs": {"linghe.facade.fp32_gemm": {"tf": 1}, "linghe.facade.fp32_gemm.fp32_gemm": {"tf": 1.4142135623730951}, "linghe.gemm": {"tf": 1}, "linghe.gemm.blockwise_fp8_gemm": {"tf": 1.4142135623730951}, "linghe.gemm.channelwise_fp8_gemm": {"tf": 1.4142135623730951}, "linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1.4142135623730951}, "linghe.gemm.fp32_gemm": {"tf": 1.4142135623730951}, "linghe.gemm.fp32_gemm.triton_fp32_gemm": {"tf": 1.7320508075688772}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_backward": {"tf": 1.7320508075688772}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_update": {"tf": 1.7320508075688772}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1.7320508075688772}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm_for_update": {"tf": 1.7320508075688772}}, "df": 12}}}, "r": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "p": {"docs": {"linghe.facade.norm.group_rms_norm_gate": {"tf": 1}, "linghe.quant.group": {"tf": 1}, "linghe.quant.group.triton_group_quant": {"tf": 1.4142135623730951}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1}}, "df": 4}}}}, "a": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "e": {"docs": {"linghe.facade.norm.group_rms_norm_gate": {"tf": 1}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1}}, "df": 2}, "h": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "r": {"docs": {"linghe.utils.gather": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}, "linghe.utils.gather.triton_index_select": {"tf": 1}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}}, "df": 9}}}}}}, "h": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "d": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "m": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "d": {"docs": {"linghe.facade.hadamard_quant_linear": {"tf": 1}, "linghe.facade.hadamard_quant_linear.HadamardQuantLinear": {"tf": 1}, "linghe.facade.hadamard_quant_linear.HadamardQuantLinear.__init__": {"tf": 1}, "linghe.quant.hadamard": {"tf": 1}, "linghe.quant.hadamard.triton_hadamard_quant": {"tf": 1.4142135623730951}}, "df": 5, "q": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "r": {"docs": {"linghe.facade.hadamard_quant_linear.HadamardQuantLinear": {"tf": 1}, "linghe.facade.hadamard_quant_linear.HadamardQuantLinear.__init__": {"tf": 1}}, "df": 2}}}}}}}}}}}}}}}}}, "l": {"docs": {}, "df": 0, "f": {"docs": {"linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.utils.rope.triton_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}}, "df": 4}}}}, "q": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "t": {"docs": {"linghe.facade.hadamard_quant_linear": {"tf": 1}, "linghe.facade.hadamard_quant_linear.HadamardQuantLinear": {"tf": 1}, "linghe.facade.hadamard_quant_linear.HadamardQuantLinear.__init__": {"tf": 1}, "linghe.facade.smooth_quant_linear": {"tf": 1}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear": {"tf": 1}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear.__init__": {"tf": 1}, "linghe.quant": {"tf": 1}, "linghe.quant.block": {"tf": 1}, "linghe.quant.block.triton_block_quant": {"tf": 1.4142135623730951}, "linghe.quant.channel": {"tf": 1}, "linghe.quant.channel.triton_row_quant": {"tf": 1.4142135623730951}, "linghe.quant.channel.triton_tokenwise_row_quant": {"tf": 1.4142135623730951}, "linghe.quant.channel.triton_transpose_row_quant": {"tf": 1.4142135623730951}, "linghe.quant.group": {"tf": 1}, "linghe.quant.group.triton_group_quant": {"tf": 1.4142135623730951}, "linghe.quant.hadamard": {"tf": 1}, "linghe.quant.hadamard.triton_hadamard_quant": {"tf": 1.4142135623730951}, "linghe.quant.smooth": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}}, "df": 23}}}}, "k": {"docs": {"linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}}, "df": 3}}, "s": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "f": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "m": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "x": {"docs": {"linghe.facade.loss.softmax_cross_entropy": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_forward": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_backward": {"tf": 1}}, "df": 3}}}}}}, "m": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "h": {"docs": {"linghe.facade.smooth_quant_linear": {"tf": 1}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear": {"tf": 1}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear.__init__": {"tf": 1}, "linghe.quant.smooth": {"tf": 1}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}}, "df": 8, "q": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "r": {"docs": {"linghe.facade.smooth_quant_linear.SmoothQuantLinear": {"tf": 1}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear.__init__": {"tf": 1}}, "df": 2}}}}}}}}}}}}}}}}, "c": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "d": {"docs": {"linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm_for_update": {"tf": 1}}, "df": 3}}}, "t": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "r": {"docs": {"linghe.utils.scatter": {"tf": 1}, "linghe.utils.scatter.triton_aligned_scatter_add": {"tf": 1.4142135623730951}, "linghe.utils.scatter.triton_scatter_add": {"tf": 1.4142135623730951}, "linghe.utils.scatter.triton_unpermute_with_mask_map": {"tf": 1}}, "df": 4}}}}}}, "e": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "c": {"docs": {}, "df": 0, "t": {"docs": {"linghe.utils.gather.triton_index_select": {"tf": 1}}, "df": 1}}}}}, "p": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "t": {"docs": {"linghe.utils.rearange.triton_split_and_cat": {"tf": 1}}, "df": 1}}}}, "u": {"docs": {}, "df": 0, "m": {"docs": {"linghe.utils.reduce.triton_batch_sum_with_ord": {"tf": 1}}, "df": 1}}, "i": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "u": {"docs": {"linghe.utils.silu": {"tf": 1}, "linghe.utils.silu.triton_weighted_silu_forward": {"tf": 1.4142135623730951}, "linghe.utils.silu.triton_weighted_silu_backward": {"tf": 1.4142135623730951}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1.4142135623730951}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 1.4142135623730951}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1.4142135623730951}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1.4142135623730951}}, "df": 7}}}}, "c": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "s": {"docs": {}, "df": 0, "s": {"docs": {"linghe.facade.loss.softmax_cross_entropy": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_forward": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_backward": {"tf": 1}}, "df": 3}}}}, "h": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "l": {"docs": {"linghe.quant.channel": {"tf": 1}, "linghe.quant.channel.triton_row_quant": {"tf": 1}, "linghe.quant.channel.triton_tokenwise_row_quant": {"tf": 1}, "linghe.quant.channel.triton_transpose_row_quant": {"tf": 1}}, "df": 4, "w": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "s": {"docs": {}, "df": 0, "e": {"docs": {"linghe.gemm.channelwise_fp8_gemm": {"tf": 1}, "linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1}}, "df": 2}}}}}}}}}}, "a": {"docs": {}, "df": 0, "t": {"docs": {"linghe.utils.rearange.triton_split_and_cat": {"tf": 1}}, "df": 1}}, "o": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "t": {"docs": {"linghe.utils.reduce.triton_batch_count_zero": {"tf": 1}}, "df": 1}}}}}, "e": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "p": {"docs": {}, "df": 0, "y": {"docs": {"linghe.facade.loss.softmax_cross_entropy": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_forward": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_backward": {"tf": 1}}, "df": 3}}}}}}}, "n": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "m": {"docs": {"linghe.facade.norm": {"tf": 1}, "linghe.facade.norm.rms_norm": {"tf": 1.4142135623730951}, "linghe.facade.norm.group_rms_norm_gate": {"tf": 1.4142135623730951}, "linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.utils.norm": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_forward": {"tf": 1.4142135623730951}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1.4142135623730951}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1.4142135623730951}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}}, "df": 10}}}}, "r": {"docs": {}, "df": 0, "m": {"docs": {}, "df": 0, "s": {"docs": {"linghe.facade.norm.rms_norm": {"tf": 1}, "linghe.facade.norm.group_rms_norm_gate": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_forward": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1}}, "df": 5}}, "o": {"docs": {}, "df": 0, "p": {"docs": {}, "df": 0, "e": {"docs": {"linghe.facade.rope": {"tf": 1}, "linghe.facade.rope.qk_norm_half_rope": {"tf": 1.4142135623730951}, "linghe.utils.rope": {"tf": 1}, "linghe.utils.rope.triton_half_rope_forward": {"tf": 1.4142135623730951}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1.4142135623730951}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1.4142135623730951}}, "df": 6}}, "w": {"docs": {"linghe.quant.channel.triton_row_quant": {"tf": 1}, "linghe.quant.channel.triton_tokenwise_row_quant": {"tf": 1}, "linghe.quant.channel.triton_transpose_row_quant": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}}, "df": 5}}, "e": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "g": {"docs": {}, "df": 0, "e": {"docs": {"linghe.utils.rearange": {"tf": 1}, "linghe.utils.rearange.triton_split_and_cat": {"tf": 1}}, "df": 2}}}}}}, "d": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "c": {"docs": {}, "df": 0, "e": {"docs": {"linghe.utils.reduce": {"tf": 1}, "linghe.utils.reduce.triton_abs_max": {"tf": 1}, "linghe.utils.reduce.triton_batch_count_zero": {"tf": 1}, "linghe.utils.reduce.triton_batch_sum_with_ord": {"tf": 1}}, "df": 4}}}}}}, "t": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "s": {"docs": {}, "df": 0, "p": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "s": {"docs": {}, "df": 0, "e": {"docs": {"linghe.facade.transpose": {"tf": 1}, "linghe.facade.transpose.transpose_dim01": {"tf": 1.4142135623730951}, "linghe.quant.channel.triton_transpose_row_quant": {"tf": 1}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.transpose": {"tf": 1}, "linghe.utils.transpose.triton_transpose": {"tf": 1.4142135623730951}, "linghe.utils.transpose.triton_transpose_and_pad": {"tf": 1.4142135623730951}, "linghe.utils.transpose.triton_batch_transpose": {"tf": 1.4142135623730951}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1.4142135623730951}}, "df": 9}}}}}}}, "i": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "n": {"docs": {"linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_backward": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_update": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm_for_update": {"tf": 1}, "linghe.quant.block.triton_block_quant": {"tf": 1}, "linghe.quant.channel.triton_row_quant": {"tf": 1}, "linghe.quant.channel.triton_tokenwise_row_quant": {"tf": 1}, "linghe.quant.channel.triton_transpose_row_quant": {"tf": 1}, "linghe.quant.group.triton_group_quant": {"tf": 1}, "linghe.quant.hadamard.triton_hadamard_quant": {"tf": 1}, "linghe.utils.add.triton_inplace_add": {"tf": 1}, "linghe.utils.dot.triton_dot": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}, "linghe.utils.gather.triton_index_select": {"tf": 1}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_forward": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_backward": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_forward": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1}, "linghe.utils.rearange.triton_split_and_cat": {"tf": 1}, "linghe.utils.reduce.triton_abs_max": {"tf": 1}, "linghe.utils.reduce.triton_batch_count_zero": {"tf": 1}, "linghe.utils.reduce.triton_batch_sum_with_ord": {"tf": 1}, "linghe.utils.rope.triton_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}, "linghe.utils.scatter.triton_aligned_scatter_add": {"tf": 1}, "linghe.utils.scatter.triton_scatter_add": {"tf": 1}, "linghe.utils.scatter.triton_unpermute_with_mask_map": {"tf": 1}, "linghe.utils.silu.triton_weighted_silu_forward": {"tf": 1}, "linghe.utils.silu.triton_weighted_silu_backward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.transpose.triton_transpose": {"tf": 1}, "linghe.utils.transpose.triton_transpose_and_pad": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1}}, "df": 47}}}}}, "o": {"docs": {}, "df": 0, "k": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "w": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "s": {"docs": {}, "df": 0, "e": {"docs": {"linghe.quant.channel.triton_tokenwise_row_quant": {"tf": 1}}, "df": 1}}}}}}}}}, "d": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "m": {"0": {"1": {"docs": {"linghe.facade.transpose.transpose_dim01": {"tf": 1}}, "df": 1}, "docs": {}, "df": 0}, "docs": {}, "df": 0}}, "o": {"docs": {}, "df": 0, "t": {"docs": {"linghe.utils.dot": {"tf": 1}, "linghe.utils.dot.triton_dot": {"tf": 1.4142135623730951}}, "df": 2}}}, "b": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "c": {"docs": {}, "df": 0, "k": {"docs": {"linghe.quant.block": {"tf": 1}, "linghe.quant.block.triton_block_quant": {"tf": 1.4142135623730951}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}}, "df": 7, "w": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "s": {"docs": {}, "df": 0, "e": {"docs": {"linghe.gemm.blockwise_fp8_gemm": {"tf": 1}}, "df": 1}}}}}}}}, "a": {"docs": {}, "df": 0, "c": {"docs": {}, "df": 0, "k": {"docs": {}, "df": 0, "w": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "d": {"docs": {"linghe.gemm.fp32_gemm.triton_fp32_gemm_for_backward": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_backward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}, "linghe.utils.silu.triton_weighted_silu_backward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}}, "df": 6}}}}}}, "t": {"docs": {}, "df": 0, "c": {"docs": {}, "df": 0, "h": {"docs": {"linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.reduce.triton_batch_count_zero": {"tf": 1}, "linghe.utils.reduce.triton_batch_sum_with_ord": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1}}, "df": 7}}}}}, "m": {"docs": {}, "df": 0, "m": {"docs": {"linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1}}, "df": 1}, "a": {"docs": {}, "df": 0, "k": {"docs": {}, "df": 0, "e": {"docs": {"linghe.utils.gather.triton_make_row_id_map": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}}, "df": 2}}, "p": {"docs": {"linghe.utils.gather.triton_make_row_id_map": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}, "linghe.utils.scatter.triton_unpermute_with_mask_map": {"tf": 1}}, "df": 5}, "s": {"docs": {}, "df": 0, "k": {"docs": {"linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}, "linghe.utils.scatter.triton_unpermute_with_mask_map": {"tf": 1}}, "df": 3}}, "x": {"docs": {"linghe.utils.reduce.triton_abs_max": {"tf": 1}}, "df": 1}}}, "u": {"docs": {}, "df": 0, "p": {"docs": {}, "df": 0, "d": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "e": {"docs": {"linghe.gemm.fp32_gemm.triton_fp32_gemm_for_update": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm_for_update": {"tf": 1}}, "df": 2}}}}}, "t": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "s": {"docs": {"linghe.utils": {"tf": 1}, "linghe.utils.add": {"tf": 1}, "linghe.utils.add.triton_inplace_add": {"tf": 1}, "linghe.utils.dot": {"tf": 1}, "linghe.utils.dot.triton_dot": {"tf": 1}, "linghe.utils.gather": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}, "linghe.utils.gather.triton_index_select": {"tf": 1}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}, "linghe.utils.loss": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_forward": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_backward": {"tf": 1}, "linghe.utils.norm": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_forward": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1}, "linghe.utils.rearange": {"tf": 1}, "linghe.utils.rearange.triton_split_and_cat": {"tf": 1}, "linghe.utils.reduce": {"tf": 1}, "linghe.utils.reduce.triton_abs_max": {"tf": 1}, "linghe.utils.reduce.triton_batch_count_zero": {"tf": 1}, "linghe.utils.reduce.triton_batch_sum_with_ord": {"tf": 1}, "linghe.utils.rope": {"tf": 1}, "linghe.utils.rope.triton_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}, "linghe.utils.scatter": {"tf": 1}, "linghe.utils.scatter.triton_aligned_scatter_add": {"tf": 1}, "linghe.utils.scatter.triton_scatter_add": {"tf": 1}, "linghe.utils.scatter.triton_unpermute_with_mask_map": {"tf": 1}, "linghe.utils.silu": {"tf": 1}, "linghe.utils.silu.triton_weighted_silu_forward": {"tf": 1}, "linghe.utils.silu.triton_weighted_silu_backward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.transpose": {"tf": 1}, "linghe.utils.transpose.triton_transpose": {"tf": 1}, "linghe.utils.transpose.triton_transpose_and_pad": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1}}, "df": 47}}}}, "n": {"docs": {}, "df": 0, "p": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "m": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "e": {"docs": {"linghe.utils.scatter.triton_unpermute_with_mask_map": {"tf": 1}}, "df": 1}}}}}}}}}, "p": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "m": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "e": {"docs": {"linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}}, "df": 5}}}}}}, "a": {"docs": {}, "df": 0, "d": {"docs": {"linghe.utils.transpose.triton_transpose_and_pad": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1}}, "df": 2}}}, "w": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "h": {"docs": {"linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}, "linghe.utils.reduce.triton_batch_sum_with_ord": {"tf": 1}, "linghe.utils.scatter.triton_unpermute_with_mask_map": {"tf": 1}}, "df": 7}}}, "e": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "g": {"docs": {}, "df": 0, "h": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "d": {"docs": {"linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.silu.triton_weighted_silu_forward": {"tf": 1}, "linghe.utils.silu.triton_weighted_silu_backward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}}, "df": 5}}}}}}}}, "z": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "o": {"docs": {"linghe.utils.reduce.triton_batch_count_zero": {"tf": 1}}, "df": 1}}}}, "o": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "d": {"docs": {"linghe.utils.reduce.triton_batch_sum_with_ord": {"tf": 1}}, "df": 1}}}}}, "annotation": {"root": {"docs": {}, "df": 0}}, "default_value": {"root": {"docs": {}, "df": 0}}, "signature": {"root": {"0": {"6": {"docs": {"linghe.facade.norm.rms_norm": {"tf": 1}, "linghe.facade.norm.group_rms_norm_gate": {"tf": 1}, "linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_forward": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}}, "df": 8}, "docs": {"linghe.utils.reduce.triton_abs_max": {"tf": 1}}, "df": 1}, "1": {"2": {"8": {"docs": {"linghe.quant.block.triton_block_quant": {"tf": 1}, "linghe.quant.group.triton_group_quant": {"tf": 1}}, "df": 2}, "docs": {}, "df": 0}, "docs": {"linghe.utils.gather.triton_make_row_id_map": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}}, "df": 2, "e": {"docs": {"linghe.facade.norm.rms_norm": {"tf": 1}, "linghe.facade.norm.group_rms_norm_gate": {"tf": 1}, "linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_forward": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1}, "linghe.utils.reduce.triton_abs_max": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}}, "df": 9}}, "2": {"docs": {"linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.reduce.triton_batch_sum_with_ord": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}}, "df": 4}, "3": {"0": {"docs": {"linghe.utils.reduce.triton_abs_max": {"tf": 1}}, "df": 1}, "2": {"docs": {"linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}}, "df": 2}, "docs": {}, "df": 0}, "4": {"docs": {"linghe.facade.norm.group_rms_norm_gate": {"tf": 1}, "linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}}, "df": 4}, "docs": {"linghe.facade.add.inplace_add": {"tf": 5.477225575051661}, "linghe.facade.fp32_gemm.fp32_gemm": {"tf": 5.477225575051661}, "linghe.facade.hadamard_quant_linear.HadamardQuantLinear.__init__": {"tf": 7.416198487095663}, "linghe.facade.loss.softmax_cross_entropy": {"tf": 6.6332495807108}, "linghe.facade.norm.rms_norm": {"tf": 6.6332495807108}, "linghe.facade.norm.group_rms_norm_gate": {"tf": 8.660254037844387}, "linghe.facade.rope.qk_norm_half_rope": {"tf": 10.14889156509222}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear.__init__": {"tf": 7.416198487095663}, "linghe.facade.transpose.transpose_dim01": {"tf": 3.1622776601683795}, "linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 9.433981132056603}, "linghe.gemm.fp32_gemm.triton_fp32_gemm": {"tf": 5.477225575051661}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_backward": {"tf": 5.477225575051661}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_update": {"tf": 5.477225575051661}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 6.48074069840786}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm_for_update": {"tf": 6.48074069840786}, "linghe.quant.block.triton_block_quant": {"tf": 5.0990195135927845}, "linghe.quant.channel.triton_row_quant": {"tf": 4.242640687119285}, "linghe.quant.channel.triton_tokenwise_row_quant": {"tf": 5.830951894845301}, "linghe.quant.channel.triton_transpose_row_quant": {"tf": 4.242640687119285}, "linghe.quant.group.triton_group_quant": {"tf": 6.164414002968976}, "linghe.quant.hadamard.triton_hadamard_quant": {"tf": 3.7416573867739413}, "linghe.utils.add.triton_inplace_add": {"tf": 6.6332495807108}, "linghe.utils.dot.triton_dot": {"tf": 3.7416573867739413}, "linghe.utils.gather.triton_make_row_id_map": {"tf": 5.656854249492381}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 6.324555320336759}, "linghe.utils.gather.triton_index_select": {"tf": 6.164414002968976}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 10.295630140987}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 8.246211251235321}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 8.717797887081348}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 8.18535277187245}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 10.583005244258363}, "linghe.utils.loss.triton_softmax_cross_entropy_forward": {"tf": 3.7416573867739413}, "linghe.utils.loss.triton_softmax_cross_entropy_backward": {"tf": 5.830951894845301}, "linghe.utils.norm.triton_rms_norm_forward": {"tf": 5.477225575051661}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 12.206555615733702}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 8.48528137423857}, "linghe.utils.rearange.triton_split_and_cat": {"tf": 5.0990195135927845}, "linghe.utils.reduce.triton_abs_max": {"tf": 6.48074069840786}, "linghe.utils.reduce.triton_batch_count_zero": {"tf": 3.1622776601683795}, "linghe.utils.reduce.triton_batch_sum_with_ord": {"tf": 4.242640687119285}, "linghe.utils.rope.triton_half_rope_forward": {"tf": 5.0990195135927845}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 8.426149773176359}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 8.246211251235321}, "linghe.utils.scatter.triton_aligned_scatter_add": {"tf": 8.306623862918075}, "linghe.utils.scatter.triton_scatter_add": {"tf": 4.242640687119285}, "linghe.utils.scatter.triton_unpermute_with_mask_map": {"tf": 6.48074069840786}, "linghe.utils.silu.triton_weighted_silu_forward": {"tf": 5.0990195135927845}, "linghe.utils.silu.triton_weighted_silu_backward": {"tf": 7.483314773547883}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 6.48074069840786}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 4.69041575982343}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 8.12403840463596}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 6.164414002968976}, "linghe.utils.transpose.triton_transpose": {"tf": 7.681145747868608}, "linghe.utils.transpose.triton_transpose_and_pad": {"tf": 5.0990195135927845}, "linghe.utils.transpose.triton_batch_transpose": {"tf": 4.242640687119285}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 5.477225575051661}}, "df": 56, "x": {"docs": {"linghe.facade.add.inplace_add": {"tf": 1}, "linghe.facade.norm.rms_norm": {"tf": 1}, "linghe.facade.transpose.transpose_dim01": {"tf": 1}, "linghe.quant.block.triton_block_quant": {"tf": 1}, "linghe.quant.channel.triton_row_quant": {"tf": 1}, "linghe.quant.channel.triton_tokenwise_row_quant": {"tf": 1}, "linghe.quant.channel.triton_transpose_row_quant": {"tf": 1}, "linghe.quant.group.triton_group_quant": {"tf": 1}, "linghe.quant.hadamard.triton_hadamard_quant": {"tf": 1}, "linghe.utils.add.triton_inplace_add": {"tf": 1}, "linghe.utils.dot.triton_dot": {"tf": 1}, "linghe.utils.gather.triton_index_select": {"tf": 1}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1.7320508075688772}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1.7320508075688772}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1.4142135623730951}, "linghe.utils.norm.triton_rms_norm_forward": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1}, "linghe.utils.rearange.triton_split_and_cat": {"tf": 1}, "linghe.utils.reduce.triton_abs_max": {"tf": 1}, "linghe.utils.scatter.triton_aligned_scatter_add": {"tf": 1}, "linghe.utils.scatter.triton_scatter_add": {"tf": 1}, "linghe.utils.silu.triton_weighted_silu_forward": {"tf": 1}, "linghe.utils.silu.triton_weighted_silu_backward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.transpose.triton_transpose": {"tf": 1}, "linghe.utils.transpose.triton_transpose_and_pad": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1.4142135623730951}}, "df": 31, "s": {"docs": {"linghe.utils.reduce.triton_batch_count_zero": {"tf": 1}, "linghe.utils.reduce.triton_batch_sum_with_ord": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose": {"tf": 1}}, "df": 3}, "t": {"docs": {}, "df": 0, "s": {"docs": {"linghe.utils.transpose.triton_batch_transpose": {"tf": 1}}, "df": 1}}}, "t": {"docs": {"linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1}}, "df": 1, "o": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "c": {"docs": {}, "df": 0, "h": {"docs": {"linghe.facade.add.inplace_add": {"tf": 1.4142135623730951}, "linghe.facade.fp32_gemm.fp32_gemm": {"tf": 1.4142135623730951}, "linghe.facade.loss.softmax_cross_entropy": {"tf": 1.4142135623730951}, "linghe.facade.norm.rms_norm": {"tf": 1.4142135623730951}, "linghe.facade.norm.group_rms_norm_gate": {"tf": 1.7320508075688772}, "linghe.facade.rope.qk_norm_half_rope": {"tf": 2}, "linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 2.23606797749979}, "linghe.gemm.fp32_gemm.triton_fp32_gemm": {"tf": 1.4142135623730951}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_backward": {"tf": 1.4142135623730951}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_update": {"tf": 1.4142135623730951}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1.7320508075688772}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm_for_update": {"tf": 1.7320508075688772}, "linghe.quant.group.triton_group_quant": {"tf": 1}, "linghe.utils.add.triton_inplace_add": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_make_row_id_map": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 2.23606797749979}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 2}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 2.23606797749979}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1.7320508075688772}, "linghe.utils.scatter.triton_aligned_scatter_add": {"tf": 2}, "linghe.utils.scatter.triton_unpermute_with_mask_map": {"tf": 1.7320508075688772}, "linghe.utils.silu.triton_weighted_silu_backward": {"tf": 1.7320508075688772}, "linghe.utils.transpose.triton_transpose": {"tf": 1}}, "df": 24}}}, "k": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "n": {"docs": {"linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1}}, "df": 3, "s": {"docs": {"linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1.4142135623730951}}, "df": 4}}}}}, "e": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "s": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "r": {"docs": {"linghe.facade.add.inplace_add": {"tf": 1.4142135623730951}, "linghe.facade.fp32_gemm.fp32_gemm": {"tf": 1.4142135623730951}, "linghe.facade.loss.softmax_cross_entropy": {"tf": 1.4142135623730951}, "linghe.facade.norm.rms_norm": {"tf": 1.4142135623730951}, "linghe.facade.norm.group_rms_norm_gate": {"tf": 1.7320508075688772}, "linghe.facade.rope.qk_norm_half_rope": {"tf": 2}, "linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 2}, "linghe.gemm.fp32_gemm.triton_fp32_gemm": {"tf": 1.4142135623730951}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_backward": {"tf": 1.4142135623730951}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_update": {"tf": 1.4142135623730951}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1.7320508075688772}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm_for_update": {"tf": 1.7320508075688772}, "linghe.utils.add.triton_inplace_add": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_make_row_id_map": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 2.23606797749979}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 2}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 2.23606797749979}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1.7320508075688772}, "linghe.utils.scatter.triton_aligned_scatter_add": {"tf": 2}, "linghe.utils.scatter.triton_unpermute_with_mask_map": {"tf": 1.7320508075688772}, "linghe.utils.silu.triton_weighted_silu_backward": {"tf": 1.7320508075688772}, "linghe.utils.transpose.triton_transpose": {"tf": 1}}, "df": 23}}}}}, "r": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "e": {"docs": {"linghe.facade.hadamard_quant_linear.HadamardQuantLinear.__init__": {"tf": 1}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear.__init__": {"tf": 1}, "linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1}, "linghe.utils.add.triton_inplace_add": {"tf": 1}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1}, "linghe.utils.rope.triton_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}, "linghe.utils.transpose.triton_transpose_and_pad": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1}}, "df": 12}}, "a": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "s": {"docs": {}, "df": 0, "p": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "s": {"docs": {}, "df": 0, "e": {"docs": {"linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1}}, "df": 1, "d": {"docs": {"linghe.utils.rope.triton_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}}, "df": 3}}}}}}}}}}, "y": {"docs": {"linghe.facade.add.inplace_add": {"tf": 1}, "linghe.utils.add.triton_inplace_add": {"tf": 1}, "linghe.utils.dot.triton_dot": {"tf": 1}}, "df": 3}, "i": {"docs": {}, "df": 0, "n": {"docs": {"linghe.facade.hadamard_quant_linear.HadamardQuantLinear.__init__": {"tf": 1}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear.__init__": {"tf": 1}}, "df": 2, "p": {"docs": {"linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}}, "df": 2, "u": {"docs": {}, "df": 0, "t": {"docs": {"linghe.facade.fp32_gemm.fp32_gemm": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_backward": {"tf": 1}}, "df": 2}}, "l": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "c": {"docs": {}, "df": 0, "e": {"docs": {"linghe.facade.loss.softmax_cross_entropy": {"tf": 1}}, "df": 1}}}}}, "t": {"docs": {"linghe.facade.hadamard_quant_linear.HadamardQuantLinear.__init__": {"tf": 1.4142135623730951}, "linghe.facade.norm.group_rms_norm_gate": {"tf": 1}, "linghe.facade.rope.qk_norm_half_rope": {"tf": 1.4142135623730951}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear.__init__": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_make_row_id_map": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 2}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.transpose.triton_transpose": {"tf": 1.4142135623730951}}, "df": 10, "e": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "v": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "d": {"docs": {"linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}}, "df": 2}}}}}}}}}, "d": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "c": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "s": {"docs": {"linghe.utils.gather.triton_index_select": {"tf": 1}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.rearange.triton_split_and_cat": {"tf": 1}, "linghe.utils.scatter.triton_aligned_scatter_add": {"tf": 1}, "linghe.utils.scatter.triton_scatter_add": {"tf": 1}}, "df": 7}}}}}}, "d": {"docs": {"linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}, "linghe.utils.scatter.triton_unpermute_with_mask_map": {"tf": 1}}, "df": 3}}, "w": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "g": {"docs": {}, "df": 0, "h": {"docs": {}, "df": 0, "t": {"docs": {"linghe.facade.fp32_gemm.fp32_gemm": {"tf": 1}, "linghe.facade.norm.rms_norm": {"tf": 1}, "linghe.facade.norm.group_rms_norm_gate": {"tf": 1}, "linghe.facade.rope.qk_norm_half_rope": {"tf": 1.4142135623730951}, "linghe.utils.norm.triton_rms_norm_forward": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1.4142135623730951}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1.4142135623730951}, "linghe.utils.silu.triton_weighted_silu_forward": {"tf": 1}, "linghe.utils.silu.triton_weighted_silu_backward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}}, "df": 13, "s": {"docs": {"linghe.utils.scatter.triton_aligned_scatter_add": {"tf": 1}}, "df": 1}}}}}}}, "f": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "s": {"docs": {"linghe.facade.hadamard_quant_linear.HadamardQuantLinear.__init__": {"tf": 1.4142135623730951}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear.__init__": {"tf": 1.4142135623730951}}, "df": 2}}}}}}}, "a": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "s": {"docs": {}, "df": 0, "e": {"docs": {"linghe.facade.loss.softmax_cross_entropy": {"tf": 1}, "linghe.quant.block.triton_block_quant": {"tf": 1}, "linghe.quant.channel.triton_row_quant": {"tf": 1}, "linghe.quant.channel.triton_tokenwise_row_quant": {"tf": 1}, "linghe.quant.channel.triton_transpose_row_quant": {"tf": 1}, "linghe.quant.group.triton_group_quant": {"tf": 1}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}}, "df": 17}}}}, "l": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "t": {"3": {"2": {"docs": {"linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1}}, "df": 1}, "docs": {}, "df": 0}, "8": {"docs": {"linghe.quant.group.triton_group_quant": {"tf": 1}}, "df": 1}, "docs": {"linghe.facade.norm.rms_norm": {"tf": 1}, "linghe.facade.norm.group_rms_norm_gate": {"tf": 1}, "linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}}, "df": 4}}}}, "r": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "q": {"docs": {}, "df": 0, "s": {"docs": {"linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.utils.rope.triton_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}}, "df": 4}}}}}, "o": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "t": {"docs": {"linghe.facade.hadamard_quant_linear.HadamardQuantLinear.__init__": {"tf": 1}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear.__init__": {"tf": 1}, "linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1}, "linghe.quant.channel.triton_tokenwise_row_quant": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}, "linghe.utils.gather.triton_index_select": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_forward": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_weighted_silu_forward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.transpose.triton_transpose_and_pad": {"tf": 1}}, "df": 14, "p": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "t": {"docs": {"linghe.facade.norm.group_rms_norm_gate": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_backward": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}}, "df": 5, "s": {"docs": {"linghe.utils.scatter.triton_aligned_scatter_add": {"tf": 1}, "linghe.utils.scatter.triton_scatter_add": {"tf": 1}}, "df": 2}}}}}}, "f": {"docs": {"linghe.utils.gather.triton_make_row_id_map": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}}, "df": 2}, "p": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "l": {"docs": {"linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1.7320508075688772}, "linghe.utils.scatter.triton_aligned_scatter_add": {"tf": 1}, "linghe.utils.silu.triton_weighted_silu_backward": {"tf": 1}, "linghe.utils.transpose.triton_transpose": {"tf": 1.4142135623730951}}, "df": 5}}}}}}}, "r": {"docs": {}, "df": 0, "g": {"docs": {"linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}}, "df": 1}, "d": {"docs": {"linghe.utils.reduce.triton_batch_sum_with_ord": {"tf": 1}}, "df": 1}}}, "b": {"docs": {"linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1.4142135623730951}, "linghe.gemm.fp32_gemm.triton_fp32_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_backward": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_update": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm_for_update": {"tf": 1}}, "df": 6, "i": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "s": {"docs": {"linghe.facade.hadamard_quant_linear.HadamardQuantLinear.__init__": {"tf": 1}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear.__init__": {"tf": 1}}, "df": 2}}}, "o": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "l": {"docs": {"linghe.facade.hadamard_quant_linear.HadamardQuantLinear.__init__": {"tf": 1}, "linghe.facade.loss.softmax_cross_entropy": {"tf": 1}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear.__init__": {"tf": 1}, "linghe.utils.add.triton_inplace_add": {"tf": 1}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}}, "df": 6}}}, "l": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "c": {"docs": {}, "df": 0, "k": {"docs": {"linghe.quant.block.triton_block_quant": {"tf": 1}}, "df": 1}}}}}, "d": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "v": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "c": {"docs": {}, "df": 0, "e": {"docs": {"linghe.facade.hadamard_quant_linear.HadamardQuantLinear.__init__": {"tf": 1}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear.__init__": {"tf": 1}}, "df": 2}}}}}, "t": {"docs": {}, "df": 0, "y": {"docs": {}, "df": 0, "p": {"docs": {}, "df": 0, "e": {"docs": {"linghe.facade.hadamard_quant_linear.HadamardQuantLinear.__init__": {"tf": 1}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear.__init__": {"tf": 1}, "linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1}, "linghe.quant.group.triton_group_quant": {"tf": 1}}, "df": 4}}}}, "a": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "a": {"docs": {"linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1}}, "df": 1}}}, "i": {"docs": {}, "df": 0, "m": {"0": {"docs": {"linghe.utils.transpose.triton_transpose": {"tf": 1}}, "df": 1}, "1": {"docs": {"linghe.utils.transpose.triton_transpose": {"tf": 1}}, "df": 1}, "docs": {}, "df": 0}}}, "n": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "e": {"docs": {"linghe.facade.hadamard_quant_linear.HadamardQuantLinear.__init__": {"tf": 1.4142135623730951}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear.__init__": {"tf": 1.4142135623730951}, "linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1}, "linghe.quant.channel.triton_tokenwise_row_quant": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_index_select": {"tf": 1.7320508075688772}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1.7320508075688772}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1.4142135623730951}, "linghe.utils.loss.triton_softmax_cross_entropy_backward": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_forward": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1.7320508075688772}, "linghe.utils.rearange.triton_split_and_cat": {"tf": 1}, "linghe.utils.reduce.triton_abs_max": {"tf": 1.4142135623730951}, "linghe.utils.scatter.triton_aligned_scatter_add": {"tf": 1}, "linghe.utils.silu.triton_weighted_silu_forward": {"tf": 1.4142135623730951}, "linghe.utils.silu.triton_weighted_silu_backward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1.4142135623730951}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1.7320508075688772}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.transpose.triton_transpose": {"tf": 1.4142135623730951}, "linghe.utils.transpose.triton_transpose_and_pad": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1}}, "df": 24}}, "r": {"docs": {}, "df": 0, "m": {"docs": {"linghe.facade.rope.qk_norm_half_rope": {"tf": 1.4142135623730951}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1.4142135623730951}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1.4142135623730951}}, "df": 3}}}, "u": {"docs": {}, "df": 0, "m": {"docs": {"linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1.7320508075688772}}, "df": 3}}}, "l": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "g": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "t": {"docs": {"linghe.utils.loss.triton_softmax_cross_entropy_backward": {"tf": 1}}, "df": 1, "s": {"docs": {"linghe.facade.loss.softmax_cross_entropy": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_forward": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_backward": {"tf": 1}}, "df": 3}}}}}, "a": {"docs": {}, "df": 0, "b": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "s": {"docs": {"linghe.facade.loss.softmax_cross_entropy": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_forward": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_backward": {"tf": 1}}, "df": 3}}}}}, "i": {"docs": {}, "df": 0, "s": {"docs": {}, "df": 0, "t": {"docs": {"linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1}}, "df": 1}}}}, "e": {"4": {"docs": {}, "df": 0, "m": {"3": {"docs": {}, "df": 0, "f": {"docs": {}, "df": 0, "n": {"docs": {"linghe.quant.group.triton_group_quant": {"tf": 1}}, "df": 1}}}, "docs": {}, "df": 0}}, "docs": {}, "df": 0, "p": {"docs": {}, "df": 0, "s": {"docs": {"linghe.facade.norm.rms_norm": {"tf": 1}, "linghe.facade.norm.group_rms_norm_gate": {"tf": 1}, "linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_forward": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}}, "df": 8}}, "x": {"docs": {}, "df": 0, "p": {"docs": {"linghe.utils.loss.triton_softmax_cross_entropy_backward": {"tf": 1}}, "df": 1, "e": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "t": {"docs": {"linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1}}, "df": 4, "s": {"docs": {"linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}}, "df": 1}}}}}}}, "a": {"docs": {"linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1.4142135623730951}, "linghe.gemm.fp32_gemm.triton_fp32_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_backward": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_update": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm_for_update": {"tf": 1}}, "df": 6, "t": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "n": {"docs": {"linghe.facade.norm.group_rms_norm_gate": {"tf": 1}}, "df": 1}}}, "c": {"docs": {}, "df": 0, "c": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "m": {"docs": {"linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1}, "linghe.utils.add.triton_inplace_add": {"tf": 1}}, "df": 2}}}}, "x": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "s": {"docs": {"linghe.utils.reduce.triton_abs_max": {"tf": 1}}, "df": 1}}}}, "g": {"docs": {"linghe.utils.silu.triton_weighted_silu_backward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}}, "df": 3, "a": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "e": {"docs": {"linghe.facade.norm.group_rms_norm_gate": {"tf": 1}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1}}, "df": 2}}}, "r": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "p": {"docs": {"linghe.facade.norm.group_rms_norm_gate": {"tf": 1}, "linghe.quant.group.triton_group_quant": {"tf": 1}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1}}, "df": 3}}}, "a": {"docs": {}, "df": 0, "d": {"docs": {"linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1.4142135623730951}, "linghe.utils.loss.triton_softmax_cross_entropy_backward": {"tf": 1.4142135623730951}, "linghe.utils.scatter.triton_unpermute_with_mask_map": {"tf": 1}}, "df": 3, "s": {"docs": {"linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}}, "df": 1}}}}, "q": {"docs": {"linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}}, "df": 1}, "k": {"docs": {"linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}}, "df": 1}, "v": {"docs": {"linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}}, "df": 1}}, "s": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "z": {"docs": {}, "df": 0, "e": {"docs": {"linghe.facade.norm.group_rms_norm_gate": {"tf": 1}, "linghe.quant.block.triton_block_quant": {"tf": 1}, "linghe.quant.group.triton_group_quant": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1}}, "df": 5}}}, "c": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "e": {"docs": {"linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1.4142135623730951}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm_for_update": {"tf": 1}, "linghe.quant.block.triton_block_quant": {"tf": 1}, "linghe.quant.channel.triton_row_quant": {"tf": 1}, "linghe.quant.channel.triton_tokenwise_row_quant": {"tf": 1.4142135623730951}, "linghe.quant.channel.triton_transpose_row_quant": {"tf": 1}, "linghe.quant.group.triton_group_quant": {"tf": 1}, "linghe.utils.gather.triton_index_select": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 2}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1.7320508075688772}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1.4142135623730951}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1.4142135623730951}, "linghe.utils.reduce.triton_abs_max": {"tf": 1.4142135623730951}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1.4142135623730951}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1.4142135623730951}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}}, "df": 20, "s": {"docs": {"linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}, "linghe.utils.rearange.triton_split_and_cat": {"tf": 1}}, "df": 5}}}}}, "m": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "h": {"docs": {"linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}, "linghe.utils.reduce.triton_abs_max": {"tf": 1}}, "df": 5}}}}}, "p": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "s": {"docs": {"linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}}, "df": 3}}}}}, "u": {"docs": {}, "df": 0, "m": {"docs": {"linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_backward": {"tf": 1}}, "df": 2}}}, "q": {"docs": {"linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.rope.triton_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}}, "df": 7, "k": {"docs": {}, "df": 0, "v": {"docs": {"linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}}, "df": 3}}}, "k": {"docs": {"linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.utils.rope.triton_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}}, "df": 4}, "h": {"docs": {"linghe.facade.rope.qk_norm_half_rope": {"tf": 1.4142135623730951}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1.4142135623730951}}, "df": 2, "m": {"docs": {"linghe.quant.hadamard.triton_hadamard_quant": {"tf": 1}}, "df": 1}, "i": {"docs": {}, "df": 0, "d": {"docs": {}, "df": 0, "d": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "n": {"docs": {"linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}}, "df": 1}}}}}}, "c": {"docs": {"linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1}}, "df": 1, "o": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "g": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "s": {"docs": {"linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}}, "df": 1}}}}}}}}, "u": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "t": {"docs": {"linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1}}, "df": 4, "s": {"docs": {"linghe.utils.rearange.triton_split_and_cat": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}}, "df": 3}}}}}}, "r": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "d": {"docs": {"linghe.quant.block.triton_block_quant": {"tf": 1}, "linghe.quant.channel.triton_row_quant": {"tf": 1}, "linghe.quant.channel.triton_tokenwise_row_quant": {"tf": 1}, "linghe.quant.channel.triton_transpose_row_quant": {"tf": 1}, "linghe.quant.group.triton_group_quant": {"tf": 1}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}}, "df": 14}}, "t": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "g": {"docs": {"linghe.utils.gather.triton_make_row_id_map": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}}, "df": 2}}}}}, "w": {"docs": {"linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}, "linghe.utils.scatter.triton_unpermute_with_mask_map": {"tf": 1}}, "df": 3}}, "e": {"docs": {}, "df": 0, "v": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "s": {"docs": {}, "df": 0, "e": {"docs": {"linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}}, "df": 3}}}}}}, "m": {"docs": {}, "df": 0, "s": {"docs": {"linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}}, "df": 1}}}, "m": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "p": {"docs": {"linghe.utils.gather.triton_make_row_id_map": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}, "linghe.utils.scatter.triton_unpermute_with_mask_map": {"tf": 1}}, "df": 5}, "x": {"docs": {"linghe.utils.loss.triton_softmax_cross_entropy_backward": {"tf": 1}}, "df": 1}}, "u": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "p": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "e": {"docs": {"linghe.utils.gather.triton_make_row_id_map": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}}, "df": 2}}}}}}}, "o": {"docs": {}, "df": 0, "d": {"docs": {}, "df": 0, "e": {"docs": {"linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}}, "df": 3}}}, "i": {"docs": {}, "df": 0, "n": {"docs": {"linghe.utils.reduce.triton_abs_max": {"tf": 1}}, "df": 1}}}, "p": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "b": {"docs": {}, "df": 0, "s": {"docs": {"linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.scatter.triton_unpermute_with_mask_map": {"tf": 1}}, "df": 2}}}}, "e": {"docs": {}, "df": 0, "r": {"docs": {"linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1}}, "df": 4}}, "a": {"docs": {}, "df": 0, "d": {"docs": {"linghe.utils.transpose.triton_transpose_and_pad": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1}}, "df": 2}}}, "v": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "e": {"docs": {"linghe.utils.reduce.triton_abs_max": {"tf": 1}}, "df": 1}}}}}}}, "bases": {"root": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "c": {"docs": {}, "df": 0, "h": {"docs": {"linghe.facade.hadamard_quant_linear.HadamardQuantLinear": {"tf": 1}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear": {"tf": 1}}, "df": 2}}}}}, "n": {"docs": {}, "df": 0, "n": {"docs": {"linghe.facade.hadamard_quant_linear.HadamardQuantLinear": {"tf": 1}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear": {"tf": 1}}, "df": 2}}, "m": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "d": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "e": {"docs": {"linghe.facade.hadamard_quant_linear.HadamardQuantLinear": {"tf": 1.4142135623730951}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear": {"tf": 1.4142135623730951}}, "df": 2, "s": {"docs": {"linghe.facade.hadamard_quant_linear.HadamardQuantLinear": {"tf": 1}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear": {"tf": 1}}, "df": 2}}}}}}}}}, "doc": {"root": {"0": {"docs": {"linghe.facade.norm.group_rms_norm_gate": {"tf": 1}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1.4142135623730951}, "linghe.utils.reduce.triton_abs_max": {"tf": 1}, "linghe.utils.scatter.triton_aligned_scatter_add": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1.4142135623730951}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1.4142135623730951}, "linghe.utils.transpose.triton_transpose": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1}}, "df": 11, "/": {"1": {"docs": {"linghe.utils.gather.triton_make_row_id_map": {"tf": 1}}, "df": 1}, "docs": {}, "df": 0}}, "1": {"1": {"docs": {"linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.scatter.triton_aligned_scatter_add": {"tf": 1}}, "df": 2}, "2": {"docs": {"linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}}, "df": 1}, "docs": {"linghe.facade.norm.group_rms_norm_gate": {"tf": 1}, "linghe.utils.dot.triton_dot": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_make_row_id_map": {"tf": 1}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1.4142135623730951}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1.4142135623730951}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1.4142135623730951}, "linghe.utils.transpose.triton_transpose": {"tf": 1}}, "df": 8, "/": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "m": {"docs": {}, "df": 0, "s": {"docs": {"linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm_for_update": {"tf": 1}}, "df": 2}}}, "s": {"docs": {}, "df": 0, "c": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "e": {"docs": {"linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}}, "df": 1}}}}}}}, "2": {"docs": {"linghe.quant.block.triton_block_quant": {"tf": 1}, "linghe.quant.channel.triton_row_quant": {"tf": 1}, "linghe.quant.channel.triton_tokenwise_row_quant": {"tf": 1.4142135623730951}, "linghe.quant.channel.triton_transpose_row_quant": {"tf": 1}, "linghe.quant.group.triton_group_quant": {"tf": 1}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1.7320508075688772}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1.7320508075688772}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1.7320508075688772}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}}, "df": 12}, "3": {"2": {"docs": {"linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.transpose.triton_transpose_and_pad": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1}}, "df": 3}, "docs": {}, "df": 0}, "4": {"docs": {"linghe.facade.transpose.transpose_dim01": {"tf": 1}}, "df": 1}, "docs": {"linghe": {"tf": 1.7320508075688772}, "linghe.facade": {"tf": 1.7320508075688772}, "linghe.facade.add": {"tf": 1.7320508075688772}, "linghe.facade.add.inplace_add": {"tf": 4.898979485566356}, "linghe.facade.fp32_gemm": {"tf": 1.7320508075688772}, "linghe.facade.fp32_gemm.fp32_gemm": {"tf": 5}, "linghe.facade.hadamard_quant_linear": {"tf": 1.7320508075688772}, "linghe.facade.hadamard_quant_linear.HadamardQuantLinear": {"tf": 1.4142135623730951}, "linghe.facade.hadamard_quant_linear.HadamardQuantLinear.__init__": {"tf": 5.0990195135927845}, "linghe.facade.loss": {"tf": 1.7320508075688772}, "linghe.facade.loss.softmax_cross_entropy": {"tf": 5.744562646538029}, "linghe.facade.norm": {"tf": 1.7320508075688772}, "linghe.facade.norm.rms_norm": {"tf": 5.291502622129181}, "linghe.facade.norm.group_rms_norm_gate": {"tf": 6.164414002968976}, "linghe.facade.rope": {"tf": 1.7320508075688772}, "linghe.facade.rope.qk_norm_half_rope": {"tf": 7.483314773547883}, "linghe.facade.smooth_quant_linear": {"tf": 1.7320508075688772}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear": {"tf": 1.4142135623730951}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear.__init__": {"tf": 5.0990195135927845}, "linghe.facade.transpose": {"tf": 1.7320508075688772}, "linghe.facade.transpose.transpose_dim01": {"tf": 4.47213595499958}, "linghe.gemm": {"tf": 1.7320508075688772}, "linghe.gemm.blockwise_fp8_gemm": {"tf": 1.7320508075688772}, "linghe.gemm.channelwise_fp8_gemm": {"tf": 1.7320508075688772}, "linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 6.6332495807108}, "linghe.gemm.fp32_gemm": {"tf": 1.7320508075688772}, "linghe.gemm.fp32_gemm.triton_fp32_gemm": {"tf": 4.898979485566356}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_backward": {"tf": 5}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_update": {"tf": 4.898979485566356}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 5.385164807134504}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm_for_update": {"tf": 5.291502622129181}, "linghe.quant": {"tf": 1.7320508075688772}, "linghe.quant.block": {"tf": 1.7320508075688772}, "linghe.quant.block.triton_block_quant": {"tf": 5.830951894845301}, "linghe.quant.channel": {"tf": 1.7320508075688772}, "linghe.quant.channel.triton_row_quant": {"tf": 4.898979485566356}, "linghe.quant.channel.triton_tokenwise_row_quant": {"tf": 4.898979485566356}, "linghe.quant.channel.triton_transpose_row_quant": {"tf": 5.477225575051661}, "linghe.quant.group": {"tf": 1.7320508075688772}, "linghe.quant.group.triton_group_quant": {"tf": 5.830951894845301}, "linghe.quant.hadamard": {"tf": 1.7320508075688772}, "linghe.quant.hadamard.triton_hadamard_quant": {"tf": 5.830951894845301}, "linghe.quant.smooth": {"tf": 1.7320508075688772}, "linghe.utils": {"tf": 1.7320508075688772}, "linghe.utils.add": {"tf": 1.7320508075688772}, "linghe.utils.add.triton_inplace_add": {"tf": 5.477225575051661}, "linghe.utils.dot": {"tf": 1.7320508075688772}, "linghe.utils.dot.triton_dot": {"tf": 5.196152422706632}, "linghe.utils.gather": {"tf": 1.7320508075688772}, "linghe.utils.gather.triton_make_row_id_map": {"tf": 4.898979485566356}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 5.385164807134504}, "linghe.utils.gather.triton_index_select": {"tf": 5.291502622129181}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 6.6332495807108}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 7.14142842854285}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 6.6332495807108}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 6.928203230275509}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 8.18535277187245}, "linghe.utils.loss": {"tf": 1.7320508075688772}, "linghe.utils.loss.triton_softmax_cross_entropy_forward": {"tf": 4.898979485566356}, "linghe.utils.loss.triton_softmax_cross_entropy_backward": {"tf": 6}, "linghe.utils.norm": {"tf": 1.7320508075688772}, "linghe.utils.norm.triton_rms_norm_forward": {"tf": 5.291502622129181}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 8.306623862918075}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 6.324555320336759}, "linghe.utils.rearange": {"tf": 1.7320508075688772}, "linghe.utils.rearange.triton_split_and_cat": {"tf": 6.164414002968976}, "linghe.utils.reduce": {"tf": 1.7320508075688772}, "linghe.utils.reduce.triton_abs_max": {"tf": 6.082762530298219}, "linghe.utils.reduce.triton_batch_count_zero": {"tf": 4.47213595499958}, "linghe.utils.reduce.triton_batch_sum_with_ord": {"tf": 5.196152422706632}, "linghe.utils.rope": {"tf": 1.7320508075688772}, "linghe.utils.rope.triton_half_rope_forward": {"tf": 5.830951894845301}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 8}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 8.366600265340756}, "linghe.utils.scatter": {"tf": 1.7320508075688772}, "linghe.utils.scatter.triton_aligned_scatter_add": {"tf": 5.656854249492381}, "linghe.utils.scatter.triton_scatter_add": {"tf": 5.291502622129181}, "linghe.utils.scatter.triton_unpermute_with_mask_map": {"tf": 5.830951894845301}, "linghe.utils.silu": {"tf": 1.7320508075688772}, "linghe.utils.silu.triton_weighted_silu_forward": {"tf": 4.898979485566356}, "linghe.utils.silu.triton_weighted_silu_backward": {"tf": 5.830951894845301}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 6.164414002968976}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 6.164414002968976}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 7.0710678118654755}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 7.211102550927978}, "linghe.utils.transpose": {"tf": 1.7320508075688772}, "linghe.utils.transpose.triton_transpose": {"tf": 5.291502622129181}, "linghe.utils.transpose.triton_transpose_and_pad": {"tf": 5.385164807134504}, "linghe.utils.transpose.triton_batch_transpose": {"tf": 4.47213595499958}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 5.291502622129181}}, "df": 90, "c": {"docs": {"linghe.facade.add": {"tf": 1}, "linghe.facade.fp32_gemm": {"tf": 1}, "linghe.facade.hadamard_quant_linear": {"tf": 1}, "linghe.facade.loss": {"tf": 1}, "linghe.facade.norm": {"tf": 1}, "linghe.facade.rope": {"tf": 1}, "linghe.facade.smooth_quant_linear": {"tf": 1}, "linghe.facade.transpose": {"tf": 1}, "linghe.gemm.blockwise_fp8_gemm": {"tf": 1}, "linghe.gemm.channelwise_fp8_gemm": {"tf": 1}, "linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 2}, "linghe.gemm.fp32_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_backward": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_update": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1}, "linghe.quant.block": {"tf": 1}, "linghe.quant.channel": {"tf": 1}, "linghe.quant.group": {"tf": 1}, "linghe.quant.hadamard": {"tf": 1}, "linghe.quant.smooth": {"tf": 1}, "linghe.utils.add": {"tf": 1}, "linghe.utils.dot": {"tf": 1}, "linghe.utils.gather": {"tf": 1}, "linghe.utils.loss": {"tf": 1}, "linghe.utils.rearange": {"tf": 1}, "linghe.utils.reduce": {"tf": 1}, "linghe.utils.rope": {"tf": 1}, "linghe.utils.scatter": {"tf": 1}, "linghe.utils.silu": {"tf": 1}, "linghe.utils.transpose": {"tf": 1}}, "df": 31, "o": {"docs": {}, "df": 0, "p": {"docs": {}, "df": 0, "y": {"docs": {"linghe.utils.add.triton_inplace_add": {"tf": 1}}, "df": 1, "r": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "g": {"docs": {}, "df": 0, "h": {"docs": {}, "df": 0, "t": {"docs": {"linghe.facade.add": {"tf": 1}, "linghe.facade.fp32_gemm": {"tf": 1}, "linghe.facade.hadamard_quant_linear": {"tf": 1}, "linghe.facade.loss": {"tf": 1}, "linghe.facade.norm": {"tf": 1}, "linghe.facade.rope": {"tf": 1}, "linghe.facade.smooth_quant_linear": {"tf": 1}, "linghe.facade.transpose": {"tf": 1}, "linghe.gemm.blockwise_fp8_gemm": {"tf": 1}, "linghe.gemm.channelwise_fp8_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm": {"tf": 1}, "linghe.quant.block": {"tf": 1}, "linghe.quant.channel": {"tf": 1}, "linghe.quant.group": {"tf": 1}, "linghe.quant.hadamard": {"tf": 1}, "linghe.quant.smooth": {"tf": 1}, "linghe.utils.add": {"tf": 1}, "linghe.utils.dot": {"tf": 1}, "linghe.utils.gather": {"tf": 1}, "linghe.utils.loss": {"tf": 1}, "linghe.utils.rearange": {"tf": 1}, "linghe.utils.reduce": {"tf": 1}, "linghe.utils.rope": {"tf": 1}, "linghe.utils.scatter": {"tf": 1}, "linghe.utils.silu": {"tf": 1}, "linghe.utils.transpose": {"tf": 1}}, "df": 26}}}}}}}, "r": {"docs": {}, "df": 0, "e": {"docs": {"linghe.facade.norm.group_rms_norm_gate": {"tf": 1}}, "df": 1}}, "l": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "m": {"docs": {}, "df": 0, "n": {"docs": {"linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1}, "linghe.utils.transpose.triton_transpose_and_pad": {"tf": 1}}, "df": 2, "w": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "s": {"docs": {}, "df": 0, "e": {"docs": {"linghe.quant.hadamard.triton_hadamard_quant": {"tf": 1.4142135623730951}, "linghe.utils.reduce.triton_abs_max": {"tf": 1}}, "df": 2}}}}}}}}, "n": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "g": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "s": {"docs": {"linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1.7320508075688772}}, "df": 1}}}}}}}}, "u": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "t": {"docs": {"linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1.7320508075688772}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.reduce.triton_batch_count_zero": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1.4142135623730951}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1.4142135623730951}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1.4142135623730951}}, "df": 8, "s": {"docs": {"linghe.utils.rearange.triton_split_and_cat": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}}, "df": 3}}}}, "m": {"docs": {}, "df": 0, "p": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "e": {"docs": {"linghe.utils.loss.triton_softmax_cross_entropy_forward": {"tf": 1}, "linghe.utils.silu.triton_weighted_silu_forward": {"tf": 1}}, "df": 2}}}}}}, "u": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "y": {"docs": {"linghe.facade.fp32_gemm.fp32_gemm": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}}, "df": 2}}}}}}}, "d": {"docs": {}, "df": 0, "a": {"docs": {"linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}}, "df": 3}}}, "r": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "s": {"docs": {}, "df": 0, "s": {"docs": {"linghe.facade.loss.softmax_cross_entropy": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_forward": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_backward": {"tf": 1}}, "df": 3}}}}, "a": {"docs": {}, "df": 0, "n": {"docs": {"linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1}}, "df": 1}, "l": {"docs": {}, "df": 0, "c": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "e": {"docs": {"linghe.utils.dot.triton_dot": {"tf": 1}}, "df": 1, "d": {"docs": {"linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.transpose.triton_transpose_and_pad": {"tf": 1}}, "df": 2}}}}}}}}, "t": {"docs": {"linghe.utils.rearange.triton_split_and_cat": {"tf": 1}}, "df": 1}}}, "a": {"docs": {"linghe.facade.hadamard_quant_linear.HadamardQuantLinear": {"tf": 1}, "linghe.facade.loss.softmax_cross_entropy": {"tf": 1}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear": {"tf": 1}, "linghe.facade.transpose.transpose_dim01": {"tf": 1.4142135623730951}, "linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1.7320508075688772}, "linghe.gemm.fp32_gemm.triton_fp32_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_backward": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_update": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1.7320508075688772}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm_for_update": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map": {"tf": 1}, "linghe.utils.reduce.triton_batch_count_zero": {"tf": 1}, "linghe.utils.reduce.triton_batch_sum_with_ord": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1}}, "df": 14, "n": {"docs": {"linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}}, "df": 1, "t": {"docs": {"linghe.facade.add": {"tf": 1}, "linghe.facade.fp32_gemm": {"tf": 1}, "linghe.facade.hadamard_quant_linear": {"tf": 1}, "linghe.facade.loss": {"tf": 1}, "linghe.facade.norm": {"tf": 1}, "linghe.facade.rope": {"tf": 1}, "linghe.facade.smooth_quant_linear": {"tf": 1}, "linghe.facade.transpose": {"tf": 1}, "linghe.gemm.blockwise_fp8_gemm": {"tf": 1}, "linghe.gemm.channelwise_fp8_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm": {"tf": 1}, "linghe.quant.block": {"tf": 1}, "linghe.quant.channel": {"tf": 1}, "linghe.quant.group": {"tf": 1}, "linghe.quant.hadamard": {"tf": 1}, "linghe.quant.smooth": {"tf": 1}, "linghe.utils.add": {"tf": 1}, "linghe.utils.dot": {"tf": 1}, "linghe.utils.gather": {"tf": 1}, "linghe.utils.loss": {"tf": 1}, "linghe.utils.rearange": {"tf": 1}, "linghe.utils.reduce": {"tf": 1}, "linghe.utils.rope": {"tf": 1}, "linghe.utils.scatter": {"tf": 1}, "linghe.utils.silu": {"tf": 1}, "linghe.utils.transpose": {"tf": 1}}, "df": 26}, "d": {"docs": {"linghe.facade.add": {"tf": 1}, "linghe.facade.fp32_gemm": {"tf": 1}, "linghe.facade.fp32_gemm.fp32_gemm": {"tf": 1}, "linghe.facade.hadamard_quant_linear": {"tf": 1}, "linghe.facade.hadamard_quant_linear.HadamardQuantLinear": {"tf": 1}, "linghe.facade.loss": {"tf": 1}, "linghe.facade.norm": {"tf": 1}, "linghe.facade.rope": {"tf": 1}, "linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.facade.smooth_quant_linear": {"tf": 1}, "linghe.facade.transpose": {"tf": 1}, "linghe.gemm.blockwise_fp8_gemm": {"tf": 1}, "linghe.gemm.channelwise_fp8_gemm": {"tf": 1}, "linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1}, "linghe.gemm.fp32_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1.4142135623730951}, "linghe.quant.block": {"tf": 1}, "linghe.quant.channel": {"tf": 1}, "linghe.quant.channel.triton_transpose_row_quant": {"tf": 1}, "linghe.quant.group": {"tf": 1}, "linghe.quant.hadamard": {"tf": 1}, "linghe.quant.hadamard.triton_hadamard_quant": {"tf": 1}, "linghe.quant.smooth": {"tf": 1}, "linghe.utils.add": {"tf": 1}, "linghe.utils.dot": {"tf": 1}, "linghe.utils.gather": {"tf": 1}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1.4142135623730951}, "linghe.utils.loss": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1.4142135623730951}, "linghe.utils.rearange": {"tf": 1}, "linghe.utils.rearange.triton_split_and_cat": {"tf": 1}, "linghe.utils.reduce": {"tf": 1}, "linghe.utils.rope": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}, "linghe.utils.scatter": {"tf": 1}, "linghe.utils.silu": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.transpose": {"tf": 1}, "linghe.utils.transpose.triton_transpose": {"tf": 1}, "linghe.utils.transpose.triton_transpose_and_pad": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1}}, "df": 50}}, "f": {"docs": {}, "df": 0, "f": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "s": {"docs": {"linghe.facade.add": {"tf": 1}, "linghe.facade.fp32_gemm": {"tf": 1}, "linghe.facade.hadamard_quant_linear": {"tf": 1}, "linghe.facade.loss": {"tf": 1}, "linghe.facade.norm": {"tf": 1}, "linghe.facade.rope": {"tf": 1}, "linghe.facade.smooth_quant_linear": {"tf": 1}, "linghe.facade.transpose": {"tf": 1}, "linghe.gemm.blockwise_fp8_gemm": {"tf": 1}, "linghe.gemm.channelwise_fp8_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm": {"tf": 1}, "linghe.quant.block": {"tf": 1}, "linghe.quant.channel": {"tf": 1}, "linghe.quant.group": {"tf": 1}, "linghe.quant.hadamard": {"tf": 1}, "linghe.quant.smooth": {"tf": 1}, "linghe.utils.add": {"tf": 1}, "linghe.utils.dot": {"tf": 1}, "linghe.utils.gather": {"tf": 1}, "linghe.utils.loss": {"tf": 1}, "linghe.utils.rearange": {"tf": 1}, "linghe.utils.reduce": {"tf": 1}, "linghe.utils.rope": {"tf": 1}, "linghe.utils.scatter": {"tf": 1}, "linghe.utils.silu": {"tf": 1}, "linghe.utils.transpose": {"tf": 1}}, "df": 26}}}}}}}}}, "d": {"docs": {}, "df": 0, "d": {"docs": {"linghe.facade.add.inplace_add": {"tf": 1.4142135623730951}, "linghe.utils.add.triton_inplace_add": {"tf": 1}, "linghe.utils.scatter.triton_aligned_scatter_add": {"tf": 1}, "linghe.utils.scatter.triton_scatter_add": {"tf": 1}, "linghe.utils.scatter.triton_unpermute_with_mask_map": {"tf": 1}}, "df": 5}}, "r": {"docs": {}, "df": 0, "g": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "m": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "s": {"docs": {"linghe.facade.add.inplace_add": {"tf": 1}, "linghe.facade.fp32_gemm.fp32_gemm": {"tf": 1}, "linghe.facade.hadamard_quant_linear.HadamardQuantLinear.__init__": {"tf": 1}, "linghe.facade.loss.softmax_cross_entropy": {"tf": 1}, "linghe.facade.norm.rms_norm": {"tf": 1}, "linghe.facade.norm.group_rms_norm_gate": {"tf": 1}, "linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear.__init__": {"tf": 1}, "linghe.facade.transpose.transpose_dim01": {"tf": 1}, "linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_backward": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_update": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm_for_update": {"tf": 1}, "linghe.quant.block.triton_block_quant": {"tf": 1}, "linghe.quant.channel.triton_row_quant": {"tf": 1}, "linghe.quant.channel.triton_tokenwise_row_quant": {"tf": 1}, "linghe.quant.channel.triton_transpose_row_quant": {"tf": 1}, "linghe.quant.group.triton_group_quant": {"tf": 1}, "linghe.quant.hadamard.triton_hadamard_quant": {"tf": 1}, "linghe.utils.add.triton_inplace_add": {"tf": 1}, "linghe.utils.dot.triton_dot": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}, "linghe.utils.gather.triton_index_select": {"tf": 1}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_forward": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_backward": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_forward": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1}, "linghe.utils.rearange.triton_split_and_cat": {"tf": 1}, "linghe.utils.reduce.triton_abs_max": {"tf": 1}, "linghe.utils.reduce.triton_batch_count_zero": {"tf": 1}, "linghe.utils.reduce.triton_batch_sum_with_ord": {"tf": 1}, "linghe.utils.rope.triton_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}, "linghe.utils.scatter.triton_aligned_scatter_add": {"tf": 1}, "linghe.utils.scatter.triton_scatter_add": {"tf": 1}, "linghe.utils.scatter.triton_unpermute_with_mask_map": {"tf": 1}, "linghe.utils.silu.triton_weighted_silu_forward": {"tf": 1}, "linghe.utils.silu.triton_weighted_silu_backward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.transpose.triton_transpose": {"tf": 1}, "linghe.utils.transpose.triton_transpose_and_pad": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1}}, "df": 56}}}}}}}, "e": {"docs": {"linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}}, "df": 3}}, "c": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "v": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "n": {"docs": {"linghe.facade.fp32_gemm.fp32_gemm": {"tf": 1}, "linghe.facade.norm.rms_norm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_backward": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_update": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1.4142135623730951}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm_for_update": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1.4142135623730951}}, "df": 7}}}}}}}}, "c": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "m": {"docs": {"linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1}, "linghe.utils.add.triton_inplace_add": {"tf": 1.4142135623730951}}, "df": 2, "u": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "g": {"docs": {"linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1}}, "df": 1}}}, "e": {"docs": {"linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1}}, "df": 1}}}}}}}}}, "t": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "n": {"docs": {"linghe.facade.norm.group_rms_norm_gate": {"tf": 1.7320508075688772}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1}}, "df": 2}, "e": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "n": {"docs": {"linghe.facade.norm.group_rms_norm_gate": {"tf": 1}, "linghe.facade.rope.qk_norm_half_rope": {"tf": 1.4142135623730951}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1.4142135623730951}}, "df": 4}}}}}}}}, "p": {"docs": {}, "df": 0, "p": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "y": {"docs": {"linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.quant.hadamard.triton_hadamard_quant": {"tf": 1}, "linghe.utils.rope.triton_half_rope_forward": {"tf": 1.4142135623730951}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}}, "df": 4}}}}, "@": {"docs": {}, "df": 0, "b": {"docs": {"linghe.gemm.fp32_gemm.triton_fp32_gemm_for_backward": {"tf": 1}}, "df": 1}}, "s": {"docs": {"linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}}, "df": 3}, "l": {"docs": {}, "df": 0, "l": {"2": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "l": {"docs": {"linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}}, "df": 1}}}}, "docs": {}, "df": 0}}, "b": {"docs": {}, "df": 0, "s": {"docs": {"linghe.utils.reduce.triton_abs_max": {"tf": 1.4142135623730951}, "linghe.utils.reduce.triton_batch_sum_with_ord": {"tf": 1}}, "df": 2}}, "x": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "s": {"docs": {"linghe.utils.reduce.triton_abs_max": {"tf": 1.4142135623730951}}, "df": 1}}}}, "f": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "c": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "l": {"docs": {"linghe.facade.add": {"tf": 1}, "linghe.facade.fp32_gemm": {"tf": 1}, "linghe.facade.hadamard_quant_linear": {"tf": 1}, "linghe.facade.loss": {"tf": 1}, "linghe.facade.norm": {"tf": 1}, "linghe.facade.rope": {"tf": 1}, "linghe.facade.smooth_quant_linear": {"tf": 1}, "linghe.facade.transpose": {"tf": 1}, "linghe.gemm.blockwise_fp8_gemm": {"tf": 1}, "linghe.gemm.channelwise_fp8_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm": {"tf": 1}, "linghe.quant.block": {"tf": 1}, "linghe.quant.channel": {"tf": 1}, "linghe.quant.group": {"tf": 1}, "linghe.quant.hadamard": {"tf": 1}, "linghe.quant.smooth": {"tf": 1}, "linghe.utils.add": {"tf": 1}, "linghe.utils.dot": {"tf": 1}, "linghe.utils.gather": {"tf": 1}, "linghe.utils.loss": {"tf": 1}, "linghe.utils.rearange": {"tf": 1}, "linghe.utils.reduce": {"tf": 1}, "linghe.utils.rope": {"tf": 1}, "linghe.utils.scatter": {"tf": 1}, "linghe.utils.silu": {"tf": 1}, "linghe.utils.transpose": {"tf": 1}}, "df": 26}}}}}}}, "r": {"docs": {}, "df": 0, "s": {"docs": {}, "df": 0, "t": {"docs": {"linghe.facade.transpose.transpose_dim01": {"tf": 1}}, "df": 1}}}, "l": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "d": {"docs": {"linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1}}, "df": 1}}}}}, "l": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "t": {"3": {"2": {"docs": {"linghe.facade.fp32_gemm.fp32_gemm": {"tf": 1}, "linghe.quant.block.triton_block_quant": {"tf": 1}, "linghe.quant.group.triton_group_quant": {"tf": 1}}, "df": 3}, "docs": {}, "df": 0}, "8": {"docs": {"linghe.quant.block.triton_block_quant": {"tf": 1}, "linghe.quant.group.triton_group_quant": {"tf": 1}}, "df": 2}, "docs": {"linghe.gemm.fp32_gemm.triton_fp32_gemm_for_backward": {"tf": 1}}, "df": 1}}}, "a": {"docs": {}, "df": 0, "s": {"docs": {}, "df": 0, "h": {"docs": {"linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}}, "df": 2}}}}, "e": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "e": {"docs": {"linghe.facade.hadamard_quant_linear.HadamardQuantLinear.__init__": {"tf": 1.4142135623730951}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear.__init__": {"tf": 1.4142135623730951}}, "df": 2, "s": {"docs": {"linghe.facade.hadamard_quant_linear.HadamardQuantLinear.__init__": {"tf": 1.4142135623730951}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear.__init__": {"tf": 1.4142135623730951}}, "df": 2}}}}}}}, "o": {"docs": {}, "df": 0, "r": {"docs": {"linghe.facade.norm.rms_norm": {"tf": 1}, "linghe.facade.norm.group_rms_norm_gate": {"tf": 1.4142135623730951}, "linghe.facade.rope.qk_norm_half_rope": {"tf": 1.7320508075688772}, "linghe.gemm.fp32_gemm.triton_fp32_gemm": {"tf": 1.4142135623730951}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_backward": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_update": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_index_select": {"tf": 1}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.rearange.triton_split_and_cat": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1.7320508075688772}, "linghe.utils.scatter.triton_aligned_scatter_add": {"tf": 1}, "linghe.utils.transpose.triton_transpose_and_pad": {"tf": 1}}, "df": 14, "m": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "t": {"docs": {"linghe.quant.group.triton_group_quant": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}}, "df": 2}}}, "w": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "d": {"docs": {"linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}, "linghe.utils.silu.triton_weighted_silu_backward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}}, "df": 5}}}}, "c": {"docs": {}, "df": 0, "e": {"docs": {"linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}}, "df": 1}}}}, "r": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "q": {"docs": {}, "df": 0, "s": {"docs": {"linghe.facade.rope.qk_norm_half_rope": {"tf": 1.4142135623730951}, "linghe.utils.rope.triton_half_rope_forward": {"tf": 1.4142135623730951}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1.4142135623730951}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}}, "df": 4}}}}, "p": {"1": {"6": {"docs": {}, "df": 0, "/": {"docs": {}, "df": 0, "b": {"docs": {}, "df": 0, "f": {"1": {"6": {"docs": {"linghe.gemm.fp32_gemm.triton_fp32_gemm": {"tf": 1.7320508075688772}}, "df": 1}, "docs": {}, "df": 0}, "docs": {}, "df": 0}}}}, "docs": {}, "df": 0}, "3": {"2": {"docs": {"linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1.4142135623730951}, "linghe.gemm.fp32_gemm.triton_fp32_gemm": {"tf": 1.4142135623730951}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_backward": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_update": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm_for_update": {"tf": 1}, "linghe.utils.reduce.triton_batch_sum_with_ord": {"tf": 1}}, "df": 7}, "docs": {}, "df": 0}, "8": {"docs": {"linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1.4142135623730951}}, "df": 1}, "docs": {}, "df": 0}, "u": {"docs": {}, "df": 0, "s": {"docs": {}, "df": 0, "e": {"docs": {"linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1.4142135623730951}}, "df": 1, "d": {"docs": {"linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1}}, "df": 3}}}}, "a": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "s": {"docs": {}, "df": 0, "e": {"docs": {"linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}}, "df": 1}}}}}, "s": {"docs": {"linghe.facade.rope.qk_norm_half_rope": {"tf": 2}, "linghe.gemm.fp32_gemm.triton_fp32_gemm": {"tf": 1}, "linghe.quant.block.triton_block_quant": {"tf": 1}, "linghe.quant.group.triton_group_quant": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 2.449489742783178}}, "df": 5, "e": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "v": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "c": {"docs": {}, "df": 0, "e": {"docs": {"linghe.facade.add": {"tf": 1}, "linghe.facade.fp32_gemm": {"tf": 1}, "linghe.facade.hadamard_quant_linear": {"tf": 1}, "linghe.facade.loss": {"tf": 1}, "linghe.facade.norm": {"tf": 1}, "linghe.facade.rope": {"tf": 1}, "linghe.facade.smooth_quant_linear": {"tf": 1}, "linghe.facade.transpose": {"tf": 1}, "linghe.gemm.blockwise_fp8_gemm": {"tf": 1}, "linghe.gemm.channelwise_fp8_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm": {"tf": 1}, "linghe.quant.block": {"tf": 1}, "linghe.quant.channel": {"tf": 1}, "linghe.quant.group": {"tf": 1}, "linghe.quant.hadamard": {"tf": 1}, "linghe.quant.smooth": {"tf": 1}, "linghe.utils.add": {"tf": 1}, "linghe.utils.dot": {"tf": 1}, "linghe.utils.gather": {"tf": 1}, "linghe.utils.loss": {"tf": 1}, "linghe.utils.rearange": {"tf": 1}, "linghe.utils.reduce": {"tf": 1}, "linghe.utils.rope": {"tf": 1}, "linghe.utils.scatter": {"tf": 1}, "linghe.utils.silu": {"tf": 1}, "linghe.utils.transpose": {"tf": 1}}, "df": 26}}}}}, "e": {"docs": {"linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm_for_update": {"tf": 1}}, "df": 1}, "l": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "c": {"docs": {}, "df": 0, "t": {"docs": {"linghe.utils.gather.triton_index_select": {"tf": 1}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1}}, "df": 3, "e": {"docs": {}, "df": 0, "d": {"docs": {"linghe.utils.gather.triton_index_select": {"tf": 1.4142135623730951}}, "df": 1}}}}}}, "t": {"docs": {"linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}}, "df": 1}}, "o": {"docs": {"linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1}}, "df": 1, "f": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "m": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "x": {"docs": {"linghe.facade.loss.softmax_cross_entropy": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_forward": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_backward": {"tf": 1}}, "df": 3}}}}}}, "h": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "p": {"docs": {}, "df": 0, "e": {"docs": {"linghe.facade.loss.softmax_cross_entropy": {"tf": 1.4142135623730951}, "linghe.facade.norm.group_rms_norm_gate": {"tf": 2}, "linghe.facade.rope.qk_norm_half_rope": {"tf": 1.7320508075688772}, "linghe.utils.gather.triton_make_row_id_map": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1.4142135623730951}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1.7320508075688772}}, "df": 6}}, "r": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "d": {"docs": {"linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1}}, "df": 1}}}}, "o": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "d": {"docs": {"linghe.facade.transpose.transpose_dim01": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1}}, "df": 2}}}}}, "i": {"docs": {}, "df": 0, "g": {"docs": {}, "df": 0, "m": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "d": {"docs": {"linghe.facade.norm.group_rms_norm_gate": {"tf": 1}}, "df": 1}}}}}, "z": {"docs": {}, "df": 0, "e": {"docs": {"linghe.facade.norm.group_rms_norm_gate": {"tf": 1.4142135623730951}, "linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm": {"tf": 1}, "linghe.quant.block.triton_block_quant": {"tf": 1}, "linghe.quant.channel.triton_tokenwise_row_quant": {"tf": 1}, "linghe.quant.group.triton_group_quant": {"tf": 1}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1.7320508075688772}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1.4142135623730951}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}, "linghe.utils.scatter.triton_unpermute_with_mask_map": {"tf": 1.4142135623730951}, "linghe.utils.transpose.triton_transpose_and_pad": {"tf": 1}}, "df": 12}}, "m": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "r": {"docs": {"linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}}, "df": 2}}}}}, "n": {"docs": {}, "df": 0, "g": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "e": {"docs": {"linghe.utils.reduce.triton_batch_count_zero": {"tf": 1}, "linghe.utils.reduce.triton_batch_sum_with_ord": {"tf": 1}}, "df": 2}}}}, "l": {"docs": {}, "df": 0, "u": {"docs": {"linghe.utils.silu.triton_weighted_silu_forward": {"tf": 1}, "linghe.utils.silu.triton_weighted_silu_backward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}}, "df": 6}}}, "p": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "t": {"docs": {"linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.utils.rearange.triton_split_and_cat": {"tf": 1.7320508075688772}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}}, "df": 3, "s": {"docs": {"linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}}, "df": 3}}}}}, "m": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "h": {"docs": {"linghe.facade.smooth_quant_linear.SmoothQuantLinear": {"tf": 1}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 2}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1.4142135623730951}, "linghe.utils.reduce.triton_abs_max": {"tf": 1.7320508075688772}}, "df": 6}}}}}, "c": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "e": {"docs": {"linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 2}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1.7320508075688772}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm_for_update": {"tf": 1}, "linghe.quant.block.triton_block_quant": {"tf": 1.7320508075688772}, "linghe.quant.channel.triton_row_quant": {"tf": 2}, "linghe.quant.channel.triton_tokenwise_row_quant": {"tf": 2}, "linghe.quant.channel.triton_transpose_row_quant": {"tf": 2}, "linghe.quant.group.triton_group_quant": {"tf": 1.7320508075688772}, "linghe.quant.hadamard.triton_hadamard_quant": {"tf": 2}, "linghe.utils.gather.triton_index_select": {"tf": 2}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 2}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 2.449489742783178}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 2}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1.7320508075688772}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 2.449489742783178}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 2.6457513110645907}, "linghe.utils.reduce.triton_abs_max": {"tf": 2}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 2.449489742783178}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 1.7320508075688772}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 2.449489742783178}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 2}}, "df": 21, "d": {"docs": {"linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm_for_update": {"tf": 1}, "linghe.utils.transpose.triton_transpose_and_pad": {"tf": 1}}, "df": 4}, "s": {"docs": {"linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.rearange.triton_split_and_cat": {"tf": 2}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 1.4142135623730951}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1.4142135623730951}}, "df": 8}}}, "t": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "r": {"docs": {"linghe.utils.scatter.triton_aligned_scatter_add": {"tf": 1}, "linghe.utils.scatter.triton_scatter_add": {"tf": 1}, "linghe.utils.scatter.triton_unpermute_with_mask_map": {"tf": 1}}, "df": 3}}}}}}, "u": {"docs": {}, "df": 0, "p": {"docs": {}, "df": 0, "p": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "t": {"docs": {"linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}}, "df": 2}}}}}, "i": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "b": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "e": {"docs": {"linghe.gemm.fp32_gemm.triton_fp32_gemm": {"tf": 1}}, "df": 1}}}}}}, "m": {"docs": {"linghe.utils.dot.triton_dot": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1.7320508075688772}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_backward": {"tf": 1}, "linghe.utils.reduce.triton_batch_sum_with_ord": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1}}, "df": 7}}, "t": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "l": {"docs": {"linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1}}, "df": 1}}}, "o": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "d": {"docs": {"linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1}}, "df": 1}}}}}, "q": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "e": {"docs": {"linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}}, "df": 1}}}}}, "l": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "w": {"docs": {"linghe.utils.scatter.triton_scatter_add": {"tf": 1}}, "df": 1}}}}, "g": {"docs": {"linghe.utils.silu.triton_weighted_silu_backward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}}, "df": 3, "r": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "p": {"docs": {"linghe.facade.add": {"tf": 1}, "linghe.facade.fp32_gemm": {"tf": 1}, "linghe.facade.hadamard_quant_linear": {"tf": 1}, "linghe.facade.loss": {"tf": 1}, "linghe.facade.norm": {"tf": 1}, "linghe.facade.norm.group_rms_norm_gate": {"tf": 2}, "linghe.facade.rope": {"tf": 1}, "linghe.facade.smooth_quant_linear": {"tf": 1}, "linghe.facade.transpose": {"tf": 1}, "linghe.gemm.blockwise_fp8_gemm": {"tf": 1}, "linghe.gemm.channelwise_fp8_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm": {"tf": 1}, "linghe.quant.block": {"tf": 1}, "linghe.quant.channel": {"tf": 1}, "linghe.quant.group": {"tf": 1}, "linghe.quant.group.triton_group_quant": {"tf": 1.7320508075688772}, "linghe.quant.hadamard": {"tf": 1}, "linghe.quant.smooth": {"tf": 1}, "linghe.utils.add": {"tf": 1}, "linghe.utils.dot": {"tf": 1}, "linghe.utils.gather": {"tf": 1}, "linghe.utils.loss": {"tf": 1}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1.7320508075688772}, "linghe.utils.rearange": {"tf": 1}, "linghe.utils.reduce": {"tf": 1}, "linghe.utils.rope": {"tf": 1}, "linghe.utils.scatter": {"tf": 1}, "linghe.utils.silu": {"tf": 1}, "linghe.utils.transpose": {"tf": 1}}, "df": 29, "w": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "s": {"docs": {}, "df": 0, "e": {"docs": {"linghe.quant.group.triton_group_quant": {"tf": 1}}, "df": 1}}}}}}}, "a": {"docs": {}, "df": 0, "d": {"docs": {"linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1.4142135623730951}, "linghe.utils.loss.triton_softmax_cross_entropy_backward": {"tf": 1.4142135623730951}, "linghe.utils.scatter.triton_unpermute_with_mask_map": {"tf": 1}}, "df": 3, "i": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "t": {"docs": {"linghe.facade.loss.softmax_cross_entropy": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_backward": {"tf": 1.4142135623730951}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_update": {"tf": 1.4142135623730951}, "linghe.utils.dot.triton_dot": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_backward": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1.4142135623730951}, "linghe.utils.reduce.triton_batch_count_zero": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 2.449489742783178}, "linghe.utils.scatter.triton_unpermute_with_mask_map": {"tf": 1}, "linghe.utils.silu.triton_weighted_silu_backward": {"tf": 1.7320508075688772}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 2.23606797749979}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 2.449489742783178}, "linghe.utils.transpose.triton_transpose_and_pad": {"tf": 1}}, "df": 13}}}}, "s": {"docs": {"linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}}, "df": 1}}}, "e": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "r": {"docs": {"linghe.facade.transpose.transpose_dim01": {"tf": 1}}, "df": 1}}}}}}, "e": {"docs": {}, "df": 0, "m": {"docs": {}, "df": 0, "m": {"docs": {"linghe.facade.fp32_gemm.fp32_gemm": {"tf": 1.7320508075688772}, "linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm": {"tf": 1.7320508075688772}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_backward": {"tf": 1.4142135623730951}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_update": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1.4142135623730951}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm_for_update": {"tf": 1}}, "df": 7}}}, "a": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "e": {"docs": {"linghe.facade.norm.group_rms_norm_gate": {"tf": 1.7320508075688772}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 2}}, "df": 2}, "h": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "r": {"docs": {"linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}, "linghe.utils.scatter.triton_aligned_scatter_add": {"tf": 1}}, "df": 3, "e": {"docs": {}, "df": 0, "d": {"docs": {"linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}}, "df": 1}}}}}}}, "t": {"docs": {"linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}}, "df": 1}, "q": {"docs": {"linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}}, "df": 1}, "k": {"docs": {"linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}}, "df": 1}, "v": {"docs": {"linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}}, "df": 1}}, "i": {"docs": {}, "df": 0, "t": {"docs": {"linghe.gemm.fp32_gemm.triton_fp32_gemm": {"tf": 1}, "linghe.utils.dot.triton_dot": {"tf": 1}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.rearange.triton_split_and_cat": {"tf": 1}, "linghe.utils.reduce.triton_abs_max": {"tf": 1}, "linghe.utils.reduce.triton_batch_count_zero": {"tf": 1}, "linghe.utils.scatter.triton_aligned_scatter_add": {"tf": 1}, "linghe.utils.transpose.triton_transpose_and_pad": {"tf": 1}}, "df": 8, "s": {"docs": {"linghe.facade.add": {"tf": 1}, "linghe.facade.fp32_gemm": {"tf": 1}, "linghe.facade.hadamard_quant_linear": {"tf": 1}, "linghe.facade.loss": {"tf": 1}, "linghe.facade.norm": {"tf": 1}, "linghe.facade.rope": {"tf": 1}, "linghe.facade.smooth_quant_linear": {"tf": 1}, "linghe.facade.transpose": {"tf": 1}, "linghe.gemm.blockwise_fp8_gemm": {"tf": 1}, "linghe.gemm.channelwise_fp8_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm": {"tf": 1}, "linghe.quant.block": {"tf": 1}, "linghe.quant.channel": {"tf": 1}, "linghe.quant.group": {"tf": 1}, "linghe.quant.hadamard": {"tf": 1}, "linghe.quant.smooth": {"tf": 1}, "linghe.utils.add": {"tf": 1}, "linghe.utils.dot": {"tf": 1}, "linghe.utils.gather": {"tf": 1}, "linghe.utils.loss": {"tf": 1}, "linghe.utils.rearange": {"tf": 1}, "linghe.utils.reduce": {"tf": 1}, "linghe.utils.rope": {"tf": 1}, "linghe.utils.scatter": {"tf": 1}, "linghe.utils.silu": {"tf": 1}, "linghe.utils.transpose": {"tf": 1}}, "df": 26}}, "n": {"docs": {"linghe.facade.fp32_gemm.fp32_gemm": {"tf": 1}, "linghe.facade.hadamard_quant_linear.HadamardQuantLinear.__init__": {"tf": 1.4142135623730951}, "linghe.facade.loss.softmax_cross_entropy": {"tf": 1}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear.__init__": {"tf": 1.4142135623730951}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1}, "linghe.quant.group.triton_group_quant": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1}, "linghe.utils.rearange.triton_split_and_cat": {"tf": 1}, "linghe.utils.reduce.triton_abs_max": {"tf": 1}, "linghe.utils.reduce.triton_batch_count_zero": {"tf": 1.4142135623730951}, "linghe.utils.scatter.triton_aligned_scatter_add": {"tf": 1}, "linghe.utils.silu.triton_weighted_silu_forward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1.4142135623730951}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1}}, "df": 20, "p": {"docs": {"linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}}, "df": 2, "l": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "c": {"docs": {}, "df": 0, "e": {"docs": {"linghe.facade.add.inplace_add": {"tf": 1}, "linghe.facade.loss.softmax_cross_entropy": {"tf": 1}, "linghe.utils.add.triton_inplace_add": {"tf": 1}}, "df": 3}}}}, "u": {"docs": {}, "df": 0, "t": {"docs": {"linghe.facade.fp32_gemm.fp32_gemm": {"tf": 1}, "linghe.facade.transpose.transpose_dim01": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_backward": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_update": {"tf": 1}, "linghe.quant.block.triton_block_quant": {"tf": 1}, "linghe.quant.channel.triton_row_quant": {"tf": 1}, "linghe.quant.channel.triton_tokenwise_row_quant": {"tf": 1}, "linghe.quant.channel.triton_transpose_row_quant": {"tf": 1}, "linghe.quant.group.triton_group_quant": {"tf": 1}, "linghe.quant.hadamard.triton_hadamard_quant": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_backward": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_forward": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1.4142135623730951}, "linghe.utils.reduce.triton_abs_max": {"tf": 1}, "linghe.utils.reduce.triton_batch_count_zero": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}, "linghe.utils.scatter.triton_aligned_scatter_add": {"tf": 1}, "linghe.utils.scatter.triton_scatter_add": {"tf": 1}, "linghe.utils.silu.triton_weighted_silu_forward": {"tf": 1}, "linghe.utils.silu.triton_weighted_silu_backward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.transpose.triton_transpose": {"tf": 1}, "linghe.utils.transpose.triton_transpose_and_pad": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose": {"tf": 1}}, "df": 26, "s": {"docs": {"linghe.facade.fp32_gemm.fp32_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm": {"tf": 1}}, "df": 2}}}}, "t": {"6": {"4": {"docs": {"linghe.utils.reduce.triton_batch_count_zero": {"tf": 1}}, "df": 1}, "docs": {}, "df": 0}, "docs": {"linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}}, "df": 2, "e": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "v": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "d": {"docs": {"linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 2.23606797749979}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}}, "df": 3}}}}}}}}}, "d": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "c": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "s": {"docs": {"linghe.utils.gather.triton_make_row_id_map": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_index_select": {"tf": 1}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}, "linghe.utils.rearange.triton_split_and_cat": {"tf": 1.4142135623730951}, "linghe.utils.scatter.triton_aligned_scatter_add": {"tf": 1.4142135623730951}, "linghe.utils.scatter.triton_scatter_add": {"tf": 1.4142135623730951}}, "df": 11}}, "a": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "s": {"docs": {"linghe.utils.gather.triton_make_row_id_map": {"tf": 1}}, "df": 1}}}}}}, "e": {"docs": {}, "df": 0, "x": {"docs": {"linghe.utils.gather.triton_index_select": {"tf": 1}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 2}}, "df": 2}}}, "c": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "d": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "g": {"docs": {"linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}}, "df": 1}}}}}}}, "v": {"docs": {"linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}}, "df": 1}}, "m": {"docs": {}, "df": 0, "p": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "m": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "n": {"docs": {"linghe.facade.hadamard_quant_linear.HadamardQuantLinear": {"tf": 1}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1.4142135623730951}}, "df": 3}}}}}}}}}}}}}, "f": {"docs": {"linghe.facade.loss.softmax_cross_entropy": {"tf": 1}, "linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1}, "linghe.utils.add.triton_inplace_add": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1.4142135623730951}, "linghe.utils.rearange.triton_split_and_cat": {"tf": 1}, "linghe.utils.reduce.triton_abs_max": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1}}, "df": 8}, "s": {"docs": {"linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1}, "linghe.quant.group.triton_group_quant": {"tf": 1}, "linghe.utils.dot.triton_dot": {"tf": 1}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1}, "linghe.utils.rearange.triton_split_and_cat": {"tf": 1.4142135623730951}, "linghe.utils.reduce.triton_abs_max": {"tf": 1.4142135623730951}, "linghe.utils.reduce.triton_batch_count_zero": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1.4142135623730951}, "linghe.utils.scatter.triton_aligned_scatter_add": {"tf": 1}, "linghe.utils.transpose.triton_transpose_and_pad": {"tf": 1}}, "df": 14}, "d": {"docs": {"linghe.utils.gather.triton_make_row_id_map": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1.7320508075688772}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}, "linghe.utils.scatter.triton_unpermute_with_mask_map": {"tf": 1.7320508075688772}}, "df": 5}, "g": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "e": {"docs": {"linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}}, "df": 1}}}}}}, "y": {"docs": {"linghe.facade.add.inplace_add": {"tf": 1.4142135623730951}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 2.449489742783178}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm_for_update": {"tf": 1}, "linghe.quant.block.triton_block_quant": {"tf": 1}, "linghe.quant.group.triton_group_quant": {"tf": 1}, "linghe.utils.add.triton_inplace_add": {"tf": 2}, "linghe.utils.dot.triton_dot": {"tf": 1}, "linghe.utils.rearange.triton_split_and_cat": {"tf": 1}}, "df": 8, "@": {"docs": {}, "df": 0, "w": {"docs": {"linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1}}, "df": 1}}}, "t": {"docs": {"linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1}}, "df": 1, "o": {"docs": {"linghe.facade.add.inplace_add": {"tf": 1.7320508075688772}, "linghe.facade.rope.qk_norm_half_rope": {"tf": 1.7320508075688772}, "linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1.4142135623730951}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1}, "linghe.quant.block.triton_block_quant": {"tf": 1}, "linghe.quant.channel.triton_row_quant": {"tf": 1}, "linghe.quant.channel.triton_tokenwise_row_quant": {"tf": 1}, "linghe.quant.channel.triton_transpose_row_quant": {"tf": 1}, "linghe.quant.group.triton_group_quant": {"tf": 1}, "linghe.utils.add.triton_inplace_add": {"tf": 1}, "linghe.utils.dot.triton_dot": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.rearange.triton_split_and_cat": {"tf": 1}, "linghe.utils.reduce.triton_batch_count_zero": {"tf": 1}, "linghe.utils.rope.triton_half_rope_forward": {"tf": 1.4142135623730951}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1.7320508075688772}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.transpose.triton_transpose_and_pad": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1}}, "df": 26, "k": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "n": {"docs": {"linghe.facade.loss.softmax_cross_entropy": {"tf": 1}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1.7320508075688772}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_forward": {"tf": 1.4142135623730951}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1.4142135623730951}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1.4142135623730951}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1}}, "df": 9, "s": {"docs": {"linghe.utils.gather.triton_make_row_id_map": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 2.449489742783178}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 2.449489742783178}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1.7320508075688772}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 2.23606797749979}, "linghe.utils.scatter.triton_unpermute_with_mask_map": {"tf": 2.23606797749979}}, "df": 7}, "w": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "s": {"docs": {}, "df": 0, "e": {"docs": {"linghe.utils.silu.triton_weighted_silu_forward": {"tf": 1}}, "df": 1}}}}}}}, "r": {"docs": {}, "df": 0, "c": {"docs": {}, "df": 0, "h": {"docs": {"linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1}, "linghe.utils.transpose.triton_transpose_and_pad": {"tf": 1}}, "df": 2}}}, "p": {"docs": {}, "df": 0, "k": {"docs": {"linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1.4142135623730951}}, "df": 1}}}, "e": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "s": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "r": {"docs": {"linghe.facade.add.inplace_add": {"tf": 1}, "linghe.facade.fp32_gemm.fp32_gemm": {"tf": 1.4142135623730951}, "linghe.facade.loss.softmax_cross_entropy": {"tf": 2}, "linghe.facade.norm.rms_norm": {"tf": 1.4142135623730951}, "linghe.facade.norm.group_rms_norm_gate": {"tf": 1}, "linghe.facade.rope.qk_norm_half_rope": {"tf": 1.4142135623730951}, "linghe.facade.transpose.transpose_dim01": {"tf": 1.7320508075688772}, "linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 2.449489742783178}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 2}, "linghe.quant.block.triton_block_quant": {"tf": 1.4142135623730951}, "linghe.quant.channel.triton_row_quant": {"tf": 1}, "linghe.quant.channel.triton_tokenwise_row_quant": {"tf": 1}, "linghe.quant.channel.triton_transpose_row_quant": {"tf": 1}, "linghe.quant.group.triton_group_quant": {"tf": 1.4142135623730951}, "linghe.quant.hadamard.triton_hadamard_quant": {"tf": 2}, "linghe.utils.add.triton_inplace_add": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_make_row_id_map": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}, "linghe.utils.gather.triton_index_select": {"tf": 1}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 2.23606797749979}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1.4142135623730951}, "linghe.utils.loss.triton_softmax_cross_entropy_forward": {"tf": 1.4142135623730951}, "linghe.utils.loss.triton_softmax_cross_entropy_backward": {"tf": 1.4142135623730951}, "linghe.utils.norm.triton_rms_norm_forward": {"tf": 1.4142135623730951}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1.7320508075688772}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1.4142135623730951}, "linghe.utils.rearange.triton_split_and_cat": {"tf": 1}, "linghe.utils.reduce.triton_abs_max": {"tf": 1.7320508075688772}, "linghe.utils.reduce.triton_batch_count_zero": {"tf": 1.7320508075688772}, "linghe.utils.reduce.triton_batch_sum_with_ord": {"tf": 1.7320508075688772}, "linghe.utils.rope.triton_half_rope_forward": {"tf": 1.4142135623730951}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1.4142135623730951}, "linghe.utils.scatter.triton_aligned_scatter_add": {"tf": 1.7320508075688772}, "linghe.utils.scatter.triton_scatter_add": {"tf": 1.7320508075688772}, "linghe.utils.scatter.triton_unpermute_with_mask_map": {"tf": 1}, "linghe.utils.silu.triton_weighted_silu_forward": {"tf": 1.4142135623730951}, "linghe.utils.silu.triton_weighted_silu_backward": {"tf": 1.7320508075688772}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 2.23606797749979}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 2.6457513110645907}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1.7320508075688772}, "linghe.utils.transpose.triton_transpose": {"tf": 1.4142135623730951}, "linghe.utils.transpose.triton_transpose_and_pad": {"tf": 1.4142135623730951}, "linghe.utils.transpose.triton_batch_transpose": {"tf": 1.4142135623730951}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1.4142135623730951}}, "df": 46, "s": {"docs": {"linghe.utils.rearange.triton_split_and_cat": {"tf": 1}, "linghe.utils.reduce.triton_batch_count_zero": {"tf": 1}}, "df": 2}}}}}}, "r": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "s": {"docs": {}, "df": 0, "f": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "m": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "n": {"docs": {"linghe.facade.hadamard_quant_linear.HadamardQuantLinear": {"tf": 1}, "linghe.quant.hadamard.triton_hadamard_quant": {"tf": 1}}, "df": 2}}}}}, "e": {"docs": {}, "df": 0, "d": {"docs": {"linghe.quant.hadamard.triton_hadamard_quant": {"tf": 1}}, "df": 1}}}}}}, "p": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "s": {"docs": {}, "df": 0, "e": {"docs": {"linghe.facade.norm.group_rms_norm_gate": {"tf": 1}, "linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.facade.transpose.transpose_dim01": {"tf": 1}, "linghe.quant.channel.triton_transpose_row_quant": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1.7320508075688772}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1.7320508075688772}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1.4142135623730951}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1.4142135623730951}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 1.4142135623730951}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1.4142135623730951}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1.4142135623730951}, "linghe.utils.transpose.triton_transpose": {"tf": 1}, "linghe.utils.transpose.triton_transpose_and_pad": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1}}, "df": 15, "d": {"docs": {"linghe.facade.transpose.transpose_dim01": {"tf": 1}, "linghe.quant.hadamard.triton_hadamard_quant": {"tf": 2}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1.7320508075688772}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1.4142135623730951}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1.7320508075688772}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 2}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 2}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 2}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 2}, "linghe.utils.transpose.triton_transpose": {"tf": 1}}, "df": 11}}}}}}, "p": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "s": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "d": {"docs": {"linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}}, "df": 2}}}}}}, "i": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "g": {"docs": {"linghe.utils.silu.triton_weighted_silu_forward": {"tf": 1}}, "df": 1}}}}}}, "u": {"docs": {}, "df": 0, "e": {"docs": {"linghe.facade.loss.softmax_cross_entropy": {"tf": 1}, "linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1}, "linghe.utils.add.triton_inplace_add": {"tf": 1}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1.4142135623730951}}, "df": 4}}, "i": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "n": {"docs": {"linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm_for_update": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}, "linghe.utils.silu.triton_weighted_silu_backward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}}, "df": 6}}}}}, "h": {"docs": {}, "df": 0, "e": {"docs": {"linghe.facade.loss.softmax_cross_entropy": {"tf": 1}, "linghe.facade.transpose.transpose_dim01": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map": {"tf": 1.7320508075688772}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1.7320508075688772}, "linghe.utils.reduce.triton_batch_sum_with_ord": {"tf": 1}, "linghe.utils.transpose.triton_transpose_and_pad": {"tf": 1}}, "df": 7, "n": {"docs": {"linghe.quant.hadamard.triton_hadamard_quant": {"tf": 1}, "linghe.utils.rope.triton_half_rope_forward": {"tf": 1}}, "df": 2}}, "a": {"docs": {}, "df": 0, "n": {"docs": {"linghe.facade.transpose.transpose_dim01": {"tf": 1}}, "df": 1}}, "i": {"docs": {}, "df": 0, "s": {"docs": {"linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}}, "df": 3}}}, "w": {"docs": {}, "df": 0, "o": {"docs": {"linghe.facade.transpose.transpose_dim01": {"tf": 1}}, "df": 1}}}, "x": {"docs": {"linghe.facade.add.inplace_add": {"tf": 2}, "linghe.facade.norm.rms_norm": {"tf": 1.4142135623730951}, "linghe.facade.transpose.transpose_dim01": {"tf": 1.4142135623730951}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1.4142135623730951}, "linghe.quant.block.triton_block_quant": {"tf": 1.4142135623730951}, "linghe.quant.channel.triton_row_quant": {"tf": 2.23606797749979}, "linghe.quant.channel.triton_tokenwise_row_quant": {"tf": 1.7320508075688772}, "linghe.quant.channel.triton_transpose_row_quant": {"tf": 2.449489742783178}, "linghe.quant.group.triton_group_quant": {"tf": 1.4142135623730951}, "linghe.quant.hadamard.triton_hadamard_quant": {"tf": 2.6457513110645907}, "linghe.utils.add.triton_inplace_add": {"tf": 2.23606797749979}, "linghe.utils.dot.triton_dot": {"tf": 1}, "linghe.utils.gather.triton_index_select": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 2}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1.7320508075688772}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1.4142135623730951}, "linghe.utils.norm.triton_rms_norm_forward": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1}, "linghe.utils.rearange.triton_split_and_cat": {"tf": 1.4142135623730951}, "linghe.utils.reduce.triton_abs_max": {"tf": 2}, "linghe.utils.reduce.triton_batch_sum_with_ord": {"tf": 1}, "linghe.utils.scatter.triton_aligned_scatter_add": {"tf": 1}, "linghe.utils.scatter.triton_scatter_add": {"tf": 1}, "linghe.utils.silu.triton_weighted_silu_forward": {"tf": 1.4142135623730951}, "linghe.utils.silu.triton_weighted_silu_backward": {"tf": 1.4142135623730951}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.transpose.triton_transpose": {"tf": 1.4142135623730951}, "linghe.utils.transpose.triton_transpose_and_pad": {"tf": 1.4142135623730951}, "linghe.utils.transpose.triton_batch_transpose": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1.7320508075688772}}, "df": 34, "/": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "m": {"docs": {}, "df": 0, "s": {"docs": {"linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1}}, "df": 1}}}}, "t": {"docs": {"linghe.quant.hadamard.triton_hadamard_quant": {"tf": 1.4142135623730951}}, "df": 1, "s": {"docs": {"linghe.utils.transpose.triton_batch_transpose": {"tf": 1}}, "df": 1}}, "*": {"docs": {}, "df": 0, "y": {"docs": {"linghe.utils.dot.triton_dot": {"tf": 1.4142135623730951}}, "df": 1}}, "s": {"docs": {"linghe.utils.reduce.triton_batch_count_zero": {"tf": 1}, "linghe.utils.reduce.triton_batch_sum_with_ord": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose": {"tf": 1}}, "df": 3}}, "w": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "h": {"docs": {"linghe.facade.add.inplace_add": {"tf": 1}, "linghe.facade.fp32_gemm.fp32_gemm": {"tf": 1}, "linghe.facade.norm.rms_norm": {"tf": 1}, "linghe.facade.norm.group_rms_norm_gate": {"tf": 1}, "linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.facade.transpose.transpose_dim01": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm": {"tf": 2}, "linghe.quant.channel.triton_tokenwise_row_quant": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.rearange.triton_split_and_cat": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}, "linghe.utils.scatter.triton_unpermute_with_mask_map": {"tf": 1}, "linghe.utils.silu.triton_weighted_silu_forward": {"tf": 1}, "linghe.utils.transpose.triton_transpose": {"tf": 1}, "linghe.utils.transpose.triton_transpose_and_pad": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1}}, "df": 18}}, "s": {"docs": {}, "df": 0, "e": {"docs": {"linghe.quant.block.triton_block_quant": {"tf": 1}, "linghe.quant.group.triton_group_quant": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_forward": {"tf": 1}}, "df": 3}}, "l": {"docs": {}, "df": 0, "l": {"docs": {"linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1}}, "df": 1}}}, "e": {"docs": {"linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1.7320508075688772}}, "df": 1, "i": {"docs": {}, "df": 0, "g": {"docs": {}, "df": 0, "h": {"docs": {}, "df": 0, "t": {"docs": {"linghe.facade.fp32_gemm.fp32_gemm": {"tf": 1.4142135623730951}, "linghe.facade.hadamard_quant_linear.HadamardQuantLinear.__init__": {"tf": 1.4142135623730951}, "linghe.facade.norm.rms_norm": {"tf": 1.7320508075688772}, "linghe.facade.norm.group_rms_norm_gate": {"tf": 1.7320508075688772}, "linghe.facade.rope.qk_norm_half_rope": {"tf": 2}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear.__init__": {"tf": 1.4142135623730951}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_backward": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_update": {"tf": 1.4142135623730951}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1}, "linghe.utils.dot.triton_dot": {"tf": 1}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_forward": {"tf": 1.4142135623730951}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1.4142135623730951}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1.4142135623730951}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 2}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 2}, "linghe.utils.scatter.triton_aligned_scatter_add": {"tf": 1}, "linghe.utils.silu.triton_weighted_silu_forward": {"tf": 1.7320508075688772}, "linghe.utils.silu.triton_weighted_silu_backward": {"tf": 1.7320508075688772}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1.4142135623730951}, "linghe.utils.transpose.triton_transpose_and_pad": {"tf": 1}}, "df": 22, "s": {"docs": {"linghe.utils.scatter.triton_aligned_scatter_add": {"tf": 1}}, "df": 1}, "e": {"docs": {}, "df": 0, "d": {"docs": {"linghe.utils.silu.triton_weighted_silu_backward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}}, "df": 2}}}}}}, "l": {"docs": {}, "df": 0, "l": {"docs": {"linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}}, "df": 1}}}, "h": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "h": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "r": {"docs": {"linghe.facade.hadamard_quant_linear.HadamardQuantLinear.__init__": {"tf": 1}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear.__init__": {"tf": 1}, "linghe.quant.block.triton_block_quant": {"tf": 1}, "linghe.quant.channel.triton_row_quant": {"tf": 1}, "linghe.quant.channel.triton_tokenwise_row_quant": {"tf": 1}, "linghe.quant.channel.triton_transpose_row_quant": {"tf": 1}, "linghe.quant.group.triton_group_quant": {"tf": 1}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1.4142135623730951}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1.4142135623730951}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.transpose.triton_transpose_and_pad": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1}}, "df": 18}}}}}}}, "m": {"docs": {"linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose": {"tf": 1}}, "df": 2, "i": {"docs": {}, "df": 0, "x": {"docs": {"linghe.facade.add.inplace_add": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_backward": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_update": {"tf": 1}}, "df": 3}, "n": {"docs": {"linghe.utils.reduce.triton_abs_max": {"tf": 1.4142135623730951}}, "df": 1}}, "o": {"docs": {}, "df": 0, "e": {"docs": {"linghe.facade.fp32_gemm.fp32_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.rearange.triton_split_and_cat": {"tf": 1}, "linghe.utils.scatter.triton_aligned_scatter_add": {"tf": 1}, "linghe.utils.silu.triton_weighted_silu_forward": {"tf": 1}}, "df": 7}, "d": {"docs": {}, "df": 0, "e": {"docs": {"linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}}, "df": 3}}, "n": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "r": {"docs": {"linghe.utils.reduce.triton_batch_count_zero": {"tf": 1}}, "df": 1}}}}}}, "m": {"docs": {"linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1}, "linghe.utils.transpose.triton_transpose_and_pad": {"tf": 1}}, "df": 2}, "a": {"docs": {}, "df": 0, "j": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "r": {"docs": {"linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1}}, "df": 1}}}, "i": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "y": {"docs": {"linghe.gemm.fp32_gemm.triton_fp32_gemm": {"tf": 1}}, "df": 1}}}}, "t": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "x": {"docs": {"linghe.gemm.fp32_gemm.triton_fp32_gemm": {"tf": 1.4142135623730951}, "linghe.quant.hadamard.triton_hadamard_quant": {"tf": 1}}, "df": 2}}}}, "k": {"docs": {}, "df": 0, "e": {"docs": {"linghe.utils.gather.triton_make_row_id_map": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}}, "df": 2}}, "p": {"docs": {"linghe.utils.gather.triton_make_row_id_map": {"tf": 1.7320508075688772}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1.7320508075688772}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1.7320508075688772}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}, "linghe.utils.scatter.triton_unpermute_with_mask_map": {"tf": 1.7320508075688772}}, "df": 5}, "y": {"docs": {"linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.reduce.triton_abs_max": {"tf": 1}}, "df": 2}, "x": {"docs": {"linghe.utils.loss.triton_softmax_cross_entropy_backward": {"tf": 1}, "linghe.utils.reduce.triton_abs_max": {"tf": 2}}, "df": 2}}, "u": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "p": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "y": {"docs": {"linghe.utils.dot.triton_dot": {"tf": 1}}, "df": 1}, "e": {"docs": {"linghe.utils.gather.triton_make_row_id_map": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.rearange.triton_split_and_cat": {"tf": 1}}, "df": 4}}}}}}, "t": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "p": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "r": {"docs": {"linghe.utils.transpose.triton_transpose_and_pad": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1}}, "df": 2}}}}}}}}, "e": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "n": {"docs": {"linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}}, "df": 1, "s": {"docs": {"linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}}, "df": 1}}}, "g": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "n": {"docs": {"linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.scatter.triton_aligned_scatter_add": {"tf": 1}}, "df": 3}}}}}}}, "]": {"docs": {}, "df": 0, "*": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "x": {"docs": {}, "df": 0, "p": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "t": {"docs": {"linghe.utils.transpose.triton_batch_transpose": {"tf": 1}}, "df": 1}}}}}}}}}, "p": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "c": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "s": {"docs": {}, "df": 0, "e": {"docs": {"linghe.facade.add.inplace_add": {"tf": 1}}, "df": 1}, "i": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "n": {"docs": {"linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm": {"tf": 1.7320508075688772}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_backward": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_update": {"tf": 1}}, "df": 4}}}}}}}, "o": {"docs": {}, "df": 0, "b": {"docs": {"linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1.4142135623730951}, "linghe.utils.scatter.triton_aligned_scatter_add": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}}, "df": 4, "s": {"docs": {"linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1.4142135623730951}, "linghe.utils.scatter.triton_unpermute_with_mask_map": {"tf": 1.4142135623730951}}, "df": 2}}}}, "e": {"docs": {}, "df": 0, "r": {"docs": {"linghe.facade.loss.softmax_cross_entropy": {"tf": 1}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 2.449489742783178}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1.4142135623730951}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1.4142135623730951}}, "df": 7, "m": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "d": {"docs": {"linghe.utils.gather.triton_permute_with_mask_map": {"tf": 2.23606797749979}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1.4142135623730951}}, "df": 2}}, "a": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "n": {"docs": {"linghe.utils.rearange.triton_split_and_cat": {"tf": 1}}, "df": 1}}}}}}}}}}, "o": {"docs": {}, "df": 0, "w": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "r": {"docs": {"linghe.quant.block.triton_block_quant": {"tf": 1}, "linghe.quant.channel.triton_row_quant": {"tf": 1}, "linghe.quant.channel.triton_tokenwise_row_quant": {"tf": 1.4142135623730951}, "linghe.quant.channel.triton_transpose_row_quant": {"tf": 1}, "linghe.quant.group.triton_group_quant": {"tf": 1}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}}, "df": 12}}}}, "a": {"docs": {}, "df": 0, "d": {"docs": {"linghe.utils.transpose.triton_transpose_and_pad": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1.7320508075688772}}, "df": 2, "d": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "g": {"docs": {"linghe.utils.gather.triton_make_row_id_map": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.transpose.triton_transpose_and_pad": {"tf": 1.4142135623730951}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1}}, "df": 5}}}, "e": {"docs": {}, "df": 0, "d": {"docs": {"linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1}}, "df": 3}}}}}, "y": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "h": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "n": {"docs": {"linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1}}, "df": 3}}}}}}, "b": {"docs": {"linghe.facade.rope.qk_norm_half_rope": {"tf": 2}, "linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1.7320508075688772}, "linghe.gemm.fp32_gemm.triton_fp32_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_backward": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_update": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1.4142135623730951}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm_for_update": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 2.449489742783178}}, "df": 8, "e": {"docs": {"linghe.facade.add.inplace_add": {"tf": 1}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1}, "linghe.utils.reduce.triton_abs_max": {"tf": 1}, "linghe.utils.transpose.triton_transpose_and_pad": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1}}, "df": 6, "f": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "e": {"docs": {"linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm_for_update": {"tf": 1}}, "df": 1}}}}}, "f": {"1": {"6": {"docs": {"linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1}}, "df": 1, "/": {"docs": {}, "df": 0, "f": {"docs": {}, "df": 0, "p": {"1": {"6": {"docs": {"linghe.facade.fp32_gemm.fp32_gemm": {"tf": 1.7320508075688772}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_backward": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_update": {"tf": 1}, "linghe.utils.silu.triton_weighted_silu_forward": {"tf": 1}}, "df": 4}, "docs": {}, "df": 0}, "docs": {}, "df": 0}}}}, "docs": {}, "df": 0}, "docs": {}, "df": 0}, "i": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "s": {"docs": {"linghe.facade.hadamard_quant_linear.HadamardQuantLinear.__init__": {"tf": 1.4142135623730951}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear.__init__": {"tf": 1.4142135623730951}}, "df": 2}}}, "s": {"docs": {"linghe.facade.norm.group_rms_norm_gate": {"tf": 1.7320508075688772}, "linghe.utils.gather.triton_index_select": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 2}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1.4142135623730951}, "linghe.utils.loss.triton_softmax_cross_entropy_backward": {"tf": 2.449489742783178}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 2.23606797749979}, "linghe.utils.rearange.triton_split_and_cat": {"tf": 1.4142135623730951}, "linghe.utils.rope.triton_half_rope_forward": {"tf": 1.4142135623730951}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1.7320508075688772}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1}}, "df": 11, "*": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "p": {"docs": {}, "df": 0, "k": {"docs": {"linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1.4142135623730951}}, "df": 2}}}}}}, "a": {"docs": {}, "df": 0, "s": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "d": {"docs": {"linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}}, "df": 2}}}, "c": {"docs": {}, "df": 0, "k": {"docs": {}, "df": 0, "w": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "d": {"docs": {"linghe.gemm.fp32_gemm.triton_fp32_gemm_for_backward": {"tf": 1}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_backward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}, "linghe.utils.silu.triton_weighted_silu_backward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}}, "df": 7}}}}}}, "t": {"docs": {}, "df": 0, "c": {"docs": {}, "df": 0, "h": {"docs": {"linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose": {"tf": 1}}, "df": 2}}}}, "l": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "c": {"docs": {}, "df": 0, "k": {"docs": {"linghe.quant.block.triton_block_quant": {"tf": 1.4142135623730951}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}}, "df": 4, "w": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "s": {"docs": {}, "df": 0, "e": {"docs": {"linghe.quant.block.triton_block_quant": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}}, "df": 3}}}}, "i": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "g": {"docs": {"linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}}, "df": 1}}}}}}}, "u": {"docs": {}, "df": 0, "t": {"docs": {"linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}}, "df": 1}}, "o": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "h": {"docs": {"linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}}, "df": 3}}}}, "u": {"docs": {}, "df": 0, "p": {"docs": {"linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}}, "df": 1, "d": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "e": {"docs": {"linghe.facade.loss.softmax_cross_entropy": {"tf": 1}}, "df": 1, "d": {"docs": {"linghe.facade.add.inplace_add": {"tf": 1.4142135623730951}, "linghe.utils.add.triton_inplace_add": {"tf": 1}}, "df": 2}}}, "i": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "g": {"docs": {"linghe.gemm.fp32_gemm.triton_fp32_gemm_for_update": {"tf": 1}}, "df": 1}}}}}}, "s": {"docs": {}, "df": 0, "e": {"docs": {"linghe.facade.hadamard_quant_linear.HadamardQuantLinear.__init__": {"tf": 1}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear.__init__": {"tf": 1}}, "df": 2, "d": {"docs": {"linghe.facade.fp32_gemm.fp32_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1}, "linghe.utils.dot.triton_dot": {"tf": 1}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.rearange.triton_split_and_cat": {"tf": 1}, "linghe.utils.reduce.triton_abs_max": {"tf": 1}, "linghe.utils.reduce.triton_batch_count_zero": {"tf": 1}, "linghe.utils.silu.triton_weighted_silu_forward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.transpose.triton_transpose_and_pad": {"tf": 1}}, "df": 13}}}, "n": {"docs": {}, "df": 0, "d": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "r": {"docs": {"linghe.quant.group.triton_group_quant": {"tf": 1}}, "df": 1}}}}}, "r": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "n": {"docs": {"linghe.facade.norm.group_rms_norm_gate": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.reduce.triton_batch_sum_with_ord": {"tf": 1}}, "df": 4, "s": {"docs": {"linghe.facade.add.inplace_add": {"tf": 1}, "linghe.facade.fp32_gemm.fp32_gemm": {"tf": 1}, "linghe.facade.loss.softmax_cross_entropy": {"tf": 1}, "linghe.facade.norm.rms_norm": {"tf": 1}, "linghe.facade.norm.group_rms_norm_gate": {"tf": 1}, "linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.facade.transpose.transpose_dim01": {"tf": 1}, "linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_backward": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_update": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm_for_update": {"tf": 1}, "linghe.quant.block.triton_block_quant": {"tf": 1}, "linghe.quant.channel.triton_row_quant": {"tf": 1}, "linghe.quant.channel.triton_tokenwise_row_quant": {"tf": 1}, "linghe.quant.channel.triton_transpose_row_quant": {"tf": 1}, "linghe.quant.group.triton_group_quant": {"tf": 1}, "linghe.quant.hadamard.triton_hadamard_quant": {"tf": 1}, "linghe.utils.add.triton_inplace_add": {"tf": 1}, "linghe.utils.dot.triton_dot": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}, "linghe.utils.gather.triton_index_select": {"tf": 1}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_forward": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_backward": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_forward": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1}, "linghe.utils.rearange.triton_split_and_cat": {"tf": 1}, "linghe.utils.reduce.triton_abs_max": {"tf": 1}, "linghe.utils.reduce.triton_batch_count_zero": {"tf": 1}, "linghe.utils.reduce.triton_batch_sum_with_ord": {"tf": 1}, "linghe.utils.rope.triton_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}, "linghe.utils.scatter.triton_aligned_scatter_add": {"tf": 1}, "linghe.utils.scatter.triton_scatter_add": {"tf": 1}, "linghe.utils.scatter.triton_unpermute_with_mask_map": {"tf": 1}, "linghe.utils.silu.triton_weighted_silu_forward": {"tf": 1}, "linghe.utils.silu.triton_weighted_silu_backward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.transpose.triton_transpose": {"tf": 1}, "linghe.utils.transpose.triton_transpose_and_pad": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1}}, "df": 54}}}}}, "s": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "t": {"docs": {"linghe.gemm.fp32_gemm.triton_fp32_gemm": {"tf": 1}}, "df": 1}}}, "t": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "e": {"docs": {"linghe.utils.scatter.triton_unpermute_with_mask_map": {"tf": 1}}, "df": 1}}}}}, "q": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "z": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "d": {"docs": {"linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}}, "df": 1}}}}}}}}}, "v": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "s": {"docs": {}, "df": 0, "e": {"docs": {"linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}}, "df": 3}}}}}, "c": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "p": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "c": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "l": {"docs": {"linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}}, "df": 1}}}}}}}}, "d": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "c": {"docs": {}, "df": 0, "e": {"docs": {"linghe.utils.reduce.triton_abs_max": {"tf": 1}}, "df": 1}}}}}, "o": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "r": {"docs": {"linghe.facade.fp32_gemm.fp32_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1.7320508075688772}, "linghe.utils.dot.triton_dot": {"tf": 1}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1.4142135623730951}, "linghe.utils.scatter.triton_aligned_scatter_add": {"tf": 1.4142135623730951}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}}, "df": 8}, "d": {"docs": {"linghe.utils.gather.triton_make_row_id_map": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1.4142135623730951}}, "df": 2}}, "i": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "g": {"docs": {"linghe.utils.gather.triton_make_row_id_map": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}}, "df": 2}}}}, "n": {"docs": {}, "df": 0, "d": {"docs": {"linghe.quant.block.triton_block_quant": {"tf": 1.4142135623730951}, "linghe.quant.channel.triton_row_quant": {"tf": 1.4142135623730951}, "linghe.quant.channel.triton_tokenwise_row_quant": {"tf": 1.4142135623730951}, "linghe.quant.channel.triton_transpose_row_quant": {"tf": 1.4142135623730951}, "linghe.quant.group.triton_group_quant": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1.4142135623730951}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 1.4142135623730951}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1.4142135623730951}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1.4142135623730951}}, "df": 15, "u": {"docs": {}, "df": 0, "p": {"docs": {"linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1.4142135623730951}}, "df": 1}}}}}, "p": {"docs": {}, "df": 0, "e": {"docs": {"linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.utils.rope.triton_half_rope_forward": {"tf": 1.4142135623730951}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}}, "df": 4}}, "w": {"docs": {"linghe.quant.channel.triton_transpose_row_quant": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map": {"tf": 1.7320508075688772}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1.7320508075688772}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 2}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}, "linghe.utils.scatter.triton_unpermute_with_mask_map": {"tf": 1.7320508075688772}}, "df": 6, "w": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "s": {"docs": {}, "df": 0, "e": {"docs": {"linghe.quant.channel.triton_row_quant": {"tf": 1}, "linghe.quant.channel.triton_tokenwise_row_quant": {"tf": 1}, "linghe.quant.group.triton_group_quant": {"tf": 1}, "linghe.quant.hadamard.triton_hadamard_quant": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1.4142135623730951}, "linghe.utils.scatter.triton_aligned_scatter_add": {"tf": 1}}, "df": 7}}}}}, "o": {"docs": {}, "df": 0, "t": {"docs": {"linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}}, "df": 1}}}, "m": {"docs": {}, "df": 0, "s": {"docs": {"linghe.facade.norm.rms_norm": {"tf": 1.7320508075688772}, "linghe.facade.norm.group_rms_norm_gate": {"tf": 2}, "linghe.facade.rope.qk_norm_half_rope": {"tf": 1.4142135623730951}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 2}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm_for_update": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_forward": {"tf": 1.7320508075688772}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1.7320508075688772}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1.7320508075688772}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1.4142135623730951}}, "df": 9, "n": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "m": {"docs": {"linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1.4142135623730951}}, "df": 2}}}}}}, "i": {"docs": {}, "df": 0, "g": {"docs": {}, "df": 0, "h": {"docs": {}, "df": 0, "t": {"docs": {"linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm": {"tf": 1}}, "df": 2}}}}}, "o": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "t": {"docs": {"linghe.facade.hadamard_quant_linear.HadamardQuantLinear.__init__": {"tf": 1.4142135623730951}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear.__init__": {"tf": 1.4142135623730951}, "linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1}, "linghe.quant.channel.triton_tokenwise_row_quant": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_index_select": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_forward": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1.4142135623730951}, "linghe.utils.scatter.triton_unpermute_with_mask_map": {"tf": 1.4142135623730951}, "linghe.utils.silu.triton_weighted_silu_forward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.transpose.triton_transpose_and_pad": {"tf": 1.4142135623730951}}, "df": 15, "p": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "t": {"docs": {"linghe.facade.fp32_gemm.fp32_gemm": {"tf": 1.4142135623730951}, "linghe.facade.norm.rms_norm": {"tf": 1}, "linghe.facade.norm.group_rms_norm_gate": {"tf": 2.23606797749979}, "linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 2.449489742783178}, "linghe.gemm.fp32_gemm.triton_fp32_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_update": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1}, "linghe.utils.dot.triton_dot": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}, "linghe.utils.gather.triton_index_select": {"tf": 1}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1.7320508075688772}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1.4142135623730951}, "linghe.utils.loss.triton_softmax_cross_entropy_backward": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_forward": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 2.6457513110645907}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1.7320508075688772}, "linghe.utils.rearange.triton_split_and_cat": {"tf": 1.7320508075688772}, "linghe.utils.reduce.triton_abs_max": {"tf": 1}, "linghe.utils.rope.triton_half_rope_forward": {"tf": 1.4142135623730951}, "linghe.utils.scatter.triton_aligned_scatter_add": {"tf": 1.4142135623730951}, "linghe.utils.scatter.triton_scatter_add": {"tf": 1.4142135623730951}, "linghe.utils.scatter.triton_unpermute_with_mask_map": {"tf": 1}, "linghe.utils.silu.triton_weighted_silu_forward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 2.6457513110645907}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 2.6457513110645907}, "linghe.utils.transpose.triton_transpose_and_pad": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1}}, "df": 28, "s": {"docs": {"linghe.utils.scatter.triton_aligned_scatter_add": {"tf": 1}, "linghe.utils.scatter.triton_scatter_add": {"tf": 1}}, "df": 2}}}}}}, "f": {"docs": {"linghe.facade.fp32_gemm.fp32_gemm": {"tf": 1}, "linghe.facade.hadamard_quant_linear.HadamardQuantLinear": {"tf": 1}, "linghe.facade.loss.softmax_cross_entropy": {"tf": 1}, "linghe.facade.norm.rms_norm": {"tf": 1}, "linghe.facade.norm.group_rms_norm_gate": {"tf": 1.7320508075688772}, "linghe.facade.rope.qk_norm_half_rope": {"tf": 1.7320508075688772}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear": {"tf": 1}, "linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1.4142135623730951}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_backward": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_update": {"tf": 1.4142135623730951}, "linghe.quant.block.triton_block_quant": {"tf": 1}, "linghe.quant.channel.triton_row_quant": {"tf": 1}, "linghe.quant.channel.triton_tokenwise_row_quant": {"tf": 1.4142135623730951}, "linghe.quant.channel.triton_transpose_row_quant": {"tf": 1}, "linghe.quant.group.triton_group_quant": {"tf": 1}, "linghe.quant.hadamard.triton_hadamard_quant": {"tf": 2}, "linghe.utils.dot.triton_dot": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_make_row_id_map": {"tf": 2}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1.7320508075688772}, "linghe.utils.gather.triton_index_select": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 2}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_forward": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_backward": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_forward": {"tf": 1.4142135623730951}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 3}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1.7320508075688772}, "linghe.utils.reduce.triton_abs_max": {"tf": 1}, "linghe.utils.reduce.triton_batch_sum_with_ord": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 2}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 2.6457513110645907}, "linghe.utils.scatter.triton_scatter_add": {"tf": 1}, "linghe.utils.silu.triton_weighted_silu_backward": {"tf": 1.7320508075688772}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 2}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 2}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 2.449489742783178}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 2.6457513110645907}, "linghe.utils.transpose.triton_transpose_and_pad": {"tf": 1.4142135623730951}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1.4142135623730951}}, "df": 40}, "n": {"docs": {"linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}}, "df": 3, "e": {"docs": {"linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}}, "df": 3}, "l": {"docs": {}, "df": 0, "y": {"docs": {"linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1.4142135623730951}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1.4142135623730951}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1.4142135623730951}}, "df": 4}}}, "r": {"docs": {}, "df": 0, "g": {"docs": {"linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}}, "df": 1}, "d": {"docs": {"linghe.utils.reduce.triton_batch_sum_with_ord": {"tf": 1.4142135623730951}}, "df": 1, "e": {"docs": {}, "df": 0, "r": {"docs": {"linghe.utils.reduce.triton_batch_sum_with_ord": {"tf": 1}}, "df": 1}}}}, "p": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "l": {"docs": {"linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1.4142135623730951}, "linghe.utils.reduce.triton_abs_max": {"tf": 1}}, "df": 2}}}}}}}, "v": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "r": {"docs": {"linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}}, "df": 1}}}}, "n": {"docs": {"linghe.facade.norm.group_rms_norm_gate": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 2.23606797749979}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1.7320508075688772}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1.7320508075688772}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1.7320508075688772}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1.7320508075688772}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1.4142135623730951}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1}, "linghe.utils.rearange.triton_split_and_cat": {"tf": 1.4142135623730951}, "linghe.utils.scatter.triton_unpermute_with_mask_map": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1}}, "df": 14, "a": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "v": {"docs": {}, "df": 0, "e": {"docs": {"linghe.facade.hadamard_quant_linear.HadamardQuantLinear": {"tf": 1}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear": {"tf": 1}, "linghe.utils.scatter.triton_scatter_add": {"tf": 1}}, "df": 3}}}, "t": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "v": {"docs": {}, "df": 0, "e": {"docs": {"linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1}}, "df": 1}}}}}, "u": {"docs": {}, "df": 0, "m": {"docs": {"linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 2.23606797749979}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 2.449489742783178}, "linghe.utils.scatter.triton_unpermute_with_mask_map": {"tf": 2.449489742783178}}, "df": 4, "b": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "r": {"docs": {"linghe.facade.hadamard_quant_linear.HadamardQuantLinear.__init__": {"tf": 1.4142135623730951}, "linghe.facade.rope.qk_norm_half_rope": {"tf": 1.4142135623730951}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear.__init__": {"tf": 1.4142135623730951}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1.4142135623730951}}, "df": 4}}}}}, "o": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "m": {"docs": {"linghe.facade.norm.rms_norm": {"tf": 1}, "linghe.facade.norm.group_rms_norm_gate": {"tf": 1.7320508075688772}, "linghe.facade.rope.qk_norm_half_rope": {"tf": 2.23606797749979}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1.7320508075688772}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm_for_update": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_forward": {"tf": 1.7320508075688772}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 2}, "linghe.utils.rope.triton_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 2.23606797749979}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 2.23606797749979}}, "df": 10, "a": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "z": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "n": {"docs": {"linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}}, "df": 3}}}}}}}}}}}, "t": {"docs": {"linghe.facade.transpose.transpose_dim01": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}, "linghe.utils.rearange.triton_split_and_cat": {"tf": 1}}, "df": 6, "e": {"docs": {"linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}}, "df": 1}}, "n": {"docs": {"linghe.quant.hadamard.triton_hadamard_quant": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1.4142135623730951}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 1.4142135623730951}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1.4142135623730951}}, "df": 8, "e": {"docs": {"linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}, "linghe.utils.rearange.triton_split_and_cat": {"tf": 1}}, "df": 3}}}, "d": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "m": {"docs": {}, "df": 0, "s": {"docs": {"linghe.facade.transpose.transpose_dim01": {"tf": 1}}, "df": 1}}}}, "e": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "d": {"docs": {"linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1.4142135623730951}, "linghe.utils.transpose.triton_transpose_and_pad": {"tf": 1}}, "df": 2}}}, "]": {"docs": {}, "df": 0, "*": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "x": {"docs": {}, "df": 0, "p": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "t": {"docs": {"linghe.utils.transpose.triton_batch_transpose": {"tf": 1}}, "df": 1}}}}}}}}}, "h": {"docs": {"linghe.facade.rope.qk_norm_half_rope": {"tf": 2.23606797749979}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 2.23606797749979}}, "df": 2, "a": {"docs": {}, "df": 0, "d": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "m": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "d": {"docs": {"linghe.facade.hadamard_quant_linear.HadamardQuantLinear": {"tf": 1}, "linghe.quant.hadamard.triton_hadamard_quant": {"tf": 1}}, "df": 2}}}}}}, "l": {"docs": {}, "df": 0, "f": {"docs": {"linghe.facade.rope.qk_norm_half_rope": {"tf": 1.4142135623730951}, "linghe.utils.rope.triton_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1.4142135623730951}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}}, "df": 4}}, "m": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "d": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "d": {"docs": {"linghe.quant.hadamard.triton_hadamard_quant": {"tf": 1}}, "df": 1}}}}}}}, "e": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "d": {"docs": {"linghe.facade.norm.group_rms_norm_gate": {"tf": 1}, "linghe.facade.rope.qk_norm_half_rope": {"tf": 1.7320508075688772}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1}, "linghe.utils.rope.triton_half_rope_forward": {"tf": 2}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 2}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 2.449489742783178}}, "df": 6, "s": {"docs": {"linghe.facade.norm.group_rms_norm_gate": {"tf": 1}, "linghe.facade.rope.qk_norm_half_rope": {"tf": 1.7320508075688772}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1.7320508075688772}}, "df": 4}}}}, "m": {"docs": {"linghe.quant.hadamard.triton_hadamard_quant": {"tf": 1}}, "df": 1}, "i": {"docs": {}, "df": 0, "d": {"docs": {}, "df": 0, "d": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "n": {"docs": {"linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1.7320508075688772}, "linghe.utils.scatter.triton_unpermute_with_mask_map": {"tf": 1.4142135623730951}}, "df": 3}}}}}, "s": {"docs": {"linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}}, "df": 1}}, "q": {"docs": {"linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1.4142135623730951}, "linghe.quant.channel.triton_row_quant": {"tf": 1}, "linghe.quant.channel.triton_transpose_row_quant": {"tf": 1}, "linghe.quant.hadamard.triton_hadamard_quant": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.rope.triton_half_rope_forward": {"tf": 1.4142135623730951}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1.7320508075688772}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 2.23606797749979}}, "df": 11, "u": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "t": {"docs": {"linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}}, "df": 5, "i": {"docs": {}, "df": 0, "z": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "n": {"docs": {"linghe.facade.hadamard_quant_linear.HadamardQuantLinear": {"tf": 1}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 2}, "linghe.quant.block.triton_block_quant": {"tf": 1}, "linghe.quant.channel.triton_row_quant": {"tf": 1}, "linghe.quant.channel.triton_tokenwise_row_quant": {"tf": 1}, "linghe.quant.channel.triton_transpose_row_quant": {"tf": 1}, "linghe.quant.group.triton_group_quant": {"tf": 1}, "linghe.quant.hadamard.triton_hadamard_quant": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1.4142135623730951}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 2.6457513110645907}, "linghe.utils.reduce.triton_abs_max": {"tf": 1.4142135623730951}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1.7320508075688772}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 1.4142135623730951}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1.4142135623730951}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1.4142135623730951}}, "df": 17}}}}}, "e": {"docs": {"linghe.quant.block.triton_block_quant": {"tf": 1}, "linghe.quant.channel.triton_row_quant": {"tf": 1}, "linghe.quant.channel.triton_tokenwise_row_quant": {"tf": 1}, "linghe.quant.channel.triton_transpose_row_quant": {"tf": 1}, "linghe.quant.group.triton_group_quant": {"tf": 1}, "linghe.quant.hadamard.triton_hadamard_quant": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}}, "df": 7, "d": {"docs": {"linghe.quant.block.triton_block_quant": {"tf": 1}, "linghe.quant.channel.triton_row_quant": {"tf": 1}, "linghe.quant.channel.triton_tokenwise_row_quant": {"tf": 1}, "linghe.quant.channel.triton_transpose_row_quant": {"tf": 1}, "linghe.quant.group.triton_group_quant": {"tf": 1}, "linghe.quant.hadamard.triton_hadamard_quant": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_index_select": {"tf": 1}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1.7320508075688772}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}, "linghe.utils.reduce.triton_abs_max": {"tf": 1.4142135623730951}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 2}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 1.4142135623730951}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 2}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1.4142135623730951}}, "df": 15}}}}}}}, "e": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "y": {"docs": {"linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.utils.rope.triton_half_rope_forward": {"tf": 1.4142135623730951}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}}, "df": 3}}}}, "k": {"docs": {"linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.utils.rope.triton_half_rope_forward": {"tf": 1.4142135623730951}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1.4142135623730951}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}}, "df": 4, "v": {"docs": {"linghe.facade.rope.qk_norm_half_rope": {"tf": 1.7320508075688772}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 2.449489742783178}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1.7320508075688772}}, "df": 3, "q": {"docs": {"linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}}, "df": 1}}}, "/": {"docs": {}, "df": 0, "k": {"docs": {"linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}}, "df": 2, "/": {"docs": {}, "df": 0, "v": {"docs": {"linghe.facade.rope.qk_norm_half_rope": {"tf": 1.4142135623730951}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1.4142135623730951}}, "df": 2}}}}, "o": {"docs": {"linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.utils.rope.triton_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}}, "df": 4}}, "d": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "v": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "c": {"docs": {}, "df": 0, "e": {"docs": {"linghe.facade.hadamard_quant_linear.HadamardQuantLinear.__init__": {"tf": 1.4142135623730951}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear.__init__": {"tf": 1.4142135623730951}}, "df": 2}}}}, "q": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "t": {"docs": {"linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 1}}, "df": 1}}}}}}, "t": {"docs": {}, "df": 0, "y": {"docs": {}, "df": 0, "p": {"docs": {}, "df": 0, "e": {"docs": {"linghe.facade.hadamard_quant_linear.HadamardQuantLinear.__init__": {"tf": 1.4142135623730951}, "linghe.facade.smooth_quant_linear.SmoothQuantLinear.__init__": {"tf": 1.4142135623730951}, "linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1.4142135623730951}}, "df": 3}}}}, "i": {"docs": {}, "df": 0, "m": {"0": {"docs": {"linghe.utils.transpose.triton_transpose": {"tf": 1.4142135623730951}}, "df": 1}, "1": {"docs": {"linghe.utils.transpose.triton_transpose": {"tf": 1.4142135623730951}}, "df": 1}, "docs": {"linghe.facade.loss.softmax_cross_entropy": {"tf": 1}, "linghe.facade.norm.group_rms_norm_gate": {"tf": 2}, "linghe.facade.rope.qk_norm_half_rope": {"tf": 2.23606797749979}, "linghe.quant.channel.triton_tokenwise_row_quant": {"tf": 1}, "linghe.utils.gather.triton_index_select": {"tf": 1}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 2}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 2}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1.7320508075688772}, "linghe.utils.loss.triton_softmax_cross_entropy_backward": {"tf": 1.7320508075688772}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 2.449489742783178}, "linghe.utils.rearange.triton_split_and_cat": {"tf": 1}, "linghe.utils.rope.triton_half_rope_forward": {"tf": 1.4142135623730951}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 2.6457513110645907}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1.7320508075688772}, "linghe.utils.transpose.triton_transpose": {"tf": 1.4142135623730951}}, "df": 15, "s": {"docs": {"linghe.facade.transpose.transpose_dim01": {"tf": 1}}, "df": 1}, "e": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "s": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "n": {"docs": {"linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}}, "df": 1}}}}}}}}, "o": {"docs": {"linghe.gemm.fp32_gemm.triton_fp32_gemm": {"tf": 1}}, "df": 1, "t": {"docs": {"linghe.utils.dot.triton_dot": {"tf": 1}}, "df": 1}}, "w": {"docs": {"linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm_for_update": {"tf": 1}, "linghe.utils.silu.triton_weighted_silu_backward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}}, "df": 3}, "y": {"docs": {"linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}}, "df": 1}, "a": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "a": {"docs": {"linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1.7320508075688772}}, "df": 2}}}, "q": {"docs": {}, "df": 0, "k": {"docs": {}, "df": 0, "v": {"docs": {"linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}}, "df": 1}}, "w": {"docs": {"linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}}, "df": 1}}, "k": {"docs": {}, "df": 0, "w": {"docs": {"linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}}, "df": 1}}, "x": {"docs": {"linghe.utils.silu.triton_weighted_silu_backward": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_backward": {"tf": 2}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 2}}, "df": 3}}, "e": {"4": {"docs": {}, "df": 0, "m": {"3": {"docs": {}, "df": 0, "f": {"docs": {}, "df": 0, "n": {"docs": {"linghe.quant.block.triton_block_quant": {"tf": 1}, "linghe.quant.group.triton_group_quant": {"tf": 1}}, "df": 2}}}, "docs": {}, "df": 0}}, "docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "p": {"docs": {}, "df": 0, "y": {"docs": {"linghe.facade.loss.softmax_cross_entropy": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_forward": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_backward": {"tf": 1}}, "df": 3}}}}}}, "p": {"docs": {}, "df": 0, "s": {"docs": {"linghe.facade.norm.rms_norm": {"tf": 1}, "linghe.facade.norm.group_rms_norm_gate": {"tf": 1}, "linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_forward": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}}, "df": 8, "i": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "n": {"docs": {"linghe.facade.norm.rms_norm": {"tf": 1}, "linghe.facade.norm.group_rms_norm_gate": {"tf": 1}, "linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_forward": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}}, "df": 7}}}}}}, "l": {"docs": {}, "df": 0, "s": {"docs": {}, "df": 0, "e": {"docs": {"linghe.utils.add.triton_inplace_add": {"tf": 1}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1.4142135623730951}}, "df": 2}}}, "a": {"docs": {}, "df": 0, "c": {"docs": {}, "df": 0, "h": {"docs": {"linghe.utils.gather.triton_make_row_id_map": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_forward": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1}}, "df": 4}}}, "x": {"docs": {}, "df": 0, "p": {"docs": {"linghe.utils.loss.triton_softmax_cross_entropy_backward": {"tf": 1}}, "df": 1, "e": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "t": {"docs": {"linghe.utils.gather.triton_make_row_id_map": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1.7320508075688772}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.silu.triton_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1.4142135623730951}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1.4142135623730951}}, "df": 9, "s": {"docs": {"linghe.utils.gather.triton_make_row_id_map": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_permute_with_mask_map": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 2.449489742783178}, "linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1.4142135623730951}, "linghe.utils.gather.triton_smooth_permute_with_mask_map": {"tf": 2}, "linghe.utils.scatter.triton_unpermute_with_mask_map": {"tf": 1.4142135623730951}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1.4142135623730951}}, "df": 9, "*": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "p": {"docs": {}, "df": 0, "k": {"docs": {"linghe.utils.gather.triton_smooth_weighted_permute_with_indices": {"tf": 1}, "linghe.utils.gather.triton_smooth_permute_with_indices": {"tf": 1}}, "df": 2}}}}}}}}}}}}, "l": {"2": {"docs": {"linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}}, "df": 3}, "docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "g": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "t": {"docs": {"linghe.utils.loss.triton_softmax_cross_entropy_backward": {"tf": 1.4142135623730951}}, "df": 1, "s": {"docs": {"linghe.facade.loss.softmax_cross_entropy": {"tf": 1.7320508075688772}, "linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1.4142135623730951}, "linghe.utils.loss.triton_softmax_cross_entropy_forward": {"tf": 1.4142135623730951}, "linghe.utils.loss.triton_softmax_cross_entropy_backward": {"tf": 1}}, "df": 4}}}}, "s": {"docs": {}, "df": 0, "s": {"docs": {"linghe.facade.loss.softmax_cross_entropy": {"tf": 1}, "linghe.utils.loss.triton_softmax_cross_entropy_forward": {"tf": 1.4142135623730951}, "linghe.utils.loss.triton_softmax_cross_entropy_backward": {"tf": 1}}, "df": 3}}, "w": {"docs": {"linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1}}, "df": 1}}, "a": {"docs": {}, "df": 0, "b": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "l": {"docs": {"linghe.utils.loss.triton_softmax_cross_entropy_backward": {"tf": 1}}, "df": 1, "s": {"docs": {"linghe.facade.loss.softmax_cross_entropy": {"tf": 1.4142135623730951}, "linghe.utils.loss.triton_softmax_cross_entropy_forward": {"tf": 1.4142135623730951}, "linghe.utils.loss.triton_softmax_cross_entropy_backward": {"tf": 1}}, "df": 3}}}}, "y": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "t": {"docs": {"linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}}, "df": 2}}}, "e": {"docs": {}, "df": 0, "r": {"docs": {"linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1}}, "df": 1}}}, "r": {"docs": {}, "df": 0, "g": {"docs": {}, "df": 0, "e": {"docs": {"linghe.gemm.fp32_gemm.triton_fp32_gemm": {"tf": 1}}, "df": 1}}}, "s": {"docs": {}, "df": 0, "t": {"docs": {"linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}}, "df": 1}}}, "e": {"docs": {}, "df": 0, "n": {"docs": {"linghe.utils.rope.triton_half_rope_forward": {"tf": 1.4142135623730951}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1.7320508075688772}}, "df": 2, "g": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "h": {"docs": {"linghe.facade.norm.group_rms_norm_gate": {"tf": 1.7320508075688772}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 2.23606797749979}}, "df": 2}}}}, "f": {"docs": {}, "df": 0, "t": {"docs": {"linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm": {"tf": 1}, "linghe.gemm.fp32_gemm.triton_fp32_gemm": {"tf": 1}}, "df": 2}}}, "i": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "r": {"docs": {"linghe.facade.smooth_quant_linear.SmoothQuantLinear": {"tf": 1}, "linghe.utils.norm.triton_group_rms_norm_gate_forward": {"tf": 1}}, "df": 2}}}}, "s": {"docs": {}, "df": 0, "t": {"docs": {"linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices": {"tf": 1}, "linghe.utils.reduce.triton_batch_count_zero": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward": {"tf": 1}, "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose": {"tf": 1.4142135623730951}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1.4142135623730951}}, "df": 6, "s": {"docs": {"linghe.utils.reduce.triton_batch_sum_with_ord": {"tf": 1}}, "df": 1}}}}}, "k": {"docs": {"linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.utils.gather.triton_index_select": {"tf": 1}, "linghe.utils.rope.triton_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1.4142135623730951}}, "df": 5, "e": {"docs": {}, "df": 0, "y": {"docs": {"linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.utils.rope.triton_half_rope_forward": {"tf": 1.4142135623730951}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}}, "df": 3, "/": {"docs": {}, "df": 0, "v": {"docs": {}, "df": 0, "a": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "e": {"docs": {"linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}}, "df": 2}}}}}}}, "r": {"docs": {}, "df": 0, "n": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "l": {"docs": {"linghe.gemm.fp32_gemm.triton_scaled_fp32_gemm": {"tf": 1.4142135623730951}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}}, "df": 2}}}}}, "o": {"docs": {"linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.utils.rope.triton_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}}, "df": 4}, "v": {"docs": {"linghe.utils.rope.triton_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}}, "df": 2}}, "v": {"docs": {"linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}}, "df": 1, "a": {"docs": {}, "df": 0, "l": {"docs": {}, "df": 0, "u": {"docs": {}, "df": 0, "e": {"docs": {"linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map": {"tf": 1}, "linghe.utils.gather.triton_make_row_id_map_and_indices": {"tf": 1}, "linghe.utils.norm.triton_rms_norm_and_block_quant_forward": {"tf": 1}, "linghe.utils.reduce.triton_abs_max": {"tf": 1.4142135623730951}, "linghe.utils.reduce.triton_batch_count_zero": {"tf": 1}, "linghe.utils.reduce.triton_batch_sum_with_ord": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}, "linghe.utils.transpose.triton_batch_transpose_and_pad": {"tf": 1}}, "df": 9, "s": {"docs": {"linghe.utils.gather.triton_make_row_id_map": {"tf": 1.4142135623730951}}, "df": 1}}}}}, "o": {"docs": {"linghe.facade.rope.qk_norm_half_rope": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_forward": {"tf": 1}, "linghe.utils.rope.triton_qk_norm_and_half_rope_backward": {"tf": 1}}, "df": 3}, "e": {"docs": {}, "df": 0, "c": {"docs": {}, "df": 0, "t": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "r": {"docs": {"linghe.utils.dot.triton_dot": {"tf": 1}}, "df": 1}}}}, "r": {"docs": {}, "df": 0, "s": {"docs": {}, "df": 0, "i": {"docs": {}, "df": 0, "o": {"docs": {}, "df": 0, "n": {"docs": {"linghe.utils.scatter.triton_scatter_add": {"tf": 1}}, "df": 1}}}}, "y": {"docs": {"linghe.utils.scatter.triton_scatter_add": {"tf": 1}}, "df": 1}}}}, "z": {"docs": {}, "df": 0, "e": {"docs": {}, "df": 0, "r": {"docs": {}, "df": 0, "o": {"docs": {"linghe.utils.reduce.triton_batch_count_zero": {"tf": 1}}, "df": 1, "s": {"docs": {"linghe.utils.reduce.triton_batch_count_zero": {"tf": 1}}, "df": 1}}}}}}}}, "pipeline": ["trimmer"], "_isPrebuiltIndex": true}; + /** pdoc search index */const docs = [{"fullname": "linghe", "modulename": "linghe", "kind": "module", "doc": "

    \n"}, {"fullname": "linghe.attn", "modulename": "linghe.attn", "kind": "module", "doc": "

    \n"}, {"fullname": "linghe.attn.la", "modulename": "linghe.attn.la", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, {"fullname": "linghe.attn.mla", "modulename": "linghe.attn.mla", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, {"fullname": "linghe.experimental", "modulename": "linghe.experimental", "kind": "module", "doc": "

    kernels should be run with torch above 2.9.0

    \n"}, {"fullname": "linghe.experimental.demb", "modulename": "linghe.experimental.demb", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, {"fullname": "linghe.experimental.demb.triton_tp_embedding_lookup_backward", "modulename": "linghe.experimental.demb", "qualname": "triton_tp_embedding_lookup_backward", "kind": "function", "doc": "

    inplace update embedding weight gradient

    \n\n
    Arguments:
    \n\n
      \n
    • y: gradient of output
    • \n
    • x: input ids Tensor
    • \n
    • g_ptr: data_ptr of embedding weight gradient
    • \n
    \n\n
    Returns:
    \n\n
    \n

    None

    \n
    \n", "signature": "(grad_output, x, g_ptr, vocab_size, hdl, group, dtype=torch.bfloat16):", "funcdef": "def"}, {"fullname": "linghe.experimental.demb.triton_sp_embedding_lookup_backward", "modulename": "linghe.experimental.demb", "qualname": "triton_sp_embedding_lookup_backward", "kind": "function", "doc": "

    inplace update embedding weight gradient

    \n\n
    Arguments:
    \n\n
      \n
    • y: gradient of output
    • \n
    • x: input ids Tensor
    • \n
    • g_ptr: data_ptr of embedding weight gradient
    • \n
    \n\n
    Returns:
    \n\n
    \n

    None

    \n
    \n", "signature": "(\tgrad_output,\tinput_ids,\tg_ptr,\tvocab_size,\thdl,\tgroup,\tdtype=torch.bfloat16,\tgathered_input_ids=None):", "funcdef": "def"}, {"fullname": "linghe.experimental.dla", "modulename": "linghe.experimental.dla", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, {"fullname": "linghe.experimental.dmm", "modulename": "linghe.experimental.dmm", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, {"fullname": "linghe.experimental.dmm.triton_split_tp_gemm", "modulename": "linghe.experimental.dmm", "qualname": "triton_split_tp_gemm", "kind": "function", "doc": "

    tensor-parallel fc2 in the shared expert, use split-k implementation\ny = all_reduce(x @ fc2)

    \n\n
    Arguments:
    \n\n
      \n
    • a: left matrix with bf16 precision
    • \n
    • b: right matrix with bf16 precision
    • \n
    \n\n
    Returns:
    \n\n
    \n

    c: all-reduced output

    \n
    \n", "signature": "(\tx: torch.Tensor,\tw: torch.Tensor,\thdl,\tgroup: torch.distributed.distributed_c10d.ProcessGroup):", "funcdef": "def"}, {"fullname": "linghe.experimental.gmem_barrier_arrive_wait", "modulename": "linghe.experimental.gmem_barrier_arrive_wait", "kind": "module", "doc": "

    copied from https://github.com/meta-pytorch/kraken/blob/main/kraken/_ptx_utils/gmem_barrier_arrive_wait.py

    \n"}, {"fullname": "linghe.experimental.gmem_barrier_arrive_wait.wait_gmem_barrier", "modulename": "linghe.experimental.gmem_barrier_arrive_wait", "qualname": "wait_gmem_barrier", "kind": "function", "doc": "

    Wait for a global memory barrier to reach the expected state.

    \n\n

    This function implements a spin-wait loop that continuously checks a memory location\nuntil it reaches the expected value, providing synchronization across GPU threads.

    \n\n
    Arguments:
    \n\n
      \n
    • addr: Memory address of the barrier to wait on (Must be a scalar)
    • \n
    • expect: Expected value to wait for (default: 1)
    • \n
    • update: Update the barrier with once acquired (default: 0)
    • \n
    • sem: Memory semantics for the atomic operation (default: \"acquire\")
    • \n
    • scope: Scope of the atomic operation. Options: \"gpu\", \"sys\" (default: \"gpu\")
    • \n
    • op: Atomic operation type (default: \"ld\", currently only supported option)
    • \n
    \n", "signature": "(\taddr,\texpect: int = 1,\tupdate: int = 0,\tsem: int = 'acquire',\tscope: int = 'gpu',\top: int = 'ld',\tskip_sync: int = False):", "funcdef": "def"}, {"fullname": "linghe.experimental.symm_mem_barrier", "modulename": "linghe.experimental.symm_mem_barrier", "kind": "module", "doc": "

    copied from https://github.com/meta-pytorch/kraken/blob/main/kraken/_ptx_utils/symm_mem_barrier.py

    \n"}, {"fullname": "linghe.experimental.symm_mem_barrier.symm_mem_sync", "modulename": "linghe.experimental.symm_mem_barrier", "qualname": "symm_mem_sync", "kind": "function", "doc": "

    Synchronizes blocks with matching block_id across participating devices.

    \n\n

    Note: the function itself is not a system level barrier/fence. It is a\nbuilding block for expressing different synchronization patterns.

    \n\n

    Pattern 0: Ensures that all writes to symm_mem buffers from previous\nkernels across all devices are visible to the current kernel:

    \n\n
    symm_mem_sync(..., hasPreviousMemAccess=False, hasSubsequentMemAccess=True)\n
    \n\n

    Pattern 1: Ensures that all writes to symm_mem buffers from the current\nblock are visible to all remote blocks with matching blockIdx:

    \n\n
    symm_mem_sync(..., hasPreviousMemAccess=True, hasSubsequentMemAccess=True)\n
    \n\n

    Pattern 2: Ensures that symm_mem buffers read by the current kernel are safe\nfor writing by subsequent kernels across all devices.

    \n\n
    symm_mem_sync(..., hasPreviousMemAccess=True, hasSubsequentMemAccess=False)\n
    \n\n
    CUDA graph friendliness:
    \n\n
    \n

    This barrier operates through atomic operations on a zero-filled signal\n pad, which resets to a zero-filled state after each successful\n synchronization. This design eliminates the need for incrementing a\n flag from host.

    \n
    \n", "signature": "(\tsignal_pad_ptrs,\tblock_id,\trank: int,\tworld_size: int,\thasPreviousMemAccess: int = False,\thasSubsequentMemAccess: int = False):", "funcdef": "def"}, {"fullname": "linghe.experimental.test_demb", "modulename": "linghe.experimental.test_demb", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, {"fullname": "linghe.experimental.test_dla", "modulename": "linghe.experimental.test_dla", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, {"fullname": "linghe.experimental.test_dmm", "modulename": "linghe.experimental.test_dmm", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, {"fullname": "linghe.facade", "modulename": "linghe.facade", "kind": "module", "doc": "

    \n"}, {"fullname": "linghe.facade.add", "modulename": "linghe.facade.add", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, {"fullname": "linghe.facade.add.inplace_add", "modulename": "linghe.facade.add", "qualname": "inplace_add", "kind": "function", "doc": "

    inplace add y to x with mix precise

    \n\n
    Arguments:
    \n\n
      \n
    • x: to be updated
    • \n
    • y: add to x
    • \n
    \n\n
    Returns:
    \n\n
    \n

    updated x tensor

    \n
    \n", "signature": "(x: torch.Tensor, y: torch.Tensor):", "funcdef": "def"}, {"fullname": "linghe.facade.emb", "modulename": "linghe.facade.emb", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, {"fullname": "linghe.facade.emb.deprecated_fused_accumulation_embedding_lookup", "modulename": "linghe.facade.emb", "qualname": "deprecated_fused_accumulation_embedding_lookup", "kind": "function", "doc": "

    embedding lookup

    \n\n
    Arguments:
    \n\n
      \n
    • x: input ids
    • \n
    • w_ptr:
    • \n
    • g_ptr:
    • \n
    • dim:
    • \n
    • dtype:
    • \n
    • grad_dtype:
    • \n
    \n\n
    Returns:
    \n\n
    \n

    lookup output

    \n
    \n", "signature": "(x: torch.Tensor, w_ptr, g_ptr, dim, dtype, grad_dtype):", "funcdef": "def"}, {"fullname": "linghe.facade.emb.fused_accumulation_embedding_lookup", "modulename": "linghe.facade.emb", "qualname": "fused_accumulation_embedding_lookup", "kind": "function", "doc": "

    embedding lookup

    \n\n
    Arguments:
    \n\n
      \n
    • x: input ids
    • \n
    • w: embedding weight, should contain a grad_name tensor
    • \n
    \n\n
    Returns:
    \n\n
    \n

    lookup output

    \n
    \n", "signature": "(\tx: torch.Tensor,\tw: torch.nn.parameter.Parameter,\tgrad_name: str = 'grad'):", "funcdef": "def"}, {"fullname": "linghe.facade.emb.embedding_lookup", "modulename": "linghe.facade.emb", "qualname": "embedding_lookup", "kind": "function", "doc": "

    embedding lookup

    \n\n
    Arguments:
    \n\n
      \n
    • x: input ids
    • \n
    • w: embedding weight
    • \n
    \n\n
    Returns:
    \n\n
    \n

    lookup output

    \n
    \n", "signature": "(x: torch.Tensor, w: torch.nn.parameter.Parameter):", "funcdef": "def"}, {"fullname": "linghe.facade.fp32_gemm", "modulename": "linghe.facade.fp32_gemm", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, {"fullname": "linghe.facade.fp32_gemm.fp32_gemm", "modulename": "linghe.facade.fp32_gemm", "qualname": "fp32_gemm", "kind": "function", "doc": "

    gemm with bf16/fp16 inputs and float32 output,\ncurrently used in MoE router gemm.

    \n\n
    Arguments:
    \n\n
      \n
    • input: bf16/fp16 activation tensor
    • \n
    • weight: bf16/fp16 weight tensor
    • \n
    \n\n
    Returns:
    \n\n
    \n

    output of gemm

    \n
    \n", "signature": "(input: torch.Tensor, weight: torch.Tensor):", "funcdef": "def"}, {"fullname": "linghe.facade.gate", "modulename": "linghe.facade.gate", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, {"fullname": "linghe.facade.gate.group_rms_norm_gate", "modulename": "linghe.facade.gate", "qualname": "group_rms_norm_gate", "kind": "function", "doc": "

    return group_rms_norm(transpose(attn_output, [0,1]), weight) * sigmoid(gate)

    \n\n
    Arguments:
    \n\n
      \n
    • attn_output: output of core attn, shape [bs, length, n_heads, head_dim]
    • \n
    • gate: gate tensor for attention output, shape [length, bs, dim]
    • \n
    • weight: weight of RMS norm, shape [dim]
    • \n
    • eps: epsilon for RMS
    • \n
    • group_size: group size of group RMS norm
    • \n
    \n\n
    Returns:
    \n\n
    \n

    output with shape [length, bs, dim]

    \n
    \n", "signature": "(\tattn_output: torch.Tensor,\tgate: torch.Tensor,\tweight: torch.Tensor,\teps: float = 1e-06,\tgroup_size: int = 4):", "funcdef": "def"}, {"fullname": "linghe.facade.hadamard_quant_linear", "modulename": "linghe.facade.hadamard_quant_linear", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, {"fullname": "linghe.facade.hadamard_quant_linear.HadamardQuantLinear", "modulename": "linghe.facade.hadamard_quant_linear", "qualname": "HadamardQuantLinear", "kind": "class", "doc": "

    a naive implementation of hadamard transformation and quantization

    \n", "bases": "torch.nn.modules.module.Module"}, {"fullname": "linghe.facade.hadamard_quant_linear.HadamardQuantLinear.__init__", "modulename": "linghe.facade.hadamard_quant_linear", "qualname": "HadamardQuantLinear.__init__", "kind": "function", "doc": "
    Arguments:
    \n\n
      \n
    • in_features: in feature number
    • \n
    • out_features: out feature number
    • \n
    • bias: whether use bias
    • \n
    • device: weight device
    • \n
    • dtype: weight dtype
    • \n
    \n", "signature": "(\tin_features: int,\tout_features: int,\tbias: bool = True,\tdevice=None,\tdtype=None)"}, {"fullname": "linghe.facade.loss", "modulename": "linghe.facade.loss", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, {"fullname": "linghe.facade.loss.softmax_cross_entropy", "modulename": "linghe.facade.loss", "qualname": "softmax_cross_entropy", "kind": "function", "doc": "

    softmax cross entropy

    \n\n
    Arguments:
    \n\n
      \n
    • logits: logits tensor, shape [...,dim]
    • \n
    • labels: labels tensor, shape [...]
    • \n
    • inplace: update gradient in the logits tensor if True
    • \n
    \n\n
    Returns:
    \n\n
    \n

    a tensor of per token loss

    \n
    \n", "signature": "(\tlogits: torch.Tensor,\tlabels: torch.Tensor,\tignore_index: int = -100,\tinplace: bool = False,\ttp_group=None):", "funcdef": "def"}, {"fullname": "linghe.facade.loss.moe_z_loss", "modulename": "linghe.facade.loss", "qualname": "moe_z_loss", "kind": "function", "doc": "

    softmax cross entropy

    \n\n
    Arguments:
    \n\n
      \n
    • logits: logits tensor, shape [...,dim]
    • \n
    • coef: z loss coef
    • \n
    \n\n
    Returns:
    \n\n
    \n

    z loss

    \n
    \n", "signature": "(logits: torch.Tensor, coef: float = 0.001):", "funcdef": "def"}, {"fullname": "linghe.facade.mla", "modulename": "linghe.facade.mla", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, {"fullname": "linghe.facade.mla.multi_latend_attention", "modulename": "linghe.facade.mla", "qualname": "multi_latend_attention", "kind": "function", "doc": "

    inplace add y to x with mix precise

    \n\n
    Arguments:
    \n\n
      \n
    • x: to be updated
    • \n
    • y: add to x
    • \n
    \n\n
    Returns:
    \n\n
    \n

    updated x tensor

    \n
    \n", "signature": "(\tq: torch.Tensor,\tk: torch.Tensor,\tv: torch.Tensor,\tcu_seqlens: Optional[torch.Tensor] = None,\tpadded_cu_seqlens: Optional[torch.Tensor] = None,\tmax_q_length: Optional[int] = None,\tcausal: bool = True,\tsafe: bool = True,\tclip_value: float = 0.0):", "funcdef": "def"}, {"fullname": "linghe.facade.norm", "modulename": "linghe.facade.norm", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, {"fullname": "linghe.facade.norm.rms_norm", "modulename": "linghe.facade.norm", "qualname": "rms_norm", "kind": "function", "doc": "

    rms norm of x with weight

    \n\n
    Arguments:
    \n\n
      \n
    • x: activation tensor
    • \n
    • weight: weight tensor
    • \n
    • eps: epsilon for RMS
    • \n
    \n\n
    Returns:
    \n\n
    \n

    rms output

    \n
    \n", "signature": "(x: torch.Tensor, weight: torch.Tensor, eps: float = 1e-06):", "funcdef": "def"}, {"fullname": "linghe.facade.norm.BlockRMSNorm", "modulename": "linghe.facade.norm", "qualname": "BlockRMSNorm", "kind": "class", "doc": "

    Base class to create custom autograd.Function.

    \n\n

    To create a custom autograd.Function, subclass this class and implement\nthe forward() and backward() static methods. Then, to use your custom\nop in the forward pass, call the class method apply. Do not call\nforward() directly.

    \n\n

    To ensure correctness and best performance, make sure you are calling the\ncorrect methods on ctx and validating your backward function using\ntorch.autograd.gradcheck().

    \n\n

    See :ref:extending-autograd for more details on how to use this class.

    \n\n

    Examples::

    \n\n
    >>> # xdoctest: +REQUIRES(env:TORCH_DOCTEST_AUTOGRAD)\n>>> class Exp(Function):\n>>>     @staticmethod\n>>>     def forward(ctx, i):\n>>>         result = i.exp()\n>>>         ctx.save_for_backward(result)\n>>>         return result\n>>>\n>>>     @staticmethod\n>>>     def backward(ctx, grad_output):\n>>>         result, = ctx.saved_tensors\n>>>         return grad_output * result\n>>>\n>>> # Use it by calling the apply method:\n>>> # xdoctest: +SKIP\n>>> output = Exp.apply(input)\n
    \n", "bases": "torch.autograd.function.Function"}, {"fullname": "linghe.facade.norm.BlockRMSNorm.forward", "modulename": "linghe.facade.norm", "qualname": "BlockRMSNorm.forward", "kind": "function", "doc": "

    Define the forward of the custom autograd Function.

    \n\n

    This function is to be overridden by all subclasses.\nThere are two ways to define forward:

    \n\n

    Usage 1 (Combined forward and ctx)::

    \n\n
    @staticmethod\ndef forward(ctx: Any, *args: Any, **kwargs: Any) -> Any:\n    pass\n
    \n\n
      \n
    • It must accept a context ctx as the first argument, followed by any\nnumber of arguments (tensors or other types).
    • \n
    • See :ref:combining-forward-context for more details
    • \n
    \n\n

    Usage 2 (Separate forward and ctx)::

    \n\n
    @staticmethod\ndef forward(*args: Any, **kwargs: Any) -> Any:\n    pass\n\n@staticmethod\ndef setup_context(ctx: Any, inputs: Tuple[Any, ...], output: Any) -> None:\n    pass\n
    \n\n
      \n
    • The forward no longer accepts a ctx argument.
    • \n
    • Instead, you must also override the torch.autograd.Function.setup_context()\nstaticmethod to handle setting up the ctx object.\noutput is the output of the forward, inputs are a Tuple of inputs\nto the forward.
    • \n
    • See :ref:extending-autograd for more details
    • \n
    \n\n

    The context can be used to store arbitrary data that can be then\nretrieved during the backward pass. Tensors should not be stored\ndirectly on ctx (though this is not currently enforced for\nbackward compatibility). Instead, tensors should be saved either with\nctx.save_for_backward() if they are intended to be used in\nbackward (equivalently, vjp) or ctx.save_for_forward()\nif they are intended to be used for in jvp.

    \n", "signature": "(ctx, input, weight, rms, eps, quantizer, cls, is_recomputing):", "funcdef": "def"}, {"fullname": "linghe.facade.norm.BlockRMSNorm.backward", "modulename": "linghe.facade.norm", "qualname": "BlockRMSNorm.backward", "kind": "function", "doc": "

    Define a formula for differentiating the operation with backward mode automatic differentiation.

    \n\n

    This function is to be overridden by all subclasses.\n(Defining this function is equivalent to defining the vjp function.)

    \n\n

    It must accept a context ctx as the first argument, followed by\nas many outputs as the forward() returned (None will be passed in\nfor non tensor outputs of the forward function),\nand it should return as many tensors, as there were inputs to\nforward(). Each argument is the gradient w.r.t the given output,\nand each returned value should be the gradient w.r.t. the\ncorresponding input. If an input is not a Tensor or is a Tensor not\nrequiring grads, you can just pass None as a gradient for that input.

    \n\n

    The context can be used to retrieve tensors saved during the forward\npass. It also has an attribute ctx.needs_input_grad as a tuple\nof booleans representing whether each input needs gradient. E.g.,\nbackward() will have ctx.needs_input_grad[0] = True if the\nfirst input to forward() needs gradient computed w.r.t. the\noutput.

    \n", "signature": "(ctx, grad_output, grad_rms):", "funcdef": "def"}, {"fullname": "linghe.facade.permutation", "modulename": "linghe.facade.permutation", "kind": "module", "doc": "

    \n"}, {"fullname": "linghe.facade.permutation.padded_permute", "modulename": "linghe.facade.permutation", "qualname": "padded_permute", "kind": "function", "doc": "

    Permute the tokens and probs based on the mask.\nTokens with the same designated expert will be grouped together.\nThe shape of mask is [tokens, num_experts], it indicates which experts were selected\nby each token.\nWhen drop_and_pad=True, in routing_map, the number of non-zeros in each column equals to\nexpert capacity. This function exploits this feature to use ops that support cuda graph.

    \n\n
    Arguments:
    \n\n
      \n
    • tokens (torch.Tensor): The input token tensor, [num_tokens, hidden].
    • \n
    • routing_map (torch.Tensor): The sparse token to expert mapping, [num_tokens, num_experts].
    • \n
    • tokens_per_expert (torch.Tensor): cpu tensor
    • \n
    \n", "signature": "(\ttokens,\trouting_map,\ttokens_per_expert_cuda_tensor,\ttokens_per_expert_list,\tprobs: Optional[torch.Tensor] = None):", "funcdef": "def"}, {"fullname": "linghe.facade.permutation.block_padded_permute", "modulename": "linghe.facade.permutation", "qualname": "block_padded_permute", "kind": "function", "doc": "

    Permute the tokens and probs based on the mask.\nTokens with the same designated expert will be grouped together.\nThe shape of mask is [tokens, num_experts], it indicates which experts were selected\nby each token.\nWhen drop_and_pad=True, in routing_map, the number of non-zeros in each column equals to\nexpert capacity. This function exploits this feature to use ops that support cuda graph.

    \n\n
    Arguments:
    \n\n
      \n
    • tokens (torch.Tensor): The input token tensor, [num_tokens, hidden].
    • \n
    • routing_map (torch.Tensor): The sparse token to expert mapping, [num_tokens, num_experts].
    • \n
    • tokens_per_expert (torch.Tensor): cpu tensor
    • \n
    \n", "signature": "(\ttokens,\trouting_map,\ttokens_per_expert_cuda_tensor,\ttokens_per_expert_list,\tquantizers,\tcls,\tprobs: Optional[torch.Tensor] = None):", "funcdef": "def"}, {"fullname": "linghe.facade.rope", "modulename": "linghe.facade.rope", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, {"fullname": "linghe.facade.rope.qk_norm_half_rope", "modulename": "linghe.facade.rope", "qualname": "qk_norm_half_rope", "kind": "function", "doc": "

    split qkv to q/k/v, apply qk norm and half rope to q/k, transpose q/k/v to flash-attention layout

    \n\n
    Arguments:
    \n\n
      \n
    • qkv: QKV tensor with size of [S, B, dim] or [T, dim] , heads are interleaved
    • \n
    • q_norm_weight: rms norm weight for query
    • \n
    • k_norm_weight: rms norm weight for key
    • \n
    • freqs: Freqs tensor based on half dim.
    • \n
    • cu_seqlens_q: accumulated query lengths, [num_seqs + 1]
    • \n
    • cu_seqlens_kv: accumulated kv lengths, [num_seqs + 1]
    • \n
    • H: Number of attention heads.
    • \n
    • h: Number of key/value heads.
    • \n
    • eps: epsilon value for L2 normalization.
    • \n
    • cp_rank: context parallel rank
    • \n
    • cp_size: context parallel size
    • \n
    • mscale: mscale for rope
    • \n
    \n\n
    Returns:
    \n\n
    \n
      \n
    • qo: shape [B, S, H, head_dim] or [T, H, head_dim]
    • \n
    • ko: shape [B, S, h, head_dim] or [T, h, head_dim]
    • \n
    • vo: shape [B, S, h, head_dim] or [T, h, head_dim]
    • \n
    \n
    \n", "signature": "(\tqkv: torch.Tensor,\tq_norm_weight: torch.Tensor,\tk_norm_weight: torch.Tensor,\tfreqs: torch.Tensor,\tcu_seqlens_q: Optional[torch.Tensor] = None,\tcu_seqlens_kv: Optional[torch.Tensor] = None,\tH: int = 32,\th: int = 4,\teps: float = 1e-06,\tcp_rank=0,\tcp_size=1,\tmscale=1.0,\tsilu=False,\treuse=False):", "funcdef": "def"}, {"fullname": "linghe.facade.rope.mla_rope", "modulename": "linghe.facade.rope", "qualname": "mla_rope", "kind": "function", "doc": "

    inplace apply rope to tail 64 dims, split kv and apply rope to k_pos_emb and copy to k

    \n\n
    Arguments:
    \n\n
      \n
    • q: query tensor with size of [S, B, H, 128] (cu_seqlens is None) \nor [N, H, 128] (cu_seqlens is not None)
    • \n
    • kv: kv tensor with size of [S, B, H, 256] (cu_seqlens is None) or \n[N, H, 256] (cu_seqlens is not None)
    • \n
    • k_pos_emb: k pos emb with size of [S, B, 1, 64] (cu_seqlens is None) or \n[N, 1, 64] (cu_seqlens is not None)
    • \n
    • freqs: Freqs tensor with size of [S, 64]
    • \n
    • cu_seqlens_q: cumulative query lengths tensor with size of [B+1]
    • \n
    • cu_seqlens_kv: cumulative kv lengths tensor with size of [B+1]
    • \n
    • mscale: mscale of rope
    • \n
    • transpose: whether transpose output layout to [B, S, H, DIM]
    • \n
    • cp_size: context-parallel size
    • \n
    • cp_rank: context-parallel rank
    • \n
    \n\n
    Returns:
    \n\n
    \n
      \n
    • qo: shape [S, B, H, 192] or [N, H, 192]
    • \n
    • ko: shape [S, B, H, 192] or [N, H, 192]
    • \n
    • vo: shape [S, B, H, 128] or [N, H, 128]
    • \n
    \n
    \n", "signature": "(\tq: torch.Tensor,\tkv: torch.Tensor,\tk_pos_emb: torch.Tensor,\tfreqs: torch.Tensor,\tcu_seqlens_q: Optional[torch.Tensor] = None,\tcu_seqlens_kv: Optional[torch.Tensor] = None,\tmscale: float = 1.0,\ttranspose: bool = False,\tcp_size: int = 1,\tcp_rank: int = 0,\treuse: bool = False):", "funcdef": "def"}, {"fullname": "linghe.facade.silu", "modulename": "linghe.facade.silu", "kind": "module", "doc": "

    \n"}, {"fullname": "linghe.facade.silu.BlockSiluFunction", "modulename": "linghe.facade.silu", "qualname": "BlockSiluFunction", "kind": "class", "doc": "

    Base class to create custom autograd.Function.

    \n\n

    To create a custom autograd.Function, subclass this class and implement\nthe forward() and backward() static methods. Then, to use your custom\nop in the forward pass, call the class method apply. Do not call\nforward() directly.

    \n\n

    To ensure correctness and best performance, make sure you are calling the\ncorrect methods on ctx and validating your backward function using\ntorch.autograd.gradcheck().

    \n\n

    See :ref:extending-autograd for more details on how to use this class.

    \n\n

    Examples::

    \n\n
    >>> # xdoctest: +REQUIRES(env:TORCH_DOCTEST_AUTOGRAD)\n>>> class Exp(Function):\n>>>     @staticmethod\n>>>     def forward(ctx, i):\n>>>         result = i.exp()\n>>>         ctx.save_for_backward(result)\n>>>         return result\n>>>\n>>>     @staticmethod\n>>>     def backward(ctx, grad_output):\n>>>         result, = ctx.saved_tensors\n>>>         return grad_output * result\n>>>\n>>> # Use it by calling the apply method:\n>>> # xdoctest: +SKIP\n>>> output = Exp.apply(input)\n
    \n", "bases": "torch.autograd.function.Function"}, {"fullname": "linghe.facade.silu.BlockSiluFunction.forward", "modulename": "linghe.facade.silu", "qualname": "BlockSiluFunction.forward", "kind": "function", "doc": "

    Define the forward of the custom autograd Function.

    \n\n

    This function is to be overridden by all subclasses.\nThere are two ways to define forward:

    \n\n

    Usage 1 (Combined forward and ctx)::

    \n\n
    @staticmethod\ndef forward(ctx: Any, *args: Any, **kwargs: Any) -> Any:\n    pass\n
    \n\n
      \n
    • It must accept a context ctx as the first argument, followed by any\nnumber of arguments (tensors or other types).
    • \n
    • See :ref:combining-forward-context for more details
    • \n
    \n\n

    Usage 2 (Separate forward and ctx)::

    \n\n
    @staticmethod\ndef forward(*args: Any, **kwargs: Any) -> Any:\n    pass\n\n@staticmethod\ndef setup_context(ctx: Any, inputs: Tuple[Any, ...], output: Any) -> None:\n    pass\n
    \n\n
      \n
    • The forward no longer accepts a ctx argument.
    • \n
    • Instead, you must also override the torch.autograd.Function.setup_context()\nstaticmethod to handle setting up the ctx object.\noutput is the output of the forward, inputs are a Tuple of inputs\nto the forward.
    • \n
    • See :ref:extending-autograd for more details
    • \n
    \n\n

    The context can be used to store arbitrary data that can be then\nretrieved during the backward pass. Tensors should not be stored\ndirectly on ctx (though this is not currently enforced for\nbackward compatibility). Instead, tensors should be saved either with\nctx.save_for_backward() if they are intended to be used in\nbackward (equivalently, vjp) or ctx.save_for_forward()\nif they are intended to be used for in jvp.

    \n", "signature": "(ctx, input, quantizer, grad_quantizer, cls):", "funcdef": "def"}, {"fullname": "linghe.facade.silu.BlockSiluFunction.backward", "modulename": "linghe.facade.silu", "qualname": "BlockSiluFunction.backward", "kind": "function", "doc": "

    Define a formula for differentiating the operation with backward mode automatic differentiation.

    \n\n

    This function is to be overridden by all subclasses.\n(Defining this function is equivalent to defining the vjp function.)

    \n\n

    It must accept a context ctx as the first argument, followed by\nas many outputs as the forward() returned (None will be passed in\nfor non tensor outputs of the forward function),\nand it should return as many tensors, as there were inputs to\nforward(). Each argument is the gradient w.r.t the given output,\nand each returned value should be the gradient w.r.t. the\ncorresponding input. If an input is not a Tensor or is a Tensor not\nrequiring grads, you can just pass None as a gradient for that input.

    \n\n

    The context can be used to retrieve tensors saved during the forward\npass. It also has an attribute ctx.needs_input_grad as a tuple\nof booleans representing whether each input needs gradient. E.g.,\nbackward() will have ctx.needs_input_grad[0] = True if the\nfirst input to forward() needs gradient computed w.r.t. the\noutput.

    \n", "signature": "(ctx, grad_output):", "funcdef": "def"}, {"fullname": "linghe.facade.silu.BlockBatchWeightedSiluFunction", "modulename": "linghe.facade.silu", "qualname": "BlockBatchWeightedSiluFunction", "kind": "class", "doc": "

    Base class to create custom autograd.Function.

    \n\n

    To create a custom autograd.Function, subclass this class and implement\nthe forward() and backward() static methods. Then, to use your custom\nop in the forward pass, call the class method apply. Do not call\nforward() directly.

    \n\n

    To ensure correctness and best performance, make sure you are calling the\ncorrect methods on ctx and validating your backward function using\ntorch.autograd.gradcheck().

    \n\n

    See :ref:extending-autograd for more details on how to use this class.

    \n\n

    Examples::

    \n\n
    >>> # xdoctest: +REQUIRES(env:TORCH_DOCTEST_AUTOGRAD)\n>>> class Exp(Function):\n>>>     @staticmethod\n>>>     def forward(ctx, i):\n>>>         result = i.exp()\n>>>         ctx.save_for_backward(result)\n>>>         return result\n>>>\n>>>     @staticmethod\n>>>     def backward(ctx, grad_output):\n>>>         result, = ctx.saved_tensors\n>>>         return grad_output * result\n>>>\n>>> # Use it by calling the apply method:\n>>> # xdoctest: +SKIP\n>>> output = Exp.apply(input)\n
    \n", "bases": "torch.autograd.function.Function"}, {"fullname": "linghe.facade.silu.BlockBatchWeightedSiluFunction.forward", "modulename": "linghe.facade.silu", "qualname": "BlockBatchWeightedSiluFunction.forward", "kind": "function", "doc": "

    Define the forward of the custom autograd Function.

    \n\n

    This function is to be overridden by all subclasses.\nThere are two ways to define forward:

    \n\n

    Usage 1 (Combined forward and ctx)::

    \n\n
    @staticmethod\ndef forward(ctx: Any, *args: Any, **kwargs: Any) -> Any:\n    pass\n
    \n\n
      \n
    • It must accept a context ctx as the first argument, followed by any\nnumber of arguments (tensors or other types).
    • \n
    • See :ref:combining-forward-context for more details
    • \n
    \n\n

    Usage 2 (Separate forward and ctx)::

    \n\n
    @staticmethod\ndef forward(*args: Any, **kwargs: Any) -> Any:\n    pass\n\n@staticmethod\ndef setup_context(ctx: Any, inputs: Tuple[Any, ...], output: Any) -> None:\n    pass\n
    \n\n
      \n
    • The forward no longer accepts a ctx argument.
    • \n
    • Instead, you must also override the torch.autograd.Function.setup_context()\nstaticmethod to handle setting up the ctx object.\noutput is the output of the forward, inputs are a Tuple of inputs\nto the forward.
    • \n
    • See :ref:extending-autograd for more details
    • \n
    \n\n

    The context can be used to store arbitrary data that can be then\nretrieved during the backward pass. Tensors should not be stored\ndirectly on ctx (though this is not currently enforced for\nbackward compatibility). Instead, tensors should be saved either with\nctx.save_for_backward() if they are intended to be used in\nbackward (equivalently, vjp) or ctx.save_for_forward()\nif they are intended to be used for in jvp.

    \n", "signature": "(\tctx,\tinput,\tweights,\tcounts,\tsplits,\tquantizers,\tgrad_quantizers,\tcls,\tis_recomputing):", "funcdef": "def"}, {"fullname": "linghe.facade.silu.BlockBatchWeightedSiluFunction.backward", "modulename": "linghe.facade.silu", "qualname": "BlockBatchWeightedSiluFunction.backward", "kind": "function", "doc": "

    Define a formula for differentiating the operation with backward mode automatic differentiation.

    \n\n

    This function is to be overridden by all subclasses.\n(Defining this function is equivalent to defining the vjp function.)

    \n\n

    It must accept a context ctx as the first argument, followed by\nas many outputs as the forward() returned (None will be passed in\nfor non tensor outputs of the forward function),\nand it should return as many tensors, as there were inputs to\nforward(). Each argument is the gradient w.r.t the given output,\nand each returned value should be the gradient w.r.t. the\ncorresponding input. If an input is not a Tensor or is a Tensor not\nrequiring grads, you can just pass None as a gradient for that input.

    \n\n

    The context can be used to retrieve tensors saved during the forward\npass. It also has an attribute ctx.needs_input_grad as a tuple\nof booleans representing whether each input needs gradient. E.g.,\nbackward() will have ctx.needs_input_grad[0] = True if the\nfirst input to forward() needs gradient computed w.r.t. the\noutput.

    \n", "signature": "(ctx, grad_output):", "funcdef": "def"}, {"fullname": "linghe.facade.smooth_quant_linear", "modulename": "linghe.facade.smooth_quant_linear", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, {"fullname": "linghe.facade.smooth_quant_linear.SmoothQuantLinear", "modulename": "linghe.facade.smooth_quant_linear", "qualname": "SmoothQuantLinear", "kind": "class", "doc": "

    a naive implementation of smooth quantization linear

    \n", "bases": "torch.nn.modules.module.Module"}, {"fullname": "linghe.facade.smooth_quant_linear.SmoothQuantLinear.__init__", "modulename": "linghe.facade.smooth_quant_linear", "qualname": "SmoothQuantLinear.__init__", "kind": "function", "doc": "
    Arguments:
    \n\n
      \n
    • in_features: in feature number
    • \n
    • out_features: out feature number
    • \n
    • bias: whether use bias
    • \n
    • device: weight device
    • \n
    • dtype: weight dtype
    • \n
    \n", "signature": "(\tin_features: int,\tout_features: int,\tbias: bool = True,\tdevice=None,\tdtype=None)"}, {"fullname": "linghe.facade.topk", "modulename": "linghe.facade.topk", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, {"fullname": "linghe.facade.topk.fused_topk", "modulename": "linghe.facade.topk", "qualname": "fused_topk", "kind": "function", "doc": "

    topk

    \n\n
    Arguments:
    \n\n
      \n
    • x: input tensor
    • \n
    • k: topk
    • \n
    • dim: dimension to apply topk, only support -1 currently
    • \n
    \n\n
    Returns:
    \n\n
    \n

    values: topk values\n indices: topk indices

    \n
    \n", "signature": "(x, k, dim=-1):", "funcdef": "def"}, {"fullname": "linghe.facade.topk.group_topk_score", "modulename": "linghe.facade.topk", "qualname": "group_topk_score", "kind": "function", "doc": "

    group topk with softmax/sigmoid function

    \n\n
    Arguments:
    \n\n
      \n
    • x: input logit tensor
    • \n
    • topk: topk
    • \n
    • expert_bias: expert bias
    • \n
    • num_groups: number of groups
    • \n
    • group_topk: group to apply topk
    • \n
    • scaling_factor: scaling factor
    • \n
    • score_function: scaling function
    • \n
    \n\n
    Returns:
    \n\n
    \n

    probs: topk probs\n routing_map: topk binary map\n counts: token count per expert

    \n
    \n", "signature": "(\tx,\ttopk,\texpert_bias=None,\tnum_groups=32,\tgroup_topk=4,\tscaling_factor=1.0,\tscore_function='sigmoid'):", "funcdef": "def"}, {"fullname": "linghe.facade.transpose", "modulename": "linghe.facade.transpose", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, {"fullname": "linghe.facade.transpose.transpose", "modulename": "linghe.facade.transpose", "qualname": "transpose", "kind": "function", "doc": "

    transpose a tensor, x.ndims should not greater than 4

    \n\n
    Arguments:
    \n\n
      \n
    • x: input tensor
    • \n
    • inner: if True, transpose the first two dimensions\nif False, transpose the last two dimensions
    • \n
    \n\n
    Returns:
    \n\n
    \n

    a transposed tensor

    \n
    \n", "signature": "(x, inner=True):", "funcdef": "def"}, {"fullname": "linghe.gemm", "modulename": "linghe.gemm", "kind": "module", "doc": "

    \n"}, {"fullname": "linghe.gemm.blockwise_fp8_gemm", "modulename": "linghe.gemm.blockwise_fp8_gemm", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, {"fullname": "linghe.gemm.channelwise_fp8_gemm", "modulename": "linghe.gemm.channelwise_fp8_gemm", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, {"fullname": "linghe.gemm.channelwise_fp8_gemm.triton_scaled_mm", "modulename": "linghe.gemm.channelwise_fp8_gemm", "qualname": "triton_scaled_mm", "kind": "function", "doc": "

    similar to torch._scaled_mm, support accumulating gemm output to c\n and low precision output tensor

    \n\n
    Arguments:
    \n\n
      \n
    • a: left fp8 tensor
    • \n
    • b: right fp8 tensor, column-major
    • \n
    • a_scale: fp32 scale of a
    • \n
    • b_scale: fp32 scale of b
    • \n
    • out_dtype: output tensor dtype
    • \n
    • c: output tensor
    • \n
    • accum: accumulate output on c if True
    • \n
    \n\n
    Returns:
    \n\n
    \n

    c: output tensor

    \n
    \n", "signature": "(\ta: torch.Tensor,\tb: torch.Tensor,\ta_scale: torch.Tensor,\tb_scale: torch.Tensor,\tout_dtype=torch.float32,\tc=None,\taccum=True):", "funcdef": "def"}, {"fullname": "linghe.gemm.fp32_gemm", "modulename": "linghe.gemm.fp32_gemm", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, {"fullname": "linghe.gemm.fp32_gemm.triton_fp32_gemm", "modulename": "linghe.gemm.fp32_gemm", "qualname": "triton_fp32_gemm", "kind": "function", "doc": "

    return fp32 gemm result with fp16/bf16 inputs,\n it's mainly used for MoE router GEMM\n and DO NOT suitable for large size GEMM

    \n\n
    Arguments:
    \n\n
      \n
    • a: left matrix with fp16/bf16 precision
    • \n
    • b: right matrix with fp16/bf16 precision
    • \n
    \n\n
    Returns:
    \n\n
    \n

    c: output with fp32 precision

    \n
    \n", "signature": "(x: torch.Tensor, w: torch.Tensor):", "funcdef": "def"}, {"fullname": "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_backward", "modulename": "linghe.gemm.fp32_gemm", "qualname": "triton_fp32_gemm_for_backward", "kind": "function", "doc": "

    mix precision gemm for backward, a@b.float()

    \n\n
    Arguments:
    \n\n
      \n
    • a: input gradient, fp32
    • \n
    • b: gemm weight, bf16/fp16
    • \n
    \n\n
    Returns:
    \n\n
    \n

    c: gradient of activation

    \n
    \n", "signature": "(y: torch.Tensor, w: torch.Tensor):", "funcdef": "def"}, {"fullname": "linghe.gemm.fp32_gemm.triton_fp32_gemm_for_update", "modulename": "linghe.gemm.fp32_gemm", "qualname": "triton_fp32_gemm_for_update", "kind": "function", "doc": "

    mix precision gemm for updaing weight

    \n\n
    Arguments:
    \n\n
      \n
    • y: gradient of output, fp32
    • \n
    • x: input activation, bf16/fp16
    • \n
    \n\n
    Returns:
    \n\n
    \n

    c: gradient of weight

    \n
    \n", "signature": "(y: torch.Tensor, x: torch.Tensor):", "funcdef": "def"}, {"fullname": "linghe.gemm.fp32_gemm.triton_split_fp32_gemm", "modulename": "linghe.gemm.fp32_gemm", "qualname": "triton_split_fp32_gemm", "kind": "function", "doc": "

    return fp32 gemm result with fp16/bf16 inputs,\n it's mainly used for MoE router GEMM\n and DO NOT suitable for large size GEMM

    \n\n
    Arguments:
    \n\n
      \n
    • a: left matrix with fp16/bf16 precision
    • \n
    • b: right matrix with fp16/bf16 precision
    • \n
    \n\n
    Returns:
    \n\n
    \n

    c: output with fp32 precision

    \n
    \n", "signature": "(x: torch.Tensor, w: torch.Tensor):", "funcdef": "def"}, {"fullname": "linghe.gemm.fp32_gemm.triton_split_fp32_gemm_for_backward", "modulename": "linghe.gemm.fp32_gemm", "qualname": "triton_split_fp32_gemm_for_backward", "kind": "function", "doc": "

    mix precision gemm for backward, a@b.float()

    \n\n
    Arguments:
    \n\n
      \n
    • a: input gradient, fp32
    • \n
    • b: gemm weight, bf16/fp16
    • \n
    \n\n
    Returns:
    \n\n
    \n

    c: gradient of activation

    \n
    \n", "signature": "(y: torch.Tensor, w: torch.Tensor):", "funcdef": "def"}, {"fullname": "linghe.gemm.fp32_gemm.triton_split_fp32_gemm_for_update", "modulename": "linghe.gemm.fp32_gemm", "qualname": "triton_split_fp32_gemm_for_update", "kind": "function", "doc": "

    mix precision gemm for updaing weight

    \n\n
    Arguments:
    \n\n
      \n
    • y: gradient of output, fp32
    • \n
    • x: input activation, bf16/fp16
    • \n
    \n\n
    Returns:
    \n\n
    \n

    c: gradient of weight

    \n
    \n", "signature": "(y: torch.Tensor, x: torch.Tensor):", "funcdef": "def"}, {"fullname": "linghe.quant", "modulename": "linghe.quant", "kind": "module", "doc": "

    \n"}, {"fullname": "linghe.quant.block", "modulename": "linghe.quant.block", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, {"fullname": "linghe.quant.block.triton_block_quant", "modulename": "linghe.quant.block", "qualname": "triton_block_quant", "kind": "function", "doc": "

    blockwise quantize x, used for blockwise recipe for weight in megatron

    \n\n
    Arguments:
    \n\n
      \n
    • x: input tensor
    • \n
    • block_size: block wise
    • \n
    • round_scale: whether round scale to power of 2
    • \n
    \n\n
    Returns:
    \n\n
    \n
      \n
    • y: quantized tensor, float8_e4m3fn
    • \n
    • s: quantization scale, float32
    • \n
    \n
    \n", "signature": "(x, block_size=128, round_scale=False):", "funcdef": "def"}, {"fullname": "linghe.quant.block.triton_blockwise_quant", "modulename": "linghe.quant.block", "qualname": "triton_blockwise_quant", "kind": "function", "doc": "

    blockwise quantization, used in blockwise recipt in megatron

    \n\n
    Arguments:
    \n\n
      \n
    • x: input tensor
    • \n
    • round_scale: whether round scale to power of 2
    • \n
    • output_mode: one of {0, 1, 2}\n0: only output non-transposed quantized tensor\n1: only output transposed quantized tensor\n2: output both
    • \n
    \n\n
    Returns:
    \n\n
    \n

    x_q: \n x_scale: \n xt_q: \n xt_scale:

    \n
    \n", "signature": "(x, round_scale=False, output_mode=2):", "funcdef": "def"}, {"fullname": "linghe.quant.block.triton_batch_blockwise_quant", "modulename": "linghe.quant.block", "qualname": "triton_batch_blockwise_quant", "kind": "function", "doc": "

    select and quant, used in megatron 0.12 flex moe

    \n\n
    Arguments:
    \n\n
      \n
    • xs: [bs, dim]
    • \n
    • token_count_per_expert: [n_experts]
    • \n
    • splits: python int list of token_count_per_expert
    • \n
    • round_scale: whether round scale to power of 2
    • \n
    \n\n
    Returns:
    \n\n
    \n
      \n
    • x_q:
    • \n
    • x_scale:
    • \n
    • xt_q:
    • \n
    • xt_scale:
    • \n
    \n
    \n", "signature": "(xs, token_count_per_expert, splits, round_scale=False):", "funcdef": "def"}, {"fullname": "linghe.quant.channel", "modulename": "linghe.quant.channel", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, {"fullname": "linghe.quant.channel.triton_row_quant", "modulename": "linghe.quant.channel", "qualname": "triton_row_quant", "kind": "function", "doc": "

    rowwise quantize x

    \n\n
    Arguments:
    \n\n
      \n
    • x: input x
    • \n
    • round_scale: whether round scale to power of 2
    • \n
    \n\n
    Returns:
    \n\n
    \n

    x_q: quantized tensor\n x_scale: quantization scale

    \n
    \n", "signature": "(x, round_scale=False):", "funcdef": "def"}, {"fullname": "linghe.quant.channel.triton_tokenwise_row_quant", "modulename": "linghe.quant.channel", "qualname": "triton_tokenwise_row_quant", "kind": "function", "doc": "

    rowwise quantize x with power of 2 dim size

    \n\n
    Arguments:
    \n\n
      \n
    • x: input x
    • \n
    • round_scale: whether round scale to power of 2
    • \n
    \n\n
    Returns:
    \n\n
    \n

    out: quantized tensor\n scale: quantization scale

    \n
    \n", "signature": "(x, out=None, scale=None, round_scale=False):", "funcdef": "def"}, {"fullname": "linghe.quant.channel.triton_transpose_row_quant", "modulename": "linghe.quant.channel", "qualname": "triton_transpose_row_quant", "kind": "function", "doc": "

    transpose x and row quantize x

    \n\n
    Arguments:
    \n\n
      \n
    • x: input x
    • \n
    • round_scale: whether round scale to power of 2
    • \n
    \n\n
    Returns:
    \n\n
    \n
      \n
    • x_q: quantized tensor
    • \n
    • x_scale: quantization scale
    • \n
    \n
    \n", "signature": "(x, round_scale=False):", "funcdef": "def"}, {"fullname": "linghe.quant.group", "modulename": "linghe.quant.group", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, {"fullname": "linghe.quant.group.triton_group_quant", "modulename": "linghe.quant.group", "qualname": "triton_group_quant", "kind": "function", "doc": "

    groupwise quantize x, group is in under rowwise format

    \n\n
    Arguments:
    \n\n
      \n
    • x: input tensor
    • \n
    • group_size: group wise
    • \n
    • round_scale: whether round scale to power of 2
    • \n
    \n\n
    Returns:
    \n\n
    \n
      \n
    • y: quantized tensor, float8_e4m3fn
    • \n
    • s: quantization scale, float32
    • \n
    \n
    \n", "signature": "(x, dtype=torch.float8_e4m3fn, group_size=128, round_scale=False):", "funcdef": "def"}, {"fullname": "linghe.quant.hadamard", "modulename": "linghe.quant.hadamard", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, {"fullname": "linghe.quant.hadamard.triton_hadamard_quant", "modulename": "linghe.quant.hadamard", "qualname": "triton_hadamard_quant", "kind": "function", "doc": "

    apply hadamard transformation and then quantize transformed tensor

    \n\n
    Arguments:
    \n\n
      \n
    • x: input tensor
    • \n
    • hm: hamadard matrix
    • \n
    \n\n
    Returns:
    \n\n
    \n
      \n
    • x_q: rowwise quantized tensor of non-transposed x
    • \n
    • x_scale: rowwise quantization scale of non-transposed x
    • \n
    • xt_q: columnwise quantized tensor of transposed x
    • \n
    • xt_scale: columnwise quantization scale of transposed x
    • \n
    \n
    \n", "signature": "(x, hm):", "funcdef": "def"}, {"fullname": "linghe.quant.smooth", "modulename": "linghe.quant.smooth", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, {"fullname": "linghe.tools", "modulename": "linghe.tools", "kind": "module", "doc": "

    \n"}, {"fullname": "linghe.tools.benchmark", "modulename": "linghe.tools.benchmark", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, {"fullname": "linghe.tools.check", "modulename": "linghe.tools.check", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, {"fullname": "linghe.tools.util", "modulename": "linghe.tools.util", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, {"fullname": "linghe.utils", "modulename": "linghe.utils", "kind": "module", "doc": "

    \n"}, {"fullname": "linghe.utils.add", "modulename": "linghe.utils.add", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, {"fullname": "linghe.utils.add.triton_inplace_add", "modulename": "linghe.utils.add", "qualname": "triton_inplace_add", "kind": "function", "doc": "

    inplace add y to x

    \n\n
    Arguments:
    \n\n
      \n
    • x: Tensor
    • \n
    • y: Tensor
    • \n
    • accum: x += y if accum=True else x.copy_(y)
    • \n
    \n\n
    Returns:
    \n\n
    \n

    updated x

    \n
    \n", "signature": "(x: torch.Tensor, y: torch.Tensor, accum: bool = True):", "funcdef": "def"}, {"fullname": "linghe.utils.emb", "modulename": "linghe.utils.emb", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, {"fullname": "linghe.utils.emb.triton_embedding_forward", "modulename": "linghe.utils.emb", "qualname": "triton_embedding_forward", "kind": "function", "doc": "

    inplace add y to x

    \n\n
    Arguments:
    \n\n
      \n
    • x: input ids Tensor
    • \n
    • w_ptr: data_ptr of embedding weight
    • \n
    \n\n
    Returns:
    \n\n
    \n

    embedding output

    \n
    \n", "signature": "(x, w_ptr, dim=4096, dtype=torch.bfloat16):", "funcdef": "def"}, {"fullname": "linghe.utils.emb.triton_atomic_embedding_backward", "modulename": "linghe.utils.emb", "qualname": "triton_atomic_embedding_backward", "kind": "function", "doc": "

    inplace update embedding weight gradient

    \n\n
    Arguments:
    \n\n
      \n
    • y: gradient of output
    • \n
    • x: input ids Tensor
    • \n
    • g_ptr: data_ptr of embedding weight gradient
    • \n
    \n\n
    Returns:
    \n\n
    \n

    None

    \n
    \n", "signature": "(y, x, g_ptr, dtype=torch.bfloat16):", "funcdef": "def"}, {"fullname": "linghe.utils.emb.triton_sync_embedding_backward", "modulename": "linghe.utils.emb", "qualname": "triton_sync_embedding_backward", "kind": "function", "doc": "

    inplace update embedding weight gradient

    \n\n
    Arguments:
    \n\n
      \n
    • y: gradient of output
    • \n
    • x: input ids Tensor
    • \n
    • g_ptr: data_ptr of embedding weight gradient
    • \n
    \n\n
    Returns:
    \n\n
    \n

    None

    \n
    \n", "signature": "(grad_output, x, g_ptr, dtype=torch.bfloat16):", "funcdef": "def"}, {"fullname": "linghe.utils.emb.triton_embedding_backward", "modulename": "linghe.utils.emb", "qualname": "triton_embedding_backward", "kind": "function", "doc": "

    inplace update embedding weight gradient

    \n\n
    Arguments:
    \n\n
      \n
    • y: gradient of output
    • \n
    • x: input ids Tensor
    • \n
    • g_ptr: data_ptr of embedding weight gradient
    • \n
    \n\n
    Returns:
    \n\n
    \n

    None

    \n
    \n", "signature": "(grad_output, x, g_ptr, dtype=torch.bfloat16):", "funcdef": "def"}, {"fullname": "linghe.utils.gate", "modulename": "linghe.utils.gate", "kind": "module", "doc": "

    \n"}, {"fullname": "linghe.utils.gate.triton_group_rms_norm_gate_forward", "modulename": "linghe.utils.gate", "qualname": "triton_group_rms_norm_gate_forward", "kind": "function", "doc": "

    norm and gate in linear attention

    \n\n
    Arguments:
    \n\n
      \n
    • x: output of attn, [bs, length, n_heads, head_dim]
    • \n
    • gate: gate tensor, [length, bs, dim] if transpose=True else [bs, length, dim]
    • \n
    • weight: rms norm weight, [dim]
    • \n
    • eps: epsilon of rms norm
    • \n
    • group_size: group size of group rms norm
    • \n
    • transpose: whether gate tensor has been transposed and output will be transposed
    • \n
    \n\n
    Returns:
    \n\n
    \n

    output tensor, [length, bs, dim] if transpose=True else [bs, length, dim]

    \n
    \n", "signature": "(\tx: torch.Tensor,\tgate: torch.Tensor,\tweight: torch.Tensor,\teps=1e-06,\tgroup_size=4,\ttranspose=True):", "funcdef": "def"}, {"fullname": "linghe.utils.gather", "modulename": "linghe.utils.gather", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, {"fullname": "linghe.utils.gather.triton_make_row_id_map", "modulename": "linghe.utils.gather", "qualname": "triton_make_row_id_map", "kind": "function", "doc": "

    make row id map, values in the tensor are the row indices

    \n\n
    Arguments:
    \n\n
      \n
    • routing_map: a tensor of 0/1 values, 1 indicates routed
    • \n
    • multiple_of: padding the tokens of each expert to multiple of this value
    • \n
    \n\n
    Returns:
    \n\n
    \n

    row id map with shape [n_tokens, n_experts]

    \n
    \n", "signature": "(routing_map: torch.Tensor, multiple_of: int = 1):", "funcdef": "def"}, {"fullname": "linghe.utils.gather.triton_make_row_id_map_and_index", "modulename": "linghe.utils.gather", "qualname": "triton_make_row_id_map_and_index", "kind": "function", "doc": "

    similar with triton_make_row_id_map, but output an indices tensor as well

    \n\n
    Arguments:
    \n\n
      \n
    • routing_map: [n_tokens, n_experts]
    • \n
    • num_out_tokens: sum(round_up_to(n_tokens, multiple_of))
    • \n
    • multiple_of: padding the tokens of each expert to this value
    • \n
    \n\n
    Returns:
    \n\n
    \n

    row_in_map: [n_tokens, n_experts]\n row_indices: [num_out_tokens]

    \n
    \n", "signature": "(routing_map: torch.Tensor, num_out_tokens: int, multiple_of: int = 1):", "funcdef": "def"}, {"fullname": "linghe.utils.gather.triton_index_select", "modulename": "linghe.utils.gather", "qualname": "triton_index_select", "kind": "function", "doc": "

    index select for quantized tensor

    \n\n
    Arguments:
    \n\n
      \n
    • x: [bs, dim]
    • \n
    • indices: [K]
    • \n
    • scale: [bs]
    • \n
    \n\n
    Returns:
    \n\n
    \n

    out: output of selected x\n scale_out: scale of selected scale

    \n
    \n", "signature": "(x, indices, scale=None, out=None, scale_out=None):", "funcdef": "def"}, {"fullname": "linghe.utils.gather.triton_permute_with_mask_map", "modulename": "linghe.utils.gather", "qualname": "triton_permute_with_mask_map", "kind": "function", "doc": "

    gather quantized tensor with row id map

    \n\n
    Arguments:
    \n\n
      \n
    • inp: [num_tokens, hidden_size], rowwise quantized tensor
    • \n
    • scale: optional, [num_tokens], quantization scale
    • \n
    • probs: optional, router prob, used as weight
    • \n
    • row_id_map: [n_experts, num_tokens]\nindex >= 0: row index of output tensor\nindex == -1: ignore\nNote: index may not be contiguous
    • \n
    • num_out_tokens: output token count, including padding tokens
    • \n
    • contiguous: whether indices in row_id_map is contiguous,\nFalse means padded
    • \n
    • tokens_per_expert: [num_experts], token count per expert,\nnon-blocking cuda tensor
    • \n
    \n\n
    Returns:
    \n\n
    \n

    output: permuted quantized tensor\n permuted_scale: permuted quantization scale\n permuted_probs: permuted router prob

    \n
    \n", "signature": "(\tinp: torch.Tensor,\tscale: torch.Tensor,\tprobs: torch.Tensor,\trow_id_map: torch.Tensor,\tnum_out_tokens: int,\tcontiguous: bool = True,\ttokens_per_expert: Optional[torch.Tensor] = None):", "funcdef": "def"}, {"fullname": "linghe.utils.gather.triton_batch_transpose_smooth_permute_with_indices", "modulename": "linghe.utils.gather", "qualname": "triton_batch_transpose_smooth_permute_with_indices", "kind": "function", "doc": "

    used for smooth quantization backward in megatron 0.12,\nx is gathered, requantized, padded to multiple of 32 and tranposed

    \n\n
    Arguments:
    \n\n
      \n
    • x: dy, [bs, dim], it is smooth quantized
    • \n
    • scale: [bs], quantized scale
    • \n
    • org_smooth_scale: [dim]
    • \n
    • smooth_scales: [n_experts, dim]
    • \n
    • indices: [sum(tokens_per_experts)]
    • \n
    • token_count_per_expert: [n_experts], tensor of token count per expert
    • \n
    • splits: [n_experts], list of token_count_per_expert
    • \n
    • round_scale: round quantization scale to power of 2
    • \n
    \n\n
    Returns:
    \n\n
    \n

    x_q: [sum(roundup(tokens_per_experts)) * dim]\n x_scale: [sum(roundup(tokens_per_experts))]

    \n
    \n", "signature": "(\tx,\tscale,\torg_smooth_scale,\tsmooth_scales,\tindices,\ttoken_count_per_expert,\tsplits,\tx_q=None,\tx_scale=None,\tround_scale=False):", "funcdef": "def"}, {"fullname": "linghe.utils.gather.triton_smooth_weighted_permute_with_indices", "modulename": "linghe.utils.gather", "qualname": "triton_smooth_weighted_permute_with_indices", "kind": "function", "doc": "

    select and smooth and quant, used in megatron 0.11 all2all moe

    \n\n
    Arguments:
    \n\n
      \n
    • grads: [bs, dim]
    • \n
    • tokens: [bs, dim]
    • \n
    • smooth_scales: [n_experts, dim]
    • \n
    • token_count_per_expert: [n_experts]
    • \n
    • indices: [n_experts*topk]
    • \n
    • reverse: whether scale is 1/scale
    • \n
    • round_scale: whether round scale to power of 2
    • \n
    \n\n
    Returns:
    \n\n
    \n

    x_q: [bs*topk, dim]\n x_scale: [bstopk]\n x_sum: [bstopk]

    \n
    \n", "signature": "(\tgrads,\ttokens,\tsmooth_scales,\ttoken_count_per_expert,\tindices,\tx_q=None,\tx_scale=None,\tx_sum=None,\treverse=False,\tround_scale=False):", "funcdef": "def"}, {"fullname": "linghe.utils.gather.triton_smooth_permute_with_indices", "modulename": "linghe.utils.gather", "qualname": "triton_smooth_permute_with_indices", "kind": "function", "doc": "

    select and smooth and quant

    \n\n
    Arguments:
    \n\n
      \n
    • grad_data: [bs, dim]
    • \n
    • grad_scale: [bs]
    • \n
    • smooth_scales: [n_experts, dim]
    • \n
    • token_count_per_expert: [n_experts]
    • \n
    • indices: [n_experts*topk]
    • \n
    • x_q: [bs*topk, dim]
    • \n
    • x_scale: [bs*topk]
    • \n
    • reverse:
    • \n
    • round_scale:
    • \n
    \n\n

    Returns:

    \n", "signature": "(\tgrad_data,\tgrad_scale,\tsmooth_scales,\ttoken_count_per_expert,\tindices,\tx_q=None,\tx_scale=None,\treverse=False,\tround_scale=False):", "funcdef": "def"}, {"fullname": "linghe.utils.gather.triton_smooth_permute_with_mask_map", "modulename": "linghe.utils.gather", "qualname": "triton_smooth_permute_with_mask_map", "kind": "function", "doc": "

    gather ( and optional dequant) and smooth quant

    \n\n
    Arguments:
    \n\n
      \n
    • inp: [num_tokens, hidden_size], rowwise quantized tensor
    • \n
    • row_id_map: [n_experts, num_tokens], indices
    • \n
    • scale: [num_tokens, hs], rowwise_scale_inv, optional
    • \n
    • num_tokens: [n_experts]
    • \n
    • num_experts:
    • \n
    • num_out_tokens:
    • \n
    • hidden_size:
    • \n
    • smooth_scales: [n_experts, hidden_size]
    • \n
    • reverse:
    • \n
    • round_scale:
    • \n
    \n\n
    Returns:
    \n\n
    \n
      \n
    • output: output tensor
    • \n
    • permuted_scale: permuted scale if scale is not None
    • \n
    \n
    \n", "signature": "(\tinp: torch.Tensor,\trow_id_map: torch.Tensor,\tscale: torch.Tensor,\tnum_tokens: int,\tnum_experts: int,\tnum_out_tokens: int,\thidden_size: int,\tsmooth_scales: torch.Tensor,\treverse=True,\tround_scale=False):", "funcdef": "def"}, {"fullname": "linghe.utils.gather.triton_batch_block_pad_permute_with_indices", "modulename": "linghe.utils.gather", "qualname": "triton_batch_block_pad_permute_with_indices", "kind": "function", "doc": "

    select and quant, used in megatron 0.12 flex moe

    \n\n
    Arguments:
    \n\n
      \n
    • xs: [bs, dim]
    • \n
    • token_count_per_expert: [n_experts]
    • \n
    • indices: [n_experts*topk]
    • \n
    • splits: python int list of token_count_per_expert
    • \n
    • probs: route weights, [bs, n_experts]
    • \n
    • round_scale: whether round scale to power of 2
    • \n
    \n\n
    Returns:
    \n\n
    \n

    x_q: \n x_scale: \n xt_q: \n xt_scale: \n prob_output:

    \n
    \n", "signature": "(\txs,\ttoken_count_per_expert,\tindices,\tsplits,\tprobs=None,\tround_scale=False):", "funcdef": "def"}, {"fullname": "linghe.utils.loss", "modulename": "linghe.utils.loss", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, {"fullname": "linghe.utils.loss.triton_softmax_cross_entropy_forward", "modulename": "linghe.utils.loss", "qualname": "triton_softmax_cross_entropy_forward", "kind": "function", "doc": "

    compute token-wise softmax cross entropy loss

    \n\n
    Arguments:
    \n\n
      \n
    • logits: logits tensor
    • \n
    • labels: labels tensor
    • \n
    \n\n
    Returns:
    \n\n
    \n

    loss of each token

    \n
    \n", "signature": "(logits, labels, ignore_index=-100):", "funcdef": "def"}, {"fullname": "linghe.utils.loss.triton_softmax_cross_entropy_backward", "modulename": "linghe.utils.loss", "qualname": "triton_softmax_cross_entropy_backward", "kind": "function", "doc": "

    backward of softmax cross entropy loss

    \n\n
    Arguments:
    \n\n
      \n
    • logits: logit tensor, [bs, dim]
    • \n
    • labels: label tensor, [bs]
    • \n
    • sum_exp: [bs]
    • \n
    • max_logit: [bs]
    • \n
    • output_grad: gradient, [bs, dim]
    • \n
    • inplace: whether to reuse logits as gradient
    • \n
    \n\n
    Returns:
    \n\n
    \n

    grad of input: [bs, dim]

    \n
    \n", "signature": "(\tlogits,\tlabels,\tsum_exp,\tmax_logit,\toutput_grad,\tignore_index=-100,\tinplace=False):", "funcdef": "def"}, {"fullname": "linghe.utils.loss.triton_parallel_softmax_cross_entropy_forward", "modulename": "linghe.utils.loss", "qualname": "triton_parallel_softmax_cross_entropy_forward", "kind": "function", "doc": "

    compute token-wise softmax cross entropy loss

    \n\n
    Arguments:
    \n\n
      \n
    • logits: logits tensor
    • \n
    • labels: labels tensor
    • \n
    \n\n
    Returns:
    \n\n
    \n

    loss of each token

    \n
    \n", "signature": "(logits, labels, group, ignore_index=-100):", "funcdef": "def"}, {"fullname": "linghe.utils.loss.triton_parallel_softmax_cross_entropy_backward", "modulename": "linghe.utils.loss", "qualname": "triton_parallel_softmax_cross_entropy_backward", "kind": "function", "doc": "

    backward of softmax cross entropy loss

    \n\n
    Arguments:
    \n\n
      \n
    • logits: logit tensor, [bs, dim]
    • \n
    • labels: label tensor, [bs]
    • \n
    • sum_exp: [bs]
    • \n
    • max_logit: [bs]
    • \n
    • output_grad: gradient, [bs, dim]
    • \n
    • inplace: whether to reuse logits as gradient
    • \n
    \n\n
    Returns:
    \n\n
    \n

    grad of input: [bs, dim]

    \n
    \n", "signature": "(\tlogits,\tlabels,\tsum_exp,\tmax_logit,\toutput_grad,\tgroup,\tignore_index=-100,\tinplace=False):", "funcdef": "def"}, {"fullname": "linghe.utils.loss.triton_moe_z_loss_forward", "modulename": "linghe.utils.loss", "qualname": "triton_moe_z_loss_forward", "kind": "function", "doc": "

    compute moe z loss,\nz_loss = torch.mean(torch.square(torch.logsumexp(logits, dim=-1))) * coef

    \n\n
    Arguments:
    \n\n
      \n
    • logits: logits tensor
    • \n
    • coef: z loss coef
    • \n
    \n\n
    Returns:
    \n\n
    \n

    z loss

    \n
    \n", "signature": "(logits, coef=1e-06):", "funcdef": "def"}, {"fullname": "linghe.utils.loss.triton_moe_z_loss_backward", "modulename": "linghe.utils.loss", "qualname": "triton_moe_z_loss_backward", "kind": "function", "doc": "

    backward of moe z loss

    \n\n
    Arguments:
    \n\n
      \n
    • grads: grad scalar tensor
    • \n
    • logits: logit tensor, [L, B, dim]
    • \n
    • coef: python scalar
    • \n
    \n\n
    Returns:
    \n\n
    \n

    output_grad: [L, B, dim]

    \n
    \n", "signature": "(grads, logits, coef=1e-06):", "funcdef": "def"}, {"fullname": "linghe.utils.mul", "modulename": "linghe.utils.mul", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, {"fullname": "linghe.utils.mul.triton_dot", "modulename": "linghe.utils.mul", "qualname": "triton_dot", "kind": "function", "doc": "

    vector dot multiply, output = sum(x*y, 1),\nit is used to calculate gradient of router weight

    \n\n
    Arguments:
    \n\n
      \n
    • x:
    • \n
    • y:
    • \n
    \n\n
    Returns:
    \n\n
    \n

    output of sum(x*y, 1)

    \n
    \n", "signature": "(x, y):", "funcdef": "def"}, {"fullname": "linghe.utils.mul.triton_inplace_scale", "modulename": "linghe.utils.mul", "qualname": "triton_inplace_scale", "kind": "function", "doc": "

    inplace scale a tensor.

    \n\n
    Arguments:
    \n\n
      \n
    • x: Tensor.
    • \n
    • scale: a python float scale
    • \n
    \n\n
    Returns:
    \n\n
    \n

    x

    \n
    \n", "signature": "(x, scale):", "funcdef": "def"}, {"fullname": "linghe.utils.mul.triton_batch_scale", "modulename": "linghe.utils.mul", "qualname": "triton_batch_scale", "kind": "function", "doc": "

    return [x*scale for x in xs],\nused to scale gradient.

    \n\n
    Arguments:
    \n\n
      \n
    • xs: Tensor lists.
    • \n
    • scale: a python float scale
    • \n
    \n\n
    Returns:
    \n\n
    \n

    xs

    \n
    \n", "signature": "(xs, scale):", "funcdef": "def"}, {"fullname": "linghe.utils.norm", "modulename": "linghe.utils.norm", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, {"fullname": "linghe.utils.norm.triton_rms_norm_forward", "modulename": "linghe.utils.norm", "qualname": "triton_rms_norm_forward", "kind": "function", "doc": "

    rms norm

    \n\n
    Arguments:
    \n\n
      \n
    • x: input tensor
    • \n
    • weight: weight of rms norm
    • \n
    • eps: epsilon of rms norm
    • \n
    • rms: use x*rms to calculate output if rms is not None, \nit will accelerate recompute of rms norm
    • \n
    \n\n
    Returns:
    \n\n
    \n

    out: output tensor\n rms: 1/rms of input tensor

    \n
    \n", "signature": "(x, weight, eps=1e-06, out=None, rms=None):", "funcdef": "def"}, {"fullname": "linghe.utils.norm.triton_rms_norm_and_block_quant_forward", "modulename": "linghe.utils.norm", "qualname": "triton_rms_norm_and_block_quant_forward", "kind": "function", "doc": "

    Fused RMSNorm forward and block quantization.

    \n\n
    Arguments:
    \n\n
      \n
    • x: Input tensor, shape [M, N]
    • \n
    • weight: RMSNorm weight, shape [N]
    • \n
    • eps: epsilon value for L2 normalization.
    • \n
    • out: output of quantization data
    • \n
    • scale: output of quantization scale.
    • \n
    • rms: output of rms
    • \n
    • round_scale: Set whether to force power of 2 scales.
    • \n
    • output_mode: one of {0, 1, 2}.\n0: only output non-transpose tensor\n1: only output transposed tensor\n2: return both
    • \n
    \n\n
    Returns:
    \n\n
    \n
      \n
    • out: quantization data.
    • \n
    • scale: quantization scale.
    • \n
    • rms: Reciprocal of the root mean square of the\n input calculated over the last dimension.
    • \n
    • transpose_output: quantization data of transposed gradient.
    • \n
    • transpose_scale: quantization scale of transposed gradient.
    • \n
    \n
    \n", "signature": "(\tx: torch.Tensor,\tweight: torch.Tensor,\teps: float = 1e-06,\tout: Optional[torch.Tensor] = None,\tscale: Optional[torch.Tensor] = None,\trms: Optional[torch.Tensor] = None,\tround_scale: bool = False,\toutput_mode: int = 2):", "funcdef": "def"}, {"fullname": "linghe.utils.norm.triton_rms_norm_fp32_gemm_block_quant_forward", "modulename": "linghe.utils.norm", "qualname": "triton_rms_norm_fp32_gemm_block_quant_forward", "kind": "function", "doc": "

    y = rms_norm(x)\nlogits = y@w_route\nx_q, x_s, xt_q, xt_s = quantization(y)

    \n\n
    Arguments:
    \n\n
      \n
    • x: input tensor
    • \n
    • norm weight: weight tensor of rms norm
    • \n
    • route_weight: moe router weight
    • \n
    • eps: epsilon of rms norm
    • \n
    • output_mode: 0 or 1\n0: only output non-transpose quantizatino tensor\n1: only output transposed quantizatino tensor
    • \n
    \n\n
    Returns:
    \n\n
    \n
      \n
    • y: rms normed tensor
    • \n
    • rms: 1/rms
    • \n
    • logits: router logit
    • \n
    • x_q:
    • \n
    • x_s:
    • \n
    • xt_q:
    • \n
    • xt_s:
    • \n
    \n
    \n", "signature": "(\tx: torch.Tensor,\tnorm_weight: torch.Tensor,\troute_weight: torch.Tensor,\trms: Optional[torch.Tensor] = None,\teps: float = 1e-06,\toutput_mode: int = 0,\tround_scale=False):", "funcdef": "def"}, {"fullname": "linghe.utils.rearange", "modulename": "linghe.utils.rearange", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, {"fullname": "linghe.utils.rearange.triton_sort_chunks_by_index", "modulename": "linghe.utils.rearange", "qualname": "triton_sort_chunks_by_index", "kind": "function", "doc": "

    split x to multiple tensors and cat with indices,\nit is used for permutation in moe with all2all communication

    \n\n
    Arguments:
    \n\n
      \n
    • x: [bs, dim]
    • \n
    • counts: [n_split]
    • \n
    • indices: [n_split]
    • \n
    • scales: [bs]
    • \n
    \n\n
    Returns:
    \n\n
    \n
      \n
    • y: output tensor
    • \n
    • output_scales: output scales if scales is not None
    • \n
    \n
    \n", "signature": "(x, counts, indices, scales=None):", "funcdef": "def"}, {"fullname": "linghe.utils.reduce", "modulename": "linghe.utils.reduce", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, {"fullname": "linghe.utils.reduce.triton_abs_max", "modulename": "linghe.utils.reduce", "qualname": "triton_abs_max", "kind": "function", "doc": "

    columnwise abs max of x, it is used in smooth quantization

    \n\n
    Arguments:
    \n\n
      \n
    • x: input tensor, may be quantized tensor
    • \n
    • scale: quantization scale if x is quantized
    • \n
    • smooth_scale: optional smooth scale
    • \n
    • min_value: output = max(max(abs(x,0)), min_value)
    • \n
    • axis: reduce axis
    • \n
    \n\n
    Returns:
    \n\n
    \n

    max tensor

    \n
    \n", "signature": "(x, scale=None, smooth_scale=None, min_value=1e-30, axis=0):", "funcdef": "def"}, {"fullname": "linghe.utils.reduce.triton_batch_count_zero", "modulename": "linghe.utils.reduce", "qualname": "triton_batch_count_zero", "kind": "function", "doc": "

    count zero in tensor list, it is used to monitor zeros in gradient tensor

    \n\n
    Arguments:
    \n\n
      \n
    • xs: input tensors
    • \n
    \n\n
    Returns:
    \n\n
    \n

    a single-value int64 tensor

    \n
    \n", "signature": "(xs):", "funcdef": "def"}, {"fullname": "linghe.utils.reduce.triton_norm", "modulename": "linghe.utils.reduce", "qualname": "triton_norm", "kind": "function", "doc": "

    calculate norm.

    \n\n
    Arguments:
    \n\n
      \n
    • x: input tensor.
    • \n
    • ord: the order of tensor. -1 means 'inf' ord.
    • \n
    • norm: only used with ord in (1, 2)\nTrue: (sum(sum(abs(x)ord) x for x in xs))(1/ord) \nFalse: sum(sum(abs(x)**ord) x for x in xs))
    • \n
    \n\n
    Returns:
    \n\n
    \n

    a scalar if scalar=True else a single-value fp32 tensor

    \n
    \n", "signature": "(x, ord=2, norm=True, scalar=True):", "funcdef": "def"}, {"fullname": "linghe.utils.reduce.triton_batch_norm", "modulename": "linghe.utils.reduce", "qualname": "triton_batch_norm", "kind": "function", "doc": "

    treat multiple tensors as a single tensor and calculate norm.

    \n\n
    Arguments:
    \n\n
      \n
    • xs: Tensor lists.
    • \n
    • ord: the order of tensor. -1 means 'inf' ord.
    • \n
    • norm: only used with ord in (1, 2)\nTrue: (sum(sum(abs(x)ord) x for x in xs))(1/ord) \nFalse: sum(sum(abs(x)**ord) x for x in xs))
    • \n
    \n\n
    Returns:
    \n\n
    \n

    a scalar if scalar=True else a single-value fp32 tensor

    \n
    \n", "signature": "(xs, ord=2, norm=True, scalar=True, high_precision=True):", "funcdef": "def"}, {"fullname": "linghe.utils.rope", "modulename": "linghe.utils.rope", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, {"fullname": "linghe.utils.rope.triton_half_rope_forward", "modulename": "linghe.utils.rope", "qualname": "triton_half_rope_forward", "kind": "function", "doc": "

    apply half rope to qk

    \n\n
    Arguments:
    \n\n
      \n
    • q: query tensor, [len, bs, q_head, head_dim]
    • \n
    • k: key tensor, [len, bs, kv_head, head_dim]
    • \n
    • freqs: rope freqs
    • \n
    \n\n
    Returns:
    \n\n
    \n
      \n
    • qo: query output
    • \n
    • ko: key output
    • \n
    \n
    \n", "signature": "(q, k, freqs, transposed=True):", "funcdef": "def"}, {"fullname": "linghe.utils.rope.triton_qk_norm_and_half_rope_forward", "modulename": "linghe.utils.rope", "qualname": "triton_qk_norm_and_half_rope_forward", "kind": "function", "doc": "

    split qkv to q/k/v, apply qk norm and half rope to q/k,\n transpose q/k/v to flash-attention layout

    \n\n
    Arguments:
    \n\n
      \n
    • qkv: QKV tensor with size of [S, B, dim], heads are interleaved
    • \n
    • q_norm_weight: rms norm weight for query
    • \n
    • k_norm_weight: rms norm weight for key
    • \n
    • freqs: Freqs tensor based on half dim.
    • \n
    • H: Number of attention heads.
    • \n
    • h: Number of key/value heads.
    • \n
    • eps: epsilon value for L2 normalization.
    • \n
    • interleaved: whether head of qkv is interleaved,\ninterleaved: [q...qkvq...qkv]\nnon-interleaved: [q...qk...kv...v]
    • \n
    • transposed: whether qkv is tranposed\ntransposed: [S, B, dim]\nnon-transposed: [B, S, dim]
    • \n
    • silu: apply silu on qkv before qk norm and rope
    • \n
    \n\n
    Returns:
    \n\n
    \n
      \n
    • qo: shape [B, S, H, head_dim]
    • \n
    • ko: shape [B, S, h, head_dim]
    • \n
    • vo: shape [B, S, h, head_dim]
    • \n
    \n
    \n", "signature": "(\tqkv,\tq_norm_weight,\tk_norm_weight,\tfreqs,\tH=32,\th=4,\teps=1e-06,\tinterleaved=True,\ttransposed=True,\tsilu=False):", "funcdef": "def"}, {"fullname": "linghe.utils.rope.triton_qk_norm_and_half_rope_backward", "modulename": "linghe.utils.rope", "qualname": "triton_qk_norm_and_half_rope_backward", "kind": "function", "doc": "

    backward kernel of triton_qk_norm_and_half_rope_forward

    \n\n
    Arguments:
    \n\n
      \n
    • gq: gradient of qo, [len, bs, q_head, head_dim]
    • \n
    • gk: gradient of ko, [len, bs, q_head, head_dim]
    • \n
    • gv: gradient of vo, [len, bs, q_head, head_dim]
    • \n
    • qkv: input qkv
    • \n
    • q_norm_weight: rms norm weight for query
    • \n
    • k_norm_weight: rms norm weight for key
    • \n
    • freqs: Freqs tensor based on half dim.
    • \n
    • eps: epsilon value for L2 normalization.
    • \n
    • interleaved: whether head of qkv is interleaved,\ninterleaved: [q...qkvq...qkv]\nnon-interleaved: [q...qk...kv...v]
    • \n
    • transposed: whether qkv is tranposed\ntransposed: [S, B, dim]\nnon-transposed: [B, S, dim]
    • \n
    • silu: whether silu is applied to qkv
    • \n
    \n\n
    Returns:
    \n\n
    \n
      \n
    • dqkv: gradient of qkv
    • \n
    • dqw: gradient of q_norm_weight
    • \n
    • dkw: gradient of k_norm_weight
    • \n
    \n
    \n", "signature": "(\tgq,\tgk,\tgv,\tqkv,\tq_norm_weight,\tk_norm_weight,\tfreqs,\teps=1e-06,\tinterleaved=True,\ttransposed=True,\tsilu=False):", "funcdef": "def"}, {"fullname": "linghe.utils.rope.triton_varlen_qk_norm_and_half_rope_forward", "modulename": "linghe.utils.rope", "qualname": "triton_varlen_qk_norm_and_half_rope_forward", "kind": "function", "doc": "

    split qkv to q/k/v, apply qk norm and half rope to q/k,\n transpose q/k/v to flash-attention layout

    \n\n
    Arguments:
    \n\n
      \n
    • qkv: QKV tensor with size of [S, B, dim], heads are interleaved
    • \n
    • q_norm_weight: rms norm weight for query
    • \n
    • k_norm_weight: rms norm weight for key
    • \n
    • freqs: Freqs tensor based on half dim.
    • \n
    • H: Number of attention heads.
    • \n
    • h: Number of key/value heads.
    • \n
    • eps: epsilon value for L2 normalization.
    • \n
    • interleaved: whether head of qkv is interleaved,\ninterleaved: [q...qkvq...qkv]\nnon-interleaved: [q...qk...kv...v]
    • \n
    • silu: apply silu on qkv before qk norm and rope
    • \n
    \n\n
    Returns:
    \n\n
    \n
      \n
    • qo: shape [B, S, H, head_dim]
    • \n
    • ko: shape [B, S, h, head_dim]
    • \n
    • vo: shape [B, S, h, head_dim]
    • \n
    \n
    \n", "signature": "(\tqkv,\tq_norm_weight,\tk_norm_weight,\tfreqs,\tcu_seqlens_q,\tcu_seqlens_kv,\tH=32,\th=4,\teps=1e-06,\tinterleaved=True,\tsilu=False,\tcp_rank=0,\tcp_size=1,\tmscale=1.0,\treuse=False):", "funcdef": "def"}, {"fullname": "linghe.utils.rope.triton_varlen_qk_norm_and_half_rope_backward", "modulename": "linghe.utils.rope", "qualname": "triton_varlen_qk_norm_and_half_rope_backward", "kind": "function", "doc": "

    backward kernel of triton_qk_norm_and_half_rope_forward

    \n\n
    Arguments:
    \n\n
      \n
    • gq: gradient of qo, [len, bs, q_head, head_dim]
    • \n
    • gk: gradient of ko, [len, bs, q_head, head_dim]
    • \n
    • gv: gradient of vo, [len, bs, q_head, head_dim]
    • \n
    • qkv: input qkv
    • \n
    • q_norm_weight: rms norm weight for query
    • \n
    • k_norm_weight: rms norm weight for key
    • \n
    • freqs: Freqs tensor based on half dim.
    • \n
    • eps: epsilon value for L2 normalization.
    • \n
    • interleaved: whether head of qkv is interleaved,\ninterleaved: [q...qkvq...qkv]\nnon-interleaved: [q...qk...kv...v]
    • \n
    • silu: whether silu is applied to qkv
    • \n
    \n\n
    Returns:
    \n\n
    \n
      \n
    • dqkv: gradient of qkv
    • \n
    • dqw: gradient of q_norm_weight
    • \n
    • dkw: gradient of k_norm_weight
    • \n
    \n
    \n", "signature": "(\tgq,\tgk,\tgv,\tqkv,\tq_norm_weight,\tk_norm_weight,\tfreqs,\tcu_seqlens_q,\tcu_seqlens_kv,\teps=1e-06,\tinterleaved=True,\tsilu=False,\tcp_rank=0,\tcp_size=1,\tmscale=1.0,\treuse=False):", "funcdef": "def"}, {"fullname": "linghe.utils.rope.triton_mla_rope_forward", "modulename": "linghe.utils.rope", "qualname": "triton_mla_rope_forward", "kind": "function", "doc": "

    apply MLA-type rope to qkv

    \n\n
    Arguments:
    \n\n
      \n
    • q: query tensor, [len, bs, n_heads, 192]
    • \n
    • kv: key-value tensor, [len, bs, n_heads, 256]
    • \n
    • k_pos_emb: k pos emb, [len, bs, 1, 64]
    • \n
    • freqs: rope freqs, [len, 64]
    • \n
    • mscale: mscale for rope
    • \n
    • transpose: whether transpose the output to [bs, len, n_heads, dim] layout
    • \n
    • cu_seqlens_q: accummulated query length
    • \n
    • cu_seqlens_kv: accummulated kv length
    • \n
    • cp_rank: rank of context parallel
    • \n
    • cp_size: size of context parallel
    • \n
    \n\n
    Returns:
    \n\n
    \n
      \n
    • qo: inplace updated query, [len, bs, n_heads, 192] if not transpose\n else [bs, len, n_heads, 192]
    • \n
    • ko: key output, [len, bs, n_heads, 192] if not transpose\n else [bs, len, n_heads, 192]
    • \n
    • vo: value output, [len, bs, n_heads, 128] if not transpose\n else [bs, len, n_heads, 128]
    • \n
    \n
    \n", "signature": "(\tq,\tkv,\tk_pos_emb,\tfreqs,\tmscale=1.0,\ttranspose=False,\tcu_seqlens_q=None,\tcu_seqlens_kv=None,\tcp_rank=0,\tcp_size=1,\treuse=False):", "funcdef": "def"}, {"fullname": "linghe.utils.scatter", "modulename": "linghe.utils.scatter", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, {"fullname": "linghe.utils.scatter.triton_aligned_scatter_add", "modulename": "linghe.utils.scatter", "qualname": "triton_aligned_scatter_add", "kind": "function", "doc": "

    scatter_add for megatron 0.11

    \n\n
    Arguments:
    \n\n
      \n
    • x: input tensor
    • \n
    • outputs: output tensor
    • \n
    • indices: gather indices
    • \n
    • weights: rowwise weight, it is router prob in MoE router
    • \n
    \n\n
    Returns:
    \n\n
    \n

    output tensor

    \n
    \n", "signature": "(\tx: torch.Tensor,\toutputs: torch.Tensor,\tindices: torch.Tensor,\tweights: Optional[torch.Tensor] = None):", "funcdef": "def"}, {"fullname": "linghe.utils.scatter.triton_scatter_add", "modulename": "linghe.utils.scatter", "qualname": "triton_scatter_add", "kind": "function", "doc": "

    naive version of scatter add, very slow

    \n\n
    Arguments:
    \n\n
      \n
    • x: input tensor
    • \n
    • outputs: output tensor
    • \n
    • indices: indices
    • \n
    \n\n
    Returns:
    \n\n
    \n

    output tensor

    \n
    \n", "signature": "(x, outputs, indices):", "funcdef": "def"}, {"fullname": "linghe.utils.scatter.triton_unpermute_with_mask_map", "modulename": "linghe.utils.scatter", "qualname": "triton_unpermute_with_mask_map", "kind": "function", "doc": "

    scatter add with row id map

    \n\n
    Arguments:
    \n\n
      \n
    • grad: gradient tensor, [num_out_tokens, hidden_size]
    • \n
    • row_id_map: row id map, [n_experts, num_tokens]
    • \n
    • probs: [num_out_tokens]
    • \n
    \n\n
    Returns:
    \n\n
    \n
      \n
    • output: [num_tokens, hidden_size]
    • \n
    • restore_probs: [num_tokens, num_experts]
    • \n
    \n
    \n", "signature": "(grad: torch.Tensor, row_id_map: torch.Tensor, probs: torch.Tensor):", "funcdef": "def"}, {"fullname": "linghe.utils.silu", "modulename": "linghe.utils.silu", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, {"fullname": "linghe.utils.silu.triton_weighted_silu_forward", "modulename": "linghe.utils.silu", "qualname": "triton_weighted_silu_forward", "kind": "function", "doc": "

    compute silu(x)*weight, used in bf16/fp16 training with MoE

    \n\n
    Arguments:
    \n\n
      \n
    • x: input tensor
    • \n
    • weight: tokenwise weight
    • \n
    \n\n
    Returns:
    \n\n
    \n

    out: output tensor

    \n
    \n", "signature": "(x, weight=None, out=None, asm=False):", "funcdef": "def"}, {"fullname": "linghe.utils.silu.triton_weighted_silu_backward", "modulename": "linghe.utils.silu", "qualname": "triton_weighted_silu_backward", "kind": "function", "doc": "

    backward of triton_weighted_silu_forward

    \n\n
    Arguments:
    \n\n
      \n
    • g: gradient tensor
    • \n
    • x: input tensor
    • \n
    • weight: weight tensor
    • \n
    \n\n
    Returns:
    \n\n
    \n
      \n
    • dx: gradient of x
    • \n
    • dw: gradient of weight
    • \n
    \n
    \n", "signature": "(\tg: torch.Tensor,\tx: torch.Tensor,\tweight: Optional[torch.Tensor] = None):", "funcdef": "def"}, {"fullname": "linghe.utils.silu.triton_silu_and_block_quant_forward", "modulename": "linghe.utils.silu", "qualname": "triton_silu_and_block_quant_forward", "kind": "function", "doc": "

    fused silu and blockwise quantization, used in shared expert

    \n\n
    Arguments:
    \n\n
      \n
    • x: input tensor
    • \n
    • round_scale: whether round scale to power of 2
    • \n
    • output_mode: one of {0, 1, 2}\n0: only output non-transposed quantized tensor\n1: only output transposed quantized tensor\n2: output both
    • \n
    \n\n
    Returns:
    \n\n
    \n
      \n
    • out: quantized tensor
    • \n
    • scale: quantization scale
    • \n
    • transpose_output: quantized tensor of transposed output
    • \n
    • transpose_scale: quantization scale of transposed output
    • \n
    \n
    \n", "signature": "(x, out=None, scale=None, round_scale=False, output_mode=2):", "funcdef": "def"}, {"fullname": "linghe.utils.silu.triton_silu_and_block_quant_backward", "modulename": "linghe.utils.silu", "qualname": "triton_silu_and_block_quant_backward", "kind": "function", "doc": "

    backward of triton_silu_and_block_quant_forward

    \n\n
    Arguments:
    \n\n
      \n
    • g: gradient
    • \n
    • x: input tensor
    • \n
    • round_scale: whether round to power of 2
    • \n
    \n\n
    Returns:
    \n\n
    \n
      \n
    • dx: quantized non-transposed gradient
    • \n
    • dx_scale: scales of quantization non-transposed gradient
    • \n
    • transpose_dx: quantized transposed gradient
    • \n
    • transpose_dx_scale: scales of quantization transposed gradient
    • \n
    \n
    \n", "signature": "(g, x, round_scale=False):", "funcdef": "def"}, {"fullname": "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_forward", "modulename": "linghe.utils.silu", "qualname": "triton_batch_weighted_silu_and_block_quant_forward", "kind": "function", "doc": "

    silu and blockwise quantize activation in routed experts

    \n\n
    Arguments:
    \n\n
      \n
    • x: activation tensor in routed experts
    • \n
    • weight: router prob tensor
    • \n
    • counts: cuda tensor of token count per expert
    • \n
    • splits: python int list of token count per expert
    • \n
    • round_scale: whether round scale to power of 2
    • \n
    • output_mode: one of {0, 1, 2}\n0: only output non-transposed quantized tensor\n1: only output transposed quantized tensor\n2: output both
    • \n
    \n\n
    Returns:
    \n\n
    \n
      \n
    • out: quantized tensor
    • \n
    • scale: quantization scale
    • \n
    • transpose_output: quantized tensor of transposed output
    • \n
    • transpose_scale: quantization scale of transposed output
    • \n
    \n
    \n", "signature": "(\tx,\tweight,\tcounts,\tsplits=None,\tout=None,\tscale=None,\tround_scale=False,\toutput_mode=2):", "funcdef": "def"}, {"fullname": "linghe.utils.silu.triton_batch_weighted_silu_and_block_quant_backward", "modulename": "linghe.utils.silu", "qualname": "triton_batch_weighted_silu_and_block_quant_backward", "kind": "function", "doc": "

    backward of triton_batch_weighted_silu_and_block_quant_forward

    \n\n
    Arguments:
    \n\n
      \n
    • g: gradient
    • \n
    • x: input tensor
    • \n
    • weight: router prob tensor
    • \n
    • counts: cuda tensor of token count per expert
    • \n
    • splits: python int list of token count per expert
    • \n
    • round_scale: whether round scale to power of 2
    • \n
    \n\n
    Returns:
    \n\n
    \n
      \n
    • dx: quantized non-transposed gradient
    • \n
    • dx_scale: scales of quantization non-transposed gradient
    • \n
    • dw: gradient of weight
    • \n
    • transpose_dx: quantized transposed gradient
    • \n
    • transpose_dx_scale: scales of quantization transposed gradient
    • \n
    \n
    \n", "signature": "(g, x, weight, counts, splits=None, round_scale=False):", "funcdef": "def"}, {"fullname": "linghe.utils.topk", "modulename": "linghe.utils.topk", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, {"fullname": "linghe.utils.topk.triton_topk_forward", "modulename": "linghe.utils.topk", "qualname": "triton_topk_forward", "kind": "function", "doc": "

    calculate topk.

    \n\n
    Arguments:
    \n\n
      \n
    • x: input tensor.
    • \n
    • k: topk
    • \n
    \n\n
    Returns:
    \n\n
    \n

    values: topk values \n indices: topk indices

    \n
    \n", "signature": "(x, k, dim=-1):", "funcdef": "def"}, {"fullname": "linghe.utils.topk.triton_topk_backward", "modulename": "linghe.utils.topk", "qualname": "triton_topk_backward", "kind": "function", "doc": "

    topk backward.

    \n\n
    Arguments:
    \n\n
      \n
    • grad_output: grad tensor of values.
    • \n
    • indices: topk indices
    • \n
    • N: dim
    • \n
    \n\n
    Returns:
    \n\n
    \n

    dx

    \n
    \n", "signature": "(grad_output, indices, N, dim=-1):", "funcdef": "def"}, {"fullname": "linghe.utils.topk.triton_group_topk_score_forward", "modulename": "linghe.utils.topk", "qualname": "triton_group_topk_score_forward", "kind": "function", "doc": "

    calculate topk.

    \n\n
    Arguments:
    \n\n
      \n
    • x: input tensor.
    • \n
    • expert_bias: expert bias
    • \n
    • k: topk
    • \n
    \n\n
    Returns:
    \n\n
    \n

    probs:
    \n routing_map: \n tokens_per_expert:

    \n
    \n", "signature": "(\tx,\tk,\texpert_bias=None,\tnum_groups=32,\tgroup_topk=4,\tscaling_factor=1.0,\tscore_function='sigmoid',\teps=1e-20):", "funcdef": "def"}, {"fullname": "linghe.utils.topk.triton_group_topk_score_backward", "modulename": "linghe.utils.topk", "qualname": "triton_group_topk_score_backward", "kind": "function", "doc": "

    topk backward.

    \n\n
    Arguments:
    \n\n
      \n
    • grad_output: grad tensor of prob.
    • \n
    • routing_map: topk indices
    • \n
    \n\n
    Returns:
    \n\n
    \n

    dx: grad of logits

    \n
    \n", "signature": "(grad_output, input, routing_map, scaling_factor=1.0, eps=1e-20):", "funcdef": "def"}, {"fullname": "linghe.utils.transpose", "modulename": "linghe.utils.transpose", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, {"fullname": "linghe.utils.transpose.triton_transpose", "modulename": "linghe.utils.transpose", "qualname": "triton_transpose", "kind": "function", "doc": "

    transpose x with dim0 and dim1

    \n\n
    Arguments:
    \n\n
      \n
    • x: input tensor
    • \n
    • inner: inner dim if True, outer dim if False
    • \n
    \n\n
    Returns:
    \n\n
    \n

    transposed tensor

    \n
    \n", "signature": "(x: torch.Tensor, inner=True):", "funcdef": "def"}, {"fullname": "linghe.utils.transpose.triton_transpose_and_pad", "modulename": "linghe.utils.transpose", "qualname": "triton_transpose_and_pad", "kind": "function", "doc": "

    transpose x and padding the column size to be mutiplier of 32,\nit is used for calculated gradient of weight with torch._scaled__mm

    \n\n
    Arguments:
    \n\n
      \n
    • x: input tensor
    • \n
    • out:
    • \n
    • pad: whether need padding
    • \n
    \n\n
    Returns:
    \n\n
    \n

    out: output tensor

    \n
    \n", "signature": "(x, out=None, pad=True):", "funcdef": "def"}, {"fullname": "linghe.utils.transpose.triton_batch_transpose", "modulename": "linghe.utils.transpose", "qualname": "triton_batch_transpose", "kind": "function", "doc": "

    batch transpose x

    \n\n
    Arguments:
    \n\n
      \n
    • xs: input tensor list, [M, N]*expert
    • \n
    \n\n
    Returns:
    \n\n
    \n

    xts: output tensor list, [N,M]*expert

    \n
    \n", "signature": "(xs, xts=None):", "funcdef": "def"}, {"fullname": "linghe.utils.transpose.triton_batch_transpose_and_pad", "modulename": "linghe.utils.transpose", "qualname": "triton_batch_transpose_and_pad", "kind": "function", "doc": "

    transpose and pad each tensor stored in x

    \n\n
    Arguments:
    \n\n
      \n
    • x: [sum(bs), N]
    • \n
    • count_list: a python list of token count
    • \n
    • pad: whether pad to mutiplier of 32,\npadding value should be filled with 0 if padded
    • \n
    \n\n
    Returns:
    \n\n
    \n

    x_t: output tensor

    \n
    \n", "signature": "(x, count_list, x_t=None, pad=True):", "funcdef": "def"}, {"fullname": "linghe.utils.unary", "modulename": "linghe.utils.unary", "kind": "module", "doc": "

    Copyright (c) Ant Financial Service Group and its affiliates.

    \n"}, {"fullname": "linghe.utils.unary.triton_batch_clip", "modulename": "linghe.utils.unary", "qualname": "triton_batch_clip", "kind": "function", "doc": "

    return [clip(x, -clip_value, clip_value) for x in xs],\nused to clip gradient.

    \n\n
    Arguments:
    \n\n
      \n
    • xs: Tensor lists.
    • \n
    • clip_value: a python float scale
    • \n
    \n\n
    Returns:
    \n\n
    \n

    updated xs

    \n
    \n", "signature": "(xs, clip_value=100.0):", "funcdef": "def"}]; // mirrored in build-search-index.js (part 1) // Also split on html tags. this is a cheap heuristic, but good enough. diff --git a/linghe/attn/__init__.py b/linghe/attn/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/linghe/attn/la.py b/linghe/attn/la.py new file mode 100644 index 0000000..2f032ba --- /dev/null +++ b/linghe/attn/la.py @@ -0,0 +1,1598 @@ +# -*- coding: utf-8 -*- +""" +Copyright (c) Ant Financial Service Group and its affiliates. +""" + +import torch +import triton +import triton.language as tl + + +@triton.jit +def fp32_lightning_attention_forward_kernel( + Q, + K, + V, + S, + Out, + softmax_scale, + stride_q, + stride_k, + stride_v, + stride_s, + decay_scales, + L, + D: tl.constexpr, + KD: tl.constexpr, + VD: tl.constexpr, + BLOCK: tl.constexpr, +): + bid = tl.program_id(0) + hid = tl.program_id(1) + kvid = tl.program_id(2) + N = D // VD + kid = kvid // N + vid = kvid % N + H = tl.num_programs(1) + + c0 = bid * L + + decay_scale = -tl.load(decay_scales + hid) + + offs_b = tl.arange(0, BLOCK) + offs_k = tl.arange(0, KD) + offs_v = tl.arange(0, VD) + + q_ptrs = ( + Q + + c0 * stride_q + + hid * D + + kid * KD + + (offs_b[:, None] * stride_q + offs_k[None, :]) + ) + k_ptrs = ( + K + + c0 * stride_k + + hid * D + + kid * KD + + (offs_b[:, None] * stride_k + offs_k[None, :]) + ) + v_ptrs = ( + V + + c0 * stride_v + + hid * D + + vid * VD + + (offs_b[:, None] * stride_v + offs_v[None, :]) + ) + out_ptrs = ( + Out + + c0 * D * H + + hid * D + + vid * VD + + (offs_b[:, None] * H * D + offs_v[None, :]) + ) + s_ptrs = ( + S + + bid * stride_s + + hid * D * D + + kid * D * KD + + vid * VD + + (offs_k[:, None] * D + offs_v[None, :]) + ) + state = tl.zeros((KD, VD), dtype=tl.float32) + block_decay = tl.exp(decay_scale * BLOCK) + + for n in range(0, L, BLOCK): + n = tl.multiple_of(n, BLOCK) + + q = tl.load(q_ptrs + n * stride_q) # .to(tl.float32) + k = tl.trans(tl.load(k_ptrs + n * stride_k)) # .to(tl.float32) + v = tl.load(v_ptrs + n * stride_v).to(tl.float32) + b_offs = BLOCK - 1 - offs_b + decays = tl.exp(decay_scale * b_offs) + inv_decays = 1 / decays + + q = q * inv_decays[:, None] + k = k * decays[None, :] + qk = tl.dot(q, k) * softmax_scale + qk = tl.where(offs_b[None, :] <= offs_b[:, None], qk, 0.0) + o = tl.dot(qk, v) + + o = tl.dot(q, state) * block_decay * softmax_scale + o + + state = state * block_decay + tl.dot(k, v) + + if KD == D: + tl.store(out_ptrs + n * H * D, o.to(Out.dtype.element_ty)) + else: + tl.atomic_add( + out_ptrs + n * H * D, o.to(Out.dtype.element_ty), sem="relaxed" + ) + + tl.store(s_ptrs, state) + + +@triton.jit +def lightning_attention_forward_kernel( + Q, + K, + V, + S, + Out, + softmax_scale, + stride_q, + stride_k, + stride_v, + stride_s, + decay_scales, + L, + D: tl.constexpr, + KD: tl.constexpr, + VD: tl.constexpr, + BLOCK: tl.constexpr, +): + bid = tl.program_id(0) + hid = tl.program_id(1) + kvid = tl.program_id(2) + N = D // VD + kid = kvid // N + vid = kvid % N + H = tl.num_programs(1) + + c0 = bid * L + + decay_scale = -tl.load(decay_scales + hid) + + offs_b = tl.arange(0, BLOCK) + offs_k = tl.arange(0, KD) + offs_v = tl.arange(0, VD) + + q_ptrs = ( + Q + + c0 * stride_q + + hid * D + + kid * KD + + (offs_b[:, None] * stride_q + offs_k[None, :]) + ) + k_ptrs = ( + K + + c0 * stride_k + + hid * D + + kid * KD + + (offs_b[:, None] * stride_k + offs_k[None, :]) + ) + v_ptrs = ( + V + + c0 * stride_v + + hid * D + + vid * VD + + (offs_b[:, None] * stride_v + offs_v[None, :]) + ) + out_ptrs = ( + Out + + c0 * D * H + + hid * D + + vid * VD + + (offs_b[:, None] * H * D + offs_v[None, :]) + ) + s_ptrs = ( + S + + bid * stride_s + + hid * D * D + + kid * D * KD + + vid * VD + + (offs_k[:, None] * D + offs_v[None, :]) + ) + state = tl.zeros((KD, VD), dtype=tl.float32) + block_decay = tl.exp(decay_scale * BLOCK) + mask = tl.exp(decay_scale * (offs_b[:, None] - offs_b[None, :])) + mask = tl.where(offs_b[None, :] <= offs_b[:, None], mask, 0.0) * softmax_scale + b_offs = BLOCK - 1 - offs_b + decays = tl.exp(decay_scale * b_offs) + inv_decays = 1 / decays * block_decay * softmax_scale + + for n in range(0, L, BLOCK): + n = tl.multiple_of(n, BLOCK) + + q = tl.load(q_ptrs + n * stride_q) + k = tl.trans(tl.load(k_ptrs + n * stride_k)) + v = tl.load(v_ptrs + n * stride_v) + + qk = tl.dot(q, k) * mask + o = tl.dot(qk.to(v.dtype), v) + + o = tl.dot((q * inv_decays[:, None]).to(q.dtype), state.to(q.dtype), o) + + state *= block_decay + state = tl.dot((k * decays[None, :]).to(v.dtype), v, state) + + if KD == D: + tl.store(out_ptrs + n * H * D, o.to(Out.dtype.element_ty)) + else: + tl.atomic_add( + out_ptrs + n * H * D, o.to(Out.dtype.element_ty), sem="relaxed" + ) + + tl.store(s_ptrs, state) + + +# (k_dim_block, length, qo_heads, d) +@triton.jit +def _output_sum_kernel(T, O, DIM: tl.constexpr, NUM_BLOCK: tl.constexpr): + pid = tl.program_id(0) + length = tl.num_programs(0) + x = tl.zeros((DIM,), dtype=tl.float32) + for i in range(NUM_BLOCK): + x += tl.load(T + i * length * DIM + pid * DIM + tl.arange(0, DIM)).to( + tl.float32 + ) + tl.store(O + pid * DIM + tl.arange(0, DIM), x) + + +def triton_lightning_attention_forward( + q, k, v, decay_scales, hpc=False, hp=False, softmax_scale=None +): + B, L, H, D = q.shape + h = k.shape[2] + assert H == h, "triton_lightning_attention_forward does NOT support GQA currently" + + if softmax_scale is None: + softmax_scale = D ** (-0.5) + + KD = 32 + VD = 128 + BLOCK = 32 + device = q.device + dtype = q.dtype + + num_warps = 2 # 2 + num_stages = 5 # 3 + + k_dim_block = D // KD + v_dim_block = D // VD + if k_dim_block == 1: + outputs = torch.empty((B, L, H, D), device=device, dtype=dtype) + else: + outputs = torch.zeros( + (B, L, H, D), device=device, dtype=torch.float32 if hpc else dtype + ) + + s = torch.empty((B, H, D, D), device=device, dtype=torch.float32) + assert L % BLOCK == 0 and BLOCK <= 64 + + kernel = ( + fp32_lightning_attention_forward_kernel + if hp + else lightning_attention_forward_kernel + ) + grid = (B, H, k_dim_block * v_dim_block) + kernel[grid]( + q, + k, + v, + s, + outputs, + softmax_scale, + q.stride(1), + k.stride(1), + v.stride(1), + s.stride(0), + decay_scales, + L, + D=D, + KD=KD, + VD=VD, + BLOCK=BLOCK, + num_warps=num_warps, + num_stages=num_stages, + ) + + if k_dim_block > 1 and hpc: + o = outputs.to(dtype) + else: + o = outputs + + return o, s + + +@triton.jit +def fp32_lightning_attention_q_backward_kernel( + Q, + K, + V, + G, + DQ, + softmax_scale, + stride_q, + stride_k, + stride_v, + stride_g, + decay_scales, + L, + D: tl.constexpr, + KD: tl.constexpr, + VD: tl.constexpr, + BLOCK: tl.constexpr, +): + bid = tl.program_id(0) + hid = tl.program_id(1) + kvid = tl.program_id(2) + N = D // VD + kid = kvid // N + vid = kvid % N + H = tl.num_programs(1) + + c0 = bid * L + + decay_scale = -tl.load(decay_scales + hid) + + offs_b = tl.arange(0, BLOCK) + offs_k = tl.arange(0, KD) + offs_v = tl.arange(0, VD) + + q_ptrs = ( + Q + + c0 * stride_q + + hid * D + + kid * KD + + (offs_b[:, None] * stride_q + offs_k[None, :]) + ) + k_ptrs = ( + K + + c0 * stride_k + + hid * D + + kid * KD + + (offs_b[:, None] * stride_k + offs_k[None, :]) + ) + v_ptrs = ( + V + + c0 * stride_v + + hid * D + + vid * VD + + (offs_b[:, None] * stride_v + offs_v[None, :]) + ) + g_ptrs = ( + G + + c0 * D * H + + hid * D + + vid * VD + + (offs_b[:, None] * stride_g + offs_v[None, :]) + ) + dq_ptrs = ( + DQ + + c0 * D * H + + hid * D + + kid * KD + + (offs_b[:, None] * H * D + offs_k[None, :]) + ) + + state = tl.zeros((KD, VD), dtype=tl.float32) + mask = tl.exp((offs_b[:, None] - offs_b[None, :]) * decay_scale) + mask = tl.where(offs_b[None, :] <= offs_b[:, None], mask, 0.0) + + n_steps = tl.cdiv(L, BLOCK) + for i in range(n_steps): + n = i * BLOCK + n = tl.multiple_of(n, BLOCK) + + q = tl.load(q_ptrs + n * stride_q).to(tl.float32) + k = tl.load(k_ptrs + n * stride_k).to(tl.float32) + v = tl.load(v_ptrs + n * stride_v).to(tl.float32) + g = tl.load(g_ptrs + n * stride_g).to(tl.float32) + + qk = tl.dot(q, tl.trans(k)) * softmax_scale + qk *= mask + + decay_offs = BLOCK - 1 - offs_b + + block_decay = tl.exp(decay_scale * BLOCK) + decays = tl.exp(decay_scale * decay_offs) # [0.01, 0.1, 1] + + state = state * block_decay + + dqk = tl.dot(g, tl.trans(v)) * mask * softmax_scale + + dq = ( + tl.dot(dqk, k) + + tl.dot(g * decays[:, None], tl.trans(state)) * softmax_scale + ) + + if VD == D: + tl.store(dq_ptrs + n * H * D, dq.to(DQ.dtype.element_ty)) + else: + tl.atomic_add( + dq_ptrs + n * H * D, dq.to(DQ.dtype.element_ty), sem="relaxed" + ) + + state = state + tl.dot(tl.trans(k * decays[:, None]), v) + + +@triton.jit +def lightning_attention_q_backward_kernel( + Q, + K, + V, + G, + DQ, + softmax_scale, + stride_q, + stride_k, + stride_v, + stride_g, + decay_scales, + L, + D: tl.constexpr, + KD: tl.constexpr, + VD: tl.constexpr, + BLOCK: tl.constexpr, +): + bid = tl.program_id(0) + hid = tl.program_id(1) + kvid = tl.program_id(2) + N = D // VD + kid = kvid // N + vid = kvid % N + H = tl.num_programs(1) + + c0 = bid * L + + decay_scale = -tl.load(decay_scales + hid) + + offs_b = tl.arange(0, BLOCK) + offs_k = tl.arange(0, KD) + offs_v = tl.arange(0, VD) + + k_ptrs = ( + K + + c0 * stride_k + + hid * D + + kid * KD + + (offs_b[:, None] * stride_k + offs_k[None, :]) + ) + v_ptrs = ( + V + + c0 * stride_v + + hid * D + + vid * VD + + (offs_b[:, None] * stride_v + offs_v[None, :]) + ) + g_ptrs = ( + G + + c0 * D * H + + hid * D + + vid * VD + + (offs_b[:, None] * stride_g + offs_v[None, :]) + ) + dq_ptrs = ( + DQ + + c0 * D * H + + hid * D + + kid * KD + + (offs_b[:, None] * H * D + offs_k[None, :]) + ) + + state = tl.zeros((KD, VD), dtype=tl.float32) + mask = tl.exp((offs_b[:, None] - offs_b[None, :]) * decay_scale) + mask = tl.where(offs_b[None, :] <= offs_b[:, None], mask, 0.0) * softmax_scale + + decay_offs = BLOCK - 1 - offs_b + + block_decay = tl.exp(decay_scale * BLOCK) + decays = tl.exp(decay_scale * decay_offs) # [0.01, 0.1, 1] + + n_steps = tl.cdiv(L, BLOCK) + for i in range(n_steps): + n = i * BLOCK + n = tl.multiple_of(n, BLOCK) + + # q = tl.load(q_ptrs + n * stride_q) + k = tl.load(k_ptrs + n * stride_k) + v = tl.load(v_ptrs + n * stride_v) + g = tl.load(g_ptrs + n * stride_g) + + # qk = tl.dot(q, tl.trans(k)) * mask + # qk *= mask + + state = state * block_decay + + dqk = tl.dot(g, tl.trans(v)) * mask + + dq = ( + tl.dot(dqk.to(k.dtype), k) + + tl.dot(g * decays[:, None], tl.trans(state)) * softmax_scale + ) + + if VD == D: + tl.store(dq_ptrs + n * H * D, dq.to(DQ.dtype.element_ty)) + else: + tl.atomic_add( + dq_ptrs + n * H * D, dq.to(DQ.dtype.element_ty), sem="relaxed" + ) + + state = state + tl.dot((tl.trans(k) * decays[None, :]).to(v.dtype), v) + + +@triton.jit +def fp32_lightning_attention_kv_backward_kernel( + Q, + K, + V, + G, + DK, + DV, + softmax_scale, + stride_q, + stride_k, + stride_v, + stride_g, + decay_scales, + L, + D: tl.constexpr, + KD: tl.constexpr, + VD: tl.constexpr, + BLOCK: tl.constexpr, +): + bid = tl.program_id(0) + hid = tl.program_id(1) + kvid = tl.program_id(2) + N = D // VD + kid = kvid // N + vid = kvid % N + H = tl.num_programs(1) + + c0 = bid * L + + decay_scale = -tl.load(decay_scales + hid) + + offs_b = tl.arange(0, BLOCK) + offs_k = tl.arange(0, KD) + offs_v = tl.arange(0, VD) + + q_ptrs = ( + Q + + c0 * stride_q + + hid * D + + kid * KD + + (offs_b[:, None] * stride_q + offs_k[None, :]) + ) + k_ptrs = ( + K + + c0 * stride_k + + hid * D + + kid * KD + + (offs_b[:, None] * stride_k + offs_k[None, :]) + ) + v_ptrs = ( + V + + c0 * stride_v + + hid * D + + vid * VD + + (offs_b[:, None] * stride_v + offs_v[None, :]) + ) + g_ptrs = ( + G + + c0 * D * H + + hid * D + + vid * VD + + (offs_b[:, None] * stride_g + offs_v[None, :]) + ) + + dk_ptrs = ( + DK + + c0 * H * D + + hid * D + + kid * KD + + (offs_b[:, None] * H * D + offs_k[None, :]) + ) + dv_ptrs = ( + DV + + c0 * H * D + + hid * D + + vid * VD + + (offs_b[:, None] * H * D + offs_v[None, :]) + ) + + gs = tl.zeros((KD, VD), dtype=tl.float32) + + n_steps = tl.cdiv(L, BLOCK) + for i in range(n_steps): + n = (n_steps - i - 1) * BLOCK + n = tl.multiple_of(n, BLOCK) + + q = tl.load(q_ptrs + n * stride_q).to(tl.float32) + k = tl.load(k_ptrs + n * stride_k).to(tl.float32) + v = tl.load(v_ptrs + n * stride_v).to(tl.float32) + g = tl.load(g_ptrs + n * stride_g).to(tl.float32) + b = BLOCK + b_offs = b - 1 - offs_b + + block_decay = tl.exp(decay_scale * b) + amps = tl.exp(-decay_scale * b_offs) # [100, 10, 1] + decays = tl.exp(decay_scale * b_offs) # [0.01, 0.1, 1] + + qs = q * amps[:, None] # [100, 10, 1] + ks = k * decays[:, None] # [0.01, 0.1, 1] + qk = tl.dot(qs, tl.trans(ks)) * softmax_scale + qk = tl.where(offs_b[None, :] <= offs_b[:, None], qk, 0.0) + + dv = tl.dot(tl.trans(qk), g) + + dv += tl.dot(ks, gs) + + dqk = tl.dot(g, tl.trans(v)) + dqk = tl.where(offs_b[None, :] <= offs_b[:, None], dqk, 0.0) * softmax_scale + dk = tl.dot(tl.trans(dqk), qs) + dk += tl.dot(v, tl.trans(gs)) + dk *= decays[:, None] + + gs *= block_decay + gs += tl.dot(tl.trans(qs), g) * softmax_scale * block_decay + + if VD == D: + tl.store(dk_ptrs + n * H * D, dk.to(DK.dtype.element_ty)) + else: + tl.atomic_add( + dk_ptrs + n * H * D, dk.to(DK.dtype.element_ty), sem="relaxed" + ) + if KD == D: + tl.store(dv_ptrs + n * H * D, dv.to(DV.dtype.element_ty)) + else: + tl.atomic_add( + dv_ptrs + n * H * D, dv.to(DV.dtype.element_ty), sem="relaxed" + ) + + +@triton.jit +def lightning_attention_kv_backward_kernel( + Q, + K, + V, + G, + DK, + DV, + softmax_scale, + stride_q, + stride_k, + stride_v, + stride_g, + decay_scales, + L, + D: tl.constexpr, + KD: tl.constexpr, + VD: tl.constexpr, + BLOCK: tl.constexpr, +): + bid = tl.program_id(0) + hid = tl.program_id(1) + kvid = tl.program_id(2) + N = D // VD + kid = kvid // N + vid = kvid % N + H = tl.num_programs(1) + + c0 = bid * L + + decay_scale = -tl.load(decay_scales + hid) + + offs_b = tl.arange(0, BLOCK) + offs_k = tl.arange(0, KD) + offs_v = tl.arange(0, VD) + + q_ptrs = ( + Q + + c0 * stride_q + + hid * D + + kid * KD + + (offs_b[:, None] * stride_q + offs_k[None, :]) + ) + k_ptrs = ( + K + + c0 * stride_k + + hid * D + + kid * KD + + (offs_b[:, None] * stride_k + offs_k[None, :]) + ) + v_ptrs = ( + V + + c0 * stride_v + + hid * D + + vid * VD + + (offs_b[:, None] * stride_v + offs_v[None, :]) + ) + g_ptrs = ( + G + + c0 * D * H + + hid * D + + vid * VD + + (offs_b[:, None] * stride_g + offs_v[None, :]) + ) + + dk_ptrs = ( + DK + + c0 * H * D + + hid * D + + kid * KD + + (offs_b[:, None] * H * D + offs_k[None, :]) + ) + dv_ptrs = ( + DV + + c0 * H * D + + hid * D + + vid * VD + + (offs_b[:, None] * H * D + offs_v[None, :]) + ) + + gs = tl.zeros((KD, VD), dtype=tl.float32) + + b_offs = BLOCK - 1 - offs_b + + block_decay = tl.exp(decay_scale * BLOCK) + amps = tl.exp(-decay_scale * b_offs) # [100, 10, 1] + # decays = tl.exp(decay_scale * b_offs) # [0.01, 0.1, 1] + decays = 1 / amps # [0.01, 0.1, 1] + sd = softmax_scale * block_decay + + mask = tl.exp((offs_b[:, None] - offs_b[None, :]) * decay_scale) + mask = tl.where(offs_b[None, :] <= offs_b[:, None], mask, 0.0) * softmax_scale + + n_steps = tl.cdiv(L, BLOCK) + for i in range(n_steps): + n = (n_steps - i - 1) * BLOCK + n = tl.multiple_of(n, BLOCK) + + q = tl.load(q_ptrs + n * stride_q) + k = tl.load(k_ptrs + n * stride_k) + v = tl.load(v_ptrs + n * stride_v) + g = tl.load(g_ptrs + n * stride_g) + + # qs = q * amps[:, None] # [100, 10, 1] + # ks = k * decays[:, None] # [0.01, 0.1, 1] + qk = tl.dot(q, tl.trans(k)) * mask + + dv = tl.dot(tl.trans(qk).to(g.dtype), g) + + dv += tl.dot(k * decays[:, None], gs) + + dqk = (tl.dot(g, tl.trans(v)) * mask).to(q.dtype) + dk = tl.dot(tl.trans(dqk), (q * amps[:, None]).to(q.dtype)) + dk = tl.dot(v, tl.trans(gs.to(v.dtype)), dk) + dk *= decays[:, None] + + gs *= block_decay + gs += tl.dot(tl.trans(q * amps[:, None]).to(g.dtype), g) * sd + + if VD == D: + tl.store(dk_ptrs + n * H * D, dk.to(DK.dtype.element_ty)) + else: + tl.atomic_add( + dk_ptrs + n * H * D, dk.to(DK.dtype.element_ty), sem="relaxed" + ) + if KD == D: + tl.store(dv_ptrs + n * H * D, dv.to(DV.dtype.element_ty)) + else: + tl.atomic_add( + dv_ptrs + n * H * D, dv.to(DV.dtype.element_ty), sem="relaxed" + ) + + +def triton_lightning_attention_backward( + output_grad, q, k, v, decay_scales, softmax_scale=None, hpc=False, hp=False +): + B, L, H, D = q.shape + if softmax_scale is None: + softmax_scale = D ** (-0.5) + + dtype = q.dtype + device = q.device + + KD = 64 + VD = 128 + BLOCK = 32 + k_dim_block = D // KD + v_dim_block = D // VD + if v_dim_block > 1: + dq = torch.zeros( + (B, L, H, D), device=device, dtype=torch.float32 if hpc else dtype + ) + else: + dq = torch.empty((B, L, H, D), device=device, dtype=dtype) + assert L % BLOCK == 0 and BLOCK <= 64 + grid = (B, H, k_dim_block * v_dim_block) + num_warps = 4 # 2 + num_stages = 3 # 5 + kernel = ( + fp32_lightning_attention_q_backward_kernel + if hp + else lightning_attention_q_backward_kernel + ) + kernel[grid]( + q, + k, + v, + output_grad, + dq, + softmax_scale, + q.stride(1), + k.stride(1), + v.stride(1), + output_grad.stride(1), + decay_scales, + L, + D=D, + KD=KD, + VD=VD, + BLOCK=BLOCK, + num_warps=num_warps, + num_stages=num_stages, + ) + + KD = 32 + VD = 128 + BLOCK = 32 + k_dim_block = D // KD + v_dim_block = D // VD + if v_dim_block > 1: + dk = torch.zeros( + (B, L, H, D), device=device, dtype=torch.float32 if hpc else dtype + ) + else: + dk = torch.empty((B, L, H, D), device=device, dtype=dtype) + if k_dim_block > 1: + dv = torch.zeros( + (B, L, H, D), device=device, dtype=torch.float32 if hpc else dtype + ) + else: + dv = torch.empty((B, L, H, D), device=device, dtype=dtype) + num_warps = 4 # 4 + num_stages = 5 # 5 + grid = (B, H, k_dim_block * v_dim_block) + kernel = ( + fp32_lightning_attention_kv_backward_kernel + if hp + else lightning_attention_kv_backward_kernel + ) + kernel[grid]( + q, + k, + v, + output_grad, + dk, + dv, + softmax_scale, + q.stride(1), + k.stride(1), + v.stride(1), + output_grad.stride(1), + decay_scales, + L, + D=D, + KD=KD, + VD=VD, + BLOCK=BLOCK, + num_warps=num_warps, + num_stages=num_stages, + ) + + if v_dim_block > 1 and hpc: + dq = dq.to(dtype) + dk = dk.to(dtype) + if k_dim_block > 1 and hpc: + dv = dv.to(dtype) + return dq, dk, dv + + +@triton.jit +def fused_lightning_attention_backward_kernel( + Q, + K, + V, + S, + G, + DQ, + DK, + DV, + softmax_scale, + stride_q, + stride_k, + stride_v, + stride_g, + decay_scales, + L, + D: tl.constexpr, + KD: tl.constexpr, + VD: tl.constexpr, + BLOCK: tl.constexpr, +): + bid = tl.program_id(0) + hid = tl.program_id(1) + kvid = tl.program_id(2) + N = D // VD + kid = kvid // N + vid = kvid % N + H = tl.num_programs(1) + + c0 = bid * L + + decay_scale = -tl.load(decay_scales + hid) + + offs_b = tl.arange(0, BLOCK) + offs_k = tl.arange(0, KD) + offs_v = tl.arange(0, VD) + + q_ptrs = ( + Q + + c0 * stride_q + + hid * D + + kid * KD + + (offs_b[:, None] * stride_q + offs_k[None, :]) + ) + k_ptrs = ( + K + + c0 * stride_k + + hid * D + + kid * KD + + (offs_b[:, None] * stride_k + offs_k[None, :]) + ) + v_ptrs = ( + V + + c0 * stride_v + + hid * D + + vid * VD + + (offs_b[:, None] * stride_v + offs_v[None, :]) + ) + g_ptrs = ( + G + + c0 * D * H + + hid * D + + vid * VD + + (offs_b[:, None] * stride_g + offs_v[None, :]) + ) + dq_ptrs = ( + DQ + + c0 * D * H + + hid * D + + kid * KD + + (offs_b[:, None] * H * D + offs_k[None, :]) + ) + dk_ptrs = ( + DK + + c0 * H * D + + hid * D + + kid * KD + + (offs_b[:, None] * H * D + offs_k[None, :]) + ) + dv_ptrs = ( + DV + + c0 * H * D + + hid * D + + vid * VD + + (offs_b[:, None] * H * D + offs_v[None, :]) + ) + s_ptrs = ( + S + + bid * H * D * D + + hid * D * D + + kid * D * KD + + vid * VD + + (offs_k[:, None] * D + offs_v[None, :]) + ) + + state = tl.load(s_ptrs).to(tl.float32) + gs = tl.zeros((KD, VD), dtype=tl.float32) + + n_steps = tl.cdiv(L, BLOCK) + for i in range(n_steps): + n = (n_steps - i - 1) * BLOCK + n = tl.multiple_of(n, BLOCK) + + q = tl.load(q_ptrs + n * stride_q).to(tl.float32) + k = tl.load(k_ptrs + n * stride_k).to(tl.float32) + v = tl.load(v_ptrs + n * stride_v).to(tl.float32) + g = tl.load(g_ptrs + n * stride_g).to(tl.float32) + b = BLOCK + b_offs = b - 1 - offs_b + + block_decay = tl.exp(decay_scale * b) + amps = tl.exp(-decay_scale * b_offs) # [100, 10, 1] + decays = tl.exp(decay_scale * b_offs) # [0.01, 0.1, 1] + + qs = q * amps[:, None] # [100, 10, 1] + ks = k * decays[:, None] # [0.01, 0.1, 1] + qk = tl.dot(qs, tl.trans(ks)) * softmax_scale + qk = tl.where(offs_b[None, :] <= offs_b[:, None], qk, 0.0) + + # state = state * block_decay + # o = tl.dot(q * decays, state) * softmax_scale + tl.dot(qk, v) + # state = state + tl.dot(tl.trans(k * decays[:, None]), v) + + state = state - tl.dot(tl.trans(ks), v) + + dv = tl.dot(tl.trans(qk), g) + + dv += tl.dot(ks, gs) + + dqk = tl.dot(g, tl.trans(v)) + dqk = tl.where(offs_b[None, :] <= offs_b[:, None], dqk, 0.0) * softmax_scale + dk = tl.dot(tl.trans(dqk), qs) + dk += tl.dot(v, tl.trans(gs)) + dk *= decays[:, None] + + dq = ( + tl.dot(dqk, ks) + + tl.dot(g, tl.trans(state)) * softmax_scale * decays[:, None] + ) + dq *= amps[:, None] + state /= block_decay + + # dq = tl.dot(dqk, ks) + tl.dot(g, tl.trans(state)) * softmax_scale + # dq *= amps[:, None] + + gs *= block_decay + gs += tl.dot(tl.trans(qs), g) * softmax_scale * block_decay + + if VD == D: + tl.store(dq_ptrs + n * H * D, dq.to(DQ.dtype.element_ty)) + tl.store(dk_ptrs + n * H * D, dk.to(DK.dtype.element_ty)) + else: + tl.atomic_add( + dq_ptrs + n * H * D, dq.to(DQ.dtype.element_ty), sem="relaxed" + ) + tl.atomic_add( + dk_ptrs + n * H * D, dk.to(DK.dtype.element_ty), sem="relaxed" + ) + if KD == D: + tl.store(dv_ptrs + n * H * D, dv.to(DV.dtype.element_ty)) + else: + tl.atomic_add( + dv_ptrs + n * H * D, dv.to(DV.dtype.element_ty), sem="relaxed" + ) + + +def triton_fused_lightning_attention_backward( + output_grad, q, k, v, s, decay_scales, softmax_scale=None, hpc=False +): + B, L, H, D = q.shape + if softmax_scale is None: + softmax_scale = D ** (-0.5) + + dtype = q.dtype + device = q.device + KD = 128 + VD = 32 + BLOCK = 32 + + k_dim_block = D // KD + v_dim_block = D // VD + if v_dim_block > 1: + dq = torch.zeros( + (B, L, H, D), device=device, dtype=torch.float32 if hpc else dtype + ) + dk = torch.zeros( + (B, L, H, D), device=device, dtype=torch.float32 if hpc else dtype + ) + else: + dq = torch.empty((B, L, H, D), device=device, dtype=dtype) + dk = torch.empty((B, L, H, D), device=device, dtype=dtype) + if k_dim_block > 1: + dv = torch.zeros( + (B, L, H, D), device=device, dtype=torch.float32 if hpc else dtype + ) + else: + dv = torch.empty((B, L, H, D), device=device, dtype=dtype) + + assert L % BLOCK == 0 and BLOCK <= 64 + grid = (B, H, k_dim_block * v_dim_block) + + num_warps = 4 # 2 + num_stages = 3 # 3 + fused_lightning_attention_backward_kernel[grid]( + q, + k, + v, + s, + output_grad, + dq, + dk, + dv, + softmax_scale, + q.stride(1), + k.stride(1), + v.stride(1), + output_grad.stride(1), + decay_scales, + L, + D=D, + KD=KD, + VD=VD, + BLOCK=BLOCK, + num_warps=num_warps, + num_stages=num_stages, + ) + if hpc: + if v_dim_block > 1: + dq = dq.to(dtype) + dk = dk.to(dtype) + if k_dim_block > 1: + dv = dv.to(dtype) + return dq, dk, dv + + +# @triton.jit +# def varlen_lightning_attention_forward_kernel( +# Q, +# K, +# V, +# S, +# Out, +# softmax_scale, +# stride_q, +# stride_k, +# stride_v, +# stride_s, +# stride_o, +# CU, +# PCU, +# decay_scales, +# D: tl.constexpr, +# KD: tl.constexpr, +# VD: tl.constexpr, +# BLOCK: tl.constexpr, +# EVEN: tl.constexpr, +# PAD: tl.constexpr +# ): +# bid = tl.program_id(0) +# hid = tl.program_id(1) +# kvid = tl.program_id(2) +# N = D // VD +# kid = kvid // N +# vid = kvid % N +# H = tl.num_programs(1) + +# if PAD: +# c01 = tl.load(CU + bid + tl.arange(0, 2)) +# c0, c1 = tl.split(c01) +# length = c1 - c0 +# pc01 = tl.load(PCU + bid + tl.arange(0, 2)) +# pc0, pc1 = tl.split(pc01) +# padded_length = pc1 - pc0 +# c0 = pc0 +# if padded_length == 0: +# return +# else: +# c01 = tl.load(CU + bid + tl.arange(0, 2)) +# c0, c1 = tl.split(c01) +# length = c1 - c0 +# padded_length = length +# if length == 0: +# return + +# decay_scale = -tl.load(decay_scales + hid) + +# offs_b = tl.arange(0, BLOCK) +# offs_k = tl.arange(0, KD) +# offs_v = tl.arange(0, VD) + +# q_ptrs = ( +# Q +# + c0 * stride_q +# + hid * D +# + kid * KD +# + (offs_b[:, None] * stride_q + offs_k[None, :]) +# ) +# k_ptrs = ( +# K +# + c0 * stride_k +# + hid * D +# + kid * KD +# + (offs_b[:, None] * stride_k + offs_k[None, :]) +# ) +# v_ptrs = ( +# V +# + c0 * stride_v +# + hid * D +# + vid * VD +# + (offs_b[:, None] * stride_v + offs_v[None, :]) +# ) +# # (num_dim_block, length, qo_heads, d) +# out_ptrs = ( +# Out +# + kid * stride_o +# + c0 * D * H +# + hid * D +# + vid * VD +# + (offs_b[:, None] * H * D + offs_v[None, :]) +# ) +# s_ptrs = ( +# S +# + bid * stride_s +# + hid * D * D +# + kid * D * KD +# + vid * VD +# + (offs_k[:, None] * D + offs_v[None, :]) +# ) +# state = tl.zeros((KD, VD), dtype=tl.float32) + +# for n in range(0, padded_length, BLOCK): +# n = tl.multiple_of(n, BLOCK) + +# if EVEN: +# q = tl.load(q_ptrs + n * stride_q).to(tl.float32) +# k = tl.trans(tl.load(k_ptrs + n * stride_k)).to(tl.float32) +# v = tl.load(v_ptrs + n * stride_v).to(tl.float32) +# b = BLOCK +# b_offs = b - 1 - offs_b +# decays = tl.exp(decay_scale * b_offs) +# inv_decays = 1 / decays +# else: +# q = tl.load( +# q_ptrs + n * stride_q, mask=(n + offs_b)[:, None] < length, other=0.0 +# ).to(tl.float32) +# k = tl.trans( +# tl.load( +# k_ptrs + n * stride_k, +# mask=(n + offs_b)[:, None] < length, +# other=0.0, +# ) +# ).to(tl.float32) +# v = tl.load( +# v_ptrs + n * stride_v, mask=(n + offs_b)[:, None] < length, other=0.0 +# ).to(tl.float32) +# b = min(BLOCK, length - n) +# b_offs = b - 1 - offs_b +# block_decays = tl.exp(decay_scale * b_offs) +# decays = tl.where(b_offs >= 0, block_decays, 0) +# inv_decays = tl.where(b_offs >= 0, 1 / block_decays, 0) + +# q = q * inv_decays[:, None] +# k = k * decays[None, :] +# qk = tl.dot(q, k) * softmax_scale +# qk = tl.where(offs_b[None, :] <= offs_b[:, None], qk, 0.0) +# o = tl.dot(qk, v) + +# block_decay = tl.exp(decay_scale * b) +# o = tl.dot(q, state) * block_decay * softmax_scale + o + +# state = state * block_decay + tl.dot(k, v) + +# if EVEN: +# tl.store(out_ptrs + n * H * D, o.to(Out.dtype.element_ty)) +# else: +# tl.store( +# out_ptrs + n * H * D, +# o.to(Out.dtype.element_ty), +# mask=(n + offs_b)[:, None] < length, +# ) +# tl.store(s_ptrs, state.to(S.dtype.element_ty)) + + +# def triton_varlen_lightning_attention_forward(q, k, v, decay_scales, cu_seqlens, padded_cu_seqlens, max_q_length, softmax_scale=None): +# length, qo_heads, D = q.shape +# _, kv_heads, _ = k.shape +# bs = cu_seqlens.size(0) - 1 +# if softmax_scale is None: +# softmax_scale = D ** (-0.5) + +# MAX_LENGTH = max_q_length + +# assert qo_heads == kv_heads, "triton_lightning_attention_forward does NOT support GQA currently" + +# KD = 32 +# VD = 32 if bs <= 2 else 64 + +# num_warps = 2 # 2 +# num_stages = 3 # 3 + +# k_dim_block = D // KD +# v_dim_block = D // VD +# tmp = torch.empty( +# (k_dim_block, length, qo_heads, D), device=q.device, dtype=q.dtype +# ) +# s = torch.empty( +# (bs, qo_heads, D, D), device=q.device, dtype=torch.float32 +# ) + +# # BLOCK should <= 64 +# BLOCK = 32 +# EVEN = MAX_LENGTH % BLOCK == 0 if bs == 1 else False +# grid = (bs, kv_heads, k_dim_block * v_dim_block) +# varlen_lightning_attention_forward_kernel[grid]( +# q, +# k, +# v, +# s, +# tmp, +# softmax_scale, +# q.stride(0), +# k.stride(0), +# v.stride(0), +# s.stride(0), +# tmp.stride(0), +# cu_seqlens, +# padded_cu_seqlens, +# decay_scales, +# D=D, +# KD=KD, +# VD=VD, +# BLOCK=BLOCK, +# EVEN=EVEN, +# num_warps=num_warps, +# num_stages=num_stages, +# ) + +# if k_dim_block > 1: +# if length < 2048: +# o = tmp.sum(0) +# else: +# o = torch.empty( +# (length, qo_heads, D), device=q.device, dtype=q.dtype +# ) +# output_sum_kernel[(length,)]( +# tmp, +# o, +# DIM=qo_heads * D, +# NUM_BLOCK=k_dim_block, +# num_warps=2, +# num_stages=3, +# ) +# else: +# o = tmp[0] + +# return o, s + + +# @triton.jit +# def varlen_lightning_attention_backward_kernel( +# Q, +# K, +# V, +# S, +# G, +# DQ, +# DK, +# DV, +# softmax_scale, +# stride_q, +# stride_k, +# stride_v, +# stride_g, +# CU, +# PCU, +# decay_scales, +# D: tl.constexpr, +# KD: tl.constexpr, +# VD: tl.constexpr, +# BLOCK: tl.constexpr, +# EVEN: tl.constexpr, +# PAD: tl.constexpr +# ): +# bid = tl.program_id(0) +# hid = tl.program_id(1) +# kvid = tl.program_id(2) +# N = D // VD +# kid = kvid // N +# vid = kvid % N +# H = tl.num_programs(1) + +# if PAD: +# c01 = tl.load(CU + bid + tl.arange(0, 2)) +# c0, c1 = tl.split(c01) +# length = c1 - c0 +# pc01 = tl.load(PCU + bid + tl.arange(0, 2)) +# pc0, pc1 = tl.split(pc01) +# padded_length = pc1 - pc0 +# c0 = pc0 +# if padded_length == 0: +# return +# else: +# c01 = tl.load(CU + bid + tl.arange(0, 2)) +# c0, c1 = tl.split(c01) +# length = c1 - c0 +# padded_length = length +# if length == 0: +# return + +# decay_scale = -tl.load(decay_scales + hid) + +# offs_b = tl.arange(0, BLOCK) +# offs_k = tl.arange(0, KD) +# offs_v = tl.arange(0, VD) + +# q_ptrs = ( +# Q +# + c0 * stride_q +# + hid * D +# + kid * KD +# + (offs_b[:, None] * stride_q + offs_k[None, :]) +# ) +# k_ptrs = ( +# K +# + c0 * stride_k +# + hid * D +# + kid * KD +# + (offs_b[:, None] * stride_k + offs_k[None, :]) +# ) +# v_ptrs = ( +# V +# + c0 * stride_v +# + hid * D +# + vid * VD +# + (offs_b[:, None] * stride_v + offs_v[None, :]) +# ) +# g_ptrs = ( +# G +# + c0 * D * H +# + hid * D +# + kid * KD +# + (offs_b[:, None] * stride_q + offs_k[None, :]) +# ) +# # (num_dim_block, length, qo_heads, d) +# dq_ptrs = ( +# DQ +# + c0 * D * H +# + hid * D +# + vid * VD +# + (offs_b[:, None] * H * D + offs_k[None, :]) +# ) +# dk_ptrs = ( +# DK +# + c0 * stride_k +# + hid * D +# + kid * KD +# + (offs_b[:, None] * stride_k + offs_k[None, :]) +# ) +# dv_ptrs = ( +# DV +# + c0 * stride_v +# + hid * D +# + vid * VD +# + (offs_b[:, None] * stride_v + offs_v[None, :]) +# ) +# s_ptrs = ( +# S +# + bid * H * D * D +# + hid * D * D +# + kid * D * KD +# + vid * VD +# + (offs_k[:, None] * D + offs_v[None, :]) +# ) +# state = tl.load(s_ptrs).to(tl.float32) +# gs = tl.zeros((KD, VD), dtype=tl.float32) +# n_steps = tl.cdiv(padded_length, BLOCK) +# for i in range(n_steps): +# n = (n_steps - i - 1) * BLOCK +# n = tl.multiple_of(n, BLOCK) + +# if EVEN: +# q = tl.load(q_ptrs + n * stride_q).to(tl.float32) +# k = tl.trans(tl.load(k_ptrs + n * stride_k)).to(tl.float32) +# v = tl.load(v_ptrs + n * stride_v).to(tl.float32) +# g = tl.load(g_ptrs + n * stride_g).to(tl.float32) +# b = BLOCK +# b_offs = b - 1 - offs_b +# decays = tl.exp(decay_scale * b_offs) +# inv_decays = 1 / decays +# else: +# q = tl.load( +# q_ptrs + n * stride_q, mask=(n + offs_b)[:, None] < length, other=0.0 +# ).to(tl.float32) +# k = tl.trans( +# tl.load( +# k_ptrs + n * stride_k, +# mask=(n + offs_b)[:, None] < length, +# other=0.0, +# ) +# ).to(tl.float32) +# v = tl.load( +# v_ptrs + n * stride_v, mask=(n + offs_b)[:, None] < length, other=0.0 +# ).to(tl.float32) +# g = tl.load( +# g_ptrs + n * stride_g, mask=(n + offs_b)[:, None] < length, other=0.0 +# ).to(tl.float32) +# b = min(BLOCK, length - n) +# b_offs = b - 1 - offs_b +# block_decays = tl.exp(decay_scale * b_offs) +# decays = tl.where(b_offs >= 0, block_decays, 0) +# inv_decays = tl.where(b_offs >= 0, 1 / block_decays, 0) +# block_decay = tl.exp(decay_scale * b) + +# q = q * inv_decays[:, None] +# k = k * decays[None, :] +# qk = tl.dot(q, k) * softmax_scale +# qk = tl.where(offs_b[None, :] <= offs_b[:, None], qk, 0.0) +# dv = tl.dot(tl.trans(qk), g) +# dqk = tl.dot(g, tl.trans(v)) +# dqk = tl.where(offs_b[None, :] <= offs_b[:, None], dqk, 0.0) +# dk = tl.dot(tl.trans(dqk), q) +# dq = tl.dot(dqk, k) + +# o = tl.dot(q, state) * block_decay * softmax_scale + o + +# state = (state - tl.dot(k, v))/block_decay +# dq += tl.dot(g, tl.trans(state)) * block_decay * softmax_scale +# dv += tl.dot(k, gs) +# dk += tl.dot(v, tl.trans(gs)) +# gs += tl.dot(tl.trans(q), g) + +# if EVEN: +# tl.store(dq_ptrs + n * H * D, dq.to(DQ.dtype.element_ty)) +# tl.store(dk_ptrs + n * H * D, dk.to(DK.dtype.element_ty)) +# tl.store(dv_ptrs + n * H * D, dv.to(DV.dtype.element_ty)) + +# else: +# tl.store( +# dq_ptrs + n * H * D, +# dq.to(DQ.dtype.element_ty), +# mask=(n + offs_b)[:, None] < length, +# ) +# tl.store( +# dk_ptrs + n * H * D, +# dk.to(DQ.dtype.element_ty), +# mask=(n + offs_b)[:, None] < length, +# ) +# tl.store( +# dv_ptrs + n * H * D, +# dv.to(DQ.dtype.element_ty), +# mask=(n + offs_b)[:, None] < length, +# ) + + +# def triton_varlen_lightning_attention_backward(output_grad, q, k, v, s, decay_scales, cu_seqlens, padded_cu_seqlens, max_q_length, softmax_scale=None): +# length, qo_heads, D = q.shape +# _, kv_heads, _ = k.shape +# bs = cu_seqlens.size(0) - 1 +# if softmax_scale is None: +# softmax_scale = D ** (-0.5) + +# MAX_LENGTH = max_q_length + +# assert qo_heads == kv_heads, "triton_lightning_attention_forward does NOT support GQA currently" + +# KD = 128 +# VD = 128 + +# num_warps = 2 # 2 +# num_stages = 3 # 3 + +# k_dim_block = D // KD +# v_dim_block = D // VD +# dq = torch.empty( +# (length, qo_heads, D), device=q.device, dtype=q.dtype +# ) +# dk = torch.empty( +# (length, qo_heads, D), device=q.device, dtype=q.dtype +# ) +# dv = torch.empty( +# (length, qo_heads, D), device=q.device, dtype=q.dtype +# ) + +# # BLOCK should <= 64 +# BLOCK = 32 +# EVEN = MAX_LENGTH % BLOCK == 0 if bs == 1 else False +# grid = (bs, kv_heads, k_dim_block * v_dim_block) +# varlen_lightning_attention_backward_kernel[grid]( +# q, +# k, +# v, +# s, +# output_grad, +# dq, +# dk, +# dv, +# softmax_scale, +# q.stride(0), +# k.stride(0), +# v.stride(0), +# output_grad.stride(0), +# cu_seqlens, +# padded_cu_seqlens, +# decay_scales, +# D=D, +# KD=KD, +# VD=VD, +# BLOCK=BLOCK, +# EVEN=EVEN, +# num_warps=num_warps, +# num_stages=num_stages, +# ) + +# return dq, dk, dv diff --git a/linghe/attn/mla.py b/linghe/attn/mla.py new file mode 100644 index 0000000..dce5eaf --- /dev/null +++ b/linghe/attn/mla.py @@ -0,0 +1,1934 @@ +# -*- coding: utf-8 -*- +""" +Copyright (c) Ant Financial Service Group and its affiliates. +""" + +import math + +import torch +import triton +import triton.language as tl + + +@triton.jit +def deprecated_mla_forward_kernel( + Q, + K, + V, + Out, + LSE, + ML, + softmax_scale, + stride_q, + stride_k, + stride_v, + L, + M: tl.constexpr, + N: tl.constexpr, + CAUSAL: tl.constexpr, +): + bid = tl.program_id(0) + hid = tl.program_id(1) + mid = tl.num_programs(2) - tl.program_id(2) - 1 + H = tl.num_programs(1) + + offs_m = tl.arange(0, M) + offs_n = tl.arange(0, N) + + offs_0 = tl.arange(0, 128) # nope + offs_1 = tl.arange(0, 64) # pe + + # [B, L, H, 192】 + q0_ptrs = ( + Q + + (bid * L + mid * M) * stride_q + + hid * 192 + + (offs_m[:, None] * stride_q + offs_0[None, :]) + ) + q1_ptrs = ( + Q + + (bid * L + mid * M) * stride_q + + hid * 192 + + 128 + + (offs_m[:, None] * stride_q + offs_1[None, :]) + ) + + k0_ptrs = ( + K + + bid * L * stride_k + + hid * 192 + + (offs_n[:, None] * stride_k + offs_0[None, :]) + ) + + k1_ptrs = ( + K + + bid * L * stride_k + + hid * 192 + + 128 + + (offs_n[:, None] * stride_k + offs_1[None, :]) + ) + + v_ptrs = ( + V + + bid * L * stride_v + + hid * 128 + + (offs_n[:, None] * stride_v + offs_0[None, :]) + ) + + q0 = tl.load(q0_ptrs) + q1 = tl.load(q1_ptrs) + + lse = tl.zeros((M,), dtype=tl.float32) + acc_o = tl.zeros((M, 128), dtype=tl.float32) + if CAUSAL: + steps = tl.cdiv(mid * M + M, N) + else: + steps = L // N + for i in range(0, steps): + n = i * N + n = tl.multiple_of(n, N) + + k1 = tl.load(k1_ptrs + n * stride_k) + + qk = tl.dot(q1, tl.trans(k1)) + + k0 = tl.load(k0_ptrs + n * stride_k) + + qk = tl.dot(q0, tl.trans(k0), qk) + + qk += tl.where((mid * M + offs_m)[:, None] >= (n + offs_n)[None, :], 0.0, -1e9) + + p = tl.exp(qk * softmax_scale) + lse += tl.sum(p, 1) + + v = tl.load(v_ptrs + n * stride_v) + p = p.to(V.dtype.element_ty) + acc_o = tl.dot(p, v, acc_o) + + acc_o = acc_o / lse[:, None] + + # [B, L, H, 128] + out_ptrs = ( + Out + + (bid * L + mid * M) * H * 128 + + hid * 128 + + (offs_m[:, None] * 128 * H + offs_0[None, :]) + ) + + tl.store(out_ptrs, acc_o) + tl.store(LSE + bid * H * L + hid * L + mid * M + tl.arange(0, M), lse) + + +@triton.jit +def mla_forward_kernel( + Q, + K, + V, + Out, + LSE, + ML, + softmax_scale, + clip_value, + stride_q, + stride_k, + stride_v, + L, + M: tl.constexpr, + N: tl.constexpr, + CAUSAL: tl.constexpr, + SAFE: tl.constexpr, + CLIP: tl.constexpr, +): + bid = tl.program_id(0) + hid = tl.program_id(1) + mid = tl.num_programs(2) - tl.program_id(2) - 1 + H = tl.num_programs(1) + + offs_m = tl.arange(0, M) + offs_n = tl.arange(0, N) + + offs_0 = tl.arange(0, 128) + offs_1 = tl.arange(0, 64) + + # [B, L, H, 192】 + q0_ptrs = ( + Q + + (bid * L + mid * M) * stride_q + + hid * 192 + + (offs_m[:, None] * stride_q + offs_1[None, :]) + ) + + k0_ptrs = ( + K + + bid * L * stride_k + + hid * 192 + + (offs_n[:, None] * stride_k + offs_1[None, :]) + ) + + v_ptrs = ( + V + + bid * L * stride_v + + hid * 128 + + (offs_n[:, None] * stride_v + offs_0[None, :]) + ) + + q0 = tl.load(q0_ptrs) + q1 = tl.load(q0_ptrs + 64) + q2 = tl.load(q0_ptrs + 128) + + acc_o = tl.zeros((M, 128), dtype=tl.float32) + if SAFE: + max_logits = tl.zeros((M,), dtype=tl.float32) - 10000.0 + lse = tl.zeros((M,), dtype=tl.float32) + else: + lse = tl.zeros((M,), dtype=tl.float32) + 1e-30 + + if CAUSAL: + steps = tl.cdiv(mid * M + M, N) + else: + steps = L // N + + for i in range(0, steps): + n = i * N + n = tl.multiple_of(n, N) + + k0 = tl.load(k0_ptrs + n * stride_k) + k1 = tl.load(k0_ptrs + n * stride_k + 64) + k2 = tl.load(k0_ptrs + n * stride_k + 128) + + qk = tl.dot(q0, tl.trans(k0)) + qk = tl.dot(q1, tl.trans(k1), qk) + qk = tl.dot(q2, tl.trans(k2), qk) + + if CAUSAL: + qk += tl.where( + (mid * M + offs_m)[:, None] >= (n + offs_n)[None, :], 0.0, -1e9 + ) + + qk *= softmax_scale + + if SAFE: + latest_max_logits = tl.maximum(max_logits, tl.max(qk, 1)) + p = tl.exp(qk - latest_max_logits[:, None]) + rescale = tl.exp(max_logits - latest_max_logits) + lse = lse * rescale + tl.sum(p, 1) + v = tl.load(v_ptrs + n * stride_v) + p = p.to(V.dtype.element_ty) + acc_o = acc_o * rescale[:, None] + acc_o = tl.dot(p, v, acc_o) + max_logits = latest_max_logits + else: + if CLIP: + p = tl.exp(tl.minimum(qk, clip_value)) + else: + p = tl.exp(qk) + lse += tl.sum(p, 1) + v = tl.load(v_ptrs + n * stride_v) + p = p.to(V.dtype.element_ty) + acc_o = tl.dot(p, v, acc_o) + + acc_o = acc_o / lse[:, None] + + # [B, L, H, 128] + out_ptrs = ( + Out + + (bid * L + mid * M) * H * 128 + + hid * 128 + + (offs_m[:, None] * 128 * H + offs_0[None, :]) + ) + + tl.store(out_ptrs, acc_o) + tl.store(LSE + bid * H * L + hid * L + mid * M + tl.arange(0, M), lse) + if SAFE: + tl.store(ML + bid * H * L + hid * L + mid * M + tl.arange(0, M), max_logits) + + +def triton_mla_forward(q, k, v, causal=True, safe=True, clip_value=None): + # q: [B, L, H, 192] + # k: [B, L, H, 192] + # v: [B, L, H, 128] + B, L, H, _ = q.shape + assert k.size(1) == L + M = 256 + N = 64 + assert L % M == 0 + assert L % N == 0 + assert M >= N + + o = torch.empty((B, L, H, 128), dtype=q.dtype, device=q.device) + lse = torch.empty((B, H, L), dtype=torch.float32, device=q.device) + max_logits = torch.empty((B, H, L), dtype=torch.float32, device=q.device) + softmax_scale = 128 ** (-0.5) + + clip = clip_value is not None and clip_value > 0.0 + clip_value = clip_value * softmax_scale if clip else 0.0 + + if clip and clip_value + math.log(L) < 88.7: + safe = False + + num_m_block = L // M + num_stages = 2 + num_warps = 8 + + grid = (B, H, num_m_block) + mla_forward_kernel[grid]( + q, + k, + v, + o, + lse, + max_logits, + softmax_scale, + clip_value, + q.stride(1), + k.stride(1), + v.stride(1), + L, + M, + N, + causal, + safe, + clip, + num_warps=num_warps, + num_stages=num_stages, + ) + return o, lse, max_logits + + +# dp and p dot sum +@triton.jit +def naive_mla_ds_kernel( + GO, + Q, + K, + V, + LSE, + ML, + DS, + softmax_scale, + stride_q, + stride_k, + stride_v, + L, + M: tl.constexpr, + N: tl.constexpr, +): + bid = tl.program_id(0) + hid = tl.program_id(1) + mid = tl.num_programs(2) - tl.program_id(2) - 1 + H = tl.num_programs(1) + + offs_m = tl.arange(0, M) + offs_n = tl.arange(0, N) + + offs_0 = tl.arange(0, 128) # nope + offs_1 = tl.arange(0, 64) # pe + + # [B, L, H, 192】 + q0_ptrs = ( + Q + + (bid * L + mid * M) * stride_q + + hid * 192 + + (offs_m[:, None] * stride_q + offs_1[None, :]) + ) + + k0_ptrs = ( + K + + bid * L * stride_k + + hid * 192 + + (offs_n[:, None] * stride_k + offs_1[None, :]) + ) + + v_ptrs = ( + V + + bid * L * stride_v + + hid * 128 + + (offs_n[:, None] * stride_v + offs_0[None, :]) + ) + + go_ptrs = ( + GO + + (bid * L + mid * M) * H * 128 + + hid * 128 + + (offs_m[:, None] * 128 * H + offs_0[None, :]) + ) + + ds = tl.zeros((M,), dtype=tl.float32) + + q0 = tl.load(q0_ptrs) + q1 = tl.load(q0_ptrs + 64) + q2 = tl.load(q0_ptrs + 128) + + go = tl.load(go_ptrs) + steps = tl.cdiv(mid * M + M, N) + + ds = tl.zeros((M, N), dtype=tl.float32) + + for i in range(0, steps): + n = i * N + n = tl.multiple_of(n, N) + + k0 = tl.load(k0_ptrs + n * stride_k) + + qk = tl.dot(q0, tl.trans(k0)) + + k1 = tl.load(k0_ptrs + n * stride_k + 64) + + qk = tl.dot(q1, tl.trans(k1), qk) + + k2 = tl.load(k0_ptrs + n * stride_k + 128) + + qk = tl.dot(q2, tl.trans(k2), qk) + + qk += tl.where((mid * M + offs_m)[:, None] >= (n + offs_n)[None, :], 0.0, -1e9) + + p = tl.exp(qk * softmax_scale) # [M, N] + v = tl.load(v_ptrs + n * stride_v) + dp = tl.dot(go, tl.trans(v)) # [M, 128]@[128, N]=[M,N] + # ds += tl.sum(p * dp, 1) # score + ds += p * dp # score + + lse = tl.load(LSE + bid * H * L + hid * L + mid * M + tl.arange(0, M)) + ds = ds.sum(1) / lse + tl.store(DS + bid * H * L + hid * L + mid * M + tl.arange(0, M), ds) + + +# dp and p dot sum +@triton.jit +def mla_ds_kernel( + G, + O, + DS, + L, + M: tl.constexpr, +): + bid = tl.program_id(0) + hid = tl.program_id(1) + mid = tl.program_id(2) + H = tl.num_programs(1) + + offs_m = tl.arange(0, M) + offs_0 = tl.arange(0, 128) # nope + + # [B, L, H, 128】 + offs = ( + (bid * L + mid * M) * H * 128 + + hid * 128 + + (offs_m[:, None] * H * 128 + offs_0[None, :]) + ) + + mask = mid * M + offs_m < L + g = tl.load(G + offs, mask=mask[:, None]).to(tl.float32) + o = tl.load(O + offs, mask=mask[:, None]).to(tl.float32) + ds = tl.sum(g * o, 1) + tl.store(DS + bid * H * L + hid * L + mid * M + tl.arange(0, M), ds, mask=mask) + + +@triton.jit +def deprecated_mla_backward_kernel( + GO, + Q, + K, + V, + GQ, + GK, + GV, + LSE, + ML, + DS, + softmax_scale, + stride_q, + stride_k, + stride_v, + L, + M: tl.constexpr, + N: tl.constexpr, + ATOMIC: tl.constexpr, # not used + CAUSAL: tl.constexpr, +): + bid = tl.program_id(0) + hid = tl.program_id(1) + nid = tl.program_id(2) + H = tl.num_programs(1).to(tl.int64) + B = tl.num_programs(0) + + offs_m = tl.arange(0, M) + offs_n = tl.arange(0, N) + + offs_1 = tl.arange(0, 64) # pe + + # [B, L, H, 192】 + q0_ptrs = ( + Q + + bid * L * stride_q + + hid * 192 + + (offs_m[:, None] * stride_q + offs_1[None, :]) + ) + + k0_ptrs = ( + K + + (bid * L + nid * N) * stride_k + + hid * 192 + + (offs_n[:, None] * stride_k + offs_1[None, :]) + ) + + k0 = tl.load(k0_ptrs) + k1 = tl.load(k0_ptrs + 64) + k2 = tl.load(k0_ptrs + 128) + + v0_ptrs = ( + V + + (bid * L + nid * N) * stride_v + + hid * 128 + + (offs_n[:, None] * stride_v + offs_1[None, :]) + ) + v0 = tl.load(v0_ptrs) + v1 = tl.load(v0_ptrs + 64) + + go_ptrs = ( + GO + + bid * L * H * 128 + + hid * 128 + + (offs_m[:, None] * 128 * H + offs_1[None, :]) + ) + + dq0_ptrs = ( + GQ + + nid * B * L * H * 192 + + bid * L * H * 192 + + hid * 192 + + (offs_m[:, None] * H * 192 + offs_1[None, :]) + ) + + dv0 = tl.zeros((N, 64), dtype=tl.float32) + dv1 = tl.zeros((N, 64), dtype=tl.float32) + dk0 = tl.zeros((N, 64), dtype=tl.float32) + dk1 = tl.zeros((N, 64), dtype=tl.float32) + dk2 = tl.zeros((N, 64), dtype=tl.float32) + if CAUSAL: + step = nid * N + else: + step = 0 + for m in range(step, L, M): + lse = tl.load(LSE + bid * H * L + hid * L + m + tl.arange(0, M)) + ds = tl.load(DS + bid * H * L + hid * L + m + tl.arange(0, M)) + + q0 = tl.load(q0_ptrs + m * stride_q) + q1 = tl.load(q0_ptrs + m * stride_q + 64) + q2 = tl.load(q0_ptrs + m * stride_q + 128) + + qk = tl.dot(q0, tl.trans(k0)) + + qk = tl.dot(q1, tl.trans(k1), qk) + + qk = tl.dot(q2, tl.trans(k2), qk) + + go0 = tl.load(go_ptrs + m * H * 128) + go1 = tl.load(go_ptrs + m * H * 128 + 64) + + if CAUSAL: + qk += tl.where( + (m + offs_m)[:, None] >= (nid * N + offs_n)[None, :], 0.0, -1e9 + ) + p = tl.exp(qk * softmax_scale) / lse[:, None] + + dp = tl.dot(go0, tl.trans(v0)) # [M, 128]@[128, N]=[M,N] + dp = tl.dot(go1, tl.trans(v1), dp) + dp = p * (dp - ds[:, None]) * softmax_scale # score + + p = p.to(V.dtype.element_ty) + dv0 = tl.dot(tl.trans(p), go0, dv0) # [N, M]@[M, 128]=[N, 128] + dv1 = tl.dot(tl.trans(p), go1, dv1) # [N, M]@[M, 128]=[N, 128] + + dp = dp.to(V.dtype.element_ty) + dk0 = tl.dot(tl.trans(dp), q0, dk0) # [N, M]@[M, 128]=[N, 128] + dk1 = tl.dot(tl.trans(dp), q1, dk1) # [N, M]@[M, 64]=[N, 64] + dk2 = tl.dot(tl.trans(dp), q2, dk2) # [N, M]@[M, 64]=[N, 64] + dq0 = tl.dot(dp, k0) # [M, N]@[N, 128]=[M, 128] + tl.store(dq0_ptrs + m * H * 192, dq0) + + dq1 = tl.dot(dp, k1) # [M, N]@[N, 64]=[M, 64] + tl.store(dq0_ptrs + m * H * 192 + 64, dq1) + + dq2 = tl.dot(dp, k2) # [M, N]@[N, 64]=[M, 64] + tl.store(dq0_ptrs + m * H * 192 + 128, dq2) + + gv_ptrs = ( + GV + + (bid * L + nid * N) * H * 128 + + hid * 128 + + (offs_n[:, None] * 128 * H + offs_1[None, :]) + ) + + tl.store(gv_ptrs, dv0) + tl.store(gv_ptrs + 64, dv1) + + gk0_ptrs = ( + GK + + (bid * L + nid * N) * H * 192 + + hid * 192 + + (offs_n[:, None] * 192 * H + offs_1[None, :]) + ) + tl.store(gk0_ptrs, dk0) + tl.store(gk0_ptrs + 64, dk1) + tl.store(gk0_ptrs + 128, dk2) + + +@triton.jit +def mla_backward_kernel( + GO, + Q, + K, + V, + GQ, + GK, + GV, + LSE, + ML, + DS, + softmax_scale, + clip_value, + stride_q, + stride_k, + stride_v, + L, + M: tl.constexpr, + N: tl.constexpr, + ATOMIC: tl.constexpr, + CAUSAL: tl.constexpr, + SAFE: tl.constexpr, + CLIP: tl.constexpr, +): + bid = tl.program_id(0) + hid = tl.program_id(1) + nid = tl.program_id(2) + H = tl.num_programs(1).to(tl.int64) + B = tl.num_programs(0) + + offs_m = tl.arange(0, M) + offs_n = tl.arange(0, N) + + offs_0 = tl.arange(0, 128) # nope + offs_1 = tl.arange(0, 64) # pe + + # [B, L, H, 192】 + q0_ptrs = ( + Q + + bid * L * stride_q + + hid * 192 + + (offs_m[:, None] * stride_q + offs_0[None, :]) + ) + q1_ptrs = ( + Q + + bid * L * stride_q + + hid * 192 + + 128 + + (offs_m[:, None] * stride_q + offs_1[None, :]) + ) + + k0_ptrs = ( + K + + (bid * L + nid * N) * stride_k + + hid * 192 + + (offs_n[:, None] * stride_k + offs_0[None, :]) + ) + k1_ptrs = ( + K + + (bid * L + nid * N) * stride_k + + hid * 192 + + 128 + + (offs_n[:, None] * stride_k + offs_1[None, :]) + ) + k0 = tl.load(k0_ptrs) + k1 = tl.load(k1_ptrs) + + v_ptrs = ( + V + + (bid * L + nid * N) * stride_v + + hid * 128 + + (offs_n[:, None] * stride_v + offs_0[None, :]) + ) + v = tl.load(v_ptrs) + + go_ptrs = ( + GO + + bid * L * H * 128 + + hid * 128 + + (offs_m[:, None] * 128 * H + offs_0[None, :]) + ) + + if ATOMIC: + # [B, L, H, 192] + dq0_ptrs = ( + GQ + + bid * L * H * 192 + + hid * 192 + + (offs_m[:, None] * H * 192 + offs_0[None, :]) + ) + dq1_ptrs = ( + GQ + + bid * L * H * 192 + + hid * 192 + + 128 + + (offs_m[:, None] * H * 192 + offs_1[None, :]) + ) + else: + dq0_ptrs = ( + GQ + + nid * B * L * H * 192 + + bid * L * H * 192 + + hid * 192 + + (offs_m[:, None] * H * 192 + offs_0[None, :]) + ) + dq1_ptrs = ( + GQ + + nid * B * L * H * 192 + + bid * L * H * 192 + + hid * 192 + + 128 + + (offs_m[:, None] * H * 192 + offs_1[None, :]) + ) + + dv = tl.zeros((N, 128), dtype=tl.float32) + dk0 = tl.zeros((N, 128), dtype=tl.float32) + dk1 = tl.zeros((N, 64), dtype=tl.float32) + if CAUSAL: + step = nid * N + n_steps = tl.cdiv(L - step, M) + else: + step = 0 + n_steps = tl.cdiv(L, M) + + for i in range(n_steps): + m = step + (n_steps - 1 - i) * M + lse = 1 / tl.load(LSE + bid * H * L + hid * L + m + tl.arange(0, M)) + if SAFE: + max_logits = tl.load(ML + bid * H * L + hid * L + m + tl.arange(0, M)) + ds = tl.load(DS + bid * H * L + hid * L + m + tl.arange(0, M)) + + q0 = tl.load(q0_ptrs + m * stride_q) + q1 = tl.load(q1_ptrs + m * stride_q) + go = tl.load(go_ptrs + m * H * 128) + + if CAUSAL: + qk = tl.where( + (m + offs_m)[:, None] >= (nid * N + offs_n)[None, :], 0.0, -10000.0 + ) + qk = tl.dot(q1, tl.trans(k1), qk) + qk = tl.dot(q0, tl.trans(k0), qk) + else: + qk = tl.dot(q1, tl.trans(k1)) + qk = tl.dot(q0, tl.trans(k0), qk) + + qk *= softmax_scale + if CLIP: + qk = tl.minimum(qk, clip_value) + + if SAFE: + p = tl.exp(qk - max_logits[:, None]) * lse[:, None] + else: + p = tl.exp(qk) * lse[:, None] + + # impl 0 + dp = tl.dot(go, tl.trans(v)) # [M, 128]@[128, N]=[M,N] + dp = p * (dp - ds[:, None]) * softmax_scale # score + # impl 1 + # dp = tl.zeros((1, N), dtype=tl.float32) - ds[:,None] + # dp = tl.dot(go, tl.trans(v), dp) # [M, 128]@[128, N]=[M,N] + # dp = softmax_scale * dp * p # score + + p = p.to(V.dtype.element_ty) + dv = tl.dot(tl.trans(p), go, dv) # [N, M]@[M, 128]=[N, 128] + + dp = dp.to(V.dtype.element_ty) + dq0 = tl.dot(dp, k0) # [M, N]@[N, 128]=[M, 128] + dq1 = tl.dot(dp, k1) # [M, N]@[N, 64]=[M, 64] + if ATOMIC: + tl.atomic_add(dq0_ptrs + m * H * 192, dq0, sem="relaxed") + tl.atomic_add(dq1_ptrs + m * H * 192, dq1, sem="relaxed") + else: + tl.store(dq0_ptrs + m * H * 192, dq0) + tl.store(dq1_ptrs + m * H * 192, dq1) + + dp = tl.trans(dp) + dk0 = tl.dot(dp, q0, dk0) # [N, M]@[M, 128]=[N, 128] + dk1 = tl.dot(dp, q1, dk1) # [N, M]@[M, 64]=[N, 64] + + gv_ptrs = ( + GV + + (bid * L + nid * N) * H * 128 + + hid * 128 + + (offs_n[:, None] * 128 * H + offs_0[None, :]) + ) + + tl.store(gv_ptrs, dv) + + gk0_ptrs = ( + GK + + (bid * L + nid * N) * H * 192 + + hid * 192 + + (offs_n[:, None] * 192 * H + offs_0[None, :]) + ) + tl.store(gk0_ptrs, dk0) + + gk1_ptrs = ( + GK + + (bid * L + nid * N) * H * 192 + + hid * 192 + + 128 + + (offs_n[:, None] * 192 * H + offs_1[None, :]) + ) + tl.store(gk1_ptrs, dk1) + + +# ragged sum +@triton.jit +def mla_rs_kernel( + Q, O, H: tl.constexpr, N: tl.constexpr, BLOCK: tl.constexpr, CAUSAL: tl.constexpr +): + bid = tl.program_id(0) + L = tl.num_programs(1).to(tl.int64) + lid = L - tl.program_id(1) - 1 + kid = tl.program_id(2) + B = tl.num_programs(0) + + offs_n = tl.arange(0, BLOCK) + + # [L//N, B, L, H, 192】 + q_ptrs = Q + bid * L * H * 192 + lid * H * 192 + kid * BLOCK + offs_n + o = tl.zeros((BLOCK,), dtype=tl.float32) + if CAUSAL: + steps = tl.cdiv(lid + 1, N) + else: + steps = L // N + + for i in range(steps): + o += tl.load(q_ptrs + i * B * L * H * 192).to(tl.float32) + + o_ptrs = O + bid * L * H * 192 + lid * H * 192 + kid * BLOCK + offs_n + + tl.store(o_ptrs, o) + + +# should use triton>=3.5.1 for better performance +# hpc: high precision cache +def triton_mla_backward( + go, + o, + q, + k, + v, + lse, + max_logits, + causal=True, + safe=True, + atomic=True, + hpc=False, + clip_value=None, +): + # q: [B, L, H, 192] + # k: [B, L, H, 192] + # v: [B, L, H, 128] + B, L, H, _ = q.shape + assert k.size(1) == L + + device = q.device + dtype = q.dtype + + ds = torch.empty((B, H, L), dtype=torch.float32, device=device) + + softmax_scale = 128 ** (-0.5) + clip = clip_value is not None and clip_value > 0.0 + clip_value = clip_value * softmax_scale if clip else 0.0 + + M = 64 + num_n_block = L // M + num_warps = 4 + num_stages = 2 + grid = (B, H, num_n_block) + mla_ds_kernel[grid](go, o, ds, L, M, num_warps=num_warps, num_stages=num_stages) + + M = 32 + N = 128 + if atomic: + gq = torch.zeros( + (B, L, H, 192), dtype=torch.float32 if hpc else dtype, device=device + ) + else: + gq = torch.empty( + (L // N, B, L, H, 192), dtype=torch.float32 if hpc else dtype, device=device + ) + + gk = torch.empty((B, L, H, 192), dtype=dtype, device=device) + gv = torch.empty((B, L, H, 128), dtype=dtype, device=device) + assert L % M == 0 + assert L % N == 0 + assert N >= M + num_n_block = L // N + num_warps = 8 + num_stages = 5 + grid = (B, H, num_n_block) + mla_backward_kernel[grid]( + go, + q, + k, + v, + gq, + gk, + gv, + lse, + max_logits, + ds, + softmax_scale, + clip_value, + q.stride(1), + k.stride(1), + v.stride(1), + L, + M, + N, + atomic, + causal, + safe, + clip, + num_warps=num_warps, + num_stages=num_stages, + ) + + if atomic: + if hpc: + gq = gq.to(q.dtype) + else: + qo = torch.empty((B, L, H, 192), dtype=dtype, device=device) + BLOCK = max([x for x in [64, 1024, 2048, 4096] if H * 192 % x == 0]) + NB = H * 192 // BLOCK + grid = (B, L, NB) + num_warps = 2 + num_stages = 3 + mla_rs_kernel[grid]( + gq, + qo, + H, + N, + BLOCK, + causal, + num_warps=num_warps, + num_stages=num_stages, + ) + gq = qo + return gq, gk, gv + + +@triton.jit +def varlen_mla_forward_kernel( + Q, + K, + V, + CU, + PCU, + Out, + LSE, + ML, + softmax_scale, + stride_q, + stride_k, + stride_v, + T, + clip_value, + M: tl.constexpr, + N: tl.constexpr, + CAUSAL: tl.constexpr, + PAD: tl.constexpr, + SAFE: tl.constexpr, + CLIP: tl.constexpr, +): + bid = tl.program_id(0) + hid = tl.program_id(1) + mid = tl.num_programs(2) - tl.program_id(2) - 1 + H = tl.num_programs(1) + + if PAD: + c01 = tl.load(CU + bid + tl.arange(0, 2)) + c0, c1 = tl.split(c01) + length = c1 - c0 + pc01 = tl.load(PCU + bid + tl.arange(0, 2)) + pc0, pc1 = tl.split(pc01) + padded_length = pc1 - pc0 + if mid + 1 > tl.cdiv(padded_length, M): + return + else: + c01 = tl.load(CU + bid + tl.arange(0, 2)) + c0, c1 = tl.split(c01) + length = c1 - c0 + if mid + 1 > tl.cdiv(length, M): + return + + offs_m = tl.arange(0, M) + offs_n = tl.arange(0, N) + + offs_0 = tl.arange(0, 128) + offs_1 = tl.arange(0, 64) + + # [T, H, 192】 + q0_ptrs = ( + Q + + (c0 + mid * M) * stride_q + + hid * 192 + + (offs_m[:, None] * stride_q + offs_1[None, :]) + ) + + k0_ptrs = ( + K + c0 * stride_k + hid * 192 + (offs_n[:, None] * stride_k + offs_1[None, :]) + ) + + v_ptrs = ( + V + c0 * stride_v + hid * 128 + (offs_n[:, None] * stride_v + offs_0[None, :]) + ) + + m_mask = (mid * M + offs_m) < length + + q0 = tl.load(q0_ptrs, mask=m_mask[:, None]) + q1 = tl.load(q0_ptrs + 64, mask=m_mask[:, None]) + q2 = tl.load(q0_ptrs + 128, mask=m_mask[:, None]) + + acc_o = tl.zeros((M, 128), dtype=tl.float32) + if SAFE: + max_logits = tl.zeros((M,), dtype=tl.float32) - 10000.0 + lse = tl.zeros((M,), dtype=tl.float32) + else: + lse = tl.zeros((M,), dtype=tl.float32) + 1e-30 + + if CAUSAL: + steps = tl.cdiv(mid * M + M, N) + else: + steps = tl.cdiv(length, N) + + for i in range(0, steps): + n = i * N + n = tl.multiple_of(n, N) + n_mask = (n + offs_n) < length + + k0 = tl.load(k0_ptrs + n * stride_k, mask=n_mask[:, None]) + k1 = tl.load(k0_ptrs + n * stride_k + 64, mask=n_mask[:, None]) + k2 = tl.load(k0_ptrs + n * stride_k + 128, mask=n_mask[:, None]) + + qk = tl.dot(q0, tl.trans(k0)) + qk = tl.dot(q1, tl.trans(k1), qk) + qk = tl.dot(q2, tl.trans(k2), qk) + + if CAUSAL: + qk += tl.where( + ((mid * M + offs_m)[:, None] >= (n + offs_n)[None, :]) + & (n_mask[None, :]), + 0.0, + -1e9, + ) + else: + qk += tl.where(n_mask[None, :], 0.0, -1e9) + + qk *= softmax_scale + + if SAFE: + latest_max_logits = tl.maximum(max_logits, tl.max(qk, 1)) + p = tl.exp(qk - latest_max_logits[:, None]) + rescale = tl.exp(max_logits - latest_max_logits) + lse = lse * rescale + tl.sum(p, 1) + v = tl.load(v_ptrs + n * stride_v, mask=n_mask[:, None]) + p = p.to(V.dtype.element_ty) + acc_o = acc_o * rescale[:, None] + acc_o = tl.dot(p, v, acc_o) + max_logits = latest_max_logits + else: + if CLIP: + p = tl.exp(tl.minimum(qk, clip_value)) + else: + p = tl.exp(qk) + lse += tl.sum(p, 1) + v = tl.load(v_ptrs + n * stride_v, mask=n_mask[:, None]) + p = p.to(V.dtype.element_ty) + acc_o = tl.dot(p, v, acc_o) + + acc_o = acc_o / lse[:, None] + + # [T, H, 128] + out_ptrs = ( + Out + + (c0 + mid * M) * H * 128 + + hid * 128 + + (offs_m[:, None] * H * 128 + offs_0[None, :]) + ) + + tl.store(out_ptrs, acc_o, mask=m_mask[:, None]) + # [H, T] + tl.store(LSE + hid * T + c0 + mid * M + tl.arange(0, M), lse, mask=m_mask) + if SAFE: + tl.store(ML + hid * T + c0 + mid * M + tl.arange(0, M), max_logits, mask=m_mask) + + +def triton_varlen_mla_forward( + q, + k, + v, + cu_seqlens, + max_q_length, + padded_cu_seqlens=None, + causal=True, + safe=True, + clip_value=None, +): + # q: [T, H, 192] + # k: [T, H, 192] + # v: [T, H, 128] + T, H, _ = q.shape + B = cu_seqlens.size(0) - 1 + assert k.size(0) == T + M = 256 + N = 64 + assert M >= N + + o = torch.empty((T, H, 128), dtype=q.dtype, device=q.device) + lse = torch.empty((H, T), dtype=torch.float32, device=q.device) + max_logits = torch.empty((H, T), dtype=torch.float32, device=q.device) + softmax_scale = 128 ** (-0.5) + clip = clip_value is not None and clip_value > 0.0 + clip_value = clip_value * softmax_scale if clip else 0.0 + if clip and clip_value + math.log(max_q_length) < 88.7: + safe = False + PAD = padded_cu_seqlens is not None + + num_m_block = triton.cdiv(max_q_length, M) + num_stages = 2 + num_warps = 8 + + grid = (B, H, num_m_block) + varlen_mla_forward_kernel[grid]( + q, + k, + v, + cu_seqlens, + padded_cu_seqlens, + o, + lse, + max_logits, + softmax_scale, + q.stride(0), + k.stride(0), + v.stride(0), + T, + clip_value, + M, + N, + causal, + PAD, + safe, + clip, + num_warps=num_warps, + num_stages=num_stages, + ) + return o, lse, max_logits + + +@triton.jit +def varlen_mla_backward_kernel( + GO, + Q, + K, + V, + CU, + PCU, + GQ, + GK, + GV, + LSE, + ML, + DS, + softmax_scale, + clip_value, + stride_q, + stride_k, + stride_v, + T, + M: tl.constexpr, + N: tl.constexpr, + ATOMIC: tl.constexpr, + CAUSAL: tl.constexpr, + PAD: tl.constexpr, + SAFE: tl.constexpr, + CLIP: tl.constexpr, +): + bid = tl.program_id(0) + hid = tl.program_id(1) + nid = tl.program_id(2) + H = tl.num_programs(1).to(tl.int64) + + if PAD: + c1 = tl.load(CU + bid + 1) + c0 = tl.load(CU + bid) + length = c1 - c0 + pc0 = tl.load(PCU + bid) + pc1 = tl.load(PCU + bid + 1) + padded_length = pc1 - pc0 + if nid + 1 > tl.cdiv(padded_length, N): + return + else: + c1 = tl.load(CU + bid + 1) + c0 = tl.load(CU + bid) + length = c1 - c0 + if nid + 1 > tl.cdiv(length, N): + return + + offs_m = tl.arange(0, M) + offs_n = tl.arange(0, N) + + offs_0 = tl.arange(0, 128) # nope + offs_1 = tl.arange(0, 64) # pe + + n_mask = (nid * N + offs_n) < length + + # [B, L, H, 192】 + q0_ptrs = ( + Q + c0 * stride_q + hid * 192 + (offs_m[:, None] * stride_q + offs_0[None, :]) + ) + q1_ptrs = ( + Q + + c0 * stride_q + + hid * 192 + + 128 + + (offs_m[:, None] * stride_q + offs_1[None, :]) + ) + + k0_ptrs = ( + K + + (c0 + nid * N) * stride_k + + hid * 192 + + (offs_n[:, None] * stride_k + offs_0[None, :]) + ) + k1_ptrs = ( + K + + (c0 + nid * N) * stride_k + + hid * 192 + + 128 + + (offs_n[:, None] * stride_k + offs_1[None, :]) + ) + k0 = tl.load(k0_ptrs, mask=n_mask[:, None]) + k1 = tl.load(k1_ptrs, mask=n_mask[:, None]) + + v_ptrs = ( + V + + (c0 + nid * N) * stride_v + + hid * 128 + + (offs_n[:, None] * stride_v + offs_0[None, :]) + ) + v = tl.load(v_ptrs, mask=n_mask[:, None]) + + go_ptrs = ( + GO + c0 * H * 128 + hid * 128 + (offs_m[:, None] * 128 * H + offs_0[None, :]) + ) + + if ATOMIC: + # [B, L, H, 192] + dq0_ptrs = ( + GQ + + c0 * H * 192 + + hid * 192 + + (offs_m[:, None] * H * 192 + offs_0[None, :]) + ) + dq1_ptrs = ( + GQ + + c0 * H * 192 + + hid * 192 + + 128 + + (offs_m[:, None] * H * 192 + offs_1[None, :]) + ) + + else: + dq0_ptrs = ( + GQ + + nid * T * H * 192 + + c0 * H * 192 + + hid * 192 + + (offs_m[:, None] * H * 192 + offs_0[None, :]) + ) + dq1_ptrs = ( + GQ + + nid * T * H * 192 + + c0 * H * 192 + + hid * 192 + + 128 + + (offs_m[:, None] * H * 192 + offs_1[None, :]) + ) + + dv = tl.zeros((N, 128), dtype=tl.float32) + dk0 = tl.zeros((N, 128), dtype=tl.float32) + dk1 = tl.zeros((N, 64), dtype=tl.float32) + if CAUSAL: + step = nid * N + n_steps = tl.cdiv(length - step, M) + else: + step = 0 + n_steps = tl.cdiv(length, M) + + for i in range(n_steps): + m = step + (n_steps - 1 - i) * M + m_mask = (m + offs_m) < length + + # [H, T] + lse = 1 / tl.load( + LSE + hid * T + c0 + m + tl.arange(0, M), mask=m_mask, other=1e30 + ) + if SAFE: + max_logits = tl.load(ML + hid * T + c0 + m + tl.arange(0, M), mask=m_mask) + ds = tl.load(DS + hid * T + c0 + m + tl.arange(0, M), mask=m_mask) + + q0 = tl.load(q0_ptrs + m * stride_q, mask=m_mask[:, None]) + q1 = tl.load(q1_ptrs + m * stride_q, mask=m_mask[:, None]) + go = tl.load(go_ptrs + m * H * 128, mask=m_mask[:, None]) + + if CAUSAL: + qk = tl.where( + ((m + offs_m)[:, None] >= (nid * N + offs_n)[None, :]) + & (n_mask[None, :]), + 0.0, + -1e9, + ) + qk = tl.dot(q1, tl.trans(k1), qk) + qk = tl.dot(q0, tl.trans(k0), qk) + else: + qk = tl.dot(q1, tl.trans(k1)) + qk = tl.dot(q0, tl.trans(k0), qk) + qk += tl.where((n_mask[None, :]), 0.0, -1e9) + + qk *= softmax_scale + if CLIP: + qk = tl.minimum(qk, clip_value) + + if SAFE: + p = tl.exp(qk - max_logits[:, None]) * lse[:, None] + else: + p = tl.exp(qk) * lse[:, None] + + # impl 0 + dp = tl.dot(go, tl.trans(v)) # [M, 128]@[128, N]=[M,N] + dp = p * (dp - ds[:, None]) * softmax_scale # score + # impl 1 + # dp = tl.zeros((1, N), dtype=tl.float32) - ds[:,None] + # dp = tl.dot(go, tl.trans(v), dp) # [M, 128]@[128, N]=[M,N] + # dp = softmax_scale * dp * p # score + + p = p.to(V.dtype.element_ty) + dv = tl.dot(tl.trans(p), go, dv) # [N, M]@[M, 128]=[N, 128] + + dp = dp.to(V.dtype.element_ty) + dq0 = tl.dot(dp, k0) # [M, N]@[N, 128]=[M, 128] + dq1 = tl.dot(dp, k1) # [M, N]@[N, 64]=[M, 64] + if ATOMIC: + tl.atomic_add( + dq0_ptrs + m * H * 192, dq0, mask=m_mask[:, None], sem="relaxed" + ) + tl.atomic_add( + dq1_ptrs + m * H * 192, dq1, mask=m_mask[:, None], sem="relaxed" + ) + else: + tl.store(dq0_ptrs + m * H * 192, dq0, mask=m_mask[:, None]) + tl.store(dq1_ptrs + m * H * 192, dq1, mask=m_mask[:, None]) + + dp = tl.trans(dp) + dk0 = tl.dot(dp, q0, dk0) # [N, M]@[M, 128]=[N, 128] + dk1 = tl.dot(dp, q1, dk1) # [N, M]@[M, 64]=[N, 64] + + gv_ptrs = ( + GV + + (c0 + nid * N) * H * 128 + + hid * 128 + + (offs_n[:, None] * 128 * H + offs_0[None, :]) + ) + tl.store(gv_ptrs, dv, mask=n_mask[:, None]) + + gk0_ptrs = ( + GK + + (c0 + nid * N) * H * 192 + + hid * 192 + + (offs_n[:, None] * 192 * H + offs_0[None, :]) + ) + tl.store(gk0_ptrs, dk0, mask=n_mask[:, None]) + + gk1_ptrs = ( + GK + + (c0 + nid * N) * H * 192 + + hid * 192 + + 128 + + (offs_n[:, None] * 192 * H + offs_1[None, :]) + ) + tl.store(gk1_ptrs, dk1, mask=n_mask[:, None]) + + +# ragged sum +@triton.jit +def varlen_mla_rs_kernel( + Q, + O, + CU, + B, + PB: tl.constexpr, + H: tl.constexpr, + N: tl.constexpr, + BLOCK: tl.constexpr, + CAUSAL: tl.constexpr, +): + tid = tl.program_id(0) + T = tl.num_programs(0).to(tl.int64) + kid = tl.program_id(1) + + # cu = tl.load(CU + tl.arange(0, PB), mask=tl.arange(0, PB) <= B) + # c0 = tl.max(tl.where(cu > tid, 0, cu), 0) + # c1 = tl.min(tl.where(cu <= c0, 2 ** 24, cu), 0) + + c0 = 0 + c1 = 1048576 + for i in range(tl.cdiv(B + 1, PB)): + cus = tl.load( + CU + i * PB + tl.arange(0, PB), mask=i * PB + tl.arange(0, PB) <= B + ) + c0 = tl.maximum(tl.max(tl.where(cus > tid, 0, cus), 0), c0) + for i in range(tl.cdiv(B, PB)): + cus = tl.load( + CU + i * PB + tl.arange(0, PB), mask=i * PB + tl.arange(0, PB) <= B + ) + c1 = tl.minimum(tl.min(tl.where(cus <= c0, 2**24, cus), 0), c1) + + length = c1 - c0 + pid = tid - c0 + + offs_n = tl.arange(0, BLOCK) + + # [max_q_length//N, T, H, 192] + q_ptrs = Q + tid * H * 192 + kid * BLOCK + offs_n + o = tl.zeros((BLOCK,), dtype=tl.float32) + if CAUSAL: + steps = tl.cdiv(pid + 1, N) + else: + steps = tl.cdiv(length, N) + + for i in range(steps): + o += tl.load(q_ptrs + i * T * H * 192).to(tl.float32) + + # [T, H, 192] + o_ptrs = O + tid * H * 192 + kid * BLOCK + offs_n + + tl.store(o_ptrs, o) + + +# should use triton>=3.5.1 for better performance +# hpc: high precision cache +def triton_varlen_mla_backward( + go, + o, + q, + k, + v, + lse, + max_logits, + cu_seqlens, + max_q_length, + padded_cu_seqlens=None, + causal=True, + safe=True, + hpc=False, + atomic=True, + clip_value=None, +): + # q: [T, H, 192] + # k: [T, H, 192] + # v: [T, H, 128] + T, H, _ = q.shape + assert k.size(0) == T + B = cu_seqlens.size(0) - 1 + PADDED = padded_cu_seqlens is not None + + device = q.device + dtype = q.dtype + + ds = torch.empty((H, T), dtype=torch.float32, device=device) + + softmax_scale = 128 ** (-0.5) + clip = clip_value is not None and clip_value > 0.0 + clip_value = clip_value * softmax_scale if clip else 0.0 + + M = 64 + num_warps = 4 + num_stages = 2 + grid = (1, H, triton.cdiv(T, M)) + mla_ds_kernel[grid](go, o, ds, T, M, num_warps=num_warps, num_stages=num_stages) + + M = 32 + N = 128 + num_n_block = triton.cdiv(T, N) + if atomic: + gq = torch.zeros( + (T, H, 192), dtype=torch.float32 if hpc else dtype, device=device + ) + else: + gq = torch.empty( + (num_n_block, T, H, 192), + dtype=torch.float32 if hpc else dtype, + device=q.device, + ) + + gk = torch.empty((T, H, 192), dtype=dtype, device=device) + gv = torch.empty((T, H, 128), dtype=dtype, device=device) + assert N >= M + num_warps = 8 + num_stages = 5 + grid = (B, H, num_n_block) + varlen_mla_backward_kernel[grid]( + go, + q, + k, + v, + cu_seqlens, + padded_cu_seqlens, + gq, + gk, + gv, + lse, + max_logits, + ds, + softmax_scale, + clip_value, + q.stride(0), + k.stride(0), + v.stride(0), + T, + M, + N, + atomic, + causal, + PADDED, + safe, + clip, + num_warps=num_warps, + num_stages=num_stages, + ) + + if atomic: + if hpc: + gq = gq.to(q.dtype) + else: + qo = torch.empty((T, H, 192), dtype=dtype, device=device) + BLOCK = max([x for x in [64, 1024, 2048, 4096] if H * 192 % x == 0]) + NB = H * 192 // BLOCK + grid = (T, NB) + num_warps = 4 + num_stages = 3 + PB = 128 + varlen_mla_rs_kernel[grid]( + gq, + qo, + padded_cu_seqlens if PADDED else cu_seqlens, + B, + PB, + H, + N, + BLOCK, + causal, + num_warps=num_warps, + num_stages=num_stages, + ) + gq = qo + return gq, gk, gv + + +@triton.jit +def deprecated_fp8_mla_forward_kernel( + Q, + K, + V, + QS, + KS, + VS, + Out, + LSE, + ML, + softmax_scale, + stride_q, + stride_k, + stride_v, + L, + M: tl.constexpr, + N: tl.constexpr, + CAUSAL: tl.constexpr, +): + bid = tl.program_id(0) + hid = tl.program_id(1) + mid = tl.num_programs(2) - tl.program_id(2) - 1 + H = tl.num_programs(1) + + offs_m = tl.arange(0, M) + offs_n = tl.arange(0, N) + + offs_0 = tl.arange(0, 128) # nope + offs_1 = tl.arange(0, 64) # pe + + # [B, L, H, 192】 + q0_ptrs = ( + Q + + (bid * L + mid * M) * stride_q + + hid * 192 + + (offs_m[:, None] * stride_q + offs_0[None, :]) + ) + q1_ptrs = ( + Q + + (bid * L + mid * M) * stride_q + + hid * 192 + + 128 + + (offs_m[:, None] * stride_q + offs_1[None, :]) + ) + + # [B, H, L] + qs_ptrs = QS + bid * H * L + hid * L + mid * M + offs_m + + k0_ptrs = ( + K + + bid * L * stride_k + + hid * 192 + + (offs_n[:, None] * stride_k + offs_0[None, :]) + ) + + k1_ptrs = ( + K + + bid * L * stride_k + + hid * 192 + + 128 + + (offs_n[:, None] * stride_k + offs_1[None, :]) + ) + + # [B, H, L] + ks_ptrs = KS + bid * H * L + hid * L + offs_n + + v_ptrs = ( + V + + bid * L * stride_v + + hid * 128 + + (offs_n[:, None] * stride_v + offs_0[None, :]) + ) + + q0 = tl.load(q0_ptrs) + q1 = tl.load(q1_ptrs) + qs = tl.load(qs_ptrs) + + lse = tl.zeros((M,), dtype=tl.float32) + acc_o = tl.zeros((M, 128), dtype=tl.float32) + if CAUSAL: + steps = tl.cdiv(mid * M + M, N) + else: + steps = L // N + for i in range(0, steps): + n = i * N + n = tl.multiple_of(n, N) + + k1 = tl.load(k1_ptrs + n * stride_k) + ks = tl.load(ks_ptrs + n) + + qk = tl.dot(q1, tl.trans(k1)) + + k0 = tl.load(k0_ptrs + n * stride_k) + + qk = tl.dot(q0, tl.trans(k0), qk) + + qk += tl.where((mid * M + offs_m)[:, None] >= (n + offs_n)[None, :], 0.0, -1e9) + qk = qk * qs[:, None] * ks[None, :] + + p = tl.exp(qk * softmax_scale) + lse += tl.sum(p, 1) + + v = tl.load(v_ptrs + n * stride_v) + p = p.to(V.dtype.element_ty) + acc_o = tl.dot(p, v, acc_o) + + acc_o = acc_o / lse[:, None] + + # [B, L, H, 128] + out_ptrs = ( + Out + + (bid * L + mid * M) * H * 128 + + hid * 128 + + (offs_m[:, None] * 128 * H + offs_0[None, :]) + ) + + tl.store(out_ptrs, acc_o) + tl.store(LSE + bid * H * L + hid * L + mid * M + tl.arange(0, M), lse) + + +@triton.jit +def padding_fp8_mla_forward_kernel( + Q, + K, + V, + QS, + KS, + VS, + Out, + LSE, + ML, + softmax_scale, + stride_q, + stride_k, + stride_v, + L, + M: tl.constexpr, + N: tl.constexpr, + CAUSAL: tl.constexpr, +): + bid = tl.program_id(0) + hid = tl.program_id(1) + mid = tl.num_programs(2) - tl.program_id(2) - 1 + H = tl.num_programs(1) + + offs_m = tl.arange(0, M) + offs_n = tl.arange(0, N) + + offs_0 = tl.arange(0, 256) + offs_1 = tl.arange(0, 128) + + # [B, L, H, 192】 + q0_ptrs = ( + Q + + (bid * L + mid * M) * stride_q + + hid * 192 + + (offs_m[:, None] * stride_q + offs_0[None, :]) + ) + + # [B, H, L] + qs_ptrs = QS + bid * H * L + hid * L + mid * M + offs_m + + k0_ptrs = ( + K + + bid * L * stride_k + + hid * 192 + + (offs_n[:, None] * stride_k + offs_0[None, :]) + ) + + # [B, H, L] + ks_ptrs = KS + bid * H * L + hid * L + offs_n + + v_ptrs = ( + V + + bid * L * stride_v + + hid * 128 + + (offs_n[:, None] * stride_v + offs_1[None, :]) + ) + mask = offs_0 < 192 + q0 = tl.load(q0_ptrs, mask=mask[None, :]) + qs = tl.load(qs_ptrs) + + lse = tl.zeros((M,), dtype=tl.float32) + acc_o = tl.zeros((M, 128), dtype=tl.float32) + if CAUSAL: + steps = tl.cdiv(mid * M + M, N) + else: + steps = L // N + for i in range(0, steps): + n = i * N + n = tl.multiple_of(n, N) + + k0 = tl.load(k0_ptrs + n * stride_k, mask=mask[None, :]) + ks = tl.load(ks_ptrs + n) + + qk = tl.dot(q0, tl.trans(k0)) + + qk += tl.where((mid * M + offs_m)[:, None] >= (n + offs_n)[None, :], 0.0, -1e9) + qk = qk * qs[:, None] * ks[None, :] + + p = tl.exp(qk * softmax_scale) + lse += tl.sum(p, 1) + + v = tl.load(v_ptrs + n * stride_v) + p = p.to(V.dtype.element_ty) + + acc_o = tl.dot(p, v, acc_o) + + acc_o = acc_o / lse[:, None] + + # [B, L, H, 128] + out_ptrs = ( + Out + + (bid * L + mid * M) * H * 128 + + hid * 128 + + (offs_m[:, None] * 128 * H + offs_1[None, :]) + ) + + tl.store(out_ptrs, acc_o) + tl.store(LSE + bid * H * L + hid * L + mid * M + tl.arange(0, M), lse) + + +@triton.jit +def fp8_mla_forward_kernel( + Q, + K, + V, + QS, + KS, + VS, + Out, + LSE, + ML, + softmax_scale, + stride_q, + stride_k, + stride_v, + L, + M: tl.constexpr, + N: tl.constexpr, + CAUSAL: tl.constexpr, +): + bid = tl.program_id(0) + hid = tl.program_id(1) + mid = tl.num_programs(2) - tl.program_id(2) - 1 + H = tl.num_programs(1) + + offs_m = tl.arange(0, M) + offs_n = tl.arange(0, N) + + offs_0 = tl.arange(0, 128) + offs_1 = tl.arange(0, 64) + + # [B, L, H, 192] + q0_ptrs = ( + Q + + (bid * L + mid * M) * stride_q + + hid * 192 + + (offs_m[:, None] * stride_q + offs_1[None, :]) + ) + + # [B, H, L] + qs_ptrs = QS + bid * H * L + hid * L + mid * M + offs_m + + k0_ptrs = ( + K + + bid * L * stride_k + + hid * 192 + + (offs_n[:, None] * stride_k + offs_1[None, :]) + ) + + # [B, H, L] + ks_ptrs = KS + bid * H * L + hid * L + offs_n + + v_ptrs = ( + V + + bid * L * stride_v + + hid * 128 + + (offs_n[:, None] * stride_v + offs_0[None, :]) + ) + + if VS is not None: + # [B, H, L] + vs_ptrs = VS + bid * H * L + hid * L + offs_n + + q0 = tl.load(q0_ptrs) + q1 = tl.load(q0_ptrs + 64) + q2 = tl.load(q0_ptrs + 128) + + qs = tl.load(qs_ptrs) * softmax_scale + + lse = tl.zeros((M,), dtype=tl.float32) + acc_o = tl.zeros((M, 128), dtype=tl.float32) + if CAUSAL: + steps = tl.cdiv(mid * M + M, N) + else: + steps = L // N + + for i in range(0, steps): + n = i * N + n = tl.multiple_of(n, N) + + k0 = tl.load(k0_ptrs + n * stride_k) + k1 = tl.load(k0_ptrs + n * stride_k + 64) + k2 = tl.load(k0_ptrs + n * stride_k + 128) + ks = tl.load(ks_ptrs + n) + + if CAUSAL: + qk = tl.where( + (mid * M + offs_m)[:, None] >= (n + offs_n)[None, :], 0.0, -1e9 + ) + qk = tl.dot(q0, tl.trans(k0), qk) + qk = tl.dot(q1, tl.trans(k1), qk) + qk = tl.dot(q2, tl.trans(k2), qk) + else: + qk = tl.dot(q0, tl.trans(k0)) + qk = tl.dot(q1, tl.trans(k1), qk) + qk = tl.dot(q2, tl.trans(k2), qk) + qk = qk * qs[:, None] * ks[None, :] + + p = tl.exp(qk) + lse += tl.sum(p, 1) + + if VS is None: + p = p.to(V.dtype.element_ty) + v = tl.load(v_ptrs + n * stride_v) + acc_o = tl.dot(p, v, acc_o) + else: + vs = tl.load(vs_ptrs) + p = p * vs + pm = tl.max(p, 1) / 448 + p = (p / pm[:, None]).to(V.dtype.element_ty) + v = tl.load(v_ptrs + n * stride_v) + acc_o = acc_o + tl.dot(p, v) * pm[:, None] + + tl.store(LSE + bid * H * L + hid * L + mid * M + tl.arange(0, M), lse) + acc_o = acc_o / lse[:, None] + # [B, L, H, 128] + out_ptrs = ( + Out + + (bid * L + mid * M) * H * 128 + + hid * 128 + + (offs_m[:, None] * 128 * H + offs_0[None, :]) + ) + tl.store(out_ptrs, acc_o) + + +def triton_fp8_mla_forward( + q, k, v, qs, ks, vs=None, causal=True, out_dtype=torch.bfloat16 +): + # q: [B, L, H, 192] + # k: [B, L, H, 192] + # v: [B, L, H, 128] + B, L, H, _ = q.shape + assert k.size(1) == L + M = 128 + N = 128 + assert L % M == 0 + assert L % N == 0 + assert M >= N + + o = torch.empty((B, L, H, 128), dtype=out_dtype, device=q.device) + lse = torch.empty((B, H, L), dtype=torch.float32, device=q.device) + max_logits = torch.empty((B, H, L), dtype=torch.float32, device=q.device) + softmax_scale = 128 ** (-0.5) + + num_m_block = L // M + num_warps = 8 + num_stages = 3 + + grid = (B, H, num_m_block) + fp8_mla_forward_kernel[grid]( + q, + k, + v, + qs, + ks, + vs, + o, + lse, + max_logits, + softmax_scale, + q.stride(1), + k.stride(1), + v.stride(1), + L, + M, + N, + causal, + num_warps=num_warps, + num_stages=num_stages, + ) + return o, lse, max_logits diff --git a/linghe/experimental/__init__.py b/linghe/experimental/__init__.py new file mode 100644 index 0000000..37910a2 --- /dev/null +++ b/linghe/experimental/__init__.py @@ -0,0 +1,3 @@ +""" +kernels should be run with torch above 2.9.0 +""" diff --git a/linghe/experimental/demb.py b/linghe/experimental/demb.py new file mode 100644 index 0000000..9d97108 --- /dev/null +++ b/linghe/experimental/demb.py @@ -0,0 +1,560 @@ +# -*- coding: utf-8 -*- +""" +Copyright (c) Ant Financial Service Group and its affiliates. +""" + +import torch +import torch.distributed as dist +import triton +import triton.language as tl + +from linghe.experimental.symm_mem_barrier import symm_mem_sync +from linghe.utils.emb import triton_scan_and_count + +""" +distributed embedding with vocab parallel +""" + + +@triton.jit +def tp_embedding_lookup_forward_kernel( + input_ids_ptr, + weights_ptr, + outputs_ptr, + buffer_ptrs, + signal_ptrs, + V, + d, + D: tl.constexpr, + SIZE: tl.constexpr, + RANK: tl.constexpr, +): + pid = tl.program_id(0) + buffer_ptrs = buffer_ptrs.to(tl.pointer_type(tl.uint64)) + + input_id = tl.load(input_ids_ptr + pid) + mask = tl.arange(0, D) < d + + if (input_id >= RANK * V) & (input_id < (RANK + 1) * V): + + buffer_ptr = tl.load(buffer_ptrs + RANK).to(tl.pointer_type(tl.bfloat16)) + buffer_ptr = tl.multiple_of(buffer_ptr, 16) + w = tl.load(weights_ptr + input_id % V * D + tl.arange(0, D), mask=mask) + tl.store(outputs_ptr + pid * D + tl.arange(0, D), w, mask=mask) + + tl.store(buffer_ptr + pid * D + tl.arange(0, D), w, mask=mask) + symm_mem_sync( + signal_ptrs, + None, + RANK, + SIZE, + hasPreviousMemAccess=True, + hasSubsequentMemAccess=True, + ) + else: + symm_mem_sync( + signal_ptrs, + None, + RANK, + SIZE, + hasPreviousMemAccess=True, + hasSubsequentMemAccess=True, + ) + buffer_ptr = tl.load(buffer_ptrs + input_id // V).to( + tl.pointer_type(tl.bfloat16) + ) + buffer_ptr = tl.multiple_of(buffer_ptr, 16) + w = tl.load(buffer_ptr + pid * D + tl.arange(0, D), mask=mask) + tl.store(outputs_ptr + pid * D + tl.arange(0, D), w, mask=mask) + + +""" +input_ids is the same cross ranks +""" + + +def triton_tp_embedding_lookup_forward(input_ids, weights, hdl, group): + group_size = hdl.world_size + group_rank = hdl.rank + shape = input_ids.shape + V, d = weights.shape + D = triton.next_power_of_2(d) + + device = weights.device + dtype = weights.dtype + assert len(shape) in (1, 2) and dtype == torch.bfloat16 + + if len(shape) == 2: + M = shape[0] * shape[1] + outputs = torch.empty((shape[0], shape[1], d), device=device, dtype=dtype) + else: + M = shape[0] + outputs = torch.empty((M, d), device=device, dtype=dtype) + + num_warps = 4 + num_stages = 3 + tp_embedding_lookup_forward_kernel[(M,)]( + input_ids, + weights, + outputs, + hdl.buffer_ptrs_dev, + hdl.signal_pad_ptrs_dev, + V, + d, + D, + group_size, + group_rank, + num_warps=num_warps, + num_stages=num_stages, + ) + return outputs + + +@triton.jit +def tp_embedding_lookup_backward_kernel( + grad_output_ptr, + sorted_ids_ptr, + sorted_indices_ptr, + accum_counts_ptr, + g_ptr, + buffer_ptrs, + signal_ptrs, + stride_0, + stride_1, + dim, + V, + B, + L, + DIM: tl.constexpr, + T: tl.constexpr, + SIZE: tl.constexpr, + RANK: tl.constexpr, +): + pid = tl.program_id(axis=0).to(tl.int64) + buffer_ptrs = buffer_ptrs.to(tl.pointer_type(tl.uint64)) + + c01 = tl.load(accum_counts_ptr + pid + tl.arange(0, 2)) + c0, c1 = tl.split(c01) + if c0 == c1: + return + + count = c1 - c0 + input_id = tl.load(sorted_ids_ptr + c0).to(tl.int64) + mask = tl.arange(0, DIM) < dim + + if T == 0: + grad_ptr = g_ptr.to(tl.pointer_type(tl.float32)) + else: + grad_ptr = g_ptr.to(tl.pointer_type(tl.bfloat16)) + + outputs = tl.zeros((DIM,), dtype=tl.float32) + + if (input_id >= RANK * V) & (input_id < (RANK + 1) * V): + for i in range(count): + pos = tl.load(sorted_indices_ptr + c0 + i) + bid = pos // L + lid = pos % L + g = tl.load( + grad_output_ptr + bid * stride_0 + lid * stride_1 + tl.arange(0, DIM), + mask=mask, + ).to(tl.float32) + outputs += g + symm_mem_sync( + signal_ptrs, + None, + RANK, + SIZE, + hasPreviousMemAccess=True, + hasSubsequentMemAccess=True, + ) + + for j in range(SIZE): + if j != RANK: + buffer_ptr = tl.load(buffer_ptrs + j).to(tl.pointer_type(tl.bfloat16)) + buffer_ptr = tl.multiple_of(buffer_ptr, 16) + g = tl.load(buffer_ptr + pid * DIM + tl.arange(0, DIM), mask=mask) + outputs += g + tl.store(grad_ptr + input_id % V * DIM + tl.arange(0, DIM), outputs, mask=mask) + + else: + + for i in range(count): + pos = tl.load(sorted_indices_ptr + c0 + i) + bid = pos // L + lid = pos % L + g = tl.load( + grad_output_ptr + bid * stride_0 + lid * stride_1 + tl.arange(0, DIM), + mask=mask, + ).to(tl.float32) + outputs += g + + buffer_ptr = tl.load(buffer_ptrs + RANK).to(tl.pointer_type(tl.bfloat16)) + buffer_ptr = tl.multiple_of(buffer_ptr, 16) + tl.store(buffer_ptr + pid * DIM + tl.arange(0, DIM), outputs, mask=mask) + symm_mem_sync( + signal_ptrs, + None, + RANK, + SIZE, + hasPreviousMemAccess=True, + hasSubsequentMemAccess=True, + ) + + +def triton_tp_embedding_lookup_backward( + grad_output, x, g_ptr, vocab_size, hdl, group, dtype=torch.bfloat16 +): + """ + inplace update embedding weight gradient + Args: + y: gradient of output + x: input ids Tensor + g_ptr: data_ptr of embedding weight gradient + Returns: + None + """ + assert dtype in (torch.bfloat16, torch.float32) + T = 0 if dtype == torch.float32 else 1 + shape = x.shape + assert len(shape) == 2 + + group_size = hdl.world_size + group_rank = hdl.rank + B, L, dim = grad_output.shape + stride_0 = grad_output.stride(0) + stride_1 = grad_output.stride(1) + + sorted_ids, sorted_indices = torch.sort(x.view(-1), stable=False) + accum_counts = triton_scan_and_count(sorted_ids) + DIM = triton.next_power_of_2(dim) + num_stages = 3 + num_warps = 2 + + grid = (B * L,) + tp_embedding_lookup_backward_kernel[grid]( + grad_output, + sorted_ids, + sorted_indices, + accum_counts, + g_ptr, + hdl.buffer_ptrs_dev, + hdl.signal_pad_ptrs_dev, + stride_0, + stride_1, + dim, + vocab_size, + B, + L, + DIM, + T, + group_size, + group_rank, + num_stages=num_stages, + num_warps=num_warps, + ) + + +@triton.jit +def sp_embedding_lookup_forward_kernel( + input_ids_ptr, + weights_ptr, + outputs_ptr, + buffer_ptrs, + signal_ptrs, + M, + V, + d, + D: tl.constexpr, + SIZE: tl.constexpr, + RANK: tl.constexpr, +): + pid = tl.program_id(0) + buffer_ptrs = buffer_ptrs.to(tl.pointer_type(tl.uint64)) + mask = tl.arange(0, D) < d + + for chunk in range(SIZE): + input_id = tl.load(input_ids_ptr + chunk * M + pid) + if chunk == RANK: + if (input_id >= RANK * V) & (input_id < (RANK + 1) * V): + w = tl.load(weights_ptr + input_id % V * D + tl.arange(0, D), mask=mask) + tl.store(outputs_ptr + pid * D + tl.arange(0, D), w, mask=mask) + else: + symm_mem_sync( + signal_ptrs, + None, + RANK, + SIZE, + hasPreviousMemAccess=True, + hasSubsequentMemAccess=True, + ) + buffer_ptr = tl.load(buffer_ptrs + input_id // V).to( + tl.pointer_type(tl.bfloat16) + ) + buffer_ptr = tl.multiple_of(buffer_ptr, 16) + w = tl.load(buffer_ptr + pid * D + tl.arange(0, D), mask=mask) + tl.store(outputs_ptr + pid % M * D + tl.arange(0, D), w, mask=mask) + else: + if (input_id >= RANK * V) & (input_id < (RANK + 1) * V): + w = tl.load(weights_ptr + input_id % V * D + tl.arange(0, D), mask=mask) + buffer_ptr = tl.load(buffer_ptrs + RANK).to( + tl.pointer_type(tl.bfloat16) + ) + buffer_ptr = tl.multiple_of(buffer_ptr, 16) + tl.store(buffer_ptr + pid * D + tl.arange(0, D), w, mask=mask) + symm_mem_sync( + signal_ptrs, + None, + RANK, + SIZE, + hasPreviousMemAccess=True, + hasSubsequentMemAccess=True, + ) + + +""" +input_ids is different cross ranks +""" + + +def triton_sp_embedding_lookup_forward(input_ids, weights, hdl, group): + group_size = hdl.world_size + group_rank = hdl.rank + shape = input_ids.shape + V, d = weights.shape + D = triton.next_power_of_2(d) + + device = weights.device + dtype = weights.dtype + assert len(shape) in (1, 2) and dtype == torch.bfloat16 + + if len(shape) == 2: + M = shape[0] * shape[1] + outputs = torch.empty((shape[0], shape[1], d), device=device, dtype=dtype) + gathered_input_ids = torch.empty( + (group_size, shape[0], shape[1]), dtype=torch.long, device=device + ) + else: + M = shape[0] + outputs = torch.empty((M, d), device=device, dtype=dtype) + gathered_input_ids = torch.empty( + (group_size, M), dtype=torch.long, device=device + ) + dist.all_gather_into_tensor(gathered_input_ids, input_ids, group=group) + + num_warps = 4 + num_stages = 3 + + sp_embedding_lookup_forward_kernel[(M,)]( + gathered_input_ids, + weights, + outputs, + hdl.buffer_ptrs_dev, + hdl.signal_pad_ptrs_dev, + M, + V, + d, + D, + group_size, + group_rank, + num_warps=num_warps, + num_stages=num_stages, + ) + return outputs + + +@triton.jit +def sp_embedding_lookup_backward_kernel( + grad_output_ptr, + sorted_ids_ptr, + sorted_indices_ptr, + accum_counts_ptr, + g_ptr, + buffer_ptrs, + signal_ptrs, + stride_0, + stride_1, + dim, + M, + V, + B, + L, + DIM: tl.constexpr, + T: tl.constexpr, + SIZE: tl.constexpr, + RANK: tl.constexpr, +): + pid = tl.program_id(axis=0).to(tl.int64) + buffer_ptrs = buffer_ptrs.to(tl.pointer_type(tl.uint64)) + mask = tl.arange(0, DIM) < dim + + for chunk in range(SIZE): + c01 = tl.load(accum_counts_ptr + chunk * (L + 1) + pid + tl.arange(0, 2)) + c0, c1 = tl.split(c01) + if c0 != c1: + + count = c1 - c0 + input_id = tl.load(sorted_ids_ptr + chunk * M + c0).to(tl.int64) + + if T == 0: + grad_ptr = g_ptr.to(tl.pointer_type(tl.float32)) + else: + grad_ptr = g_ptr.to(tl.pointer_type(tl.bfloat16)) + + outputs = tl.zeros((DIM,), dtype=tl.float32) + + if chunk == RANK: + if (input_id >= RANK * V) & (input_id < (RANK + 1) * V): + for i in range(count): + pos = tl.load(sorted_indices_ptr + chunk * M + c0 + i) + bid = pos // L + lid = pos % L + g = tl.load( + grad_output_ptr + + bid * stride_0 + + lid * stride_1 + + tl.arange(0, DIM), + mask=mask, + ).to(tl.float32) + outputs += g + + tl.atomic_add( + grad_ptr + input_id % V * DIM + tl.arange(0, DIM), + outputs, + mask=mask, + sem="relaxed", + ) + + else: + + for i in range(count): + pos = tl.load(sorted_indices_ptr + chunk * M + c0 + i) + bid = pos // L + lid = pos % L + g = tl.load( + grad_output_ptr + + bid * stride_0 + + lid * stride_1 + + tl.arange(0, DIM), + mask=mask, + ).to(tl.float32) + outputs += g + + # save to dst addr + buffer_ptr = tl.load(buffer_ptrs + input_id // V).to( + tl.pointer_type(tl.float32) + ) + buffer_ptr = tl.multiple_of(buffer_ptr, 16) + tl.store( + buffer_ptr + pid * DIM + tl.arange(0, DIM), outputs, mask=mask + ) + symm_mem_sync( + signal_ptrs, + None, + RANK, + SIZE, + hasPreviousMemAccess=True, + hasSubsequentMemAccess=True, + ) + else: + if (input_id >= RANK * V) & (input_id < (RANK + 1) * V): + symm_mem_sync( + signal_ptrs, + None, + RANK, + SIZE, + hasPreviousMemAccess=True, + hasSubsequentMemAccess=True, + ) + + buffer_ptr = tl.load(buffer_ptrs + RANK).to( + tl.pointer_type(tl.float32) + ) + buffer_ptr = tl.multiple_of(buffer_ptr, 16) + g = tl.load(buffer_ptr + pid * DIM + tl.arange(0, DIM), mask=mask) + tl.atomic_add( + grad_ptr + input_id % V * DIM + tl.arange(0, DIM), + g, + mask=mask, + sem="relaxed", + ) + + +def triton_sp_embedding_lookup_backward( + grad_output, + input_ids, + g_ptr, + vocab_size, + hdl, + group, + dtype=torch.bfloat16, + gathered_input_ids=None, +): + """ + inplace update embedding weight gradient + Args: + y: gradient of output + x: input ids Tensor + g_ptr: data_ptr of embedding weight gradient + Returns: + None + """ + + assert dtype in (torch.bfloat16, torch.float32) + T = 0 if dtype == torch.float32 else 1 + group_size = hdl.world_size + group_rank = hdl.rank + + shape = input_ids.shape + device = input_ids.device + dim = grad_output.size(-1) + + if len(shape) == 2: + B, L = shape + M = B * L + if gathered_input_ids is None: + gathered_input_ids = torch.empty( + (group_size, B, L), dtype=torch.long, device=device + ) + dist.all_gather_into_tensor(gathered_input_ids, input_ids, group=group) + else: + M = shape[0] + if gathered_input_ids is None: + gathered_input_ids = torch.empty( + (group_size, M), dtype=torch.long, device=device + ) + dist.all_gather_into_tensor(gathered_input_ids, input_ids, group=group) + + stride_0 = grad_output.stride(0) + stride_1 = grad_output.stride(1) + + sorted_ids, sorted_indices = torch.sort( + gathered_input_ids.view(group_size, M), stable=False, dim=-1 + ) + accum_counts = triton_scan_and_count(sorted_ids) + + DIM = triton.next_power_of_2(dim) + num_stages = 3 + num_warps = 2 + grid = (M,) + sp_embedding_lookup_backward_kernel[grid]( + grad_output, + sorted_ids, + sorted_indices, + accum_counts, + g_ptr, + hdl.buffer_ptrs_dev, + hdl.signal_pad_ptrs_dev, + stride_0, + stride_1, + dim, + M, + vocab_size, + B, + L, + DIM, + T, + group_size, + group_rank, + num_stages=num_stages, + num_warps=num_warps, + ) diff --git a/linghe/experimental/dla.py b/linghe/experimental/dla.py new file mode 100644 index 0000000..5fcea16 --- /dev/null +++ b/linghe/experimental/dla.py @@ -0,0 +1,1214 @@ +# -*- coding: utf-8 -*- +""" +Copyright (c) Ant Financial Service Group and its affiliates. +""" + +import torch +import triton +import triton.language as tl + +from linghe.experimental.symm_mem_barrier import symm_mem_sync + +""" +context-parallel lightning attention +""" + + +@triton.jit +def cp_lightning_attention_forward_kernel( + Q, + K, + V, + S, + Out, + buffer_ptrs, + signal_ptrs, + softmax_scale, + stride_q, + stride_k, + stride_v, + stride_s, + decay_scales, + L, + D: tl.constexpr, + KD: tl.constexpr, + VD: tl.constexpr, + BLOCK: tl.constexpr, + SIZE: tl.constexpr, + RANK: tl.constexpr, +): + bid = tl.program_id(0) + hid = tl.program_id(1) + kvid = tl.program_id(2) + N = D // VD + kid = kvid // N + vid = kvid % N + H = tl.num_programs(1) + + buffer_ptrs = buffer_ptrs.to(tl.pointer_type(tl.uint64)) + + c0 = bid * L + + decay_scale = -tl.load(decay_scales + hid) + + offs_b = tl.arange(0, BLOCK) + offs_k = tl.arange(0, KD) + offs_v = tl.arange(0, VD) + + q_ptrs = ( + Q + + c0 * stride_q + + hid * D + + kid * KD + + (offs_b[:, None] * stride_q + offs_k[None, :]) + ) + k_ptrs = ( + K + + c0 * stride_k + + hid * D + + kid * KD + + (offs_b[:, None] * stride_k + offs_k[None, :]) + ) + v_ptrs = ( + V + + c0 * stride_v + + hid * D + + vid * VD + + (offs_b[:, None] * stride_v + offs_v[None, :]) + ) + out_ptrs = ( + Out + + c0 * D * H + + hid * D + + vid * VD + + (offs_b[:, None] * H * D + offs_v[None, :]) + ) + s_ptrs = ( + S + + bid * stride_s + + hid * 2 * D * D + + kid * D * KD + + vid * VD + + (offs_k[:, None] * D + offs_v[None, :]) + ) + + buffer_offs = ( + bid * stride_s + + hid * 2 * D * D + + kid * D * KD + + vid * VD + + (offs_k[:, None] * D + offs_v[None, :]) + ) + + block_decay = tl.exp(decay_scale * BLOCK) + mask = tl.exp(decay_scale * (offs_b[:, None] - offs_b[None, :])) + mask = tl.where(offs_b[None, :] <= offs_b[:, None], mask, 0.0) * softmax_scale + b_offs = BLOCK - 1 - offs_b + decays = tl.exp(decay_scale * b_offs) + amps = block_decay * softmax_scale / decays + + state0 = tl.zeros((KD, VD), dtype=tl.float32) + + # cid = 0 + for n in range(0, L // 2, BLOCK): + n = tl.multiple_of(n, BLOCK) + + q = tl.load(q_ptrs + n * stride_q) + k = tl.trans(tl.load(k_ptrs + n * stride_k)) + v = tl.load(v_ptrs + n * stride_v) + + qk = tl.dot(q, k) * mask + o = tl.dot(qk.to(v.dtype), v) + + # o = tl.dot((q * amps[:, None]).to(q.dtype), state.to(q.dtype), o) + o += tl.dot((q * amps[:, None]), state0) + + state0 *= block_decay + # state = tl.dot((k * decays[None, :]).to(v.dtype), v, state) + state0 += tl.dot((k * decays[None, :]), v.to(tl.float32)) + + if KD == D: + tl.store(out_ptrs + n * H * D, o) + else: + tl.atomic_add(out_ptrs + n * H * D, o, sem="relaxed") + + buffer_ptr = tl.load(buffer_ptrs + RANK).to(tl.pointer_type(tl.float32)) + buffer_ptr = tl.multiple_of(buffer_ptr, 16) + tl.store(buffer_ptr + buffer_offs, state0) + + # c1 + state1 = tl.zeros((KD, VD), dtype=tl.float32) + DD = D * D + for n in range(L // 2, L, BLOCK): + n = tl.multiple_of(n, BLOCK) + + q = tl.load(q_ptrs + n * stride_q) + k = tl.trans(tl.load(k_ptrs + n * stride_k)) + v = tl.load(v_ptrs + n * stride_v) + + qk = tl.dot(q, k) * mask + o = tl.dot(qk.to(v.dtype), v) + + # o = tl.dot((q * amps[:, None]).to(q.dtype), state.to(q.dtype), o) + o += tl.dot((q * amps[:, None]), state1) + + state1 *= block_decay + # state = tl.dot((k * decays[None, :]).to(v.dtype), v, state) + state1 += tl.dot((k * decays[None, :]), v.to(tl.float32)) + + if KD == D: + tl.store(out_ptrs + n * H * D, o) + else: + tl.atomic_add(out_ptrs + n * H * D, o, sem="relaxed") + + tl.store(buffer_ptr + DD + buffer_offs, state1) + symm_mem_sync( + signal_ptrs, + None, + RANK, + SIZE, + hasPreviousMemAccess=True, + hasSubsequentMemAccess=True, + ) + + # accumulate c0 + gcid = RANK + if (gcid + 1) % 2 == 0: + pre_rank = RANK - 1 + chunk_decay = tl.exp(L // 2 * decay_scale) + pre_buffer_ptr = tl.load(buffer_ptrs + pre_rank).to(tl.pointer_type(tl.float32)) + state0 += tl.load(pre_buffer_ptr + buffer_offs) * chunk_decay + tl.store(buffer_ptr + buffer_offs, state0) + + # accumulate c1 + gcid = 2 * SIZE - 1 - RANK + if (gcid + 1) % 2 == 0: + pre_rank = RANK + 1 + chunk_decay = tl.exp(L // 2 * decay_scale) + pre_buffer_ptr = tl.load(buffer_ptrs + pre_rank).to(tl.pointer_type(tl.float32)) + state1 += tl.load(pre_buffer_ptr + DD + buffer_offs) * chunk_decay + tl.store(buffer_ptr + DD + buffer_offs, state1) + + symm_mem_sync( + signal_ptrs, + None, + RANK, + SIZE, + hasPreviousMemAccess=True, + hasSubsequentMemAccess=True, + ) + + if SIZE >= 4: + gcid = RANK + if (gcid + 1) % 4 == 0: + pre_rank = RANK - 2 + chunk_decay = tl.exp(L * decay_scale) + pre_buffer_ptr = tl.load(buffer_ptrs + pre_rank).to( + tl.pointer_type(tl.float32) + ) + state0 += tl.load(pre_buffer_ptr + buffer_offs) * chunk_decay + tl.store(buffer_ptr + buffer_offs, state0) + + gcid = 2 * SIZE - 1 - RANK + if (gcid + 1) % 4 == 0: + pre_rank = RANK + 2 + chunk_decay = tl.exp(L * decay_scale) + pre_buffer_ptr = tl.load(buffer_ptrs + pre_rank).to( + tl.pointer_type(tl.float32) + ) + state1 += tl.load(pre_buffer_ptr + DD + buffer_offs) * chunk_decay + tl.store(buffer_ptr + DD + buffer_offs, state1) + + symm_mem_sync( + signal_ptrs, + None, + RANK, + SIZE, + hasPreviousMemAccess=True, + hasSubsequentMemAccess=True, + ) + + if SIZE >= 8: + gcid = RANK + if (gcid + 1) % 8 == 0: + pre_rank = RANK - 4 + chunk_decay = tl.exp(L * 2 * decay_scale) + pre_buffer_ptr = tl.load(buffer_ptrs + pre_rank).to( + tl.pointer_type(tl.float32) + ) + state0 += tl.load(pre_buffer_ptr + buffer_offs) * chunk_decay + tl.store(buffer_ptr + buffer_offs, state0) + + symm_mem_sync( + signal_ptrs, + None, + RANK, + SIZE, + hasPreviousMemAccess=True, + hasSubsequentMemAccess=True, + ) + + # scatter + if SIZE >= 8: + gcid = 2 * SIZE - 1 - RANK + + if (gcid + 1) % 8 == 4: + pre_gcid = gcid - 4 + pre_rank = pre_gcid if pre_gcid < SIZE else 2 * SIZE - 1 - pre_gcid + pre_offs = 0 if pre_gcid < SIZE else DD + chunk_decay = tl.exp(L * decay_scale) + pre_buffer_ptr = tl.load(buffer_ptrs + pre_rank).to( + tl.pointer_type(tl.float32) + ) + state1 += tl.load(pre_buffer_ptr + pre_offs + buffer_offs) * chunk_decay + tl.store(buffer_ptr + buffer_offs, state1) + + symm_mem_sync( + signal_ptrs, + None, + RANK, + SIZE, + hasPreviousMemAccess=True, + hasSubsequentMemAccess=True, + ) + + # accumulate c0 + + if SIZE >= 4: + gcid = RANK + if ((gcid + 1) % 4 == 2) and (gcid > 4): + pre_gcid = gcid - 2 + pre_rank = pre_gcid if pre_gcid < SIZE else 2 * SIZE - 1 - pre_gcid + pre_offs = 0 if pre_gcid < SIZE else DD + chunk_decay = tl.exp(L * decay_scale) + pre_buffer_ptr = tl.load(buffer_ptrs + pre_rank).to( + tl.pointer_type(tl.float32) + ) + state0 += tl.load(pre_buffer_ptr + pre_offs + buffer_offs) * chunk_decay + tl.store(buffer_ptr + buffer_offs, state0) + + gcid = 2 * SIZE - 1 - RANK + if (gcid + 1) % 4 == 2: + pre_gcid = gcid - 2 + pre_rank = pre_gcid if pre_gcid < SIZE else 2 * SIZE - 1 - pre_gcid + pre_offs = 0 if pre_gcid < SIZE else DD + chunk_decay = tl.exp(L * decay_scale) + pre_buffer_ptr = tl.load(buffer_ptrs + pre_rank).to( + tl.pointer_type(tl.float32) + ) + state1 += tl.load(pre_buffer_ptr + pre_offs + buffer_offs) * chunk_decay + tl.store(buffer_ptr + DD + buffer_offs, state1) + + symm_mem_sync( + signal_ptrs, + None, + RANK, + SIZE, + hasPreviousMemAccess=True, + hasSubsequentMemAccess=True, + ) + + gcid = RANK + if ((gcid + 1) % 2 == 1) and (gcid > 0): + chunk_decay = tl.exp(L // 2 * decay_scale) + pre_gcid = gcid - 1 + pre_rank = pre_gcid if pre_gcid < SIZE else 2 * SIZE - 1 - pre_gcid + pre_offs = 0 if pre_gcid < SIZE else DD + pre_buffer_ptr = tl.load(buffer_ptrs + pre_rank).to(tl.pointer_type(tl.float32)) + state0 += tl.load(pre_buffer_ptr + pre_offs + buffer_offs) * chunk_decay + tl.store(buffer_ptr + buffer_offs, state0) + + gcid = 2 * SIZE - 1 - RANK + if (gcid + 1) % 2 == 1: + chunk_decay = tl.exp(L // 2 * decay_scale) + pre_gcid = gcid - 1 + pre_rank = pre_gcid if pre_gcid < SIZE else 2 * SIZE - 1 - pre_gcid + pre_offs = 0 if pre_gcid < SIZE else DD + pre_buffer_ptr = tl.load(buffer_ptrs + pre_rank).to(tl.pointer_type(tl.float32)) + state1 += tl.load(pre_buffer_ptr + pre_offs + buffer_offs) * chunk_decay + tl.store(buffer_ptr + DD + buffer_offs, state1) + + symm_mem_sync( + signal_ptrs, + None, + RANK, + SIZE, + hasPreviousMemAccess=True, + hasSubsequentMemAccess=True, + ) + + tl.store(s_ptrs, state0) + tl.store(s_ptrs + DD, state1) + + # read pre state + # chunk_decay = tl.exp(L // 2 * decay_scale) + if RANK > 0: + gcid = RANK + pre_gcid = gcid - 1 + pre_rank = pre_gcid if pre_gcid < SIZE else 2 * SIZE - 1 - pre_gcid + pre_offs = 0 if pre_gcid < SIZE else DD + buffer_ptr = tl.load(buffer_ptrs + pre_rank).to(tl.pointer_type(tl.float32)) + buffer_ptr = tl.multiple_of(buffer_ptr, 16) + # pre_state = tl.load(buffer_ptr + pre_offs + buffer_offs).to(Q.dtype.element_ty) + pre_state = tl.load( + buffer_ptr + pre_offs + buffer_offs + ) # .to(Q.dtype.element_ty) + + iter_amps = amps * block_decay + for n in range(0, L // 2, BLOCK): + n = tl.multiple_of(n, BLOCK) + + q = tl.load(q_ptrs + n * stride_q) + # q = (q * iter_amps[:, None]).to(q.dtype) + q = q * iter_amps[:, None] + + # o = tl.dot(q, pre_state.to(q.dtype)) + o = tl.dot(q, pre_state) + + iter_amps *= block_decay + tl.atomic_add(out_ptrs + n * H * D, o, sem="relaxed") + + gcid = 2 * SIZE - 1 - RANK + pre_gcid = gcid - 1 + pre_rank = pre_gcid if pre_gcid < SIZE else 2 * SIZE - 1 - pre_gcid + pre_offs = 0 if pre_gcid < SIZE else DD + buffer_ptr = tl.load(buffer_ptrs + pre_rank).to(tl.pointer_type(tl.float32)) + buffer_ptr = tl.multiple_of(buffer_ptr, 16) + # pre_state = tl.load(buffer_ptr + pre_offs + buffer_offs).to(Q.dtype.element_ty) + pre_state = tl.load(buffer_ptr + pre_offs + buffer_offs) # .to(Q.dtype.element_ty) + + iter_amps = amps * block_decay + for n in range(L // 2, L, BLOCK): + n = tl.multiple_of(n, BLOCK) + + q = tl.load(q_ptrs + n * stride_q) + # q = (q * iter_amps[:, None]).to(q.dtype) + q = q * iter_amps[:, None] + + # o = tl.dot(q, pre_state.to(q.dtype)) + o = tl.dot(q, pre_state) + + iter_amps *= block_decay + tl.atomic_add(out_ptrs + n * H * D, o, sem="relaxed") + + +def triton_cp_lightning_attention_forward( + q, k, v, decay_scales, hdl, group, hpc=True, softmax_scale=None +): + B, L, H, D = q.shape + h = k.shape[2] + assert H == h, "triton_lightning_attention_forward does NOT support GQA currently" + + if softmax_scale is None: + softmax_scale = D ** (-0.5) + + KD = 32 + VD = 128 + BLOCK = 32 + device = q.device + dtype = q.dtype + + num_warps = 2 # 2 + num_stages = 5 # 3 + + k_dim_block = D // KD + v_dim_block = D // VD + if k_dim_block == 1: + outputs = torch.empty( + (B, L, H, D), device=device, dtype=torch.float32 if hpc else dtype + ) + else: + outputs = torch.zeros( + (B, L, H, D), device=device, dtype=torch.float32 if hpc else dtype + ) + + s = torch.empty((B, H, 2, D, D), device=device, dtype=torch.float32) + assert L % BLOCK == 0 and BLOCK <= 64 + group_size = hdl.world_size + group_rank = hdl.rank + + grid = (B, H, k_dim_block * v_dim_block) + cp_lightning_attention_forward_kernel[grid]( + q, + k, + v, + s, + outputs, + hdl.buffer_ptrs_dev, + hdl.signal_pad_ptrs_dev, + softmax_scale, + q.stride(1), + k.stride(1), + v.stride(1), + s.stride(0), + decay_scales, + L, + D=D, + KD=KD, + VD=VD, + BLOCK=BLOCK, + SIZE=group_size, + RANK=group_rank, + num_warps=num_warps, + num_stages=num_stages, + ) + + o = outputs.to(dtype) + + return o, s + + +@triton.jit +def cp_lightning_attention_q_backward_kernel( + Q, + K, + V, + S, + G, + DQ, + buffer_ptrs, + signal_ptrs, + softmax_scale, + stride_q, + stride_k, + stride_v, + stride_s, + stride_g, + decay_scales, + L, + D: tl.constexpr, + KD: tl.constexpr, + VD: tl.constexpr, + BLOCK: tl.constexpr, + SIZE: tl.constexpr, + RANK: tl.constexpr, +): + bid = tl.program_id(0) + hid = tl.program_id(1) + kvid = tl.program_id(2) + N = D // VD + kid = kvid // N + vid = kvid % N + H = tl.num_programs(1) + DD = D * D + + c0 = bid * L + + buffer_ptrs = buffer_ptrs.to(tl.pointer_type(tl.uint64)) + decay_scale = -tl.load(decay_scales + hid) + + offs_b = tl.arange(0, BLOCK) + offs_k = tl.arange(0, KD) + offs_v = tl.arange(0, VD) + + k_ptrs = ( + K + + c0 * stride_k + + hid * D + + kid * KD + + (offs_b[:, None] * stride_k + offs_k[None, :]) + ) + v_ptrs = ( + V + + c0 * stride_v + + hid * D + + vid * VD + + (offs_b[:, None] * stride_v + offs_v[None, :]) + ) + s_ptrs = ( + S + + bid * stride_s + + hid * 2 * D * D + + kid * D * KD + + vid * VD + + (offs_k[:, None] * D + offs_v[None, :]) + ) + g_ptrs = ( + G + + c0 * D * H + + hid * D + + vid * VD + + (offs_b[:, None] * stride_g + offs_v[None, :]) + ) + dq_ptrs = ( + DQ + + c0 * D * H + + hid * D + + kid * KD + + (offs_b[:, None] * H * D + offs_k[None, :]) + ) + + buffer_offs = ( + bid * stride_s + + hid * 2 * D * D + + kid * D * KD + + vid * VD + + (offs_k[:, None] * D + offs_v[None, :]) + ) + + # store state to buffer + state = tl.load(s_ptrs) + buffer_ptr = tl.load(buffer_ptrs + RANK).to(tl.pointer_type(tl.float32)) + buffer_ptr = tl.multiple_of(buffer_ptr, 16) + tl.store(buffer_ptr + buffer_offs, state) + + state = tl.load(s_ptrs + DD) + buffer_ptr = tl.load(buffer_ptrs + RANK).to(tl.pointer_type(tl.float32)) + buffer_ptr = tl.multiple_of(buffer_ptr, 16) + tl.store(buffer_ptr + DD + buffer_offs, state) + + symm_mem_sync( + signal_ptrs, + None, + RANK, + SIZE, + hasPreviousMemAccess=True, + hasSubsequentMemAccess=True, + ) + + mask = tl.exp((offs_b[:, None] - offs_b[None, :]) * decay_scale) + mask = tl.where(offs_b[None, :] <= offs_b[:, None], mask, 0.0) * softmax_scale + decay_offs = BLOCK - 1 - offs_b + block_decay = tl.exp(decay_scale * BLOCK) + decays = tl.exp(decay_scale * decay_offs) # [0.01, 0.1, 1] + + # read state from buffer, c0 + state = tl.zeros((KD, VD), dtype=tl.float32) + if RANK > 0: + gcid = RANK + pre_gcid = gcid - 1 + pre_rank = pre_gcid if pre_gcid < SIZE else 2 * SIZE - 1 - pre_gcid + pre_offs = 0 if pre_gcid < SIZE else DD + buffer_ptr = tl.load(buffer_ptrs + pre_rank).to(tl.pointer_type(tl.float32)) + buffer_ptr = tl.multiple_of(buffer_ptr, 16) + state += tl.load(buffer_ptr + pre_offs + buffer_offs) # .to(Q.dtype.element_ty) + + for n in range(0, L // 2, BLOCK): + n = tl.multiple_of(n, BLOCK) + + k = tl.load(k_ptrs + n * stride_k) + v = tl.load(v_ptrs + n * stride_v) + g = tl.load(g_ptrs + n * stride_g) + + state = state * block_decay + + dqk = tl.dot(g, tl.trans(v)) * mask + + dq = ( + tl.dot(dqk.to(k.dtype), k) + + tl.dot(g * decays[:, None], tl.trans(state)) * softmax_scale + ) + + if VD == D: + tl.store(dq_ptrs + n * H * D, dq.to(DQ.dtype.element_ty)) + else: + tl.atomic_add( + dq_ptrs + n * H * D, dq.to(DQ.dtype.element_ty), sem="relaxed" + ) + + state = state + tl.dot((tl.trans(k) * decays[None, :]).to(v.dtype), v) + + symm_mem_sync( + signal_ptrs, + None, + RANK, + SIZE, + hasPreviousMemAccess=True, + hasSubsequentMemAccess=True, + ) + + # c1 + state = tl.zeros((KD, VD), dtype=tl.float32) + gcid = 2 * SIZE - 1 - RANK + pre_gcid = gcid - 1 + pre_rank = pre_gcid if pre_gcid < SIZE else 2 * SIZE - 1 - pre_gcid + pre_offs = 0 if pre_gcid < SIZE else DD + buffer_ptr = tl.load(buffer_ptrs + pre_rank).to(tl.pointer_type(tl.float32)) + buffer_ptr = tl.multiple_of(buffer_ptr, 16) + state += tl.load(buffer_ptr + pre_offs + buffer_offs) # .to(Q.dtype.element_ty) + + for n in range(L // 2, L, BLOCK): + n = tl.multiple_of(n, BLOCK) + + k = tl.load(k_ptrs + n * stride_k) + v = tl.load(v_ptrs + n * stride_v) + g = tl.load(g_ptrs + n * stride_g) + + state = state * block_decay + + dqk = tl.dot(g, tl.trans(v)) * mask + + dq = ( + tl.dot(dqk.to(k.dtype), k) + + tl.dot(g * decays[:, None], tl.trans(state)) * softmax_scale + ) + + if VD == D: + tl.store(dq_ptrs + n * H * D, dq.to(DQ.dtype.element_ty)) + else: + tl.atomic_add( + dq_ptrs + n * H * D, dq.to(DQ.dtype.element_ty), sem="relaxed" + ) + + state = state + tl.dot((tl.trans(k) * decays[None, :]).to(v.dtype), v) + + +@triton.jit +def cp_lightning_attention_kv_backward_kernel( + Q, + K, + V, + G, + DK, + DV, + buffer_ptrs, + signal_ptrs, + softmax_scale, + stride_q, + stride_k, + stride_v, + stride_g, + decay_scales, + L, + D: tl.constexpr, + KD: tl.constexpr, + VD: tl.constexpr, + BLOCK: tl.constexpr, + SIZE: tl.constexpr, + RANK: tl.constexpr, +): + bid = tl.program_id(0) + hid = tl.program_id(1) + kvid = tl.program_id(2) + N = D // VD + kid = kvid // N + vid = kvid % N + H = tl.num_programs(1) + DD = D * D + c0 = bid * L + + buffer_ptrs = buffer_ptrs.to(tl.pointer_type(tl.uint64)) + + decay_scale = -tl.load(decay_scales + hid) + + offs_b = tl.arange(0, BLOCK) + offs_k = tl.arange(0, KD) + offs_v = tl.arange(0, VD) + + q_ptrs = ( + Q + + c0 * stride_q + + hid * D + + kid * KD + + (offs_b[:, None] * stride_q + offs_k[None, :]) + ) + k_ptrs = ( + K + + c0 * stride_k + + hid * D + + kid * KD + + (offs_b[:, None] * stride_k + offs_k[None, :]) + ) + v_ptrs = ( + V + + c0 * stride_v + + hid * D + + vid * VD + + (offs_b[:, None] * stride_v + offs_v[None, :]) + ) + g_ptrs = ( + G + + c0 * D * H + + hid * D + + vid * VD + + (offs_b[:, None] * stride_g + offs_v[None, :]) + ) + + dk_ptrs = ( + DK + + c0 * H * D + + hid * D + + kid * KD + + (offs_b[:, None] * H * D + offs_k[None, :]) + ) + dv_ptrs = ( + DV + + c0 * H * D + + hid * D + + vid * VD + + (offs_b[:, None] * H * D + offs_v[None, :]) + ) + buffer_offs = ( + bid * H * 2 * D * D + + hid * 2 * D * D + + kid * D * KD + + vid * VD + + (offs_k[:, None] * D + offs_v[None, :]) + ) + + b_offs = BLOCK - 1 - offs_b + + block_decay = tl.exp(decay_scale * BLOCK) + amps = tl.exp(-decay_scale * b_offs) # [100, 10, 1] + decays = 1 / amps # [0.01, 0.1, 1] + sd = softmax_scale * block_decay + + mask = tl.exp((offs_b[:, None] - offs_b[None, :]) * decay_scale) + mask = tl.where(offs_b[None, :] <= offs_b[:, None], mask, 0.0) * softmax_scale + + gs0 = tl.zeros((KD, VD), dtype=tl.float32) + n_steps = tl.cdiv(L // 2, BLOCK) + for i in range(n_steps): + n = (n_steps - i - 1) * BLOCK + n = tl.multiple_of(n, BLOCK) + + q = tl.load(q_ptrs + n * stride_q) + k = tl.load(k_ptrs + n * stride_k) + v = tl.load(v_ptrs + n * stride_v) + g = tl.load(g_ptrs + n * stride_g) + + # qs = q * amps[:, None] # [100, 10, 1] + # ks = k * decays[:, None] # [0.01, 0.1, 1] + qk = tl.dot(q, tl.trans(k)) * mask + + dv = tl.dot(tl.trans(qk).to(g.dtype), g) + + dv += tl.dot(k * decays[:, None], gs0) + + dqk = (tl.dot(g, tl.trans(v)) * mask).to(q.dtype) + dk = tl.dot(tl.trans(dqk), (q * amps[:, None]).to(q.dtype)) + dk = tl.dot(v, tl.trans(gs0.to(v.dtype)), dk) + dk *= decays[:, None] + + gs0 *= block_decay + gs0 += tl.dot(tl.trans(q * amps[:, None]).to(g.dtype), g) * sd + + if VD == D: + tl.store(dk_ptrs + n * H * D, dk.to(DK.dtype.element_ty)) + else: + tl.atomic_add( + dk_ptrs + n * H * D, dk.to(DK.dtype.element_ty), sem="relaxed" + ) + if KD == D: + tl.store(dv_ptrs + n * H * D, dv.to(DV.dtype.element_ty)) + else: + tl.atomic_add( + dv_ptrs + n * H * D, dv.to(DV.dtype.element_ty), sem="relaxed" + ) + + buffer_ptr = tl.load(buffer_ptrs + RANK).to(tl.pointer_type(tl.float32)) + buffer_ptr = tl.multiple_of(buffer_ptr, 16) + tl.store(buffer_ptr + buffer_offs, gs0) + + gs1 = tl.zeros((KD, VD), dtype=tl.float32) + n_steps = tl.cdiv(L // 2, BLOCK) + for i in range(n_steps): + n = (n_steps - i - 1) * BLOCK + L // 2 + n = tl.multiple_of(n, BLOCK) + + q = tl.load(q_ptrs + n * stride_q) + k = tl.load(k_ptrs + n * stride_k) + v = tl.load(v_ptrs + n * stride_v) + g = tl.load(g_ptrs + n * stride_g) + + # qs = q * amps[:, None] # [100, 10, 1] + # ks = k * decays[:, None] # [0.01, 0.1, 1] + qk = tl.dot(q, tl.trans(k)) * mask + + dv = tl.dot(tl.trans(qk).to(g.dtype), g) + + dv += tl.dot(k * decays[:, None], gs1) + + dqk = (tl.dot(g, tl.trans(v)) * mask).to(q.dtype) + dk = tl.dot(tl.trans(dqk), (q * amps[:, None]).to(q.dtype)) + dk = tl.dot(v, tl.trans(gs1.to(v.dtype)), dk) + dk *= decays[:, None] + + gs1 *= block_decay + gs1 += tl.dot(tl.trans(q * amps[:, None]).to(g.dtype), g) * sd + + if VD == D: + tl.store(dk_ptrs + n * H * D, dk.to(DK.dtype.element_ty)) + else: + tl.atomic_add( + dk_ptrs + n * H * D, dk.to(DK.dtype.element_ty), sem="relaxed" + ) + if KD == D: + tl.store(dv_ptrs + n * H * D, dv.to(DV.dtype.element_ty)) + else: + tl.atomic_add( + dv_ptrs + n * H * D, dv.to(DV.dtype.element_ty), sem="relaxed" + ) + + tl.store(buffer_ptr + DD + buffer_offs, gs1) + + symm_mem_sync( + signal_ptrs, + None, + RANK, + SIZE, + hasPreviousMemAccess=True, + hasSubsequentMemAccess=True, + ) + + # accumulate c0 + gcid = RANK + if RANK > 0: + if (gcid + 1) % 2 == 1: + next_rank = RANK + 1 + chunk_decay = tl.exp(L // 2 * decay_scale) + next_buffer_ptr = tl.load(buffer_ptrs + next_rank).to( + tl.pointer_type(tl.float32) + ) + gs0 += tl.load(next_buffer_ptr + buffer_offs) * chunk_decay + tl.store(buffer_ptr + buffer_offs, gs0) + + # accumulate c1 + gcid = 2 * SIZE - 1 - RANK + if (gcid + 1) % 2 == 1: + next_rank = RANK - 1 + chunk_decay = tl.exp(L // 2 * decay_scale) + next_buffer_ptr = tl.load(buffer_ptrs + next_rank).to( + tl.pointer_type(tl.float32) + ) + gs1 += tl.load(next_buffer_ptr + DD + buffer_offs) * chunk_decay + tl.store(buffer_ptr + DD + buffer_offs, gs1) + + symm_mem_sync( + signal_ptrs, + None, + RANK, + SIZE, + hasPreviousMemAccess=True, + hasSubsequentMemAccess=True, + ) + + if RANK >= 4: + # accumulate c0 + gcid = RANK + if RANK > 0: + if (gcid + 1) % 4 == 1: + next_rank = RANK + 2 + chunk_decay = tl.exp(L * decay_scale) + next_buffer_ptr = tl.load(buffer_ptrs + next_rank).to( + tl.pointer_type(tl.float32) + ) + gs0 += tl.load(next_buffer_ptr + buffer_offs) * chunk_decay + tl.store(buffer_ptr + buffer_offs, gs0) + + # accumulate c1 + gcid = 2 * SIZE - 1 - RANK + if (gcid + 1) % 4 == 1: + next_rank = RANK - 2 + chunk_decay = tl.exp(L * decay_scale) + next_buffer_ptr = tl.load(buffer_ptrs + next_rank).to( + tl.pointer_type(tl.float32) + ) + gs1 += tl.load(next_buffer_ptr + DD + buffer_offs) * chunk_decay + tl.store(buffer_ptr + DD + buffer_offs, gs1) + + symm_mem_sync( + signal_ptrs, + None, + RANK, + SIZE, + hasPreviousMemAccess=True, + hasSubsequentMemAccess=True, + ) + + if SIZE >= 8: + # accumulate c1 + gcid = 2 * SIZE - 1 - RANK + if (gcid + 1) % 8 == 1: + next_rank = RANK - 4 + chunk_decay = tl.exp(L * 2 * decay_scale) + next_buffer_ptr = tl.load(buffer_ptrs + next_rank).to( + tl.pointer_type(tl.float32) + ) + gs1 += tl.load(next_buffer_ptr + DD + buffer_offs) * chunk_decay + tl.store(buffer_ptr + DD + buffer_offs, gs1) + + symm_mem_sync( + signal_ptrs, + None, + RANK, + SIZE, + hasPreviousMemAccess=True, + hasSubsequentMemAccess=True, + ) + + # scatter + if SIZE >= 8: + # accumulate c0 + gcid = RANK + if (gcid + 1) % 8 == 5: + next_gcid = gcid + 4 + next_rank = next_gcid if next_gcid < SIZE else 2 * SIZE - 1 - next_gcid + next_offs = 0 if next_gcid < SIZE else DD + chunk_decay = tl.exp(L * 2 * decay_scale) + next_buffer_ptr = tl.load(buffer_ptrs + next_rank).to( + tl.pointer_type(tl.float32) + ) + gs0 += tl.load(next_buffer_ptr + next_offs + buffer_offs) * chunk_decay + tl.store(buffer_ptr + buffer_offs, gs0) + + symm_mem_sync( + signal_ptrs, + None, + RANK, + SIZE, + hasPreviousMemAccess=True, + hasSubsequentMemAccess=True, + ) + + # scatter + if SIZE >= 4: + # accumulate c0 + gcid = RANK + if (gcid + 1) % 4 == 3: + next_gcid = gcid + 2 + next_rank = next_gcid if next_gcid < SIZE else 2 * SIZE - 1 - next_gcid + next_offs = 0 if next_gcid < SIZE else DD + chunk_decay = tl.exp(L * decay_scale) + next_buffer_ptr = tl.load(buffer_ptrs + next_rank).to( + tl.pointer_type(tl.float32) + ) + gs0 += tl.load(next_buffer_ptr + next_offs + buffer_offs) * chunk_decay + tl.store(buffer_ptr + buffer_offs, gs0) + + gcid = 2 * SIZE - 1 - RANK + if ((gcid + 1) % 4 == 3) and (gcid < 14): + next_gcid = gcid + 2 + next_rank = next_gcid if next_gcid < SIZE else 2 * SIZE - 1 - next_gcid + next_offs = 0 if next_gcid < SIZE else DD + chunk_decay = tl.exp(L * decay_scale) + next_buffer_ptr = tl.load(buffer_ptrs + next_rank).to( + tl.pointer_type(tl.float32) + ) + gs1 += tl.load(next_buffer_ptr + next_offs + buffer_offs) * chunk_decay + tl.store(buffer_ptr + buffer_offs, gs1) + + symm_mem_sync( + signal_ptrs, + None, + RANK, + SIZE, + hasPreviousMemAccess=True, + hasSubsequentMemAccess=True, + ) + + # accumulate c0 + gcid = RANK + if (gcid + 1) % 2 == 0: + next_gcid = gcid + 1 + next_rank = next_gcid if next_gcid < SIZE else 2 * SIZE - 1 - next_gcid + next_offs = 0 if next_gcid < SIZE else DD + chunk_decay = tl.exp(L // 2 * decay_scale) + next_buffer_ptr = tl.load(buffer_ptrs + next_rank).to( + tl.pointer_type(tl.float32) + ) + gs0 += tl.load(next_buffer_ptr + next_offs + buffer_offs) * chunk_decay + tl.store(buffer_ptr + buffer_offs, gs0) + + # accumulate c1 + gcid = 2 * SIZE - 1 - RANK + if RANK > 0: + if (gcid + 1) % 2 == 0: + next_gcid = gcid + 1 + next_rank = next_gcid if next_gcid < SIZE else 2 * SIZE - 1 - next_gcid + next_offs = 0 if next_gcid < SIZE else DD + chunk_decay = tl.exp(L * decay_scale) + next_buffer_ptr = tl.load(buffer_ptrs + next_rank).to( + tl.pointer_type(tl.float32) + ) + gs1 += tl.load(next_buffer_ptr + DD + buffer_offs) * chunk_decay + tl.store(buffer_ptr + DD + buffer_offs, gs1) + + symm_mem_sync( + signal_ptrs, + None, + RANK, + SIZE, + hasPreviousMemAccess=True, + hasSubsequentMemAccess=True, + ) + + gcid = RANK + next_gcid = gcid + 1 + next_rank = next_gcid if next_gcid < SIZE else 2 * SIZE - 1 - next_gcid + next_offs = 0 if next_gcid < SIZE else DD + chunk_decay = tl.exp(L // 2 * decay_scale) + next_buffer_ptr = tl.load(buffer_ptrs + next_rank).to(tl.pointer_type(tl.float32)) + gs = tl.load(next_buffer_ptr + next_offs + buffer_offs) + + n_steps = tl.cdiv(L // 2, BLOCK) + for i in range(n_steps): + n = (n_steps - i - 1) * BLOCK + n = tl.multiple_of(n, BLOCK) + + k = tl.load(k_ptrs + n * stride_k) + v = tl.load(v_ptrs + n * stride_v) + + dv = tl.dot(k * decays[:, None], gs) + + dk = tl.dot(v.to(tl.float32), tl.trans(gs)) + + gs *= block_decay + + tl.atomic_add(dk_ptrs + n * H * D, dk.to(DK.dtype.element_ty), sem="relaxed") + tl.atomic_add(dv_ptrs + n * H * D, dv.to(DV.dtype.element_ty), sem="relaxed") + + if RANK > 0: + gcid = 2 * SIZE - 1 - RANK + next_gcid = gcid + 1 + next_rank = next_gcid if next_gcid < SIZE else 2 * SIZE - 1 - next_gcid + next_offs = 0 if next_gcid < SIZE else DD + chunk_decay = tl.exp(L // 2 * decay_scale) + next_buffer_ptr = tl.load(buffer_ptrs + next_rank).to( + tl.pointer_type(tl.float32) + ) + gs = tl.load(next_buffer_ptr + next_offs + buffer_offs) + + n_steps = tl.cdiv(L // 2, BLOCK) + for i in range(n_steps): + n = (n_steps - i - 1) * BLOCK + L // 2 + n = tl.multiple_of(n, BLOCK) + + k = tl.load(k_ptrs + n * stride_k) + v = tl.load(v_ptrs + n * stride_v) + + dv = tl.dot(k * decays[:, None], gs) + + dk = tl.dot(v, tl.trans(gs.to(v.dtype))) + + gs *= block_decay + + tl.atomic_add( + dk_ptrs + n * H * D, dk.to(DK.dtype.element_ty), sem="relaxed" + ) + tl.atomic_add( + dv_ptrs + n * H * D, dv.to(DV.dtype.element_ty), sem="relaxed" + ) + + +def triton_cp_lightning_attention_backward( + output_grad, + q, + k, + v, + s, + decay_scales, + hdl, + group, + softmax_scale=None, + hpc=False, + hp=False, +): + B, L, H, D = q.shape + if softmax_scale is None: + softmax_scale = D ** (-0.5) + + dtype = q.dtype + device = q.device + + KD = 64 + VD = 128 + BLOCK = 32 + k_dim_block = D // KD + v_dim_block = D // VD + if v_dim_block > 1: + dq = torch.zeros( + (B, L, H, D), device=device, dtype=torch.float32 if hpc else dtype + ) + else: + dq = torch.empty( + (B, L, H, D), device=device, dtype=torch.float32 if hpc else dtype + ) + assert L % BLOCK == 0 and BLOCK <= 64 + group_size = hdl.world_size + group_rank = hdl.rank + + grid = (B, H, k_dim_block * v_dim_block) + num_warps = 4 # 2 + num_stages = 3 # 5 + cp_lightning_attention_q_backward_kernel[grid]( + q, + k, + v, + s, + output_grad, + dq, + hdl.buffer_ptrs_dev, + hdl.signal_pad_ptrs_dev, + softmax_scale, + q.stride(1), + k.stride(1), + v.stride(1), + s.stride(0), + output_grad.stride(1), + decay_scales, + L, + D=D, + KD=KD, + VD=VD, + BLOCK=BLOCK, + SIZE=group_size, + RANK=group_rank, + num_warps=num_warps, + num_stages=num_stages, + ) + + KD = 32 + VD = 128 + BLOCK = 32 + k_dim_block = D // KD + v_dim_block = D // VD + if v_dim_block > 1: + dk = torch.zeros( + (B, L, H, D), device=device, dtype=torch.float32 if hpc else dtype + ) + else: + dk = torch.empty( + (B, L, H, D), device=device, dtype=torch.float32 if hpc else dtype + ) + if k_dim_block > 1: + dv = torch.zeros( + (B, L, H, D), device=device, dtype=torch.float32 if hpc else dtype + ) + else: + dv = torch.empty( + (B, L, H, D), device=device, dtype=torch.float32 if hpc else dtype + ) + num_warps = 4 # 4 + num_stages = 5 # 5 + grid = (B, H, k_dim_block * v_dim_block) + cp_lightning_attention_kv_backward_kernel[grid]( + q, + k, + v, + output_grad, + dk, + dv, + hdl.buffer_ptrs_dev, + hdl.signal_pad_ptrs_dev, + softmax_scale, + q.stride(1), + k.stride(1), + v.stride(1), + output_grad.stride(1), + decay_scales, + L, + D=D, + KD=KD, + VD=VD, + BLOCK=BLOCK, + SIZE=group_size, + RANK=group_rank, + num_warps=num_warps, + num_stages=num_stages, + ) + + dq = dq.to(dtype) + dk = dk.to(dtype) + dv = dv.to(dtype) + return dq, dk, dv diff --git a/linghe/experimental/dmm.py b/linghe/experimental/dmm.py new file mode 100644 index 0000000..c700aab --- /dev/null +++ b/linghe/experimental/dmm.py @@ -0,0 +1,194 @@ +# -*- coding: utf-8 -*- +""" +Copyright (c) Ant Financial Service Group and its affiliates. +""" + +import torch +import torch.distributed as dist +import triton +import triton.language as tl + +from linghe.experimental.symm_mem_barrier import symm_mem_sync + + +@triton.jit +def split_tp_mm_kernel( + a_ptr, + b_ptr, + c_ptr, + split_atomic_ptr, + buffer_ptrs, + signal_ptrs, + M, + N: tl.constexpr, + K: tl.constexpr, + BLOCK_SIZE_K: tl.constexpr, + BLOCK_SIZE_M: tl.constexpr, + BLOCK_SIZE_N: tl.constexpr, + SPLIT_COUNT: tl.constexpr, + SIZE: tl.constexpr, + RANK: tl.constexpr, +): + pid_m = tl.program_id(axis=0) + pid_n = tl.program_id(axis=1) + nbn = tl.num_programs(1) + pid_k = tl.program_id(axis=2) + buffer_ptrs = buffer_ptrs.to(tl.pointer_type(tl.uint64)) + + k = tl.cdiv(K, BLOCK_SIZE_K * SPLIT_COUNT) + offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) + offs_k = tl.arange(0, BLOCK_SIZE_K) + a_ptrs = a_ptr + pid_k * K // SPLIT_COUNT + offs_m[:, None] * K + offs_k[None, :] + b_ptrs = b_ptr + pid_k * K // SPLIT_COUNT + offs_n[None, :] * K + offs_k[:, None] + + c = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32) + for i in range(k): + a = tl.load(a_ptrs) + b = tl.load(b_ptrs) + c = tl.dot(a, b, c) + a_ptrs += BLOCK_SIZE_K + b_ptrs += BLOCK_SIZE_K + offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) + c_ptrs = c_ptr + offs_m[:, None] * N + offs_n[None, :] + + if SPLIT_COUNT == 1: + buffer_ptr = tl.load(buffer_ptrs + RANK).to(tl.pointer_type(tl.float32)) + buffer_ptr = tl.multiple_of(buffer_ptr, 16) + + tl.store(buffer_ptr + offs_m[:, None] * N + offs_n[None, :], c) + symm_mem_sync( + signal_ptrs, + None, + RANK, + SIZE, + hasPreviousMemAccess=True, + hasSubsequentMemAccess=True, + ) + + outputs = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32) + for i in tl.static_range(SIZE): + buffer_ptr = tl.load(buffer_ptrs + i).to(tl.pointer_type(tl.float32)) + buffer_ptr = tl.multiple_of(buffer_ptr, 16) + outputs += tl.load(buffer_ptr + offs_m[:, None] * N + offs_n[None, :]) + tl.store(c_ptrs, outputs) + + # for j in range(0, RANK): + # buffer_ptr = tl.load(buffer_ptrs + j).to(tl.pointer_type(tl.float32)) + # buffer_ptr = tl.multiple_of(buffer_ptr, 16) + # c += tl.load(buffer_ptr + offs_m[:, None] * N + offs_n[None, :]) + # for j in range(RANK+1, SIZE): + # buffer_ptr = tl.load(buffer_ptrs + j).to(tl.pointer_type(tl.float32)) + # buffer_ptr = tl.multiple_of(buffer_ptr, 16) + # c += tl.load(buffer_ptr + offs_m[:, None] * N + offs_n[None, :]) + # tl.store(c_ptrs, c) + + else: + tl.atomic_add(c_ptrs, c, sem="relaxed") + # tl.atomic_add(c_ptrs, c) + atomic_index = pid_m * nbn + pid_n + tl.atomic_add(split_atomic_ptr + atomic_index, 1) + + if pid_k == SPLIT_COUNT - 1: # pid_k == 0 will result in error output + tl.debug_barrier() + while tl.load(split_atomic_ptr + atomic_index) < SPLIT_COUNT: + pass + + c = tl.load(c_ptrs) + + buffer_ptr = tl.load(buffer_ptrs + RANK).to(tl.pointer_type(tl.float32)) + buffer_ptr = tl.multiple_of(buffer_ptr, 16) + tl.store(buffer_ptr + offs_m[:, None] * N + offs_n[None, :], c) + + symm_mem_sync( + signal_ptrs, + None, + RANK, + SIZE, + hasPreviousMemAccess=True, + hasSubsequentMemAccess=True, + ) + + outputs = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32) + for j in tl.static_range(SIZE): + buffer_ptr = tl.load(buffer_ptrs + j).to(tl.pointer_type(tl.float32)) + buffer_ptr = tl.multiple_of(buffer_ptr, 16) + outputs += tl.load(buffer_ptr + offs_m[:, None] * N + offs_n[None, :]) + tl.store(c_ptrs, outputs) + + # for j in range(0, RANK): + # buffer_ptr = tl.load(buffer_ptrs + j).to(tl.pointer_type(tl.float32)) + # buffer_ptr = tl.multiple_of(buffer_ptr, 16) + # c += tl.load(buffer_ptr + offs_m[:, None] * N + offs_n[None, :]) + # for j in range(RANK+1, SIZE): + # buffer_ptr = tl.load(buffer_ptrs + j).to(tl.pointer_type(tl.float32)) + # buffer_ptr = tl.multiple_of(buffer_ptr, 16) + # c += tl.load(buffer_ptr + offs_m[:, None] * N + offs_n[None, :]) + # tl.store(c_ptrs, c) + + +def triton_split_tp_gemm( + x: torch.Tensor, w: torch.Tensor, hdl, group: dist.ProcessGroup +): + """ + tensor-parallel fc2 in the shared expert, use split-k implementation + y = all_reduce(x @ fc2) + Args: + a: left matrix with bf16 precision + b: right matrix with bf16 precision + + Returns: + c: all-reduced output + """ + assert x.is_contiguous() and w.is_contiguous() + M, K = x.size() + N, K = w.size() + BLOCK_SIZE_K = 128 + BLOCK_SIZE_M = 128 + BLOCK_SIZE_N = max([x for x in [16, 32, 64, 128] if N % x == 0]) + SPLIT_COUNT = 1 # min(triton.cdiv(K, 2048), 4) + assert M % BLOCK_SIZE_M == 0 and K % BLOCK_SIZE_K == 0 + assert K % (BLOCK_SIZE_K * SPLIT_COUNT) == 0 + + device = x.device + if SPLIT_COUNT == 1: + c = torch.empty(M, N, dtype=x.dtype, device=device) + split_atomic_signal = None + else: + c = torch.zeros(M, N, dtype=torch.float32, device=device) + split_atomic_signal = torch.zeros( + M // BLOCK_SIZE_M * N // BLOCK_SIZE_N, dtype=torch.int32, device=device + ) + group_size = hdl.world_size + group_rank = hdl.rank + + grid = lambda META: ( + triton.cdiv(M, META["BLOCK_SIZE_M"]), + triton.cdiv(N, META["BLOCK_SIZE_N"]), + SPLIT_COUNT, + ) # noqa + num_warps = 4 + num_stages = 3 + split_tp_mm_kernel[grid]( + x, + w, + c, + split_atomic_signal, + hdl.buffer_ptrs_dev, + hdl.signal_pad_ptrs_dev, + M, + N, + K, + BLOCK_SIZE_K, + BLOCK_SIZE_M, + BLOCK_SIZE_N, + SPLIT_COUNT, + group_size, + group_rank, + num_warps=num_warps, + num_stages=num_stages, + ) + if SPLIT_COUNT > 1: + c = c.to(x.dtype) + return c diff --git a/linghe/experimental/gemm.py b/linghe/experimental/gemm.py new file mode 100644 index 0000000..adea8f8 --- /dev/null +++ b/linghe/experimental/gemm.py @@ -0,0 +1,113 @@ + +# -*- coding: utf-8 -*- +""" +Copyright (c) Ant Financial Service Group and its affiliates. +""" + +import os +import torch +import triton +import triton.language as tl +from triton import Config + + +@triton.jit +def _compute_pid(tile_id, num_pid_in_group, num_pid_m, GROUP_SIZE_M): + group_id = tile_id // num_pid_in_group + first_pid_m = group_id * GROUP_SIZE_M + group_size_m = min(num_pid_m - first_pid_m, GROUP_SIZE_M) + pid_m = first_pid_m + (tile_id % group_size_m) + pid_n = (tile_id % num_pid_in_group) // group_size_m + return pid_m, pid_n + + +@triton.jit +def tma_persistent_matmul_kernel( + a_desc, + b_desc, + c_desc, + M, + N, + K, + BLOCK_SIZE_M: tl.constexpr, + BLOCK_SIZE_N: tl.constexpr, + BLOCK_SIZE_K: tl.constexpr, + GROUP_SIZE_M: tl.constexpr, + SM: tl.constexpr, ): + start_pid = tl.program_id(axis=0) + num_pid_m = tl.cdiv(M, BLOCK_SIZE_M) + num_pid_n = tl.cdiv(N, BLOCK_SIZE_N) + k_tiles = tl.cdiv(K, BLOCK_SIZE_K) + num_tiles = num_pid_m * num_pid_n + + tid_c = start_pid - SM + num_pid_in_group = GROUP_SIZE_M * num_pid_n + + for tid in tl.range(start_pid, num_tiles, SM, flatten=True): + pid_m, pid_n = _compute_pid(tid, num_pid_in_group, num_pid_m, + GROUP_SIZE_M) + offs_a = pid_m * BLOCK_SIZE_M + offs_b = pid_n * BLOCK_SIZE_N + + accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32) + for k in range(k_tiles): + offs_k = k * BLOCK_SIZE_K + a = a_desc.load([offs_a, offs_k]) + b = b_desc.load([offs_b, offs_k]) + accumulator = tl.dot(a, b.T, accumulator) + + tid_c += SM + pid_m, pid_n = _compute_pid(tid_c, num_pid_in_group, num_pid_m, GROUP_SIZE_M) + offs_a_acc = pid_m * BLOCK_SIZE_M + offs_b_acc = pid_n * BLOCK_SIZE_N + + acc = tl.reshape(accumulator, (BLOCK_SIZE_M, 2, BLOCK_SIZE_N // 2)) + acc = tl.permute(acc, (0, 2, 1)) + acc0, acc1 = tl.split(acc) + c_desc.store([offs_a_acc, offs_b_acc], acc0) + c_desc.store([offs_a_acc, offs_b_acc + BLOCK_SIZE_N // 2], acc1) + + +def triton_tma_persistent_matmul(a, b): + M, K = a.shape + N, K = b.shape + dtype = torch.float32 + + c = torch.empty((M, N), device=a.device, dtype=dtype) + + SM = torch.cuda.get_device_properties("cuda").multi_processor_count + + BLOCK_M = 128 + BLOCK_K = 64 + BLOCK_N = 64 + GROUP_SIZE_M = 8 + + a_desc = triton.tools.tensor_descriptor.TensorDescriptor(a, a.shape, + a.stride(), + [BLOCK_M, BLOCK_K]) + b_desc = triton.tools.tensor_descriptor.TensorDescriptor(b, b.shape, + b.stride(), + [BLOCK_N, BLOCK_K]) + c_desc = triton.tools.tensor_descriptor.TensorDescriptor(c, c.shape, + c.stride(), + [BLOCK_M, + BLOCK_N // 2]) + + def grid(META): + nonlocal a_desc, b_desc, c_desc + return (min(SM, + triton.cdiv(M, BLOCK_M) * triton.cdiv(N, BLOCK_N), ),) + + tma_persistent_matmul_kernel[grid]( + a_desc, + b_desc, + c_desc, + M, + N, + K, + BLOCK_M, + BLOCK_K, + BLOCK_N, + GROUP_SIZE_M, + SM=SM) + return c diff --git a/linghe/experimental/gmem_barrier_arrive_wait.py b/linghe/experimental/gmem_barrier_arrive_wait.py new file mode 100644 index 0000000..b50bad6 --- /dev/null +++ b/linghe/experimental/gmem_barrier_arrive_wait.py @@ -0,0 +1,73 @@ +""" +copied from https://github.com/meta-pytorch/kraken/blob/main/kraken/_ptx_utils/gmem_barrier_arrive_wait.py +""" + +import triton +import triton.language as tl + + +@triton.jit +def arrive_gmem_barrier( + addr, + update: tl.constexpr = 1, # set the lock value to + sem: tl.constexpr = "release", + scope: tl.constexpr = "gpu", + op: tl.constexpr = "atomic_xchg", + skip_sync: tl.constexpr = False, +): + tl.static_assert( + op == "atomic_xchg", + "Currently only support atomic_xchg wait on gmem_barriers. ", + ) + + if not skip_sync: + tl.inline_asm_elementwise( + "bar.sync 0;", "=r", [], dtype=tl.int32, is_pure=False, pack=1 + ) + return tl.atomic_xchg(addr, update, sem=sem, scope=scope) + + +@triton.jit +def wait_gmem_barrier( + addr, + expect: tl.constexpr = 1, # wait until lock is set to expect + update: tl.constexpr = 0, # update the lock once it is aquired. + sem: tl.constexpr = "acquire", + scope: tl.constexpr = "gpu", + op: tl.constexpr = "ld", + skip_sync: tl.constexpr = False, +): + """ + 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. + + Args: + addr: Memory address of the barrier to wait on (Must be a scalar) + expect: Expected value to wait for (default: 1) + update: Update the barrier with once acquired (default: 0) + sem: Memory semantics for the atomic operation (default: "acquire") + scope: Scope of the atomic operation. Options: "gpu", "sys" (default: "gpu") + op: Atomic operation type (default: "ld", currently only supported option) + """ + tl.static_assert( + op == "ld" and update == 0, "Currently only support ld wait on gmem_barriers. " + ) + # TODO(joydddd): add support for cas barriers. + + tl.static_assert(addr.type.is_ptr(), "Barrier address must be a scalar.") + # TODO(joydddd): add wait_gmem_multi_barrier. (each thread waits on a different barrier). + + # Spin-wait loop: + # Uses atomic_add with update=0 for ld.global.{sem}.{scope} + # Triton generates smem broadcasting of tl.atomic_add return value in ptx, + # but it is optimized away by ptxas in SASS, hence no performance overhead. + while tl.atomic_add(addr, update, sem=sem, scope=scope) != expect: + pass + + if not skip_sync: + tl.inline_asm_elementwise( + "bar.sync 0;", "=r", [], dtype=tl.int32, is_pure=False, pack=1 + ) + # tl.debug_barrier() cause significant performance loss. (Perhaps breaks triton prefetching?) diff --git a/linghe/experimental/symm_mem_barrier.py b/linghe/experimental/symm_mem_barrier.py new file mode 100644 index 0000000..04b4eb8 --- /dev/null +++ b/linghe/experimental/symm_mem_barrier.py @@ -0,0 +1,163 @@ +# -*- coding: utf-8 -*- +""" +copied from https://github.com/meta-pytorch/kraken/blob/main/kraken/_ptx_utils/symm_mem_barrier.py +""" + +import triton +import triton.language as tl + + +@triton.jit +def _get_tid(): + return tl.inline_asm_elementwise( + """ + mov.u32 $0, %tid.x; + mov.u32 $1, %tid.y; + mov.u32 $2, %tid.z; + """, + "=r,=r,=r", + [], + dtype=(tl.uint32, tl.uint32, tl.uint32), + is_pure=True, + pack=1, + ) + + +@triton.jit +def _get_ntid(): + return tl.inline_asm_elementwise( + """ + mov.u32 $0, %ntid.x; + mov.u32 $1, %ntid.y; + mov.u32 $2, %ntid.z; + """, + "=r,=r,=r", + [], + dtype=(tl.uint32, tl.uint32, tl.uint32), + is_pure=True, + pack=1, + ) + + +@triton.jit +def _get_flat_tid(): + tid_x, tid_y, tid_z = _get_tid() + ntid_x, ntid_y, _ = _get_ntid() + return tid_z * ntid_y * ntid_x + tid_y * ntid_x + tid_x + + +@triton.jit +def _get_flat_bid(): + return ( + tl.program_id(2) * tl.num_programs(1) * tl.num_programs(0) + + tl.program_id(1) * tl.num_programs(0) + + tl.program_id(0) + ) + + +@triton.jit +def _send_signal(addrs, sem: tl.constexpr): + tl.inline_asm_elementwise( + f""" + {{ + .reg .u32 %tmp32_<1>; + .reg .pred %p<1>; + + send_signal: + atom.global.{sem}.sys.cas.b32 %tmp32_0, [$1], 0, 1; + setp.eq.u32 %p0, %tmp32_0, 0; + @!%p0 bra send_signal; + }} + """, + "=r, l", + [addrs], + dtype=addrs.dtype, + is_pure=False, + pack=1, + ) + + +@triton.jit +def _wait_signal(addrs, sem: tl.constexpr): + tl.inline_asm_elementwise( + f""" + {{ + .reg .u32 %tmp32_<1>; + .reg .pred %p<1>; + + wait_signal: + atom.global.sys.{sem}.cas.b32 %tmp32_0, [$1], 1, 0; + setp.eq.u32 %p0, %tmp32_0, 1; + @!%p0 bra wait_signal; + }} + """, + "=r, l", + [addrs], + dtype=tl.int32, + is_pure=False, + pack=1, + ) + + +@triton.jit +def symm_mem_sync( + signal_pad_ptrs, + block_id, + rank: tl.constexpr, + world_size: tl.constexpr, + hasPreviousMemAccess: tl.constexpr = False, + hasSubsequentMemAccess: tl.constexpr = False, +): + """ + 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) + + CUDA graph friendliness: + + 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. + """ + if block_id is None: + block_id = _get_flat_bid() + flat_tid = _get_flat_tid() + + remote_ranks = tl.arange(0, world_size) + signal_pad_ptrs = signal_pad_ptrs.to(tl.pointer_type(tl.uint64)) + remote_signal_pad_addrs = tl.load(signal_pad_ptrs + remote_ranks).to( + tl.pointer_type(tl.uint32) + ) + send_addrs = remote_signal_pad_addrs + block_id * world_size + rank + + local_signal_pad_addr = tl.load(signal_pad_ptrs + rank).to( + tl.pointer_type(tl.uint32) + ) + wait_addrs = local_signal_pad_addr + block_id * world_size + remote_ranks + + if hasPreviousMemAccess: + tl.debug_barrier() + + if flat_tid < world_size: + _send_signal(send_addrs, "release" if hasPreviousMemAccess else "relaxed") + _wait_signal(wait_addrs, "acquire" if hasSubsequentMemAccess else "relaxed") + + if hasSubsequentMemAccess: + tl.debug_barrier() diff --git a/linghe/experimental/test_demb.py b/linghe/experimental/test_demb.py new file mode 100644 index 0000000..91f1660 --- /dev/null +++ b/linghe/experimental/test_demb.py @@ -0,0 +1,269 @@ +# -*- coding: utf-8 -*- +""" +Copyright (c) Ant Financial Service Group and its affiliates. +""" + +import os +from datetime import timedelta + +import torch +import torch.distributed as dist +import torch.distributed._symmetric_memory as symm_mem +import torch.nn.functional as F + +from linghe.experimental.demb import ( + triton_tp_embedding_lookup_forward, + triton_tp_embedding_lookup_backward, + triton_sp_embedding_lookup_forward, + triton_sp_embedding_lookup_backward, +) +from linghe.tools.benchmark import benchmark_func +from linghe.tools.check import output_check + + +def torch_tp_emb(input_ids, weights, group): + V, D = weights.shape + group_rank = group.rank() + ids = input_ids % V + output = weights[ids] + mask = torch.logical_and( + input_ids >= group_rank * V, input_ids < (group_rank + 1) * V + ) + output = torch.where(mask[:, :, None], output, 0.0 * output) + dist.all_reduce(output, op=dist.ReduceOp.SUM) + return output + + +def test_tp_emb( + B=1, M=4096, N=157184, D=4096, coef=1.0, grad_coef=1.0, group=None, bench=False +): + group_size = group.size() + group_rank = group.rank() + + device_module = torch.get_device_module("cuda") + device_module.set_device(torch.device(f"cuda:{group_rank}")) + + device = "cuda" + dtype = torch.bfloat16 + buffers = symm_mem.empty( + (B * M, D), + dtype=torch.bfloat16, + device=device, + ) + hdl = symm_mem.rendezvous(buffers, dist.group.WORLD) + + local_weights = torch.randn((N, D), dtype=dtype, device=device, requires_grad=False) + + local_weights = (local_weights * coef).detach().clone().requires_grad_() + + global_weights = torch.empty((group_size, N, D), dtype=dtype, device=device) + dist.all_gather_into_tensor(global_weights, local_weights.detach(), group=group) + global_weights = ( + torch.reshape(global_weights, (group_size * N, D)).contiguous().requires_grad_() + ) + + local_ids = torch.randint( + 0, N * group_size, (B, M), dtype=torch.long, device=device + ) + global_ids = torch.empty((group_size, B, M), dtype=torch.long, device=device) + dist.all_gather_into_tensor(global_ids, local_ids, group=group) + global_ids = global_ids[0] + + local_grad = torch.randn((B, M, D), dtype=dtype, device=device, requires_grad=False) + global_grad = torch.empty((group_size, B, M, D), dtype=dtype, device=device) + dist.all_gather_into_tensor(global_grad, local_grad, group=group) + global_grad = global_grad.sum(0) + + output_ref = global_weights[global_ids] + output_ref.backward(global_grad) + global_weight_grad_ref = global_weights.grad + local_weight_grad_ref = global_weight_grad_ref[ + group_rank * N : (group_rank + 1) * N + ] + global_weights.grad = None + # dist_output = torch_tp_emb(global_ids, local_weights, group) + # output_check(output_ref, dist_output, name=f'output:{group_rank}', atol=1e-4, rtol=1e-5) + + output = triton_tp_embedding_lookup_forward( + global_ids, + local_weights, + hdl, + group, + ) + output_check(output_ref, output, name=f"output:{group_rank}", atol=1e-4, rtol=1e-5) + + weight_grad = torch.zeros((N, D), dtype=torch.float32, device=device) + triton_tp_embedding_lookup_backward( + local_grad, + global_ids, + weight_grad.data_ptr(), + N, + hdl, + group, + dtype=weight_grad.dtype, + ) + output_check( + local_weight_grad_ref, + weight_grad.to(dtype), + name=f"grad:{group_rank}", + atol=-1e-4, + rtol=1e-2, + ) + + if bench: + benchmark_func( + F.embedding, global_ids, global_weights, ref_bytes=M * D * group_size * 2 + ) + benchmark_func( + torch_tp_emb, + global_ids, + local_weights, + group, + ref_bytes=M * D * group_size * 2, + ) + benchmark_func( + triton_tp_embedding_lookup_forward, + global_ids, + local_weights, + hdl, + group, + ref_bytes=M * D * group_size * 2, + ) + benchmark_func( + triton_tp_embedding_lookup_backward, + local_grad, + global_ids, + weight_grad.data_ptr(), + N, + hdl, + group, + dtype=weight_grad.dtype, + ref_bytes=M * D * group_size * 2, + ) + + +def test_sp_emb( + B=1, M=4096, N=157184, D=4096, coef=1.0, grad_coef=1.0, group=None, bench=False +): + group_size = group.size() + group_rank = group.rank() + + device_module = torch.get_device_module("cuda") + device_module.set_device(torch.device(f"cuda:{group_rank}")) + + device = "cuda" + dtype = torch.bfloat16 + buffers = symm_mem.empty( + (B * M, D), + dtype=torch.float32, + device=device, + ) + hdl = symm_mem.rendezvous(buffers, dist.group.WORLD) + + local_weights = torch.randn((N, D), dtype=dtype, device=device, requires_grad=False) + + local_weights = (local_weights * coef).detach().clone().requires_grad_() + + global_weights = torch.empty((group_size, N, D), dtype=dtype, device=device) + dist.all_gather_into_tensor(global_weights, local_weights.detach(), group=group) + global_weights = ( + torch.reshape(global_weights, (group_size * N, D)).contiguous().requires_grad_() + ) + + local_ids = torch.randint( + 0, N * group_size, (B, M), dtype=torch.long, device=device + ) + global_ids = torch.empty((group_size, B, M), dtype=torch.long, device=device) + dist.all_gather_into_tensor(global_ids, local_ids, group=group) + global_ids = torch.reshape(global_ids, (group_size * B, M)) + + local_grad = torch.randn((B, M, D), dtype=dtype, device=device, requires_grad=False) + global_grad = torch.empty((group_size, B, M, D), dtype=dtype, device=device) + dist.all_gather_into_tensor(global_grad, local_grad, group=group) + global_grad = torch.reshape(global_grad, (group_size * B, M, D)) + + output_ref = global_weights[global_ids] + output_ref.backward(global_grad) + global_weight_grad_ref = global_weights.grad + local_weight_grad_ref = global_weight_grad_ref[ + group_rank * N : (group_rank + 1) * N + ] + global_weights.grad = None + output_ref = output_ref[group_rank * B : (group_rank + 1) * B] + # dist_output = torch_sp_emb(global_ids, local_weights, group) + # output_check(output_ref, dist_output, name=f'output:{group_rank}', atol=1e-4, rtol=1e-5) + + output = triton_sp_embedding_lookup_forward( + local_ids, + local_weights, + hdl, + group, + ) + output_check(output_ref, output, name=f"output:{group_rank}", atol=1e-4, rtol=1e-5) + + weight_grad = torch.zeros((N, D), dtype=torch.float32, device=device) + triton_sp_embedding_lookup_backward( + local_grad, + local_ids, + weight_grad.data_ptr(), + N, + hdl, + group, + dtype=weight_grad.dtype, + ) + output_check( + local_weight_grad_ref, + weight_grad.to(dtype), + name=f"grad:{group_rank}", + atol=0.02, + rtol=0.03, + ) + + if bench: + benchmark_func( + F.embedding, global_ids, global_weights, ref_bytes=M * D * group_size * 2 + ) + # benchmark_func(torch_sp_emb, global_ids, local_weights, group, + # ref_bytes=M * D * group_size * 2) + benchmark_func( + triton_sp_embedding_lookup_forward, + local_ids, + local_weights, + hdl, + group, + ref_bytes=M * D * group_size * 2, + ) + benchmark_func( + triton_sp_embedding_lookup_backward, + local_grad, + local_ids, + weight_grad.data_ptr(), + N, + hdl, + group, + dtype=weight_grad.dtype, + ref_bytes=M * D * group_size * 2, + ) + + +if __name__ == "__main__": + # torchrun --nproc_per_node=2 test_demb.py + world_size = int(os.environ["WORLD_SIZE"]) + local_rank = int(os.environ["LOCAL_RANK"]) + os.environ["TORCH_NCCL_AVOID_RECORD_STREAMS"] = "1" + print(f"{world_size=} {local_rank=}") + dist.init_process_group( + backend="nccl", + init_method="env://", + world_size=world_size, + rank=local_rank, + timeout=timedelta(seconds=10), + ) + group = dist.distributed_c10d._get_default_group() + torch.distributed.distributed_c10d._set_pg_timeout( + timedelta(seconds=10), dist.group.WORLD + ) + # test_tp_emb(M=8192, N=157184, D=4096, coef=1.0, grad_coef=1.0, group=group, bench=True) + test_sp_emb( + M=8192, N=157184, D=4096, coef=1.0, grad_coef=1.0, group=group, bench=True + ) diff --git a/linghe/experimental/test_dla.py b/linghe/experimental/test_dla.py new file mode 100644 index 0000000..1de0151 --- /dev/null +++ b/linghe/experimental/test_dla.py @@ -0,0 +1,202 @@ +# -*- coding: utf-8 -*- +""" +Copyright (c) Ant Financial Service Group and its affiliates. +""" + +import math +import os +from datetime import timedelta + +import torch +import torch.distributed as dist +import torch.distributed._symmetric_memory as symm_mem + +from linghe.experimental.dla import ( + triton_cp_lightning_attention_forward, + triton_cp_lightning_attention_backward, +) +from linghe.tools.benchmark import benchmark_func +from linghe.tools.check import output_check + + +def torch_la(q, k, v, decay_scales, s=None, hp=False): + dtype = q.dtype + if hp: + q = q.double() + k = k.double() + v = v.double() + if s is not None: + s = s.double() + else: + q = q.float() + k = k.float() + v = v.float() + if s is not None: + s = s.float() + B, L, H, D = q.shape + h = k.shape[2] + assert L == k.shape[1] + softmax_scale = 1.0 / math.sqrt(D) + query = q.transpose(1, 2) # [B, head, len, D] + key = torch.permute(k, (0, 2, 3, 1)) # [B, head, D, len] + value = v.transpose(1, 2) # [B, head, len, D] + if h != H: + g = H // h + key = torch.repeat_interleave(key, g, D=1) + value = torch.repeat_interleave(value, g, D=1) + + arr = torch.arange(L, dtype=torch.float64 if hp else torch.float32, device=q.device) + decay_matrix = arr.view(-1, 1) - arr.view(1, -1) + decay_matrix = torch.exp(-decay_scales[:, None, None] * decay_matrix[None]) + decay_matrix = torch.tril(decay_matrix, 0) + + score = torch.matmul(query, key) * softmax_scale + score *= decay_matrix[None] + att = torch.matmul(score, value) + + decay_arr = torch.exp(-decay_scales[:, None, None] * (arr[:, None] + 1)) + if s is not None: + att = att + torch.matmul(query * decay_arr, s) + + att = torch.reshape(att.transpose(1, 2), [B, L, H, D]).contiguous() + + decay_key = key * torch.exp(-decay_scales[:, None, None] * (L - 1 - arr)) + state = decay_key @ value + if s is not None: + state += s * torch.exp(-decay_scales[:, None, None]) + + return att.to(dtype), state.to(torch.float32) + + +def rearange(x, group): + B, L, H, D = x.shape + group_size = group.size() + X = torch.empty((group_size, B, L, H, D), dtype=x.dtype, device=x.device) + dist.all_gather_into_tensor(X, x.detach(), group=group) + X = torch.permute( + torch.reshape(X, (group_size, B, 2, L // 2, H, D)), (1, 2, 0, 3, 4, 5) + ) + X = torch.reshape( + torch.cat([X[:, 0], torch.flip(X[:, 1], (1,))], 1), (B, L * group_size, H, D) + ) + X = X.contiguous().requires_grad_() + return X + + +def select(x, group): + group_size = group.size() + group_rank = group.rank() + B, L, H, D = x.shape + l = L // (2 * group_size) + x1 = x[:, group_rank * l : (group_rank + 1) * l] + x2 = x[:, (group_size * 2 - group_rank - 1) * l : (group_size * 2 - group_rank) * l] + return torch.cat([x1, x2], 1) + + +def test_dist_la( + B=1, + L=4096, + H=16, + D=128, + group=None, + hpc=True, + digest=False, + coef=1.0, + grad_coef=1.0, + bench=False, +): + group_size = group.size() + group_rank = group.rank() + + device_module = torch.get_device_module("cuda") + device_module.set_device(torch.device(f"cuda:{group_rank}")) + + device = torch.device("cuda") + dtype = torch.bfloat16 + + buffers = symm_mem.empty((B, H, 2, D, D), dtype=torch.float32, device=device) + hdl = symm_mem.rendezvous(buffers, group) + + q = torch.randn(B, L, H, D, dtype=dtype, device=device) * coef + q = q.requires_grad_() + + k = torch.randn(B, L, H, D, dtype=dtype, device=device) * coef + k = k.requires_grad_() + + v = torch.randn(B, L, H, D, dtype=dtype, device=device) * coef + v = v.requires_grad_() + + g = torch.randn(B, L, H, D, dtype=dtype, device=device) * grad_coef + + decay_scales = 2 ** ( + -0.5 * torch.arange(1, H + 1, dtype=torch.float32, device=device) + ) + # decay_scales = 0.0 * torch.arange(1, H+1, dtype=torch.float32, device=device) + + Q = rearange(q.detach(), group).requires_grad_() + K = rearange(k.detach(), group).requires_grad_() + V = rearange(v.detach(), group).requires_grad_() + G = rearange(g, group) + + global_output_ref, global_state_ref = torch_la(Q, K, V, decay_scales, hp=False) + global_output_ref.backward(G) + DQ_ref = Q.grad + DK_ref = K.grad + DV_ref = V.grad + Q.grad = None + K.grad = None + V.grad = None + output_ref = select(global_output_ref, group) + dq_ref = select(DQ_ref, group) + dk_ref = select(DK_ref, group) + dv_ref = select(DV_ref, group) + + output, state = triton_cp_lightning_attention_forward( + q, k, v, decay_scales, hdl, group, hpc=hpc + ) + output_check(output_ref, output, atol=-0.2, rtol=0.05, name=f"output:{group_rank}") + + dq, dk, dv = triton_cp_lightning_attention_backward( + g, q, k, v, state, decay_scales, hdl, group, hpc=hpc + ) + output_check(dq_ref, dq, name="dq", rtol=-0.1, atol=1.0) + output_check(dk_ref, dk, name="dk", rtol=-0.1, atol=1.0) + output_check(dv_ref, dv, name="dv", rtol=-0.1, atol=1.0) + + if bench: + ref_bytes = (B * L * H * D * 8 + B * H * D * D * 8) * group_size + benchmark_func(torch_la, Q, K, V, decay_scales, ref_bytes=ref_bytes) + benchmark_func( + triton_cp_lightning_attention_forward, + q, + k, + v, + decay_scales, + hdl, + group, + hpc=hpc, + ref_bytes=ref_bytes, + ) + + +if __name__ == "__main__": + # torchrun --nproc_per_node=2 test_dla.py + world_size = int(os.environ["WORLD_SIZE"]) + local_rank = int(os.environ["LOCAL_RANK"]) + os.environ["TORCH_NCCL_AVOID_RECORD_STREAMS"] = "1" + print(f"{world_size=} {local_rank=}") + dist.init_process_group( + backend="nccl", + init_method="env://", + world_size=world_size, + rank=local_rank, + timeout=timedelta(seconds=10), + ) + group = dist.distributed_c10d._get_default_group() + torch.distributed.distributed_c10d._set_pg_timeout( + timedelta(seconds=10), dist.group.WORLD + ) + test_dist_la( + B=1, L=4096, H=64, D=128, group=group, hpc=True, digest=False, bench=False + ) + # test_dist_la(B=1, L=4096, H=64, D=128, group=group, hpc=True, digest=False, bench=True) diff --git a/linghe/experimental/test_dmm.py b/linghe/experimental/test_dmm.py new file mode 100644 index 0000000..ef31bb1 --- /dev/null +++ b/linghe/experimental/test_dmm.py @@ -0,0 +1,115 @@ +# -*- coding: utf-8 -*- +""" +Copyright (c) Ant Financial Service Group and its affiliates. +""" + +import os +from datetime import timedelta + +import torch +import torch.distributed as dist +import torch.distributed._symmetric_memory as symm_mem +import torch.nn.functional as F + +from linghe.experimental.dmm import triton_split_tp_gemm +from linghe.tools.benchmark import benchmark_func +from linghe.tools.check import output_check + + +def torch_dist_mm(x, w, group): + dtype = x.dtype + output = x @ w.t() + output = output.float() + dist.all_reduce(output, op=dist.ReduceOp.SUM) + return output.to(dtype) + + +def test_dist_mm( + M=4096, N=2048, K=4096, coef=1.0, grad_coef=1.0, group=None, bench=False +): + group_size = group.size() + group_rank = group.rank() + + device_module = torch.get_device_module("cuda") + device_module.set_device(torch.device(f"cuda:{group_rank}")) + + device = "cuda" + dtype = torch.bfloat16 + + buffers = symm_mem.empty((M, N), dtype=torch.float32, device=device) + hdl = symm_mem.rendezvous(buffers, dist.group.WORLD) + + # hdl = symm_mem.get_symm_mem_workspace(group.group_name, min_size=M * N * 4) + # buf_list = [ + # hdl.get_buffer(i, [M, N], torch.float32, 0) + # for i in range(hdl.world_size) + # ] + # buffer_tuple = tuple(buf_list) + + local_weights = torch.randn((N, K), dtype=dtype, device=device, requires_grad=False) + local_weights = (local_weights * coef).detach().clone().requires_grad_() + global_weights = torch.empty((group_size, N, K), dtype=dtype, device=device) + dist.all_gather_into_tensor(global_weights, local_weights.detach(), group=group) + global_weights = ( + torch.reshape(torch.permute(global_weights, (1, 0, 2)), (N, group_size * K)) + .contiguous() + .requires_grad_() + ) + + local_states = torch.randn((M, K), dtype=dtype, device=device) + global_states = torch.empty((group_size, M, K), dtype=dtype, device=device) + dist.all_gather_into_tensor(global_states, local_states, group=group) + global_states = torch.reshape( + torch.permute(global_states, (1, 0, 2)), (M, group_size * K) + ).contiguous() + + output_ref = global_states @ global_weights.t() + + dist_output = torch_dist_mm(local_states, local_weights, group) + output_check(output_ref, dist_output, name=f"output:{group_rank}", atol=10.0) + + output = triton_split_tp_gemm( + local_states, + local_weights, + hdl, + group, + ) + output_check(output_ref, output, name=f"output:{group_rank}", atol=10.0) + + if bench: + ref_flops = M * N * K * 2 * group_size + benchmark_func(F.linear, global_states, global_weights, ref_flops=ref_flops) + benchmark_func( + torch_dist_mm, local_states, local_weights, group, ref_flops=ref_flops + ) + benchmark_func( + triton_split_tp_gemm, + local_states, + local_weights, + hdl, + group, + ref_flops=ref_flops, + ) + + +if __name__ == "__main__": + # torchrun --nproc_per_node=2 test_dmm.py + world_size = int(os.environ["WORLD_SIZE"]) + local_rank = int(os.environ["LOCAL_RANK"]) + os.environ["TORCH_NCCL_AVOID_RECORD_STREAMS"] = "1" + print(f"{world_size=} {local_rank=}") + dist.init_process_group( + backend="nccl", + init_method="env://", + world_size=world_size, + rank=local_rank, + timeout=timedelta(seconds=10), + ) + group = dist.distributed_c10d._get_default_group() + torch.distributed.distributed_c10d._set_pg_timeout( + timedelta(seconds=10), dist.group.WORLD + ) + # test_dist_mm(M=1024, N=8192, K=8192, group=group, bench=True) + # test_dist_mm(M=8192, N=1024, K=8192, group=group, bench=True) + # test_dist_mm(M=8192, N=8192, K=1024, group=group, bench=True) + test_dist_mm(M=8192, N=8192, K=8192, group=group, bench=True) diff --git a/linghe/facade/add.py b/linghe/facade/add.py index cac4306..f289cbe 100644 --- a/linghe/facade/add.py +++ b/linghe/facade/add.py @@ -10,6 +10,7 @@ class InplaceAddFunction(torch.autograd.Function): """""" + @staticmethod def forward(ctx, x: torch.Tensor, y: torch.Tensor): return triton_inplace_add(x, y) @@ -28,4 +29,4 @@ def inplace_add(x: torch.Tensor, y: torch.Tensor): Returns: updated x tensor """ - return InplaceAddFunction.apply(x, y) \ No newline at end of file + return InplaceAddFunction.apply(x, y) diff --git a/linghe/facade/emb.py b/linghe/facade/emb.py new file mode 100644 index 0000000..55fb5c6 --- /dev/null +++ b/linghe/facade/emb.py @@ -0,0 +1,110 @@ +# -*- coding: utf-8 -*- +""" +Copyright (c) Ant Financial Service Group and its affiliates. +""" + +import torch + +from linghe.utils.emb import triton_embedding_forward, triton_embedding_backward + + +class DeprecatedFusedAccumulationEmbeddingLookup(torch.autograd.Function): + """""" + + @staticmethod + def forward(ctx, x, w_ptr, g_ptr, dim, dtype, grad_dtype): + x = x.long() + ctx.grad_dtype = grad_dtype + ctx.g_ptr = g_ptr + ctx.save_for_backward(x) + return triton_embedding_forward(x, w_ptr, dim, dtype) + + @staticmethod + def backward(ctx, grad_output): + (x,) = ctx.saved_tensors + triton_embedding_backward(grad_output, x, ctx.g_ptr, ctx.grad_dtype) + return None, None, None, None, None, None + + +def deprecated_fused_accumulation_embedding_lookup( + x: torch.Tensor, w_ptr, g_ptr, dim, dtype, grad_dtype +): + """ + embedding lookup + Args: + x: input ids + w_ptr: + g_ptr: + dim: + dtype: + grad_dtype: + Returns: + lookup output + """ + x = x.double().requires_grad_() + return DeprecatedFusedAccumulationEmbeddingLookup.apply( + x, w_ptr, g_ptr, dim, dtype, grad_dtype + ) + + +class FusedAccumulationEmbeddingLookup(torch.autograd.Function): + """""" + + @staticmethod + def forward(ctx, x, w, grad_name): + dim = w.size(1) + dtype = w.dtype + ctx.save_for_backward(x, w) + ctx.grad_name = grad_name + return triton_embedding_forward(x, w.data_ptr(), dim, dtype) + + @staticmethod + def backward(ctx, grad_output): + x, w = ctx.saved_tensors + grad = getattr(w, ctx.grad_name) + triton_embedding_backward(grad_output, x, grad.data_ptr(), grad.dtype) + return None, None, None + + +def fused_accumulation_embedding_lookup( + x: torch.Tensor, w: torch.nn.Parameter, grad_name: str = "grad" +): + """ + embedding lookup + Args: + x: input ids + w: embedding weight, should contain a `grad_name` tensor + Returns: + lookup output + """ + return FusedAccumulationEmbeddingLookup.apply(x, w, grad_name) + + +class EmbeddingLookup(torch.autograd.Function): + """""" + + @staticmethod + def forward(ctx, x, w): + dim = w.size(1) + dtype = w.dtype + ctx.save_for_backward(x, w) + return triton_embedding_forward(x, w.data_ptr(), dim, dtype) + + @staticmethod + def backward(ctx, grad_output): + x, w = ctx.saved_tensors + grad = torch.zeros_like(w) + triton_embedding_backward(grad_output, x, grad.data_ptr(), grad.dtype) + return None, grad + + +def embedding_lookup(x: torch.Tensor, w: torch.nn.Parameter): + """ + embedding lookup + Args: + x: input ids + w: embedding weight + Returns: + lookup output + """ + return EmbeddingLookup.apply(x, w) diff --git a/linghe/facade/fp32_gemm.py b/linghe/facade/fp32_gemm.py index 2f6ff34..bcd6237 100644 --- a/linghe/facade/fp32_gemm.py +++ b/linghe/facade/fp32_gemm.py @@ -3,52 +3,5 @@ Copyright (c) Ant Financial Service Group and its affiliates. """ -import torch - -from linghe.gemm.fp32_gemm import (triton_fp32_gemm, - triton_fp32_gemm_for_backward, - triton_fp32_gemm_for_update) - - -class Fp32GEMM(torch.autograd.Function): - """""" - @staticmethod - def forward(ctx, input: torch.Tensor, weight: torch.Tensor): - shape = input.shape - assert len(shape) == 3 - input = input.view(shape[0] * shape[1], shape[2]) - - logits = triton_fp32_gemm(input, weight.data) - - ctx.input_requires_grad = input.requires_grad - ctx.weight_requires_grad = weight.requires_grad - ctx.shape = shape - ctx.save_for_backward(input, weight.data) - - return logits.view(shape[0], shape[1], weight.shape[0]) - - @staticmethod - def backward(ctx, grad_output): - shape = grad_output.shape - grad_output = grad_output.view(shape[0] * shape[1], shape[2]) - input, weight = ctx.saved_tensors - - dx = triton_fp32_gemm_for_backward(grad_output, weight) - dx = dx.view(*ctx.shape) - - dw = triton_fp32_gemm_for_update(grad_output, input) - - return dx, dw - - -def fp32_gemm(input: torch.Tensor, weight: torch.Tensor): - """ - gemm with bf16/fp16 inputs and float32 output, - currently used in MoE router gemm. - Args: - input: bf16/fp16 activation tensor - weight: bf16/fp16 weight tensor - Returns: - output of gemm - """ - return Fp32GEMM.apply(input, weight) \ No newline at end of file +# the file will be removed in future, please import functions in gemm.py +from linghe.facade.gemm import fp32_gemm diff --git a/linghe/facade/fp8_gemm.py b/linghe/facade/fp8_gemm.py new file mode 100644 index 0000000..2367f42 --- /dev/null +++ b/linghe/facade/fp8_gemm.py @@ -0,0 +1,7 @@ +# -*- coding: utf-8 -*- +""" +Copyright (c) Ant Financial Service Group and its affiliates. +""" + +# the file will be removed in future, please import functions in gemm.py +from linghe.facade.gemm import smooth_gemm, smooth_groued_gemm diff --git a/linghe/facade/gate.py b/linghe/facade/gate.py new file mode 100644 index 0000000..5843f1f --- /dev/null +++ b/linghe/facade/gate.py @@ -0,0 +1,134 @@ +# -*- coding: utf-8 -*- +""" +Copyright (c) Ant Financial Service Group and its affiliates. +""" + +import torch + +from linghe.utils.gate import ( + triton_group_rms_norm_gate_forward, + triton_group_rms_norm_gate_backward, + triton_group_rms_norm_gate_and_mxfp8_quant_forward, + triton_group_rms_norm_gate_and_mxfp8_quant_backward, +) + + +class GroupRMSNormGateFunction(torch.autograd.Function): + """""" + + @staticmethod + def forward(ctx, attn_output, gate, weight, eps=1e-6, group_size=4): + + output = triton_group_rms_norm_gate_forward( + attn_output, gate, weight, eps=eps, group_size=group_size + ) + + ctx.save_for_backward(attn_output, gate, weight) + ctx.eps = eps + ctx.group_size = group_size + + return output + + @staticmethod + def backward(ctx, dy): + attn_output, gate, weight = ctx.saved_tensors + + dx, dg, dw = triton_group_rms_norm_gate_backward( + dy, attn_output, gate, weight, eps=ctx.eps, group_size=ctx.group_size + ) + + return dx, dg, dw, None, None + + +def group_rms_norm_gate( + attn_output: torch.Tensor, + gate: torch.Tensor, + weight: torch.Tensor, + eps: float = 1e-6, + group_size: int = 4, + transpose: bool = True, +): + """ + return group_rms_norm(transpose(attn_output, [0,1]), weight) * sigmoid(gate) + Args: + attn_output: output of core attn, shape [bs, length, n_heads, head_dim] + gate: gate tensor for attention output, shape [length, bs, dim] + weight: weight of RMS norm, shape [dim] + eps: epsilon for RMS + group_size: group size of group RMS norm + transpose: whether gate is transposed + Returns: + output with shape [length, bs, dim] + """ + assert transpose + return GroupRMSNormGateFunction.apply(attn_output, gate, weight, eps, group_size) + + +class Mxfp8GroupRMSNormGateFunction(torch.autograd.Function): + """""" + + @staticmethod + def forward( + ctx, + attn_output, + gate, + weight, + quantizer, + grad_quantizer, + cls, + eps=1e-6, + group_size=4, + ): + + shape = attn_output.shape + assert len(shape) == 3 + + x_q, x_s, xt_q, xt_s = triton_group_rms_norm_gate_and_mxfp8_quant_forward( + attn_output, gate, weight, eps=eps, group_size=group_size + ) + + output = cls( + shape=x_q.shape, + dtype=attn_output.dtype, + fp8_dtype=quantizer.dtype, + rowwise_data=x_q.view(shape), + rowwise_scale_inv=x_s, + columnwise_data=xt_q.view(shape), + columnwise_scale_inv=xt_s, + quantizer=quantizer, + requires_grad=attn_output.requires_grad, + ) + + ctx.save_for_backward(attn_output, gate, weight) + ctx.eps = eps + ctx.group_size = group_size + ctx.shape = shape + ctx.grad_quantizer = grad_quantizer + ctx.cls = cls + + return output + + @staticmethod + def backward(ctx, dy): + attn_output, gate, weight = ctx.saved_tensors + grad_quantizer = ctx.grad_quantizer + + dx, dg_q, dg_s, dgt_q, dgt_s, dw = ( + triton_group_rms_norm_gate_and_mxfp8_quant_backward( + dy, attn_output, gate, weight, eps=ctx.eps, group_size=ctx.group_size + ) + ) + + dg_out = ctx.cls( + shape=ctx.shape, + dtype=dy.dtype, + fp8_dtype=grad_quantizer.dtype, + rowwise_data=dg_q.view(ctx.shape) if dg_q is not None else None, + rowwise_scale_inv=dg_s, + columnwise_data=dgt_q.view(ctx.shape) if dgt_q is not None else None, + columnwise_scale_inv=dgt_s, + quantizer=grad_quantizer, + requires_grad=attn_output.requires_grad, + ) + + return dx, dg_out, dw, None, None, None, None, None diff --git a/linghe/facade/gemm.py b/linghe/facade/gemm.py new file mode 100644 index 0000000..85f982a --- /dev/null +++ b/linghe/facade/gemm.py @@ -0,0 +1,218 @@ +# -*- coding: utf-8 -*- +""" +Copyright (c) Ant Financial Service Group and its affiliates. +""" + +import torch + +from linghe.gemm.fp32_gemm import ( + triton_fp32_gemm, + triton_fp32_gemm_for_backward, + triton_fp32_gemm_for_update, + triton_split_fp32_gemm, + triton_split_fp32_gemm_for_backward, + triton_split_fp32_gemm_for_update, +) +from linghe.experimental.gemm import triton_tma_persistent_matmul +from linghe.utils.add import triton_inplace_add +from linghe.utils.transpose import triton_pad_transpose + + +# the function is used in transformer_engine/pytorch/cpp_extensions/gemm.py +def smooth_gemm(A, B, layout="TN", out=None, accumulate=True, online_transpose=False): + if 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 + out = torch._scaled_mm( + x_q, + w_q.t(), + scale_a=x_scale.view(-1, 1), + scale_b=w_scale.view(1, -1), + out_dtype=torch.bfloat16, + use_fast_accum=True, + ) + out = out.view(*B.shape[:-1], A.shape[0]) + elif layout == "NN": # backward, dx=dy@wT + y_q = B._rowwise_data + y_scale = B._rowwise_scale_inv + w_q = A._columnwise_data + if w_q is None: + w_q = triton_pad_transpose( + A._rowwise_data, out=A._columnwise_data, multiple=32 + ) + if not online_transpose: + A._columnwise_data = w_q + w_scale = A._columnwise_scale_inv + out = torch._scaled_mm( + y_q, + w_q.t(), + scale_a=y_scale.view(-1, 1), + scale_b=w_scale.view(1, -1), + out_dtype=torch.bfloat16, + use_fast_accum=True, + ) + out = out.view(*B.shape[:-1], A.shape[1]) + elif layout == "NT": # update, dw=dyT@dx + y_q = B._columnwise_data + y_scale = B._columnwise_scale_inv + if A._columnwise_data is None: + x_q = triton_pad_transpose(A._rowwise_data, multiple=32) + else: + x_q = A._columnwise_data + x_scale = A._columnwise_scale_inv + o = torch._scaled_mm( + y_q, + x_q.t(), + scale_a=y_scale.view(-1, 1), + scale_b=x_scale.view(1, -1), + out_dtype=torch.bfloat16, + use_fast_accum=True, + ) + triton_inplace_add(out, o, accum=accumulate) + A._columnwise_data = None + else: + raise ValueError(f"layout {layout} is not supported") + return out + + +def smooth_groued_gemm( + A, B, out, m_splits, layout="TN", accumulate=True, online_transpose=False +): + s = 0 + if layout == "TN": # forward, y=x@w + for i, m in enumerate(m_splits): + if m == 0: + continue + + x_q = B[i]._rowwise_data + x_scale = B[i]._rowwise_scale_inv + w_q = A[i]._rowwise_data + w_scale = A[i]._rowwise_scale_inv + torch._scaled_mm( + x_q, + w_q.t(), + scale_a=x_scale.view(-1, 1), + scale_b=w_scale.view(1, -1), + out_dtype=torch.bfloat16, + use_fast_accum=True, + out=out[0][s : s + m], + ) + s += m + elif layout == "NN": # backward, dx=dy@wT + for i, m in enumerate(m_splits): + if m == 0: + continue + + y_q = B[i]._rowwise_data + y_scale = B[i]._rowwise_scale_inv + w_q = A[i]._columnwise_data + if w_q is None: + w_q = triton_pad_transpose( + A[i]._rowwise_data, out=A[i]._columnwise_data, multiple=32 + ) + if not online_transpose: + A[i]._columnwise_data = w_q + + w_scale = A[i]._columnwise_scale_inv + torch._scaled_mm( + y_q, + w_q.t(), + scale_a=y_scale.view(-1, 1), + scale_b=w_scale.view(1, -1), + out_dtype=torch.bfloat16, + use_fast_accum=True, + out=out[0][s : s + m], + ) + s += m + elif layout == "NT": # update, dw=dyT@dx + for i, m in enumerate(m_splits): + if m == 0: + continue + + y_q = B[i]._columnwise_data + y_scale = B[i]._columnwise_scale_inv + if A[i]._columnwise_data is None: + x_q = triton_pad_transpose(A[i]._rowwise_data, multiple=32) + else: + x_q = A[i]._columnwise_data + x_scale = A[i]._columnwise_scale_inv + # out is float32 + o = torch._scaled_mm( + y_q, + x_q.t(), + scale_a=y_scale.view(-1, 1), + scale_b=x_scale.view(1, -1), + out_dtype=torch.bfloat16, + use_fast_accum=True, + ) + + triton_inplace_add(out[i], o, accum=accumulate) + A[i]._columnwise_data = None + s += m + else: + raise ValueError(f"layout {layout} is not supported") + return out + + +class Fp32GEMM(torch.autograd.Function): + """""" + + @staticmethod + def forward(ctx, input: torch.Tensor, weight: torch.Tensor, impl: str): + shape = input.shape + if len(shape) == 3: + input = input.view(shape[0] * shape[1], shape[2]) + if impl == "native": + logits = triton_fp32_gemm(input, weight) + elif impl == "tma": + logits = triton_tma_persistent_matmul(input, weight) + elif impl == "split": + logits = triton_split_fp32_gemm(input, weight) + + ctx.input_requires_grad = input.requires_grad + ctx.weight_requires_grad = weight.requires_grad + ctx.shape = shape + ctx.impl = impl + ctx.save_for_backward(input, weight) + if len(shape) == 3: + logits = logits.view(shape[0], shape[1], weight.shape[0]) + return logits + + @staticmethod + def backward(ctx, grad_output): + grad_shape = grad_output.shape + if len(grad_shape) == 3: + grad_output = grad_output.view(grad_shape[0] * grad_shape[1], grad_shape[2]) + + input, weight = ctx.saved_tensors + + if ctx.impl == "split": + dx = triton_split_fp32_gemm_for_backward(grad_output, weight) + else: + dx = triton_fp32_gemm_for_backward(grad_output, weight) + if len(grad_shape) == 3: + dx = dx.view(*ctx.shape) + + if ctx.impl == "split": + dw = triton_split_fp32_gemm_for_update(grad_output, input) + else: + dw = triton_fp32_gemm_for_update(grad_output, input) + + return dx, dw, None + + +def fp32_gemm(input: torch.Tensor, weight: torch.Tensor, impl="native"): + """ + gemm with bf16/fp16 inputs and float32 output, + currently used in MoE router gemm. + Args: + input: bf16/fp16 activation tensor + weight: bf16/fp16 weight tensor + Returns: + output of gemm + """ + assert impl in ("native", "split", "tma") + assert input.dtype == weight.dtype, f"{input.dtype=} {weight.dtype=}" + return Fp32GEMM.apply(input, weight, impl) diff --git a/linghe/facade/hadamard_quant_linear.py b/linghe/facade/hadamard_quant_linear.py index 5b3dd45..4bc7414 100644 --- a/linghe/facade/hadamard_quant_linear.py +++ b/linghe/facade/hadamard_quant_linear.py @@ -11,15 +11,14 @@ from linghe.quant.hadamard import triton_hadamard_quant - class _HadamardQuantLinear(torch.autograd.Function): @staticmethod def forward( - ctx, - input: torch.Tensor, - weight: torch.Tensor, - bias: Optional[torch.Tensor], - hadamard_matrix: torch.Tensor + ctx, + input: torch.Tensor, + weight: torch.Tensor, + bias: Optional[torch.Tensor], + hadamard_matrix: torch.Tensor, ): ctx.input_requires_grad = input.requires_grad ctx.weight_requires_grad = weight.requires_grad @@ -32,13 +31,14 @@ def forward( x_q, x_scale, xt_q, xt_scale = triton_hadamard_quant(input, hadamard_matrix) w_q, w_scale, wt_q, wt_scale = triton_hadamard_quant(weight, hadamard_matrix) - output = torch._scaled_mm(x_q, - w_q.t(), - scale_a=x_scale.view(-1,1), - scale_b=w_scale.view(1,-1), - out_dtype=ctx.out_dtype, - use_fast_accum=True - ) + output = torch._scaled_mm( + x_q, + w_q.t(), + scale_a=x_scale.view(-1, 1), + scale_b=w_scale.view(1, -1), + out_dtype=ctx.out_dtype, + use_fast_accum=True, + ) if bias is not None: output += bias @@ -48,7 +48,11 @@ def forward( xt_scale if ctx.weight_requires_grad else None, wt_q if ctx.input_requires_grad else None, wt_scale if ctx.input_requires_grad else None, - hadamard_matrix if ctx.weight_requires_grad or ctx.weight_requires_grad else None + ( + hadamard_matrix + if ctx.weight_requires_grad or ctx.weight_requires_grad + else None + ), ] ctx.save_for_backward(*saved_tensors) @@ -57,32 +61,36 @@ def forward( @staticmethod def backward( - ctx, - output_grad: torch.Tensor, + ctx, + output_grad: torch.Tensor, ): xt_q, xt_scale, wt_q, wt_scale, hadamard_matrix = ctx.saved_tensors output_grad = output_grad.view(-1, output_grad.shape[-1]) - y_q, y_scale, yt_q, yt_scale = triton_hadamard_quant(output_grad, hadamard_matrix) + y_q, y_scale, yt_q, yt_scale = triton_hadamard_quant( + output_grad, hadamard_matrix + ) - dx = torch._scaled_mm(y_q, - wt_q.t(), - scale_a=y_scale.view(-1,1), - scale_b=wt_scale.view(1,-1), - out_dtype=ctx.out_dtype, - use_fast_accum=True - ) + dx = torch._scaled_mm( + y_q, + wt_q.t(), + scale_a=y_scale.view(-1, 1), + scale_b=wt_scale.view(1, -1), + out_dtype=ctx.out_dtype, + use_fast_accum=True, + ) dx = dx.view(ctx.input_shape) - dw = torch._scaled_mm(yt_q, - xt_q.t(), - scale_a=yt_scale.view(-1,1), - scale_b=xt_scale.view(1,-1), - out_dtype=ctx.out_dtype, - use_fast_accum=True - ) + dw = torch._scaled_mm( + yt_q, + xt_q.t(), + scale_a=yt_scale.view(-1, 1), + scale_b=xt_scale.view(1, -1), + out_dtype=ctx.out_dtype, + use_fast_accum=True, + ) db = None if ctx.bias_requires_grad: @@ -95,13 +103,14 @@ class HadamardQuantLinear(torch.nn.Module): """ a naive implementation of hadamard transformation and quantization """ + def __init__( - self, - in_features: int, - out_features: int, - bias: bool = True, - device=None, - dtype=None + self, + in_features: int, + out_features: int, + bias: bool = True, + device=None, + dtype=None, ): """ Args: @@ -115,19 +124,18 @@ def __init__( self.in_features = in_features self.out_features = out_features self.weight = torch.nn.parameter.Parameter( - torch.empty((out_features, in_features), device=device, - dtype=dtype)) + torch.empty((out_features, in_features), device=device, dtype=dtype) + ) if bias: self.bias = torch.nn.parameter.Parameter( - torch.empty(out_features, device=device, dtype=dtype)) + torch.empty(out_features, device=device, dtype=dtype) + ) else: self.bias = None - size = 32 if 'H20' in torch.cuda.get_device_properties(0).name else 64 - data = self._hadamard_matrix(size, device=device, dtype=dtype, - norm=True) - self.hadamard_matrix = torch.nn.parameter.Parameter(data, - requires_grad=False) + size = 32 if "H20" in torch.cuda.get_device_properties(0).name else 64 + data = self._hadamard_matrix(size, device=device, dtype=dtype, norm=True) + self.hadamard_matrix = torch.nn.parameter.Parameter(data, requires_grad=False) self.reset_parameters() def _hadamard_matrix(self, size, device=None, dtype=None, norm=False): @@ -137,7 +145,7 @@ def _hadamard_matrix(self, size, device=None, dtype=None, norm=False): for _ in range(int(math.log2(size)) - 1): m = torch.kron(m, m2) if norm: - m = m / size ** 0.5 + m = m / size**0.5 if dtype is not None: m = m.to(dtype) return m @@ -145,8 +153,9 @@ def _hadamard_matrix(self, size, device=None, dtype=None, norm=False): def forward(self, input: torch.Tensor) -> torch.Tensor: """""" if self.training: - return _HadamardQuantLinear.apply(input, self.weight, self.bias, - self.hadamard_matrix) + return _HadamardQuantLinear.apply( + input, self.weight, self.bias, self.hadamard_matrix + ) else: output = input @ self.weight.t() if self.bias is not None: diff --git a/linghe/facade/linear.py b/linghe/facade/linear.py new file mode 100644 index 0000000..403aa98 --- /dev/null +++ b/linghe/facade/linear.py @@ -0,0 +1,343 @@ +# -*- coding: utf-8 -*- +""" +Copyright (c) Ant Financial Service Group and its affiliates. +""" + +import math +from typing import Optional + +import torch + +from linghe.quant.hadamard import triton_hadamard_quant +from linghe.quant.smooth import triton_smooth_quant, triton_transpose_smooth_quant +from linghe.utils.reduce import triton_abs_max +from linghe.utils.transpose import triton_pad_transpose + + +class _HadamardQuantLinear(torch.autograd.Function): + @staticmethod + def forward( + ctx, + input: torch.Tensor, + weight: torch.Tensor, + bias: Optional[torch.Tensor], + hadamard_matrix: torch.Tensor, + ): + ctx.input_requires_grad = input.requires_grad + ctx.weight_requires_grad = weight.requires_grad + ctx.bias_requires_grad = bias is not None and bias.requires_grad + + ctx.out_dtype = input.dtype + ctx.input_shape = input.shape + input = input.view(-1, input.shape[-1]) + + x_q, x_scale, xt_q, xt_scale = triton_hadamard_quant(input, hadamard_matrix) + w_q, w_scale, wt_q, wt_scale = triton_hadamard_quant(weight, hadamard_matrix) + + output = torch._scaled_mm( + x_q, + w_q.t(), + scale_a=x_scale.view(-1, 1), + scale_b=w_scale.view(1, -1), + out_dtype=ctx.out_dtype, + use_fast_accum=True, + ) + + if bias is not None: + output += bias + + saved_tensors = [ + xt_q if ctx.weight_requires_grad else None, + xt_scale if ctx.weight_requires_grad else None, + wt_q if ctx.input_requires_grad else None, + wt_scale if ctx.input_requires_grad else None, + ( + hadamard_matrix + if ctx.weight_requires_grad or ctx.weight_requires_grad + else None + ), + ] + + ctx.save_for_backward(*saved_tensors) + out_shape = (*ctx.input_shape[0:-1], -1) + return output.view(out_shape) + + @staticmethod + def backward( + ctx, + output_grad: torch.Tensor, + ): + xt_q, xt_scale, wt_q, wt_scale, hadamard_matrix = ctx.saved_tensors + + output_grad = output_grad.view(-1, output_grad.shape[-1]) + + y_q, y_scale, yt_q, yt_scale = triton_hadamard_quant( + output_grad, hadamard_matrix + ) + + dx = torch._scaled_mm( + y_q, + wt_q.t(), + scale_a=y_scale.view(-1, 1), + scale_b=wt_scale.view(1, -1), + out_dtype=ctx.out_dtype, + use_fast_accum=True, + ) + + dx = dx.view(ctx.input_shape) + + dw = torch._scaled_mm( + yt_q, + xt_q.t(), + scale_a=yt_scale.view(-1, 1), + scale_b=xt_scale.view(1, -1), + out_dtype=ctx.out_dtype, + use_fast_accum=True, + ) + + db = None + if ctx.bias_requires_grad: + db = torch.sum(output_grad, dim=0) + + return dx, dw, db, None + + +class HadamardQuantLinear(torch.nn.Module): + """ + a naive implementation of hadamard transformation and quantization + """ + + def __init__( + self, + in_features: int, + out_features: int, + bias: bool = True, + device=None, + dtype=None, + ): + """ + Args: + in_features: in feature number + out_features: out feature number + bias: whether use bias + device: weight device + dtype: weight dtype + """ + super().__init__() + self.in_features = in_features + self.out_features = out_features + self.weight = torch.nn.parameter.Parameter( + torch.empty((out_features, in_features), device=device, dtype=dtype) + ) + if bias: + self.bias = torch.nn.parameter.Parameter( + torch.empty(out_features, device=device, dtype=dtype) + ) + else: + self.bias = None + + size = 32 if "H20" in torch.cuda.get_device_properties(0).name else 64 + data = self._hadamard_matrix(size, device=device, dtype=dtype, norm=True) + self.hadamard_matrix = torch.nn.parameter.Parameter(data, requires_grad=False) + self.reset_parameters() + + def _hadamard_matrix(self, size, device=None, dtype=None, norm=False): + assert 2 ** int(math.log2(size)) == size + m2 = torch.tensor([[1, 1], [1, -1]], device=device, dtype=torch.float32) + m = m2 + for _ in range(int(math.log2(size)) - 1): + m = torch.kron(m, m2) + if norm: + m = m / size**0.5 + if dtype is not None: + m = m.to(dtype) + return m + + def forward(self, input: torch.Tensor) -> torch.Tensor: + """""" + if self.training: + return _HadamardQuantLinear.apply( + input, self.weight, self.bias, self.hadamard_matrix + ) + else: + output = input @ self.weight.t() + if self.bias is not None: + output = output + self.bias + return output + + def extra_repr(self) -> str: + """""" + return f"in_features={self.in_features}, out_features={self.out_features}, bias={self.bias is not None}" + + def reset_parameters(self): + """""" + self.weight.data.normal_(mean=0.0, std=0.02) + if self.bias is not None: + self.bias.data.zero_() + + +class _SmoothQuantLinear(torch.autograd.Function): + @staticmethod + def forward( + ctx, + input: torch.Tensor, + weight: torch.Tensor, + bias: Optional[torch.Tensor], + smooth_scale: torch.Tensor, + ): + ctx.input_requires_grad = input.requires_grad + ctx.weight_requires_grad = weight.requires_grad + ctx.bias_requires_grad = bias is not None and bias.requires_grad + ctx.out_dtype = input.dtype + ctx.input_shape = input.shape + + round_scale = True + ctx.round_scale = round_scale + + input = input.view(-1, input.shape[-1]) + + x_q, x_scale = triton_smooth_quant( + input, 1 / smooth_scale, round_scale=round_scale + ) + w_q, w_scale = triton_smooth_quant( + weight, smooth_scale, round_scale=round_scale + ) + + output = torch._scaled_mm( + x_q, + w_q.t(), + scale_a=x_scale.view(-1, 1), + scale_b=w_scale.view(1, -1), + out_dtype=ctx.out_dtype, + use_fast_accum=True, + ) + + if bias is not None: + output += bias + + saved_tensors = [ + x_q if ctx.weight_requires_grad else None, + x_scale if ctx.weight_requires_grad else None, + w_q if ctx.input_requires_grad else None, + w_scale if ctx.input_requires_grad else None, + ( + smooth_scale + if ctx.weight_requires_grad or ctx.weight_requires_grad + else None + ), + ] + + ctx.save_for_backward(*saved_tensors) + out_shape = (*ctx.input_shape[0:-1], -1) + return output.view(out_shape) + + @staticmethod + def backward(ctx, output_grad: torch.Tensor): + + x_q, x_s, w_q, w_s, smooth_scale = ctx.saved_tensors + + output_grad = output_grad.view(-1, output_grad.shape[-1]) + round_scale = ctx.round_scale + y_q, y_scale = triton_smooth_quant( + output_grad, w_s, reverse=True, round_scale=round_scale + ) + + wt_q = triton_pad_transpose(w_q, multiple=32) + dx = torch._scaled_mm( + y_q, + wt_q.t(), + scale_a=y_scale.view(-1, 1), + scale_b=smooth_scale.view(1, -1), + out_dtype=ctx.out_dtype, + use_fast_accum=True, + ) + + yt_q, yt_scale = triton_transpose_smooth_quant( + output_grad, x_s, reverse=True, round_scale=round_scale + ) + + xt_q = triton_pad_transpose(x_q, multiple=32) + dw = torch._scaled_mm( + yt_q, + xt_q.t(), + scale_a=yt_scale.view(-1, 1), + scale_b=1.0 / smooth_scale.view(1, -1), + out_dtype=ctx.out_dtype, + use_fast_accum=True, + ) + + db = None + if ctx.bias_requires_grad: + db = torch.sum(output_grad, dim=0) + + return dx, dw, db, None + + +class SmoothQuantLinear(torch.nn.Module): + """ + a naive implementation of smooth quantization linear + """ + + def __init__( + self, + in_features: int, + out_features: int, + bias: bool = True, + device=None, + dtype=None, + ): + """ + Args: + in_features: in feature number + out_features: out feature number + bias: whether use bias + device: weight device + dtype: weight dtype + """ + super().__init__() + self.in_features = in_features + self.out_features = out_features + self.weight = torch.nn.parameter.Parameter( + torch.empty((out_features, in_features), device=device, dtype=dtype) + ) + if bias: + self.bias = torch.nn.parameter.Parameter( + torch.empty(out_features, device=device, dtype=dtype) + ) + else: + self.bias = None + + self.gap_step = 16 + self.smooth_scale = None + self.smooth_update_step = 0 + + self.reset_parameters() + + def forward(self, input: torch.Tensor) -> torch.Tensor: + """""" + if self.training: + + if self.smooth_update_step % self.gap_step == 0: + input_maxs = triton_abs_max(input) + weight_maxs = triton_abs_max(self.weight) + self.smooth_scale = torch.sqrt(input_maxs * weight_maxs) + + output = _SmoothQuantLinear.apply( + input, self.weight, self.bias, self.smooth_scale + ) + self.smooth_update_step += 1 + else: + output = input @ self.weight.t() + if self.bias is not None: + output = output + self.bias + return output + + def extra_repr(self) -> str: + """""" + return f"in_features={self.in_features}, out_features={self.out_features}, bias={self.bias is not None}" + + def reset_parameters(self): + """""" + self.weight.data.normal_(mean=0.0, std=0.02) + if self.bias is not None: + self.bias.data.zero_() diff --git a/linghe/facade/loss.py b/linghe/facade/loss.py index 0feac8a..b588871 100644 --- a/linghe/facade/loss.py +++ b/linghe/facade/loss.py @@ -5,22 +5,38 @@ import torch -from linghe.utils.loss import triton_softmax_cross_entropy_forward, \ - triton_softmax_cross_entropy_backward +from linghe.utils.loss import ( + triton_softmax_cross_entropy_forward, + triton_softmax_cross_entropy_backward, + triton_parallel_softmax_cross_entropy_forward, + triton_parallel_softmax_cross_entropy_backward, + triton_moe_z_loss_forward, + triton_moe_z_loss_backward, +) class SoftmaxCrossEntropyFunction(torch.autograd.Function): """""" + @staticmethod - def forward(ctx, logits, labels, inplace=False): + def forward(ctx, logits, labels, ignore_index=-100, inplace=False, tp_group=None): shape = logits.shape - if len(shape) == 3: - logits = logits.view(-1, shape[-1]) - loss, sum_exp, max_logit = triton_softmax_cross_entropy_forward(logits, - labels) + logits_view = logits.view(-1, shape[-1]) if len(shape) == 3 else logits + parallel = tp_group is not None and tp_group.size() > 1 + if parallel: + loss, sum_exp, max_logit = triton_parallel_softmax_cross_entropy_forward( + logits_view, labels, tp_group, ignore_index=ignore_index + ) + else: + loss, sum_exp, max_logit = triton_softmax_cross_entropy_forward( + logits_view, labels, ignore_index=ignore_index + ) ctx.save_for_backward(logits, labels, sum_exp, max_logit) + ctx.ignore_index = ignore_index ctx.inplace = inplace ctx.shape = shape + ctx.tp_group = tp_group + ctx.parallel = parallel if len(shape) == 3: loss = loss.view(shape[0], shape[1]) return loss @@ -29,17 +45,43 @@ def forward(ctx, logits, labels, inplace=False): def backward(ctx, grad_output): logits, labels, sum_exp, max_logit = ctx.saved_tensors shape = ctx.shape - grad = logits if ctx.inplace else None - grad = triton_softmax_cross_entropy_backward(logits, labels, sum_exp, - max_logit, - grad_output, - output_grad=grad) + if len(shape) == 3: + logits = logits.view(-1, shape[-1]) + grad_output = torch.reshape(grad_output, (-1,)) + if ctx.parallel: + grad = triton_parallel_softmax_cross_entropy_backward( + logits, + labels, + sum_exp, + max_logit, + grad_output, + ctx.tp_group, + ignore_index=ctx.ignore_index, + inplace=ctx.inplace, + ) + else: + + grad = triton_softmax_cross_entropy_backward( + logits, + labels, + sum_exp, + max_logit, + grad_output, + ignore_index=ctx.ignore_index, + inplace=ctx.inplace, + ) if len(shape) == 3: grad = grad.view(shape) - return grad, None, None, None + return grad, None, None, None, None -def softmax_cross_entropy(logits: torch.Tensor, labels: torch.Tensor, inplace: bool = False): +def softmax_cross_entropy( + logits: torch.Tensor, + labels: torch.Tensor, + ignore_index: int = -100, + inplace: bool = False, + tp_group=None, +): """ softmax cross entropy Args: @@ -51,22 +93,62 @@ def softmax_cross_entropy(logits: torch.Tensor, labels: torch.Tensor, inplace: b """ assert logits.is_contiguous() assert labels.is_contiguous() - return SoftmaxCrossEntropyFunction.apply(logits, labels, inplace) + return SoftmaxCrossEntropyFunction.apply( + logits, labels, ignore_index, inplace, tp_group + ) -class GradScalingFunction(torch.autograd.Function): +class GradientScalingFunction(torch.autograd.Function): """""" + @staticmethod - def forward(ctx, x, coef=0.2): + def forward(ctx, x, coef=1.0): ctx.coef = coef return x @staticmethod def backward(ctx, grad_output): - shape = grad_output.shape - assert len(shape) == 2 - bs, length = grad_output.shape - array = length - torch.arange(0, length, device=grad_output.device) - scale = 1 / torch.pow(array.float(), ctx.coef) - grad = grad_output * scale + grad = grad_output * ctx.coef + return grad, None + + +def gradient_scaling(x: torch.Tensor, coef: float = 1.0): + """ + scale gradient + Args: + x: input tensor + coef: scale coefficient + """ + return GradientScalingFunction.apply(x, coef) + + +class MoeZLossFunction(torch.autograd.Function): + """""" + + @staticmethod + def forward(ctx, logits, coef): + loss = triton_moe_z_loss_forward(logits, coef=coef) + ctx.save_for_backward( + logits, + ) + ctx.coef = coef + return loss + + @staticmethod + def backward(ctx, grad_output): + (logits,) = ctx.saved_tensors + grad = triton_moe_z_loss_backward(grad_output, logits, coef=ctx.coef) return grad, None + + +def moe_z_loss(logits: torch.Tensor, coef: float = 1e-3): + """ + softmax cross entropy + Args: + logits: logits tensor, shape [...,dim] + coef: z loss coef + Returns: + z loss + """ + assert logits.is_contiguous() + return MoeZLossFunction.apply(logits, coef) diff --git a/linghe/facade/mla.py b/linghe/facade/mla.py new file mode 100644 index 0000000..fc440a1 --- /dev/null +++ b/linghe/facade/mla.py @@ -0,0 +1,127 @@ +# -*- coding: utf-8 -*- +""" +Copyright (c) Ant Financial Service Group and its affiliates. +""" + +from typing import Optional + +import torch + +from linghe.attn.mla import ( + triton_mla_forward, + triton_mla_backward, + triton_varlen_mla_forward, + triton_varlen_mla_backward, +) + + +class MultiLatentAttention(torch.autograd.Function): + """""" + + @staticmethod + def forward( + ctx, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + cu_seqlens: Optional[torch.Tensor] = None, + padded_cu_seqlens: Optional[torch.Tensor] = None, + max_q_length: Optional[int] = None, + causal: bool = True, + safe: bool = True, + clip_value: Optional[float] = None, + ): + ctx.cu_seqlens = cu_seqlens + ctx.padded_cu_seqlens = padded_cu_seqlens + ctx.max_q_length = max_q_length + ctx.causal = causal + ctx.safe = safe + ctx.clip_value = clip_value + VARLEN = cu_seqlens is not None + ctx.VARLEN = VARLEN + if VARLEN: + output, lse, max_logits = triton_varlen_mla_forward( + q, + k, + v, + cu_seqlens, + padded_cu_seqlens=None, + max_q_length=max_q_length, + causal=causal, + safe=safe, + clip_value=clip_value, + ) + else: + output, lse, max_logits = triton_mla_forward( + q, k, v, causal=causal, safe=safe, clip_value=clip_value + ) + ctx.save_for_backward(q, k, v, output, lse, max_logits) + return output + + @staticmethod + def backward(ctx, grad_output): + q, k, v, output, lse, max_logits = ctx.saved_tensors + if ctx.VARLEN: + dq, dk, dv = triton_varlen_mla_backward( + grad_output, + output, + q, + k, + v, + lse, + max_logits, + ctx.cu_seqlens, + ctx.max_q_length, + padded_cu_seqlens=ctx.padded_cu_seqlens, + causal=ctx.causal, + safe=ctx.safe, + clip_value=ctx.clip_value, + ) + else: + dq, dk, dv = triton_mla_backward( + grad_output, + output, + q, + k, + v, + lse, + max_logits, + causal=ctx.causal, + safe=ctx.safe, + clip_value=ctx.clip_value, + ) + return ( + dq, + dk, + dv, + None, + None, + None, + None, + None, + None, + ) + + +def multi_latend_attention( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + cu_seqlens: Optional[torch.Tensor] = None, + padded_cu_seqlens: Optional[torch.Tensor] = None, + max_q_length: Optional[int] = None, + causal: bool = True, + safe: bool = True, + clip_value: float = 0.0, +): + """ + inplace add y to x with mix precise + Args: + x: to be updated + y: add to x + Returns: + updated x tensor + """ + return MultiLatentAttention.apply( + q, k, v, cu_seqlens, padded_cu_seqlens, max_q_length, causal, safe, clip_value + ) diff --git a/linghe/facade/norm.py b/linghe/facade/norm.py index 38d7fd3..58fdf40 100644 --- a/linghe/facade/norm.py +++ b/linghe/facade/norm.py @@ -5,36 +5,29 @@ import torch -from linghe.utils.norm import triton_rms_norm_forward, triton_rms_norm_backward, \ - triton_group_rms_norm_gate_forward, triton_group_rms_norm_gate_backward +from linghe.utils.norm import ( + triton_rms_norm_forward, + triton_rms_norm_backward, + triton_rms_norm_and_block_quant_forward, + triton_rms_norm_and_mxfp8_quant_forward, + triton_rms_norm_and_smooth_quant_forward, +) class RMSNormFunction(torch.autograd.Function): """""" + @staticmethod def forward(ctx, x, weight, eps=1e-6): - output = triton_rms_norm_forward( - x, - weight, - eps - ) - # ctx.save_for_backward(x, weight, norm) + output, rms = triton_rms_norm_forward(x, weight, eps) ctx.save_for_backward(x, weight) ctx.eps = eps - return output @staticmethod def backward(ctx, dy): x, weight = ctx.saved_tensors - - dx, dw = triton_rms_norm_backward( - dy, - x, - weight, - ctx.eps - ) - + dx, dw = triton_rms_norm_backward(dy, x, weight, ctx.eps) return dx, dw, None @@ -49,58 +42,185 @@ def rms_norm(x: torch.Tensor, weight: torch.Tensor, eps: float = 1e-6): Returns: rms output """ - assert x.contiguous() - assert weight.contiguous() + assert x.is_contiguous() + assert weight.is_contiguous() return RMSNormFunction.apply(x, weight, eps) -class GroupRMSNormGateFunction(torch.autograd.Function): - """""" + +# used in attention rms norm +class BlockRMSNorm(torch.autograd.Function): @staticmethod - def forward(ctx, attn_output, gate, weight, eps=1e-6, group_size=4): - output = triton_group_rms_norm_gate_forward( - attn_output, - gate, - weight.data, - eps=eps, - group_size=group_size + def forward(ctx, input, weight, rms, eps, quantizer, cls, is_recomputing): + shape = input.shape + assert len(shape) == 3 + + if is_recomputing is None: + output_mode = 2 + elif is_recomputing: + output_mode = 1 + else: + output_mode = 0 + + input_view = input.view(shape[0] * shape[1], shape[2]) + x_q, x_scale, output_rms, xt_q, xt_scale = ( + triton_rms_norm_and_block_quant_forward( + input_view, + weight, + rms=rms, + eps=eps, + round_scale=quantizer.force_pow_2_scales, + output_mode=output_mode, + ) ) - ctx.save_for_backward(attn_output, gate, weight.data) + + transpose_shape = (shape[2], shape[0], shape[1]) + output = cls( + shape=shape, + dtype=input.dtype, + fp8_dtype=quantizer.dtype, + rowwise_data=x_q.view(shape) if x_q is not None else None, + rowwise_scale_inv=x_scale, + columnwise_data=xt_q.view(transpose_shape) if xt_q is not None else None, + columnwise_scale_inv=xt_scale, + quantizer=quantizer, + requires_grad=input.requires_grad, + is_2D_scaled=False, + ) + ctx.input_requires_grad = input.requires_grad + ctx.weight_requires_grad = weight.requires_grad + ctx.shape = shape ctx.eps = eps - ctx.group_size = group_size + ctx.save_for_backward(input, weight) - return output + return output, output_rms @staticmethod - def backward(ctx, dy): - attn_output, gate, weight = ctx.saved_tensors - - dx, dg, dw = triton_group_rms_norm_gate_backward( - dy, - attn_output, - gate, - weight, - ctx.eps, - ctx.group_size + def backward(ctx, grad_output, grad_rms): + shape = grad_output.shape + grad_output = grad_output.view(shape[0] * shape[1], shape[2]) + input, weight = ctx.saved_tensors + input = input.view(shape[0] * shape[1], shape[2]) + dx, dw = triton_rms_norm_backward(grad_output, input, weight, eps=ctx.eps) + dx = dx.view(*shape) + + return dx, dw, None, None, None, None, None + + +def block_rms_norm(input, weight, rms, quantizer, cls, eps=1e-6, is_recomputing=None): + output, output_rms = BlockRMSNorm.apply( + input, weight, rms, eps, quantizer, cls, is_recomputing + ) + output_rms = output_rms.detach() + return output, output_rms + + +class MXFP8RMSNorm(torch.autograd.Function): + @staticmethod + def forward(ctx, input, weight, rms, eps, quantizer, cls, is_recomputing): + shape = input.shape + assert len(shape) == 3 + input = input.view(shape[0] * shape[1], shape[2]) + rms = None + if is_recomputing is None: + output_mode = 2 + elif is_recomputing: + output_mode = 1 + else: + output_mode = 0 + x_q, x_scale, output_rms, xt_q, xt_scale = ( + triton_rms_norm_and_mxfp8_quant_forward( + input, weight.data, rms=rms, eps=eps, output_mode=output_mode + ) ) - return dx, dg, dw, None, None + # transpose_shape = (shape[2], shape[0], shape[1]) + output = cls( + shape=shape, + dtype=input.dtype, + fp8_dtype=quantizer.dtype, + rowwise_data=x_q.view(shape) if x_q is not None else None, + rowwise_scale_inv=x_scale, + columnwise_data=xt_q.view(shape) if xt_q is not None else None, + columnwise_scale_inv=xt_scale, + quantizer=quantizer, + requires_grad=input.requires_grad, + ) + ctx.input_requires_grad = input.requires_grad + ctx.weight_requires_grad = weight.requires_grad + ctx.shape = shape + ctx.eps = eps + ctx.save_for_backward(input, weight) + return output, output_rms -def group_rms_norm_gate(attn_output: torch.Tensor, - gate: torch.Tensor, - weight: torch.Tensor, - eps: float = 1e-6, - group_size: int = 4): - """ - return group_rms_norm(transpose(attn_output, [0,1]), weight) * sigmoid(gate) - Args: - attn_output: output of core attn, shape [bs, length, n_heads, head_dim] - gate: gate tensor for attention output, shape [length, bs, dim] - weight: weight of RMS norm, shape [dim] - eps: epsilon for RMS - group_size: group size of group RMS norm - Returns: - output with shape [length, bs, dim] - """ - return GroupRMSNormGateFunction.apply(attn_output, gate, weight, eps, group_size) \ No newline at end of file + @staticmethod + def backward(ctx, grad_output, grad_rms): + shape = grad_output.shape + grad_output = grad_output.view(shape[0] * shape[1], shape[2]) + input, weight = ctx.saved_tensors + dx, dw = triton_rms_norm_backward(grad_output, input, weight, eps=ctx.eps) + dx = dx.view(*shape) + + return dx, dw, None, None, None, None, None + + +def mxfp8_rms_norm(input, weight, rms, quantizer, cls, eps=1e-6, is_recomputing=None): + # input: [length,bs,dim] + output, output_rms = MXFP8RMSNorm.apply( + input, weight, rms, eps, quantizer, cls, is_recomputing + ) + output_rms = output_rms.detach() + return output, output_rms + + +# used in attention rms norm +class SmoothRMSNorm(torch.autograd.Function): + @staticmethod + def forward(ctx, input, weight, quantizer, cls, eps, is_first_microbatch): + shape = input.shape + assert len(shape) == 3 + input = input.view(shape[0] * shape[1], shape[2]) + ctx.input_requires_grad = input.requires_grad + ctx.weight_requires_grad = weight.requires_grad + ctx.shape = shape + ctx.eps = eps + ctx.save_for_backward(input, weight) + x_q, x_scale, x_maxs, rms = triton_rms_norm_and_smooth_quant_forward( + input, + weight.data, + smooth_scale=quantizer.smooth_scale, + eps=eps, + calibrate=is_first_microbatch, + output_rms=False, + round_scale=quantizer.force_pow_2_scales, + ) + output = cls( + shape=(shape[0], shape[1], shape[2]), + dtype=input.dtype, + fp8_dtype=quantizer.dtype, + rowwise_data=x_q, + rowwise_scale_inv=x_scale, + columnwise_data=None, + columnwise_scale_inv=quantizer.smooth_scale, + quantizer=quantizer, + requires_grad=input.requires_grad, + ) + return output, x_maxs + + @staticmethod + def backward(ctx, grad_output, grad_max): + shape = grad_output.shape + grad_output = grad_output.view(shape[0] * shape[1], shape[2]) + input, weight = ctx.saved_tensors + dx, dw = triton_rms_norm_backward(grad_output, input, weight, eps=ctx.eps) + dx = dx.view(*shape) + return dx, dw, None, None, None, None + + +def smooth_rms_norm(input, weight, quantizer, cls, eps=1e-6, is_first_microbatch=False): + # input: [length,bs,dim] + output, x_maxs = SmoothRMSNorm.apply( + input, weight, quantizer, cls, eps, is_first_microbatch + ) + return output, x_maxs diff --git a/linghe/facade/permutation.py b/linghe/facade/permutation.py new file mode 100644 index 0000000..27cdadb --- /dev/null +++ b/linghe/facade/permutation.py @@ -0,0 +1,1220 @@ +# -*- coding: utf-8 -*- +""" +Copyright (c) Ant Financial Service Group and its affiliates. +""" + +from typing import Optional, List + +import torch + +from linghe.quant.mxfp8 import triton_batch_mxfp8_quant +from linghe.utils.gather import ( + triton_permute_with_mask_map, + triton_make_row_id_map, + triton_make_row_id_map_and_index, + triton_batch_block_pad_permute_with_indices, + triton_batch_mxfp8_permute_with_indices, + triton_batch_smooth_permute_with_indices, + triton_batch_transpose_smooth_permute_with_indices, + triton_batch_smooth_fused_permute_with_indices, + triton_batch_transpose_smooth_fused_permute_with_indices, + triton_make_chunk_sort_map, +) +from linghe.utils.scatter import ( + triton_unpermute_with_mask_map, + triton_unpermute_with_reverse_map, +) + + +class _PaddedPermute(torch.autograd.Function): + @staticmethod + def forward( + ctx, + tokens, + probs, + routing_map, + tokens_per_expert_cuda_tensor, + tokens_per_expert_list, + multiple, + ): + """Forward function.""" + num_tokens, hidden_dim = tokens.shape + + row_id_map = triton_make_row_id_map(routing_map, multiple_of=multiple) + num_out_tokens = sum( + [((x - 1) // multiple + 1) * multiple for x in tokens_per_expert_list] + ) + + ctx.num_tokens = num_tokens + ctx.hidden_dim = hidden_dim + ctx.prob_shape = probs.shape + ctx.shape = tokens.shape + ctx.row_id_map = row_id_map + permuted_tokens, _, permuted_probs = triton_permute_with_mask_map( + tokens, + None, + probs, + row_id_map, + num_out_tokens, + contiguous=False, + tokens_per_expert=tokens_per_expert_cuda_tensor, + ) + ctx.save_for_backward(row_id_map) + return permuted_tokens, permuted_probs, row_id_map + + @staticmethod + def backward(ctx, grad_output, grad_prob, grad_map): + """Backward function.""" + (row_id_map,) = ctx.saved_tensors + output, prob_output = triton_unpermute_with_mask_map( + grad_output, row_id_map, grad_prob + ) + return ( + output.view(ctx.shape), + prob_output.view(ctx.prob_shape), + None, + None, + None, + None, + ) + + +def padded_permute( + tokens, + routing_map, + tokens_per_expert_cuda_tensor, + tokens_per_expert_list, + probs: Optional[torch.Tensor] = None, + multiple: int = 32, +): + """Permute the tokens and probs based on the mask. + Tokens with the same designated expert will be grouped together. + The shape of mask is [tokens, num_experts], it indicates which experts were selected + by each token. + When drop_and_pad=True, in routing_map, the number of non-zeros in each column equals to + expert capacity. This function exploits this feature to use ops that support cuda graph. + Args: + tokens (torch.Tensor): The input token tensor, [num_tokens, hidden]. + routing_map (torch.Tensor): The sparse token to expert mapping, [num_tokens, num_experts]. + tokens_per_expert (torch.Tensor): cpu tensor + """ + + permuted_input, permuted_probs, row_id_map = _PaddedPermute.apply( + tokens, + probs, + routing_map, + tokens_per_expert_cuda_tensor, + tokens_per_expert_list, + multiple, + ) + return permuted_input, permuted_probs, row_id_map + + +class _PaddedUnpermute(torch.autograd.Function): + @staticmethod + def forward(ctx, permuted_tokens, row_id_map, tokens_per_expert, restore_shape): + """Forward function.""" + num_tokens, hidden_size = restore_shape + num_out_tokens = permuted_tokens.shape[0] + n_experts = row_id_map.size(1) + ctx.save_for_backward(row_id_map) + ctx.input_requires_grad = permuted_tokens.requires_grad + ctx.num_experts = n_experts + ctx.restore_shape = restore_shape + ctx.num_tokens = num_tokens + ctx.num_out_tokens = num_out_tokens + ctx.hidden_size = hidden_size + ctx.tokens_per_expert = tokens_per_expert + + output, _ = triton_unpermute_with_mask_map(permuted_tokens, row_id_map, None) + return output + + @staticmethod + def backward(ctx, grad_output): + """Backward function.""" + (row_id_map,) = ctx.saved_tensors + permuted_tokens, _, _ = triton_permute_with_mask_map( + grad_output, + None, + None, + row_id_map, + ctx.num_out_tokens, + contiguous=False, + tokens_per_expert=ctx.tokens_per_expert, + ) + + return permuted_tokens, None, None, None + + +def padded_unpermute( + permuted_tokens: torch.Tensor, + row_id_map: torch.Tensor, + tokens_per_expert: torch.Tensor, + restore_shape: torch.Size, +): + output = _PaddedUnpermute.apply( + permuted_tokens, row_id_map, tokens_per_expert, restore_shape + ) + return output + + +class _BlockPaddedPermute(torch.autograd.Function): + @staticmethod + def forward( + ctx, + tokens, + probs, + routing_map, + tokens_per_expert_cuda_tensor, + tokens_per_expert_list, + quantizers, + cls, + ): + """Forward function.""" + num_tokens, hidden_dim = tokens.shape + + num_out_tokens = sum([(x + 15) // 16 * 16 for x in tokens_per_expert_list]) + row_id_map, row_id_index = triton_make_row_id_map_and_index( + routing_map, num_out_tokens, multiple_of=16 + ) + + ctx.num_tokens = num_tokens + ctx.hidden_dim = hidden_dim + ctx.prob_shape = probs.shape + ctx.shape = tokens.shape + ctx.cls = cls + x_q, x_scale, xt_q, xt_scale, permuted_probs = ( + triton_batch_block_pad_permute_with_indices( + tokens, + tokens_per_expert_cuda_tensor, + row_id_index, + tokens_per_expert_list, + probs=probs, + round_scale=quantizers[0].force_pow_2_scales, + ) + ) + + output = cls( + shape=x_q.shape, + dtype=tokens.dtype, + fp8_dtype=quantizers[0].dtype, + rowwise_data=x_q, + rowwise_scale_inv=x_scale, + columnwise_data=xt_q, + columnwise_scale_inv=xt_scale, + quantizer=quantizers, + requires_grad=tokens.requires_grad, + is_2D_scaled=False, + ) + ctx.save_for_backward(row_id_map) + return output, permuted_probs, row_id_map, row_id_index + + @staticmethod + def backward(ctx, grad_output, grad_prob, grad_map, grad_index): + """Backward function.""" + (row_id_map,) = ctx.saved_tensors + output, prob_output = triton_unpermute_with_mask_map( + grad_output, row_id_map, grad_prob + ) + return ( + output.view(ctx.shape), + prob_output.view(ctx.prob_shape), + None, + None, + None, + None, + None, + ) + + +def block_padded_permute( + tokens, + routing_map, + tokens_per_expert_cuda_tensor, + tokens_per_expert_list, + quantizers, + cls, + probs: Optional[torch.Tensor] = None, +): + """Permute the tokens and probs based on the mask. + Tokens with the same designated expert will be grouped together. + The shape of mask is [tokens, num_experts], it indicates which experts were selected + by each token. + When drop_and_pad=True, in routing_map, the number of non-zeros in each column equals to + expert capacity. This function exploits this feature to use ops that support cuda graph. + Args: + tokens (torch.Tensor): The input token tensor, [num_tokens, hidden]. + routing_map (torch.Tensor): The sparse token to expert mapping, [num_tokens, num_experts]. + tokens_per_expert (torch.Tensor): cpu tensor + """ + + permuted_input, permuted_probs, row_id_map, row_id_index = ( + _BlockPaddedPermute.apply( + tokens, + probs, + routing_map, + tokens_per_expert_cuda_tensor, + tokens_per_expert_list, + quantizers, + cls, + ) + ) + return permuted_input, permuted_probs, row_id_map, row_id_index + + +class _BlockPaddedUnpermute(torch.autograd.Function): + @staticmethod + def forward( + ctx, + permuted_tokens, + row_id_map, + row_id_index, + tokens_per_expert, + splits, + restore_shape, + quantizers, + cls, + ): + """Forward function.""" + num_tokens, hidden_size = restore_shape + num_out_tokens = permuted_tokens.shape[0] + n_experts = row_id_map.size(1) + ctx.save_for_backward(row_id_index) + ctx.input_requires_grad = permuted_tokens.requires_grad + ctx.num_experts = n_experts + ctx.restore_shape = restore_shape + ctx.num_tokens = num_tokens + ctx.num_out_tokens = num_out_tokens + ctx.hidden_size = hidden_size + ctx.tokens_per_expert = tokens_per_expert + ctx.splits = splits + ctx.quantizers = quantizers + ctx.cls = cls + + output, _ = triton_unpermute_with_mask_map(permuted_tokens, row_id_map, None) + return output + + @staticmethod + def backward(ctx, grad_output): + """Backward function.""" + (row_id_index,) = ctx.saved_tensors + + quantizers = ctx.quantizers + x_q, x_scale, xt_q, xt_scale, _ = triton_batch_block_pad_permute_with_indices( + grad_output, + ctx.tokens_per_expert, + row_id_index, + ctx.splits, + round_scale=quantizers[0].force_pow_2_scales, + ) + + output = ctx.cls( + shape=x_q.shape, + dtype=grad_output.dtype, + fp8_dtype=quantizers[0].dtype, + rowwise_data=x_q, + rowwise_scale_inv=x_scale, + columnwise_data=xt_q, + columnwise_scale_inv=xt_scale, + quantizer=quantizers, + requires_grad=False, + is_2D_scaled=False, + ) + + return output, None, None, None, None, None, None, None + + +def block_padded_unpermute( + permuted_tokens: torch.Tensor, + row_id_map: torch.Tensor, + row_id_index: torch.Tensor, + tokens_per_expert: torch.Tensor, + splits: List, + restore_shape: torch.Size, + quantizers, + cls, +): + output = _BlockPaddedUnpermute.apply( + permuted_tokens, + row_id_map, + row_id_index, + tokens_per_expert, + splits, + restore_shape, + quantizers, + cls, + ) + return output + + +class _MXFP8PermuteDP(torch.autograd.Function): + @staticmethod + def forward( + ctx, + tokens, + probs, + routing_map, + tokens_per_expert_cuda_tensor, + tokens_per_expert_list, + quantizers, + cls, + ): + """Forward function.""" + num_tokens, hidden_dim = tokens.shape + + num_out_tokens = sum(tokens_per_expert_list) + row_id_map, row_id_index = triton_make_row_id_map_and_index( + routing_map, num_out_tokens + ) + + ctx.num_tokens = num_tokens + ctx.hidden_dim = hidden_dim + ctx.prob_shape = probs.shape + ctx.shape = tokens.shape + ctx.cls = cls + x_q, x_scale, xt_q, xt_scale, permuted_probs = ( + triton_batch_mxfp8_permute_with_indices( + tokens, + tokens_per_expert_cuda_tensor, + row_id_index, + tokens_per_expert_list, + probs=probs, + ) + ) + + output = cls( + shape=x_q.shape, + dtype=tokens.dtype, + fp8_dtype=quantizers[0].dtype, + rowwise_data=x_q, + rowwise_scale_inv=x_scale, + columnwise_data=xt_q, + columnwise_scale_inv=xt_scale, + quantizer=quantizers, + requires_grad=tokens.requires_grad, + ) + ctx.save_for_backward(row_id_map) + return output, permuted_probs, row_id_map, row_id_index + + @staticmethod + def backward(ctx, grad_output, grad_prob, grad_map, grad_index): + """Backward function.""" + (row_id_map,) = ctx.saved_tensors + output, prob_output = triton_unpermute_with_mask_map( + grad_output, row_id_map, grad_prob + ) + return ( + output.view(ctx.shape), + prob_output.view(ctx.prob_shape), + None, + None, + None, + None, + None, + ) + + +def mxfp8_permute( + tokens, + routing_map, + tokens_per_expert_cuda_tensor, + tokens_per_expert_list, + quantizers, + cls, + probs: Optional[torch.Tensor] = None, +): + permuted_input, permuted_probs, row_id_map, row_id_index = _MXFP8PermuteDP.apply( + tokens, + probs, + routing_map, + tokens_per_expert_cuda_tensor, + tokens_per_expert_list, + quantizers, + cls, + ) + return permuted_input, permuted_probs, row_id_map, row_id_index + + +class _MXFP8UnpermuteDP(torch.autograd.Function): + @staticmethod + def forward( + ctx, + permuted_tokens, + row_id_map, + row_id_index, + tokens_per_expert, + splits, + restore_shape, + quantizers, + cls, + ): + """Forward function.""" + num_tokens, hidden_size = restore_shape + num_out_tokens = permuted_tokens.shape[0] + n_experts = row_id_map.size(1) + ctx.save_for_backward(row_id_index) + ctx.input_requires_grad = permuted_tokens.requires_grad + ctx.num_experts = n_experts + ctx.restore_shape = restore_shape + ctx.num_tokens = num_tokens + ctx.num_out_tokens = num_out_tokens + ctx.hidden_size = hidden_size + ctx.tokens_per_expert = tokens_per_expert + ctx.splits = splits + ctx.quantizers = quantizers + ctx.cls = cls + + output, _ = triton_unpermute_with_mask_map(permuted_tokens, row_id_map, None) + return output + + @staticmethod + def backward(ctx, grad_output): + """Backward function.""" + (row_id_index,) = ctx.saved_tensors + + quantizers = ctx.quantizers + x_q, x_scale, xt_q, xt_scale, _ = triton_batch_mxfp8_permute_with_indices( + grad_output, + ctx.tokens_per_expert, + row_id_index, + ctx.splits, + ) + + output = ctx.cls( + shape=x_q.shape, + dtype=grad_output.dtype, + fp8_dtype=quantizers[0].dtype, + rowwise_data=x_q, + rowwise_scale_inv=x_scale, + columnwise_data=xt_q, + columnwise_scale_inv=xt_scale, + quantizer=quantizers, + requires_grad=False, + ) + + return output, None, None, None, None, None, None, None + + +def mxfp8_unpermute( + permuted_tokens: torch.Tensor, + row_id_map: torch.Tensor, + row_id_index: torch.Tensor, + tokens_per_expert: torch.Tensor, + splits: List, + restore_shape: torch.Size, + quantizers, + cls, +): + output = _MXFP8UnpermuteDP.apply( + permuted_tokens, + row_id_map, + row_id_index, + tokens_per_expert, + splits, + restore_shape, + quantizers, + cls, + ) + return output + + +class _MXFP8QuantDispatch(torch.autograd.Function): + @staticmethod + def forward( + ctx, tokens, tokens_per_expert_cuda, tokens_per_expert, quantizers, cls + ): + """Forward function.""" + num_tokens, hidden_dim = tokens.shape + + ctx.num_tokens = num_tokens + ctx.hidden_dim = hidden_dim + ctx.shape = tokens.shape + ctx.cls = cls + + inp_q, inp_scale, inpt_q, inpt_scale = triton_batch_mxfp8_quant( + tokens, tokens_per_expert_cuda, tokens_per_expert.tolist(), output_mode=2 + ) + + output = cls( + shape=inp_q.size(), + dtype=tokens.dtype, + fp8_dtype=quantizers[0].dtype, + rowwise_data=inp_q, + rowwise_scale_inv=inp_scale, + columnwise_data=inpt_q, + columnwise_scale_inv=inpt_scale, + quantizer=None, + requires_grad=tokens.requires_grad, + ) + + return output + + @staticmethod + def backward(ctx, grad_output): + return grad_output, None, None, None, None + + +def mxfp8_quant_dispatch( + tokens, tokens_per_expert_cuda, tokens_per_expert, quantizers, cls +): + output = _MXFP8QuantDispatch.apply( + tokens, + tokens_per_expert_cuda, + tokens_per_expert, + quantizers, + cls, + ) + return output + + +class _MXFP8QuantCombine(torch.autograd.Function): + @staticmethod + def forward( + ctx, tokens, tokens_per_expert_cuda, tokens_per_expert, quantizers, cls + ): + """Forward function.""" + num_tokens, hidden_dim = tokens.shape + + ctx.num_tokens = num_tokens + ctx.hidden_dim = hidden_dim + ctx.quantizers = quantizers + ctx.tokens_per_expert_cuda = tokens_per_expert_cuda + ctx.tokens_per_expert = tokens_per_expert + ctx.cls = cls + + return tokens + + @staticmethod + def backward(ctx, grad_output): + quantizers = ctx.quantizers + tokens_per_expert = ctx.tokens_per_expert + tokens_per_expert_cuda = ctx.tokens_per_expert_cuda + + inp_q, inp_scale, inpt_q, inpt_scale = triton_batch_mxfp8_quant( + grad_output, + tokens_per_expert_cuda, + tokens_per_expert.tolist(), + output_mode=2, + ) + + grad_output = ctx.cls( + shape=inp_q.size(), + dtype=grad_output.dtype, + fp8_dtype=quantizers[0].dtype, + rowwise_data=inp_q, + rowwise_scale_inv=inp_scale, + columnwise_data=inpt_q, + columnwise_scale_inv=inpt_scale, + quantizer=None, + requires_grad=grad_output.requires_grad, + ) + + return grad_output, None, None, None, None + + +def mxfp8_quant_combine( + tokens, tokens_per_expert_cuda, tokens_per_expert, quantizers, cls +): + output = _MXFP8QuantCombine.apply( + tokens, + tokens_per_expert_cuda, + tokens_per_expert, + quantizers, + cls, + ) + return output + + +# bf16 forward and bf16 backward +class _SmoothPermute(torch.autograd.Function): + @staticmethod + def forward( + ctx, tokens, probs, routing_map, tokens_per_expert, splits, quantizers, cls + ): + """Forward function.""" + num_tokens, hidden_dim = tokens.shape + + smooth_scales = torch.stack([x.smooth_scale for x in quantizers], 0) + # num_out_tokens should including padding tokens + row_id_map, row_id_indices = triton_make_row_id_map_and_index( + routing_map, sum(splits) + ) + ctx.num_tokens = num_tokens + ctx.hidden_dim = hidden_dim + ctx.prob_shape = probs.shape + ctx.shape = tokens.shape + ctx.row_id_map = row_id_map + permuted_input_data, permuted_input_scales, permuted_probs = ( + triton_batch_smooth_permute_with_indices( + tokens, + smooth_scales, + tokens_per_expert, + row_id_indices, + probs=probs, + reverse=False, + round_scale=False, + ) + ) + permuted_input = cls( + shape=permuted_input_data.shape, + dtype=tokens.dtype, + fp8_dtype=quantizers[0].dtype, + rowwise_data=permuted_input_data, + rowwise_scale_inv=permuted_input_scales, + columnwise_data=None, + columnwise_scale_inv=smooth_scales, + quantizer=quantizers, + requires_grad=tokens.requires_grad, + ) + return permuted_input, permuted_probs, row_id_map, row_id_indices + + @staticmethod + def backward(ctx, grad_output, grad_prob, grad_map, grad_indices): + """Backward function.""" + output, prob_output = triton_unpermute_with_mask_map( + grad_output, ctx.row_id_map, grad_prob + ) + return ( + output.view(ctx.shape), + prob_output.view(ctx.prob_shape), + None, + None, + None, + None, + None, + ) + + +def smooth_permute( + tokens, + routing_map, + tokens_per_expert, + splits, + quantizers, + cls, + probs: Optional[torch.Tensor] = None, +): + """Permute the tokens and probs based on the mask. + Tokens with the same designated expert will be grouped together. + The shape of mask is [tokens, num_experts], it indicates which experts were selected + by each token. + When drop_and_pad=True, in routing_map, the number of non-zeros in each column equals to + expert capacity. This function exploits this feature to use ops that support cuda graph. + Args: + tokens (torch.Tensor): The fp8 input token tensor, [num_tokens, hidden]. + routing_map (torch.Tensor): The sparse token to expert mapping, [num_tokens, num_experts]. + num_out_tokens (int, optional): The number of output tokens. If None, it's set to + the number of input tokens. + fused (bool, optional): Whether use the fused permute function. + drop_and_pad (bool, optional): Whether or not the token dispatcher uses token-drop + and pads the number of tokens to the expert capacity. + If set to true, routing_map has a fixed number of non-zeros + in each column. + """ + (permuted_input, permuted_probs, row_id_map, row_id_indices) = _SmoothPermute.apply( + tokens, probs, routing_map, tokens_per_expert, splits, quantizers, cls + ) + return permuted_input, permuted_probs, row_id_map, row_id_indices + + +# bf16 forward and bf16 backward +class _SmoothUnpermute(torch.autograd.Function): + @staticmethod + def forward( + ctx, + permuted_tokens, + row_id_map, + row_id_indices, + token_count_per_expert, + splits, + restore_shape, + quantizers, + cls, + ): + """Forward function.""" + num_tokens, hidden_size = restore_shape + num_out_tokens = permuted_tokens.shape[0] + n_experts = row_id_map.size(1) + ctx.save_for_backward(row_id_map, row_id_indices, token_count_per_expert) + ctx.num_experts = n_experts + ctx.restore_shape = restore_shape + ctx.num_tokens = num_tokens + ctx.num_out_tokens = num_out_tokens + ctx.hidden_size = hidden_size + ctx.splits = splits + ctx.quantizers = quantizers + ctx.cls = cls + output, _ = triton_unpermute_with_mask_map(permuted_tokens, row_id_map, None) + return output + + @staticmethod + def backward(ctx, grad_output): + """Backward function.""" + row_id_map, row_id_indices, token_count_per_expert = ctx.saved_tensors + quantizers = ctx.quantizers + # TODO(nanxiao): smooth_scale_inv will updated in every forward, it will cause error with PP + grad_smooth_scales = torch.stack([x.smooth_scale_inv for x in quantizers], 0) + transpose_grad_smooth_scales = [ + x.transpose_smooth_scale_inv + for x in quantizers + if x.transpose_smooth_scale_inv.numel() > 0 + ] + if len(transpose_grad_smooth_scales) > 0: + transpose_grad_smooth_scales = torch.cat(transpose_grad_smooth_scales, 0) + else: + transpose_grad_smooth_scales = None + round_scale = quantizers[0].force_pow_2_scales + + # import pydevd + # pydevd.settrace(suspend=False, trace_only_current_thread=True) + permuted_grad_data, permuted_grad_scales, _ = ( + triton_batch_smooth_permute_with_indices( + grad_output, + grad_smooth_scales, + token_count_per_expert, + row_id_indices, + probs=None, + reverse=True, + round_scale=round_scale, + ) + ) + permuted_grad_data_t, permuted_grad_scales_t = ( + triton_batch_transpose_smooth_permute_with_indices( + grad_output, + transpose_grad_smooth_scales, + row_id_indices, + token_count_per_expert, + ctx.splits, + round_scale=round_scale, + ) + ) + + input_grad = ctx.cls( + shape=permuted_grad_data.shape, + dtype=grad_output.dtype, + fp8_dtype=quantizers[0].dtype, + rowwise_data=permuted_grad_data, + rowwise_scale_inv=permuted_grad_scales, + columnwise_data=permuted_grad_data_t, + columnwise_scale_inv=permuted_grad_scales_t, + quantizer=quantizers, + requires_grad=False, + ) + return input_grad, None, None, None, None, None, None, None + + +def smooth_unpermute( + permuted_tokens: torch.Tensor, + row_id_map: torch.Tensor, + row_id_indices: torch.Tensor, + token_count_per_expert: torch.Tensor, + splits: List[int], + restore_shape: torch.Size, + quantizers: List, + cls, +): + output = _SmoothUnpermute.apply( + permuted_tokens, + row_id_map, + row_id_indices, + token_count_per_expert, + splits, + restore_shape, + quantizers, + cls, + ) + return output + + +# fp8 forward and bf16 backward +class _SmoothFusedPermute(torch.autograd.Function): + @staticmethod + def forward(ctx, tokens, probs, routing_map, tokens_per_expert, splits, cls): + """Forward function.""" + num_tokens, hidden_dim = tokens._rowwise_data.shape + # num_experts = routing_map.shape[1] + counts = routing_map.sum(-1) + + # num_out_tokens should including padding tokens + row_id_map, row_id_indices = triton_make_row_id_map_and_index( + routing_map, sum(splits) + ) + ctx.num_tokens = num_tokens + ctx.hidden_dim = hidden_dim + ctx.prob_shape = probs.shape + ctx.shape = tokens.shape + ctx.counts = counts + ctx.row_id_map = row_id_map + permuted_input_data, permuted_input_scales, permuted_probs = ( + triton_permute_with_mask_map( + tokens._rowwise_data, + tokens._rowwise_scale_inv, + probs, + row_id_map, + sum(splits), + ) + ) + permuted_input = cls( + shape=permuted_input_data.shape, + dtype=tokens.dtype, + fp8_dtype=tokens._fp8_dtype, + rowwise_data=permuted_input_data, + rowwise_scale_inv=permuted_input_scales, + columnwise_data=None, + columnwise_scale_inv=tokens._columnwise_scale_inv, + quantizer=tokens._quantizer, + requires_grad=tokens.requires_grad, + ) + return permuted_input, permuted_probs, row_id_map, row_id_indices + + @staticmethod + def backward(ctx, grad_output, grad_prob, grad_map, grad_indices): + """Backward function.""" + output, prob_output = triton_unpermute_with_mask_map( + grad_output, ctx.row_id_map, grad_prob + ) + return ( + output.view(ctx.shape), + prob_output.view(ctx.prob_shape), + None, + None, + None, + None, + ) + + +def smooth_fused_permute( + tokens, + routing_map, + probs: Optional[torch.Tensor] = None, + tokens_per_expert: Optional[torch.Tensor] = None, + splits: Optional[List[int]] = None, +): + """Permute the tokens and probs based on the mask. + Tokens with the same designated expert will be grouped together. + The shape of mask is [tokens, num_experts], it indicates which experts were selected + by each token. + When drop_and_pad=True, in routing_map, the number of non-zeros in each column equals to + expert capacity. This function exploits this feature to use ops that support cuda graph. + Args: + tokens (torch.Tensor): The fp8 input token tensor, [num_tokens, hidden]. + routing_map (torch.Tensor): The sparse token to expert mapping, [num_tokens, num_experts]. + num_out_tokens (int, optional): The number of output tokens. If None, it's set to + the number of input tokens. + fused (bool, optional): Whether use the fused permute function. + drop_and_pad (bool, optional): Whether or not the token dispatcher uses token-drop + and pads the number of tokens to the expert capacity. + If set to true, routing_map has a fixed number of non-zeros + in each column. + """ + (permuted_input, permuted_probs, row_id_map, row_id_indices) = ( + _SmoothFusedPermute.apply(tokens, probs, routing_map, tokens_per_expert, splits) + ) + return permuted_input, permuted_probs, row_id_map, row_id_indices + + +# bf16 forward and fp8 backward +# fused unpermute and quant +class _SmoothFusedUnpermute(torch.autograd.Function): + @staticmethod + def forward( + ctx, + permuted_tokens, + row_id_map, + row_id_indices, + org_smooth_scale, + token_count_per_expert, + splits, + quantizers, + restore_shape, + cls, + ): + """Forward function.""" + num_tokens, hidden_size = restore_shape + num_out_tokens = permuted_tokens.shape[0] + n_experts = row_id_map.size(1) + ctx.save_for_backward( + row_id_map, row_id_indices, org_smooth_scale, token_count_per_expert + ) + ctx.input_requires_grad = permuted_tokens.requires_grad + ctx.num_experts = n_experts + ctx.restore_shape = restore_shape + ctx.num_tokens = num_tokens + ctx.num_out_tokens = num_out_tokens + ctx.hidden_size = hidden_size + ctx.splits = splits + ctx.quantizers = quantizers + ctx.cls = cls + output, _ = triton_unpermute_with_mask_map(permuted_tokens, row_id_map, None) + return output + + @staticmethod + def backward(ctx, grad_output): + """Backward function. grad_output is smooth quantized""" + row_id_map, row_id_indices, org_smooth_scale, token_count_per_expert = ( + ctx.saved_tensors + ) + quantizers = grad_output._quantizer + # smooth_scale_inv will updated in every forward, it will cause error with PP + smooth_scales = torch.stack([x.smooth_scale_inv for x in quantizers], 0) + transpose_smooth_scales = torch.cat( + [x.transpose_smooth_scale_inv for x in quantizers], 0 + ) + # todo(nanxiao): smooth + round_scale = quantizers[0].force_pow_2_scales + + update_smooth_scales = smooth_scales * org_smooth_scale + for i, quantizer in enumerate(quantizers): + quantizer.smooth_scale_inv = update_smooth_scales[i] + # import pydevd + # pydevd.settrace(suspend=False, trace_only_current_thread=True) + permuted_grad_data, permuted_grad_scales = ( + triton_batch_smooth_fused_permute_with_indices( + grad_output._rowwise_data, + grad_output._rowwise_scale_inv, + update_smooth_scales, + token_count_per_expert, + row_id_indices, + reverse=True, + round_scale=round_scale, + ) + ) + permuted_grad_data_t, permuted_grad_scales_t = ( + triton_batch_transpose_smooth_fused_permute_with_indices( + grad_output._rowwise_data, + grad_output._rowwise_scale_inv, + grad_smooth_scale, + transpose_smooth_scales, + row_id_indices, + token_count_per_expert, + ctx.splits, + round_scale=round_scale, + ) + ) + + input_grad = ctx.cls( + shape=permuted_grad_data.shape, + dtype=grad_output.dtype, + fp8_dtype=quantizers[0].dtype, + rowwise_data=permuted_grad_data, + rowwise_scale_inv=permuted_grad_scales, + columnwise_data=permuted_grad_data_t, + columnwise_scale_inv=permuted_grad_scales_t, + quantizer=quantizers, + requires_grad=ctx.input_requires_grad, + ) + return input_grad, None, None, None, None, None, None, None, None + + +def smooth_fused_unpermute( + permuted_tokens: torch.Tensor, + row_id_map: torch.Tensor, + row_id_indices: torch.Tensor, + org_smooth_scale: torch.Tensor, + token_count_per_expert: torch.Tensor, + splits: List[int], + quantizers: List, + restore_shape: torch.Size, + cls, +): + output = _SmoothFusedUnpermute.apply( + permuted_tokens, + row_id_map, + row_id_indices, + org_smooth_scale, + token_count_per_expert, + splits, + quantizers, + restore_shape, + cls, + ) + return output + + +class _MXFP8Permute(torch.autograd.Function): + @staticmethod + def forward( + ctx, + inp: torch.Tensor, + split_sizes: torch.Tensor, + sorted_idxs: torch.Tensor, + probs: torch.Tensor, + tokens_per_expert, + tokens_per_expert_cpu, + quantizers, + cls, + ): + """Forward for alltoall dispatch.""" + if not inp.numel(): + return inp, probs + + row_id_map, reverse_row_id_map = triton_make_chunk_sort_map( + split_sizes.reshape(-1, tokens_per_expert.size(0)), + tokens_per_expert_cpu.tolist(), + ) + + x_q, x_s, xt_q, xt_s, permuted_probs = triton_batch_mxfp8_permute_with_indices( + inp, + tokens_per_expert, + row_id_map, + tokens_per_expert_cpu.tolist(), + probs=probs, + ) + + ctx.save_for_backward(reverse_row_id_map) + ctx.prob_shape = probs.shape + ctx.cls = cls + + output = cls( + shape=x_q.shape, + dtype=inp.dtype, + fp8_dtype=quantizers[0].dtype, + rowwise_data=x_q, + rowwise_scale_inv=x_s, + columnwise_data=xt_q, + columnwise_scale_inv=xt_s, + quantizer=quantizers, + requires_grad=inp.requires_grad, + ) + return output, permuted_probs, row_id_map, reverse_row_id_map + + @staticmethod + def backward( + ctx, + permuted_act_grad, + permuted_probs_grad, + row_id_map_grad, + reverse_row_id_map_grad, + ): + (reverse_row_id_map,) = ctx.saved_tensors + if not permuted_act_grad.numel(): + return permuted_act_grad, None, None, permuted_probs_grad + act_grad, probs_grad = triton_unpermute_with_reverse_map( + permuted_act_grad, reverse_row_id_map, permuted_probs_grad + ) + return act_grad, None, None, probs_grad, None, None, None, None + + +def mxfp8_dispatch( + inp, + split_sizes, + sorted_idxs, + probs, + tokens_per_expert, + tokens_per_expert_cpu, + quantizers, + cls, +): + output, permuted_probs, row_id_map, reverse_row_id_map = _MXFP8Permute.apply( + inp, + split_sizes, + sorted_idxs, + probs, + tokens_per_expert, + tokens_per_expert_cpu, + quantizers, + cls, + ) + return output, permuted_probs, row_id_map, reverse_row_id_map + + +class _MXFP8Unpermute(torch.autograd.Function): + @staticmethod + def forward( + ctx, + inp: torch.Tensor, + ori_size, + split_sizes: torch.Tensor, + sorted_idxs: torch.Tensor, + probs: torch.Tensor, + tokens_per_expert, + tokens_per_expert_cpu, + quantizers, + cls, + row_id_map, + reverse_row_id_map, + ): + """Forward for alltoall dispatch.""" + if not inp.numel(): + return inp, probs + + output, unpermuted_probs = triton_unpermute_with_reverse_map( + inp, reverse_row_id_map, probs + ) + + ctx.save_for_backward(row_id_map) + ctx.cls = cls + ctx.tokens_per_expert = tokens_per_expert + ctx.tokens_per_expert_cpu = tokens_per_expert_cpu + ctx.quantizers = quantizers + ctx.sorted_idxs = sorted_idxs + ctx.split_sizes = split_sizes + + return output, unpermuted_probs + + @staticmethod + def backward(ctx, permuted_act_grad, permuted_probs_grad): + (row_id_map,) = ctx.saved_tensors + + x_q, x_s, xt_q, xt_s, permuted_probs = triton_batch_mxfp8_permute_with_indices( + permuted_act_grad, + ctx.tokens_per_expert, + row_id_map, + ctx.tokens_per_expert_cpu.tolist(), + probs=permuted_probs_grad, + ) + + act_grad = ctx.cls( + shape=x_q.shape, + dtype=permuted_act_grad.dtype, + fp8_dtype=ctx.quantizers[0].dtype, + rowwise_data=x_q, + rowwise_scale_inv=x_s, + columnwise_data=xt_q, + columnwise_scale_inv=xt_s, + quantizer=ctx.quantizers, + requires_grad=permuted_act_grad.requires_grad, + ) + + return ( + act_grad, + None, + None, + None, + permuted_probs, + None, + None, + None, + None, + None, + None, + ) + + +def mxfp8_combine( + inp, + ori_size, + split_sizes, + sorted_idxs, + tokens_per_expert, + tokens_per_expert_cpu, + quantizers, + cls, + row_id_map, + reverse_row_id_map, + probs=None, +): + output = _MXFP8Unpermute.apply( + inp, + ori_size, + split_sizes, + sorted_idxs, + probs, + tokens_per_expert, + tokens_per_expert_cpu, + quantizers, + cls, + row_id_map, + reverse_row_id_map, + ) + return output diff --git a/linghe/facade/quantization.py b/linghe/facade/quantization.py new file mode 100644 index 0000000..8433dd4 --- /dev/null +++ b/linghe/facade/quantization.py @@ -0,0 +1,359 @@ +# -*- coding: utf-8 -*- +""" +Copyright (c) Ant Financial Service Group and its affiliates. +""" + +import torch + +from linghe.quant.smooth import ( + triton_smooth_quant, + triton_transpose_smooth_quant, + triton_batch_smooth_quant, + triton_batch_transpose_smooth_quant, + triton_subrow_smooth_quant, +) +from linghe.utils.transpose import triton_pad_transpose + +""" +smooth quantization v2 for fp8 training +1.1 calculate M = max(abs(w)). +1.2 calculate smooth_scale = sqrt(M) + +2.1 quantize weight w_q, w_s = quant(w / smooth_scale) +2.2 set weight quantizer.smooth_scale = smooth_scale +2.3 set weight quantizer.smooth_scale_inv = w_s +3.4 set quantized weight._columnwise_scale_inv = smooth_scale + +3.1 set activation quantizer.smooth_scale = 1/smooth_scale +3.1 quantize activation a_q, a_s = quant(a * smooth_scale) +3.3 set activation quantizer.transpose_smooth_scale_inv = a_s + +4.1 set grad quantizer.smooth_scale_inv = w_s +4.2 set grad quantizer.transpose_smooth_scale_inv = a_s +4.3 quantize grad g_q, g_s = quant(g * smooth_scale_inv) +4.4 quantize transposed grad gt_q, gt_s = quant(g * transpose_smooth_scale_inv) +""" + + +class SmoothQuantize(torch.autograd.Function): + @staticmethod + def forward(ctx, hidden_states, quantizer, grad_quantizer, cls): + ctx.grad_quantizer = grad_quantizer + ctx.cls = cls + + shape = hidden_states.shape + if len(shape) == 3: + hidden_states = hidden_states.view(-1, shape[-1]) + + if hasattr(hidden_states, "_quantizer") or quantizer is None: + return hidden_states + + ctx.grad_quantizer = grad_quantizer + x_q, x_scale = triton_smooth_quant( + hidden_states, + quantizer.smooth_scale, + reverse=False, + round_scale=quantizer.force_pow_2_scales, + ) + output = cls( + shape=shape, + dtype=hidden_states.dtype, + fp8_dtype=quantizer.dtype, + rowwise_data=x_q, + rowwise_scale_inv=x_scale, + columnwise_data=None, + columnwise_scale_inv=quantizer.smooth_scale, + quantizer=quantizer, + requires_grad=hidden_states.requires_grad, + ) + return output + + @staticmethod + def backward(ctx, grad_output): + if hasattr(grad_output, "_quantizer") or ctx.grad_quantizer is None: + return (grad_output,) + shape = grad_output.shape # rank-3 tensor + grad_output = grad_output.view(-1, shape[-1]) + # import pydevd + # pydevd.settrace(suspend=False, trace_only_current_thread=True) + grad_quantizer = ctx.grad_quantizer + y_q, y_scale = triton_smooth_quant( + grad_output, + grad_quantizer.smooth_scale_inv, + reverse=True, + round_scale=False, + ) + yt_q, yt_scale = triton_transpose_smooth_quant( + grad_output, + grad_quantizer.transpose_smooth_scale_inv, + reverse=True, + pad=True, + round_scale=False, + ) + # import math + # if math.isnan(y_q.float().max()) or math.isnan(yt_q.float().max()): + # print(f'ReverseSmoothQuantize {grad_output.max()=} {y_scale.max()=} {yt_scale.max()=}') + output = ctx.cls( + shape=shape, + dtype=grad_output.dtype, + fp8_dtype=grad_quantizer.dtype, + rowwise_data=y_q, + rowwise_scale_inv=y_scale, + columnwise_data=yt_q, + columnwise_scale_inv=yt_scale, + quantizer=grad_quantizer, + requires_grad=False, + ) + return output, None, None, None + + +def smooth_quantize(hidden_states, quantizer, grad_quantizer, cls): + return SmoothQuantize.apply(hidden_states, quantizer, grad_quantizer, cls) + + +class BatchSmoothQuantize(torch.autograd.Function): + @staticmethod + def forward( + ctx, + hidden_states, + token_count_per_expert, + quantizers, + grad_quantizers, + splits, + cls, + ): + ctx.grad_quantizers = grad_quantizers + ctx.splits = splits + ctx.cls = cls + ctx.token_count_per_expert = None + + if hasattr(hidden_states, "_quantizer") or quantizers is None: + return hidden_states + + shape = hidden_states.shape + if len(shape) == 3: + hidden_states = hidden_states.view(-1, shape[-1]) + + if token_count_per_expert is None: + token_count_per_expert = torch.tensor(splits).cuda(non_blocking=True) + ctx.token_count_per_expert = token_count_per_expert + + smooth_scales = [x.smooth_scale for x in quantizers] + if any([x is None for x in smooth_scales]): + smooth_scales = torch.ones( + (len(smooth_scales), hidden_states.shape[-1]), + dtype=torch.float32, + device=hidden_states.device, + ) + for i, x in enumerate(quantizers): + x.smooth_scale = smooth_scales[i] + else: + smooth_scales = torch.stack(smooth_scales, 0) + x_q, x_scale = triton_batch_smooth_quant( + hidden_states, + smooth_scales, + token_count_per_expert, + reverse=False, + round_scale=quantizers[0].force_pow_2_scales, + ) + output = cls( + shape=shape, + dtype=hidden_states.dtype, + fp8_dtype=quantizers[0].dtype, + rowwise_data=x_q, + rowwise_scale_inv=x_scale, + columnwise_data=None, + columnwise_scale_inv=smooth_scales, + quantizer=quantizers, + requires_grad=hidden_states.requires_grad, + ) + return output + + @staticmethod + def backward(ctx, grad_output): + if hasattr(grad_output, "_quantizer") or ctx.grad_quantizers is None: + return grad_output, None, None, None, None, None + + token_count_per_expert = ctx.token_count_per_expert + if token_count_per_expert is None: + token_count_per_expert = torch.tensor(ctx.splits).cuda(non_blocking=True) + + shape = grad_output.shape # rank-3 tensor + grad_output = grad_output.view(-1, shape[-1]) + # import pydevd + # pydevd.settrace(suspend=False, trace_only_current_thread=True) + + grad_quantizers = ctx.grad_quantizers + smooth_scale_invs = [x.smooth_scale_inv for x in grad_quantizers] + smooth_scale_invs = torch.stack(smooth_scale_invs, 0) + + y_q, y_scale = triton_batch_smooth_quant( + grad_output, + smooth_scale_invs, + token_count_per_expert, + reverse=True, + round_scale=grad_quantizers[0].force_pow_2_scales, + ) + + transpose_smooth_scale_invs = [ + x.transpose_smooth_scale_inv for x in grad_quantizers + ] + transpose_smooth_scale_invs = torch.cat(transpose_smooth_scale_invs, 0) + + yt_q, yt_scale = triton_batch_transpose_smooth_quant( + grad_output, + transpose_smooth_scale_invs, + token_count_per_expert, + ctx.splits, + reverse=True, + round_scale=grad_quantizers[0].force_pow_2_scales, + ) + + # import math + # if math.isnan(yt_q.float().max()): + # print(f'BatchReverseSmoothQuantize {grad_output.max()=} {yt_scale.max()=}') + output = ctx.cls( + shape=shape, + dtype=grad_output.dtype, + fp8_dtype=ctx.grad_quantizers[0].dtype, + rowwise_data=y_q, + rowwise_scale_inv=y_scale, + columnwise_data=yt_q, + columnwise_scale_inv=yt_scale, + quantizer=ctx.grad_quantizers, + requires_grad=False, + ) + return output, None, None, None, None, None + + +def batch_smooth_quantize( + hidden_states, token_count_per_expert, quantizers, grad_quantizers, splits, cls +): + return BatchSmoothQuantize.apply( + hidden_states, token_count_per_expert, quantizers, grad_quantizers, splits, cls + ) + + +# y = x @ w +# dx = y @ wT +# dwT = yT @ x +def triton_smooth_quant_activation( + x, + smooth_scale, + x_q=None, + x_scale=None, + xt_q=None, + transpose=True, + pad=True, + round_scale=False, +): + """""" + x_q, x_scale = triton_smooth_quant( + x, + smooth_scale, + x_q=x_q, + x_scale=x_scale, + reverse=False, + round_scale=round_scale, + ) + + if transpose: + xt_q = triton_pad_transpose(x_q, out=xt_q, multiple=32) + else: + xt_q = None + xt_scale = smooth_scale + + return x_q, x_scale, xt_q, xt_scale + + +# y = x @ w +# dx = y @ wT +# dwT = yT @ x +def triton_smooth_quant_gradient( + y, + smooth_scale, + transpose_smooth_scale, + reverse=True, + transpose=True, + pad=True, + round_scale=False, +): + """""" + assert reverse, ( + "args `smooth_scale` and/or `transpose_smooth_scale` " + "must be in reciprocal format in triton_smooth_quant_grad" + ) + y_q, y_scale = triton_smooth_quant( + y, smooth_scale, reverse=True, round_scale=round_scale + ) + if transpose: + yt_q, yt_scale = triton_transpose_smooth_quant( + y, transpose_smooth_scale, reverse=True, pad=pad, round_scale=round_scale + ) + else: + yt_q, yt_scale = None, None + + return y_q, y_scale, yt_q, yt_scale + + +def triton_smooth_quant_weight( + w, smooth_scale, w_q, quant_scale, subrow_scales, offset=0, round_scale=False +): + """""" + assert w.ndim == 1 + assert w_q.size(1) == smooth_scale.size(0) + + size = w.numel() + M, N = w_q.shape + + if size == M * N: + triton_smooth_quant( + w.view(M, N), + smooth_scale, + x_q=w_q, + x_scale=quant_scale, + round_scale=round_scale, + ) + elif offset % N == 0 and size % N == 0: + n_row = size // N + row_id = offset // N + w_q_slice = w_q[row_id : row_id + n_row] + quant_scale_slice = quant_scale[row_id : row_id + n_row] + triton_smooth_quant( + w.view(n_row, N), + smooth_scale, + x_q=w_q_slice, + x_scale=quant_scale_slice, + round_scale=round_scale, + ) + else: + row_si = (offset - 1) // N + 1 + row_ei = (offset + size) // N + col_si = offset % N + col_ei = (offset + size) % N + n_row = row_ei - row_si + mw_offset = 0 if col_si == 0 else N - col_si + w_q_slice = w_q[row_si:row_ei] + quant_scale_slice = quant_scale[row_si:row_ei] + w_slice = w[mw_offset : mw_offset + n_row * N].view(n_row, N) + triton_smooth_quant( + w_slice, + smooth_scale, + x_q=w_q_slice, + x_scale=quant_scale_slice, + round_scale=round_scale, + ) + + # subrow scale is writed by the row with leading master weights + if col_si > 0 or col_ei > 0: + triton_subrow_smooth_quant( + w, + smooth_scale, + w_q, + quant_scale, + subrow_scales, + offset, + size, + reverse=False, + round_scale=round_scale, + ) diff --git a/linghe/facade/rope.py b/linghe/facade/rope.py index c31f1c8..b516cc6 100644 --- a/linghe/facade/rope.py +++ b/linghe/facade/rope.py @@ -3,80 +3,301 @@ Copyright (c) Ant Financial Service Group and its affiliates. """ +from typing import Optional + import torch -from linghe.utils.rope import triton_qk_norm_and_half_rope_forward, \ - triton_qk_norm_and_half_rope_backward +from linghe.utils.rope import ( + triton_qk_norm_and_half_rope_forward, + triton_qk_norm_and_half_rope_backward, + triton_varlen_qk_norm_and_half_rope_forward, + triton_varlen_qk_norm_and_half_rope_backward, + triton_mla_rope_forward, + triton_mla_rope_backward, +) class QkNormHalfRopeFunction(torch.autograd.Function): """""" + @staticmethod - def forward(ctx, qkv, q_norm_weight, k_norm_weight, freqs, H=32, h=4, - eps=1e-6): - shape = qkv.shape - qo, ko, vo = triton_qk_norm_and_half_rope_forward(qkv, - q_norm_weight.data, - k_norm_weight.data, - freqs, - H=H, - h=h, - eps=eps, - interleaved=True, - transposed=True) - - ctx.save_for_backward(qkv, q_norm_weight.data, k_norm_weight.data, - freqs) + def forward( + ctx, + qkv, + q_norm_weight, + k_norm_weight, + freqs, + cu_seqlens_q, + cu_seqlens_kv, + H=32, + h=4, + eps=1e-6, + cp_rank=0, + cp_size=1, + mscale=1.0, + silu=False, + reuse=False, + ): + if cu_seqlens_q is None: + qo, ko, vo = triton_qk_norm_and_half_rope_forward( + qkv, + q_norm_weight, + k_norm_weight, + freqs, + H=H, + h=h, + eps=eps, + interleaved=True, + transposed=True, + silu=silu, + ) + else: + qo, ko, vo = triton_varlen_qk_norm_and_half_rope_forward( + qkv, + q_norm_weight, + k_norm_weight, + freqs, + cu_seqlens_q, + cu_seqlens_kv, + H=H, + h=h, + eps=eps, + interleaved=True, + cp_rank=cp_rank, + cp_size=cp_size, + mscale=mscale, + silu=silu, + reuse=reuse, + ) + ctx.save_for_backward(qkv, q_norm_weight, k_norm_weight, freqs) ctx.H = H ctx.h = h ctx.eps = eps - ctx.shape = shape + ctx.cp_rank = cp_rank + ctx.cp_size = cp_size + ctx.mscale = mscale + ctx.silu = silu + ctx.reuse = reuse + ctx.cu_seqlens_q = cu_seqlens_q + ctx.cu_seqlens_kv = cu_seqlens_kv return qo, ko, vo @staticmethod def backward(ctx, grad_q, grad_k, grad_v): qkv, q_norm_weight, k_norm_weight, freqs = ctx.saved_tensors - dqkv, dqw, dkw = triton_qk_norm_and_half_rope_backward(grad_q, - grad_k, - grad_v, - qkv, - q_norm_weight, - k_norm_weight, - freqs, - eps=ctx.eps, - transposed=True, - interleaved=True) - return dqkv, dqw, dkw, None, None, None, None - - -def qk_norm_half_rope(qkv: torch.Tensor, - q_norm_weight: torch.Tensor, - k_norm_weight: torch.Tensor, - freqs: torch.Tensor, - H: int = 32, - h: int = 4, - eps: float = 1e-6): + if ctx.cu_seqlens_q is None: + dqkv, dqw, dkw = triton_qk_norm_and_half_rope_backward( + grad_q, + grad_k, + grad_v, + qkv, + q_norm_weight, + k_norm_weight, + freqs, + eps=ctx.eps, + transposed=True, + interleaved=True, + silu=ctx.silu, + ) + else: + dqkv, dqw, dkw = triton_varlen_qk_norm_and_half_rope_backward( + grad_q, + grad_k, + grad_v, + qkv, + q_norm_weight, + k_norm_weight, + freqs, + ctx.cu_seqlens_q, + ctx.cu_seqlens_kv, + eps=ctx.eps, + interleaved=True, + cp_rank=ctx.cp_rank, + cp_size=ctx.cp_size, + mscale=ctx.mscale, + silu=ctx.silu, + reuse=ctx.reuse, + ) + return ( + dqkv, + dqw, + dkw, + None, + None, + None, + None, + None, + None, + None, + None, + None, + None, + None, + ) + + +def qk_norm_half_rope( + qkv: torch.Tensor, + q_norm_weight: torch.Tensor, + k_norm_weight: torch.Tensor, + freqs: torch.Tensor, + cu_seqlens_q: Optional[torch.Tensor] = None, + cu_seqlens_kv: Optional[torch.Tensor] = None, + H: int = 32, + h: int = 4, + eps: float = 1e-6, + cp_rank=0, + cp_size=1, + mscale=1.0, + silu=False, + reuse=False, +): """ split qkv to q/k/v, apply qk norm and half rope to q/k, transpose q/k/v to flash-attention layout Args: - qkv: QKV tensor with size of [S, B, dim], heads are interleaved + qkv: QKV tensor with size of [S, B, dim] or [T, dim] , heads are interleaved q_norm_weight: rms norm weight for query k_norm_weight: rms norm weight for key freqs: Freqs tensor based on half dim. + cu_seqlens_q: accumulated query lengths, [num_seqs + 1] + cu_seqlens_kv: accumulated kv lengths, [num_seqs + 1] H: Number of attention heads. h: Number of key/value heads. eps: epsilon value for L2 normalization. + cp_rank: context parallel rank + cp_size: context parallel size + mscale: mscale for rope + + Returns: + - qo: shape [B, S, H, head_dim] or [T, H, head_dim] + - ko: shape [B, S, h, head_dim] or [T, h, head_dim] + - vo: shape [B, S, h, head_dim] or [T, h, head_dim] + """ + return QkNormHalfRopeFunction.apply( + qkv, + q_norm_weight, + k_norm_weight, + freqs, + cu_seqlens_q, + cu_seqlens_kv, + H, + h, + eps, + cp_rank, + cp_size, + mscale, + silu, + reuse, + ) + + +class MLARopeFunction(torch.autograd.Function): + """""" + @staticmethod + def forward( + ctx, + q, + kv, + k_pos_emb, + freqs, + mscale, + transpose, + cu_seqlens_q, + cu_seqlens_kv, + cp_size, + cp_rank, + reuse, + ): + qo, ko, vo = triton_mla_rope_forward( + q, + kv, + k_pos_emb, + freqs, + mscale=mscale, + transpose=transpose, + cu_seqlens_q=cu_seqlens_q, + cu_seqlens_kv=cu_seqlens_kv, + cp_size=cp_size, + cp_rank=cp_rank, + reuse=reuse, + ) + + ctx.save_for_backward(freqs) + ctx.mscale = mscale + ctx.cp_size = cp_size + ctx.cp_rank = cp_rank + ctx.transpose = transpose + ctx.reuse = reuse + ctx.cu_seqlens_q = cu_seqlens_q + ctx.cu_seqlens_kv = cu_seqlens_kv + return qo, ko, vo + + @staticmethod + def backward(ctx, grad_q, grad_k, grad_v): + (freqs,) = ctx.saved_tensors + dq, dkv, dp = triton_mla_rope_backward( + grad_q, + grad_k, + grad_v, + freqs, + mscale=ctx.mscale, + transposed=ctx.transpose, + cu_seqlens_q=ctx.cu_seqlens_q, + cu_seqlens_kv=ctx.cu_seqlens_kv, + cp_size=ctx.cp_size, + cp_rank=ctx.cp_rank, + reuse=ctx.reuse, + ) + return dq, dkv, dp, None, None, None, None, None, None, None, None + + +def mla_rope( + q: torch.Tensor, + kv: torch.Tensor, + k_pos_emb: torch.Tensor, + freqs: torch.Tensor, + cu_seqlens_q: Optional[torch.Tensor] = None, + cu_seqlens_kv: Optional[torch.Tensor] = None, + mscale: float = 1.0, + transpose: bool = False, + cp_size: int = 1, + cp_rank: int = 0, + reuse: bool = False, +): + """ + inplace apply rope to tail 64 dims, split kv and apply rope to k_pos_emb and copy to k + Args: + q: query tensor with size of [S, B, H, 128] (cu_seqlens is None) + or [N, H, 128] (cu_seqlens is not None) + kv: kv tensor with size of [S, B, H, 256] (cu_seqlens is None) or + [N, H, 256] (cu_seqlens is not None) + k_pos_emb: k pos emb with size of [S, B, 1, 64] (cu_seqlens is None) or + [N, 1, 64] (cu_seqlens is not None) + freqs: Freqs tensor with size of [S, 64] + cu_seqlens_q: cumulative query lengths tensor with size of [B+1] + cu_seqlens_kv: cumulative kv lengths tensor with size of [B+1] + mscale: mscale of rope + transpose: whether transpose output layout to [B, S, H, DIM] + cp_size: context-parallel size + cp_rank: context-parallel rank Returns: - - qo: shape [B, S, H, head_dim] - - ko: shape [B, S, h, head_dim] - - vo: shape [B, S, h, head_dim] + - qo: shape [S, B, H, 192] or [N, H, 192] + - ko: shape [S, B, H, 192] or [N, H, 192] + - vo: shape [S, B, H, 128] or [N, H, 128] """ - return QkNormHalfRopeFunction.apply(qkv, - q_norm_weight, - k_norm_weight, - freqs, - H, - h, - eps) \ No newline at end of file + q, k, v = MLARopeFunction.apply( + q, + kv, + k_pos_emb, + freqs, + mscale, + transpose, + cu_seqlens_q, + cu_seqlens_kv, + cp_size, + cp_rank, + reuse, + ) + return q, k, v diff --git a/linghe/facade/silu.py b/linghe/facade/silu.py new file mode 100644 index 0000000..3f4fe91 --- /dev/null +++ b/linghe/facade/silu.py @@ -0,0 +1,509 @@ +import torch + +from linghe.utils.silu import ( + triton_silu_and_block_quant_forward, + triton_silu_and_block_quant_backward, + triton_batch_weighted_silu_and_block_quant_forward, + triton_batch_weighted_silu_and_block_quant_backward, + triton_silu_and_mxfp8_quant_forward, + triton_silu_and_mxfp8_quant_backward, + triton_batch_weighted_silu_and_mxfp8_quant_forward, + triton_batch_weighted_silu_and_mxfp8_quant_backward, + triton_silu_and_smooth_quant_forward, + triton_silu_and_smooth_quant_backward, + triton_batch_weighted_silu_and_smooth_quant_forward, + triton_batch_weighted_silu_and_smooth_quant_backward, +) + + +class BlockSiluFunction(torch.autograd.Function): + @staticmethod + def forward(ctx, input, quantizer, grad_quantizer, cls, limit): + shape = input.shape + assert len(shape) == 3 + input_view = input.view(shape[0] * shape[1], shape[2]) + ctx.grad_quantizer = grad_quantizer + ctx.input_requires_grad = input.requires_grad + ctx.shape = shape + ctx.cls = cls + ctx.limit = limit + ctx.save_for_backward(input) + + round_scale = quantizer.force_pow_2_scales + x_q, x_scale, xt_q, xt_scale = triton_silu_and_block_quant_forward( + input_view, round_scale=round_scale, limit=limit + ) + output_shape = (shape[0], shape[1], shape[2] // 2) + transpose_shape = (shape[2] // 2, shape[0], shape[1]) + output = cls( + shape=output_shape, + dtype=input.dtype, + fp8_dtype=quantizer.dtype, + rowwise_data=x_q.view(output_shape), + rowwise_scale_inv=x_scale, + columnwise_data=xt_q.view(transpose_shape), + columnwise_scale_inv=xt_scale, + quantizer=quantizer, + requires_grad=input.requires_grad, + is_2D_scaled=False, + ) + return output + + @staticmethod + def backward(ctx, grad_output): + shape = grad_output.shape + grad_output_view = grad_output.view(shape[0] * shape[1], shape[2]) + (input,) = ctx.saved_tensors + grad_quantizer = ctx.grad_quantizer + input_view = input.view(shape[0] * shape[1], shape[2] * 2) + round_scale = grad_quantizer.force_pow_2_scales + x_q, x_scale, xt_q, xt_scale = triton_silu_and_block_quant_backward( + grad_output_view, input_view, round_scale=round_scale, limit=ctx.limit + ) + output = ctx.cls( + shape=ctx.shape, + dtype=grad_output.dtype, + fp8_dtype=grad_quantizer.dtype, + rowwise_data=x_q.view(ctx.shape), + rowwise_scale_inv=x_scale, + columnwise_data=xt_q.view(ctx.shape[2], shape[0], shape[1]), + columnwise_scale_inv=xt_scale, + quantizer=grad_quantizer, + requires_grad=ctx.input_requires_grad, + is_2D_scaled=False, + ) + + return output, None, None, None, None + + +def block_silu_impl(input, quantizer, grad_quantizer, cls, limit=None): + output = BlockSiluFunction.apply(input, quantizer, grad_quantizer, cls, limit) + return output + + +class BlockBatchWeightedSiluFunction(torch.autograd.Function): + @staticmethod + def forward( + ctx, + input, + weights, + counts, + splits, + quantizers, + grad_quantizers, + cls, + limit, + is_recomputing, + ): + shape = input.shape + ctx.grad_quantizers = grad_quantizers + ctx.input_requires_grad = input.requires_grad + ctx.shape = shape + ctx.splits = splits + ctx.cls = cls + ctx.limit = limit + ctx.save_for_backward(input, weights, counts) + + if is_recomputing is None: + output_mode = 2 + elif is_recomputing: + output_mode = 1 + else: + output_mode = 0 + + round_scale = quantizers[0].force_pow_2_scales + + (x_q, x_scale, xt_q, xt_scale) = ( + triton_batch_weighted_silu_and_block_quant_forward( + input, + weights, + counts, + splits=splits, + limit=limit, + round_scale=round_scale, + output_mode=output_mode, + ) + ) + + output = cls( + shape=x_q.shape, + dtype=input.dtype, + fp8_dtype=quantizers[0].dtype, + rowwise_data=x_q, + rowwise_scale_inv=x_scale, + columnwise_data=xt_q, + columnwise_scale_inv=xt_scale, + quantizer=quantizers, + requires_grad=input.requires_grad, + is_2D_scaled=False, + ) + return output + + @staticmethod + def backward(ctx, grad_output): + input, weights, counts = ctx.saved_tensors + grad_quantizers = ctx.grad_quantizers + round_scale = grad_quantizers[0].force_pow_2_scales + + (x_q, x_scale, wgrad, xt_q, xt_scale) = ( + triton_batch_weighted_silu_and_block_quant_backward( + grad_output, + input, + weights, + counts, + splits=ctx.splits, + round_scale=round_scale, + limit=ctx.limit, + ) + ) + output = ctx.cls( + shape=ctx.shape, + dtype=grad_output.dtype, + fp8_dtype=grad_quantizers[0].dtype, + rowwise_data=x_q, + rowwise_scale_inv=x_scale, + columnwise_data=xt_q, + columnwise_scale_inv=xt_scale, + quantizer=grad_quantizers, + requires_grad=ctx.input_requires_grad, + is_2D_scaled=False, + ) + + return output, wgrad, None, None, None, None, None, None, None + + +def block_batch_weighted_silu_impl( + input, + weights, + counts, + splits, + quantizers, + grad_quantizers, + cls, + limit=None, + is_recomputing=None, +): + assert input.ndim == 2 + output = BlockBatchWeightedSiluFunction.apply( + input, + weights, + counts, + splits, + quantizers, + grad_quantizers, + cls, + limit, + is_recomputing, + ) + return output + + +class MXFP8SiluFunction(torch.autograd.Function): + @staticmethod + # bias is an optional argument + def forward(ctx, input, quantizer, grad_quantizer, cls, limit): + shape = input.shape + assert len(shape) == 3 + input = input.view(shape[0] * shape[1], shape[2]) + ctx.grad_quantizer = grad_quantizer + ctx.input_requires_grad = input.requires_grad + ctx.shape = shape + ctx.cls = cls + ctx.limit = limit + ctx.save_for_backward(input) + x_q, x_scale, xt_q, xt_scale = triton_silu_and_mxfp8_quant_forward( + input, limit=limit + ) + + output_shape = (shape[0], shape[1], shape[2] // 2) + # transpose_shape = (shape[2]//2, shape[0], shape[1]) + output = cls( + shape=output_shape, + dtype=input.dtype, + fp8_dtype=quantizer.dtype, + rowwise_data=x_q.view(output_shape), + rowwise_scale_inv=x_scale, + columnwise_data=xt_q.view(output_shape), + columnwise_scale_inv=xt_scale, + quantizer=quantizer, + requires_grad=input.requires_grad, + ) + return output + + @staticmethod + def backward(ctx, grad_output): + shape = grad_output.shape + grad_output = grad_output.view(shape[0] * shape[1], shape[2]) + (input,) = ctx.saved_tensors + grad_quantizer = ctx.grad_quantizer + x_q, x_scale, xt_q, xt_scale = triton_silu_and_mxfp8_quant_backward( + grad_output, input, limit=ctx.limit + ) + output = ctx.cls( + shape=ctx.shape, + dtype=grad_output.dtype, + fp8_dtype=grad_quantizer.dtype, + rowwise_data=x_q.view(ctx.shape) if x_q is not None else None, + rowwise_scale_inv=x_scale, + columnwise_data=xt_q.view(ctx.shape) if xt_q is not None else None, + columnwise_scale_inv=xt_scale, + quantizer=grad_quantizer, + requires_grad=ctx.input_requires_grad, + ) + + return output, None, None, None, None + + +def mxfp8_silu_impl(input, quantizer, grad_quantizer, cls, limit=None): + # input: [length,bs,dim] + output = MXFP8SiluFunction.apply(input, quantizer, grad_quantizer, cls, limit) + return output + + +class MXFP8BatchWeightedSiluFunction(torch.autograd.Function): + @staticmethod + def forward( + ctx, + input, + weights, + counts, + splits, + quantizers, + grad_quantizers, + cls, + limit, + is_recomputing, + ): + shape = input.shape + ctx.grad_quantizers = grad_quantizers + ctx.input_requires_grad = input.requires_grad + ctx.shape = shape + ctx.splits = splits + ctx.cls = cls + ctx.limit = limit + ctx.save_for_backward(input, weights, counts) + + if is_recomputing is None: + output_mode = 2 + elif is_recomputing: + output_mode = 1 + else: + output_mode = 0 + + (x_q, x_scale, xt_q, xt_scale) = ( + triton_batch_weighted_silu_and_mxfp8_quant_forward( + input, + weights, + counts, + splits=splits, + limit=limit, + output_mode=output_mode, + ) + ) + + output = cls( + shape=x_q.shape, + dtype=input.dtype, + fp8_dtype=quantizers[0].dtype, + rowwise_data=x_q, + rowwise_scale_inv=x_scale, + columnwise_data=xt_q, + columnwise_scale_inv=xt_scale, + quantizer=quantizers, + requires_grad=input.requires_grad, + ) + return output + + @staticmethod + def backward(ctx, grad_output): + input, weights, counts = ctx.saved_tensors + quantizers = ctx.grad_quantizers + (x_q, x_scale, wgrad, xt_q, xt_scale) = ( + triton_batch_weighted_silu_and_mxfp8_quant_backward( + grad_output, input, weights, counts, splits=ctx.splits, limit=ctx.limit + ) + ) + output = ctx.cls( + shape=ctx.shape, + dtype=grad_output.dtype, + fp8_dtype=quantizers[0].dtype, + rowwise_data=x_q, + rowwise_scale_inv=x_scale, + columnwise_data=xt_q, + columnwise_scale_inv=xt_scale, + quantizer=quantizers, + requires_grad=ctx.input_requires_grad, + ) + + return output, wgrad, None, None, None, None, None, None, None + + +def mxfp8_batch_weighted_silu_impl( + input, + weights, + counts, + splits, + quantizers, + grad_quantizers, + cls, + limit=None, + is_recomputing=None, +): + assert input.ndim == 2 + output = MXFP8BatchWeightedSiluFunction.apply( + input, + weights, + counts, + splits, + quantizers, + grad_quantizers, + cls, + limit, + is_recomputing, + ) + return output + + +class SmoothSiluFunction(torch.autograd.Function): + @staticmethod + # bias is an optional argument + def forward(ctx, input, quantizer, grad_quantizer, cls): + shape = input.shape + assert len(shape) == 3 + input = input.view(shape[0] * shape[1], shape[2]) + ctx.grad_quantizer = grad_quantizer + ctx.shape = shape + ctx.cls = cls + ctx.save_for_backward(input) + x_q, x_scale = triton_silu_and_smooth_quant_forward( + input, + smooth_scale=quantizer.smooth_scale, + round_scale=quantizer.force_pow_2_scales, + ) + output = cls( + shape=(shape[0], shape[1], shape[2] // 2), + dtype=input.dtype, + fp8_dtype=quantizer.dtype, + rowwise_data=x_q, + rowwise_scale_inv=x_scale, + columnwise_data=None, + columnwise_scale_inv=quantizer.smooth_scale, + quantizer=quantizer, + requires_grad=input.requires_grad, + ) + return output + + @staticmethod + def backward(ctx, grad_output): + shape = grad_output.shape + grad_output = grad_output.view(shape[0] * shape[1], shape[2]) + (input,) = ctx.saved_tensors + grad_quantizer = ctx.grad_quantizer + # we use requant implementation, + # so must use round_scale=True to avoid second quantization error + round_scale = True # quantizer.force_pow_2_scales + x_q, x_scale, xt_q, xt_scale = triton_silu_and_smooth_quant_backward( + grad_output, + input, + smooth_scale=grad_quantizer.smooth_scale_inv, + transpose_smooth_scale=grad_quantizer.transpose_smooth_scale_inv, + reverse=True, + round_scale=round_scale, + ) + + output = ctx.cls( + shape=ctx.shape, + dtype=grad_output.dtype, + fp8_dtype=grad_quantizer.dtype, + rowwise_data=x_q, + rowwise_scale_inv=x_scale, + columnwise_data=xt_q, + columnwise_scale_inv=xt_scale, + quantizer=grad_quantizer, + requires_grad=False, + ) + + return output, None, None, None + + +def smooth_silu_impl(input, quantizer, grad_quantizer, cls): + # input: [length, bs, dim] + output = SmoothSiluFunction.apply(input, quantizer, grad_quantizer, cls) + return output + + +class SmoothBatchWeightedSiluFunction(torch.autograd.Function): + @staticmethod + def forward(ctx, inputs, weights, counts, splits, quantizers, grad_quantizers, cls): + shape = inputs.shape + ctx.grad_quantizers = grad_quantizers + ctx.shape = shape + ctx.save_for_backward(inputs, weights, counts) + ctx.splits = splits + ctx.cls = cls + smooth_scales = torch.stack([x.smooth_scale for x in quantizers], 0) + round_scale = quantizers[0].force_pow_2_scales + x_q, x_scale = triton_batch_weighted_silu_and_smooth_quant_forward( + inputs, weights, counts, smooth_scale=smooth_scales, round_scale=round_scale + ) + output = cls( + shape=x_q.shape, + dtype=inputs.dtype, + fp8_dtype=quantizers[0].dtype, + rowwise_data=x_q, + rowwise_scale_inv=x_scale, + columnwise_data=None, + columnwise_scale_inv=smooth_scales, + quantizer=quantizers, + requires_grad=inputs.requires_grad, + ) + return output + + @staticmethod + def backward(ctx, grad_output): + inputs, weights, counts = ctx.saved_tensors + grad_quantizers = ctx.grad_quantizers + # smooth_scale_inv may not exist in forward step + smooth_scales = torch.stack([x.smooth_scale_inv for x in grad_quantizers], 0) + transpose_smooth_scales = torch.cat( + [x.transpose_smooth_scale_inv for x in grad_quantizers], 0 + ) + # we use requant implementation, so must use round_scale=True to avoid second quantization error + round_scale = True # grad_quantizers[0].force_pow_2_scales + x_q, x_scale, wgrad, xt_q, xt_scale = ( + triton_batch_weighted_silu_and_smooth_quant_backward( + grad_output, + inputs, + weights, + counts, + smooth_scale=smooth_scales, + transpose_smooth_scale=transpose_smooth_scales, + splits=ctx.splits, + reverse=True, + round_scale=round_scale, + ) + ) + output = ctx.cls( + shape=ctx.shape, + dtype=grad_output.dtype, + fp8_dtype=grad_quantizers[0].dtype, + rowwise_data=x_q, + rowwise_scale_inv=x_scale, + columnwise_data=xt_q, + columnwise_scale_inv=xt_scale, + quantizer=grad_quantizers, + requires_grad=False, + ) + + return output, wgrad, None, None, None, None, None + + +def smooth_batch_weighted_silu_impl( + inputs, weights, counts, splits, quantizers, grad_quantizers, cls +): + # TODO(nanxiao): support recomputation + assert inputs.ndim == 2 + output = SmoothBatchWeightedSiluFunction.apply( + inputs, weights, counts, splits, quantizers, grad_quantizers, cls + ) + return output diff --git a/linghe/facade/smooth_quant_linear.py b/linghe/facade/smooth_quant_linear.py index dd284e2..e05c756 100644 --- a/linghe/facade/smooth_quant_linear.py +++ b/linghe/facade/smooth_quant_linear.py @@ -7,20 +7,19 @@ import torch - -from linghe.quant.smooth import triton_smooth_quant, \ - triton_transpose_smooth_quant -from linghe.utils.transpose import triton_transpose_and_pad +from linghe.quant.smooth import triton_smooth_quant, triton_transpose_smooth_quant from linghe.utils.reduce import triton_abs_max +from linghe.utils.transpose import triton_transpose_and_pad + class _SmoothQuantLinear(torch.autograd.Function): @staticmethod def forward( - ctx, - input: torch.Tensor, - weight: torch.Tensor, - bias: Optional[torch.Tensor], - smooth_scale: torch.Tensor, + ctx, + input: torch.Tensor, + weight: torch.Tensor, + bias: Optional[torch.Tensor], + smooth_scale: torch.Tensor, ): ctx.input_requires_grad = input.requires_grad ctx.weight_requires_grad = weight.requires_grad @@ -33,15 +32,21 @@ def forward( input = input.view(-1, input.shape[-1]) - x_q, x_scale, x_maxs = triton_smooth_quant(input, 1 / smooth_scale, round_scale=round_scale) - w_q, w_scale, w_maxs = triton_smooth_quant(weight, smooth_scale, round_scale=round_scale) - - output = torch._scaled_mm(x_q, - w_q.t(), - scale_a=x_scale.view(-1, 1), - scale_b=w_scale.view(1, -1), - out_dtype=ctx.out_dtype, - use_fast_accum=True) + x_q, x_scale, x_maxs = triton_smooth_quant( + input, 1 / smooth_scale, round_scale=round_scale + ) + w_q, w_scale, w_maxs = triton_smooth_quant( + weight, smooth_scale, round_scale=round_scale + ) + + output = torch._scaled_mm( + x_q, + w_q.t(), + scale_a=x_scale.view(-1, 1), + scale_b=w_scale.view(1, -1), + out_dtype=ctx.out_dtype, + use_fast_accum=True, + ) if bias is not None: output += bias @@ -51,7 +56,11 @@ def forward( x_scale if ctx.weight_requires_grad else None, w_q if ctx.input_requires_grad else None, w_scale if ctx.input_requires_grad else None, - smooth_scale if ctx.weight_requires_grad or ctx.weight_requires_grad else None + ( + smooth_scale + if ctx.weight_requires_grad or ctx.weight_requires_grad + else None + ), ] ctx.save_for_backward(*saved_tensors) @@ -59,40 +68,39 @@ def forward( return output.view(out_shape) @staticmethod - def backward( - ctx, - output_grad: torch.Tensor - ): + def backward(ctx, output_grad: torch.Tensor): x_q, x_s, w_q, w_s, smooth_scale = ctx.saved_tensors output_grad = output_grad.view(-1, output_grad.shape[-1]) round_scale = ctx.round_scale - y_q, y_scale, y_maxs = triton_smooth_quant(output_grad, - w_s, - reverse=True, - round_scale=round_scale) + y_q, y_scale, y_maxs = triton_smooth_quant( + output_grad, w_s, reverse=True, round_scale=round_scale + ) wt_q = triton_transpose_and_pad(w_q, pad=True) - dx = torch._scaled_mm(y_q, - wt_q.t(), - scale_a=y_scale.view(-1, 1), - scale_b=smooth_scale.view(1, -1), - out_dtype=ctx.out_dtype, - use_fast_accum=True) - - yt_q, yt_scale = triton_transpose_smooth_quant(output_grad, - x_s, - reverse=True , - round_scale=round_scale) + dx = torch._scaled_mm( + y_q, + wt_q.t(), + scale_a=y_scale.view(-1, 1), + scale_b=smooth_scale.view(1, -1), + out_dtype=ctx.out_dtype, + use_fast_accum=True, + ) + + yt_q, yt_scale = triton_transpose_smooth_quant( + output_grad, x_s, reverse=True, round_scale=round_scale + ) xt_q = triton_transpose_and_pad(x_q, pad=True) - dw = torch._scaled_mm(yt_q, - xt_q.t(), - scale_a=yt_scale.view(-1, 1), - scale_b=1/smooth_scale.view(1, -1), - out_dtype=ctx.out_dtype, - use_fast_accum=True) + dw = torch._scaled_mm( + yt_q, + xt_q.t(), + scale_a=yt_scale.view(-1, 1), + scale_b=1 / smooth_scale.view(1, -1), + out_dtype=ctx.out_dtype, + use_fast_accum=True, + ) db = None if ctx.bias_requires_grad: @@ -105,13 +113,14 @@ class SmoothQuantLinear(torch.nn.Module): """ a naive implementation of smooth quantization linear """ + def __init__( - self, - in_features: int, - out_features: int, - bias: bool = True, - device=None, - dtype=None + self, + in_features: int, + out_features: int, + bias: bool = True, + device=None, + dtype=None, ): """ Args: @@ -125,11 +134,12 @@ def __init__( self.in_features = in_features self.out_features = out_features self.weight = torch.nn.parameter.Parameter( - torch.empty((out_features, in_features), device=device, - dtype=dtype)) + torch.empty((out_features, in_features), device=device, dtype=dtype) + ) if bias: self.bias = torch.nn.parameter.Parameter( - torch.empty(out_features, device=device, dtype=dtype)) + torch.empty(out_features, device=device, dtype=dtype) + ) else: self.bias = None @@ -145,13 +155,12 @@ def forward(self, input: torch.Tensor) -> torch.Tensor: if self.smooth_update_step % self.gap_step == 0: input_maxs = triton_abs_max(input) - weight_maxs = triton_abs_max(self.weight.data) + weight_maxs = triton_abs_max(self.weight) self.smooth_scale = torch.sqrt(input_maxs * weight_maxs) - output = _SmoothQuantLinear.apply(input, - self.weight, - self.bias, - self.smooth_scale) + output = _SmoothQuantLinear.apply( + input, self.weight, self.bias, self.smooth_scale + ) self.smooth_update_step += 1 else: output = input @ self.weight.t() diff --git a/linghe/facade/topk.py b/linghe/facade/topk.py new file mode 100644 index 0000000..c41a520 --- /dev/null +++ b/linghe/facade/topk.py @@ -0,0 +1,113 @@ +# -*- coding: utf-8 -*- +""" +Copyright (c) Ant Financial Service Group and its affiliates. +""" + +import torch + +from linghe.utils.topk import ( + triton_topk_forward, + triton_topk_backward, + triton_group_topk_score_forward, + triton_group_topk_score_backward, +) + + +class TopkFunction(torch.autograd.Function): + """""" + + @staticmethod + def forward(ctx, x, k, dim, sorted, impl): + values, indices = triton_topk_forward(x, k, dim=dim, sorted=sorted, impl=impl) + ctx.dim = dim + ctx.shape = x.shape + ctx.save_for_backward(indices) + return values, indices + + @staticmethod + def backward(ctx, grad_output, grad_indices): + (indices,) = ctx.saved_tensors + grad_input = triton_topk_backward( + grad_output, indices, ctx.shape[-1], dim=ctx.dim + ) + return grad_input, None, None, None, None + + +def fused_topk(x, k, dim=-1, sorted=True, impl="iter"): + """ + topk + Args: + x: input tensor + k: topk + dim: dimension to apply topk, only support -1 currently + Returns: + values: topk values + indices: topk indices + """ + return TopkFunction.apply(x, k, dim, sorted, impl) + + +class GroupTopkScoreFunction(torch.autograd.Function): + """""" + + @staticmethod + def forward( + ctx, + x, + topk, + expert_bias, + num_groups, + group_topk, + scaling_factor, + score_function, + ): + probs, routing_map, counts = triton_group_topk_score_forward( + x, + topk, + expert_bias=expert_bias, + num_groups=num_groups, + group_topk=group_topk, + scaling_factor=scaling_factor, + score_function=score_function, + ) + ctx.save_for_backward(x, routing_map) + ctx.scaling_factor = scaling_factor + ctx.score_function = score_function + return probs, routing_map, counts + + @staticmethod + def backward(ctx, grad_output, grad_map, grad_counts): + x, routing_map = ctx.saved_tensors + grad_input = triton_group_topk_score_backward( + grad_output, x, routing_map, scaling_factor=ctx.scaling_factor + ) + return grad_input, None, None, None, None, None, None + + +def group_topk_score( + x, + topk, + expert_bias=None, + num_groups=32, + group_topk=4, + scaling_factor=1.0, + score_function="sigmoid", +): + """ + group topk with softmax/sigmoid function + Args: + x: input logit tensor + topk: topk + expert_bias: expert bias + num_groups: number of groups + group_topk: group to apply topk + scaling_factor: scaling factor + score_function: scaling function + Returns: + probs: topk probs + routing_map: topk binary map + counts: token count per expert + """ + return GroupTopkScoreFunction.apply( + x, topk, expert_bias, num_groups, group_topk, scaling_factor, score_function + ) diff --git a/linghe/facade/transpose.py b/linghe/facade/transpose.py index 9332e18..961babf 100644 --- a/linghe/facade/transpose.py +++ b/linghe/facade/transpose.py @@ -8,24 +8,28 @@ from linghe.utils.transpose import triton_transpose -class TransposeDim01Function(torch.autograd.Function): +class TransposeFunction(torch.autograd.Function): """""" + @staticmethod - def forward(ctx, x): - return triton_transpose(x, dim0=0, dim1=1) + def forward(ctx, x, inner): + ctx.inner = inner + return triton_transpose(x, inner=inner) @staticmethod def backward(ctx, grad_output): - return triton_transpose(grad_output, dim0=0, dim1=1) + return triton_transpose(grad_output, inner=ctx.inner) -def transpose_dim01(x): +def transpose(x, inner=True): """ - transpose a tensor with the first two dims, x.ndims should not greater than 4 + transpose a tensor, x.ndims should not greater than 4 Args: x: input tensor - + inner: + if True, transpose the first two dimensions + if False, transpose the last two dimensions Returns: a transposed tensor """ - return TransposeDim01Function.apply(x) \ No newline at end of file + return TransposeFunction.apply(x, inner) diff --git a/linghe/gemm/blockwise_fp8_gemm.py b/linghe/gemm/blockwise_fp8_gemm.py index 1d396b7..b2c0903 100644 --- a/linghe/gemm/blockwise_fp8_gemm.py +++ b/linghe/gemm/blockwise_fp8_gemm.py @@ -3,32 +3,27 @@ Copyright (c) Ant Financial Service Group and its affiliates. """ -import os import torch import triton import triton.language as tl -from triton import Config -# adapt from deepseek # os.environ["TRITON_PRINT_AUTOTUNING"] = "1" - - @triton.jit -def fp8_gemm_bb_kernel( - a_ptr, - b_ptr, - c_ptr, - a_s_ptr, - b_s_ptr, - M, - N: tl.constexpr, - K: tl.constexpr, - BLOCK_SIZE_K: tl.constexpr, - BLOCK_SIZE_M: tl.constexpr, - BLOCK_SIZE_N: tl.constexpr, +def fp8_blockwise_gemm_kernel( + a_ptr, + b_ptr, + c_ptr, + a_s_ptr, + b_s_ptr, + M, + N: tl.constexpr, + K: tl.constexpr, + BLOCK_SIZE_K: tl.constexpr, + BLOCK_SIZE_M: tl.constexpr, + BLOCK_SIZE_N: tl.constexpr, ): pid_m = tl.program_id(axis=0) pid_n = tl.program_id(axis=1) @@ -45,10 +40,8 @@ def fp8_gemm_bb_kernel( for i in range(0, k): a_s = tl.load(a_s_ptr + pid_m * nb + i) b_s = tl.load(b_s_ptr + pid_n * nb + i) - a = tl.load(a_ptrs, mask=offs_k[None, :] < K - i * BLOCK_SIZE_K, - other=0.0) - b = tl.load(b_ptrs, mask=offs_k[None, :] < K - i * BLOCK_SIZE_K, - other=0.0) + a = tl.load(a_ptrs, mask=offs_k[None, :] < K - i * BLOCK_SIZE_K, other=0.0) + b = tl.load(b_ptrs, mask=offs_k[None, :] < K - i * BLOCK_SIZE_K, other=0.0) accumulator += tl.dot(a, tl.trans(b)) * (a_s * b_s) # accumulator = tl.dot(a, tl.trans(b), accumulator) # accumulator += (accumulators-accumulator) * scale @@ -62,108 +55,39 @@ def fp8_gemm_bb_kernel( tl.store(c_ptrs, c, mask=mask) -# use for hadamard quantization, too slow on H800 -def triton_bb_fp8_gemm(a: torch.Tensor, - b: torch.Tensor, - a_s: torch.Tensor, - b_s: torch.Tensor, - out_dtype=torch.bfloat16, - block_size=128): +# use for hadamard transform and blockwise quantization, too slow on H800 +def triton_blockwise_fp8_gemm( + a: torch.Tensor, + b: torch.Tensor, + a_s: torch.Tensor, + b_s: torch.Tensor, + out_dtype=torch.bfloat16, + block_size=128, +): assert a.is_contiguous() and b.is_contiguous() assert a_s.is_contiguous() and b_s.is_contiguous() K = a.size(-1) M = a.numel() // K N = b.size(0) c = torch.empty(M, N, dtype=out_dtype, device=a.device) - grid = lambda META: (triton.cdiv(M, META["BLOCK_SIZE_M"]), - triton.cdiv(N, META["BLOCK_SIZE_N"])) # noqa - - fp8_gemm_bb_kernel[grid](a, b, c, a_s, b_s, - M, N, K, - BLOCK_SIZE_K=block_size, - BLOCK_SIZE_M=block_size, - BLOCK_SIZE_N=block_size, - num_warps=8, - num_stages=4 - ) - return c - - - -fp8_gemm_configs = [ - Config({"BLOCK_SIZE_M": block_m, "BLOCK_SIZE_N": block_n}, - num_stages=num_stages, num_warps=8) - for block_m in [32, 64, 128] - for block_n in [32, 64, 128] - for num_stages in [3, 4, 5, 6] -] - -@triton.autotune(configs=fp8_gemm_configs, key=["N", "K"]) -@triton.jit -def fp8_gemm_tt_kernel( - a_ptr, - b_ptr, - c_ptr, - a_s_ptr, - b_s_ptr, + grid = lambda META: ( + triton.cdiv(M, META["BLOCK_SIZE_M"]), + triton.cdiv(N, META["BLOCK_SIZE_N"]), + ) # noqa + + fp8_blockwise_gemm_kernel[grid]( + a, + b, + c, + a_s, + b_s, M, - N: tl.constexpr, - K: tl.constexpr, - BLOCK_SIZE_K: tl.constexpr, - BLOCK_SIZE_M: tl.constexpr, - BLOCK_SIZE_N: tl.constexpr, -): - # a and b all tilewise quantization. - pid_m = tl.program_id(axis=0) - pid_n = tl.program_id(axis=1) - k = tl.cdiv(K, BLOCK_SIZE_K) - offs_m = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M - offs_n = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N - offs_k = tl.arange(0, BLOCK_SIZE_K) - a_ptrs = a_ptr + offs_m[:, None] * K + offs_k[None, :] - b_ptrs = b_ptr + offs_n[None, :] * K + offs_k[:, None] - a_s_ptrs = a_s_ptr + offs_m * k - b_s_ptrs = b_s_ptr + offs_n * k - - accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32) - for i in range(k): - a = tl.load(a_ptrs, mask=offs_k[None, :] < K - i * BLOCK_SIZE_K, - other=0.0) - b = tl.load(b_ptrs, mask=offs_k[:, None] < K - i * BLOCK_SIZE_K, - other=0.0) - a_s = tl.load(a_s_ptrs) - b_s = tl.load(b_s_ptrs) - accumulator += tl.dot(a, b) * a_s[:, None] * b_s[None, :] - a_ptrs += BLOCK_SIZE_K - b_ptrs += BLOCK_SIZE_K - a_s_ptrs += 1 - b_s_ptrs += 1 - - c = accumulator.to(c_ptr.dtype.element_ty) - offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) - offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) - c_ptrs = c_ptr + offs_m[:, None] * N + offs_n[None, :] - mask = (offs_m[:, None] < M) & (offs_n[None, :] < N) - tl.store(c_ptrs, c, mask=mask) - - -# use to mock mxfp8 gemm, too slow on H800 -def triton_tt_fp8_gemm(a: torch.Tensor, - b: torch.Tensor, - a_s: torch.Tensor, - b_s: torch.Tensor, - out_dtype=torch.bfloat16, - block_size=128): - assert a.is_contiguous() and b.is_contiguous() - assert a_s.is_contiguous() and b_s.is_contiguous() - K = a.size(-1) - M = a.numel() // K - N = b.size(0) - c = torch.empty(*a.size()[:-1], N, dtype=out_dtype, device=a.device) - grid = lambda META: (triton.cdiv(M, META["BLOCK_SIZE_M"]), - triton.cdiv(N, META["BLOCK_SIZE_N"])) # noqa - fp8_gemm_tt_kernel[grid](a, b, c, - a_s, b_s, - M, N, K, - BLOCK_SIZE_K=block_size) + N, + K, + BLOCK_SIZE_K=block_size, + BLOCK_SIZE_M=block_size, + BLOCK_SIZE_N=block_size, + num_warps=8, + num_stages=4, + ) return c diff --git a/linghe/gemm/channelwise_fp8_gemm.py b/linghe/gemm/channelwise_fp8_gemm.py index a6413b1..fa4684f 100644 --- a/linghe/gemm/channelwise_fp8_gemm.py +++ b/linghe/gemm/channelwise_fp8_gemm.py @@ -7,7 +7,6 @@ import triton import triton.language as tl - # os.environ["TRITON_PRINT_AUTOTUNING"] = "1" @@ -22,27 +21,28 @@ # # for num_warps in [8] # ] + # @triton.autotune(configs=fp8_gemm_configs, key=["M", "N", "K"]) @triton.jit def scaled_mm_kernel( - a_ptr, - b_ptr, - c_ptr, - a_scale_ptr, - b_scale_ptr, - N, - K, - ACCUM: tl.constexpr, - EVEN: tl.constexpr, - BLOCK_SIZE_K: tl.constexpr, - BLOCK_SIZE_M: tl.constexpr, - BLOCK_SIZE_N: tl.constexpr, + a_ptr, + b_ptr, + c_ptr, + a_scale_ptr, + b_scale_ptr, + N, + K, + ACCUM: tl.constexpr, + EVEN: tl.constexpr, + BLOCK_SIZE_K: tl.constexpr, + BLOCK_SIZE_M: tl.constexpr, + BLOCK_SIZE_N: tl.constexpr, ): pid_m = tl.program_id(axis=0) pid_n = tl.program_id(axis=1) k = tl.cdiv(K, BLOCK_SIZE_K) - offs_m = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) - offs_n = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) + offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) offs_k = tl.arange(0, BLOCK_SIZE_K) a_ptrs = a_ptr + offs_m[:, None] * K + offs_k[None, :] b_ptrs = b_ptr + offs_n[None, :] * K + offs_k[:, None] @@ -80,13 +80,15 @@ def scaled_mm_kernel( tl.store(c_ptrs, accumulator) -def triton_scaled_mm(a: torch.Tensor, - b: torch.Tensor, - a_scale: torch.Tensor, - b_scale: torch.Tensor, - out_dtype=torch.float32, - c=None, - accum=True): +def triton_scaled_mm( + a: torch.Tensor, + b: torch.Tensor, + a_scale: torch.Tensor, + b_scale: torch.Tensor, + out_dtype=torch.float32, + c=None, + accum=True, +): """ similar to torch._scaled_mm, support accumulating gemm output to c and low precision output tensor @@ -113,17 +115,21 @@ def triton_scaled_mm(a: torch.Tensor, BLOCK_SIZE_N = 256 EVEN = K % BLOCK_SIZE_K == 0 grid = (M // BLOCK_SIZE_M, N // BLOCK_SIZE_N) # noqa - scaled_mm_kernel[grid](a, b, c, - a_scale, - b_scale, - N, K, - ACCUM, - EVEN, - BLOCK_SIZE_K, - BLOCK_SIZE_M, - BLOCK_SIZE_N, - num_stages=3, - num_warps=8 - ) + scaled_mm_kernel[grid]( + a, + b, + c, + a_scale, + b_scale, + N, + K, + ACCUM, + EVEN, + BLOCK_SIZE_K, + BLOCK_SIZE_M, + BLOCK_SIZE_N, + num_stages=3, + num_warps=8, + ) return c diff --git a/linghe/gemm/fp32_gemm.py b/linghe/gemm/fp32_gemm.py index 5ad778c..6e0dac6 100644 --- a/linghe/gemm/fp32_gemm.py +++ b/linghe/gemm/fp32_gemm.py @@ -3,53 +3,40 @@ Copyright (c) Ant Financial Service Group and its affiliates. """ -from typing import Optional - +import os import torch import triton import triton.language as tl - +from triton import Config # os.environ["TRITON_PRINT_AUTOTUNING"] = "1" -# fp32_gemm_configs = [ -# Config({"BLOCK_SIZE_K": block_k, "BLOCK_SIZE_M": block_m, "BLOCK_SIZE_N": block_n}, num_stages=num_stages, num_warps=num_warps) -# for block_k in [64, 128, 256] -# for block_m in [32, 64, 128] -# for block_n in [32, 64, 128] -# for num_stages in [2, 3, 4, 5] -# for num_warps in [4, 8] -# ] - - -# @triton.autotune(configs=fp32_gemm_configs, key=["M", "N", "K"]) @triton.jit def fp32_gemm_kernel( - a_ptr, - b_ptr, - c_ptr, - M, - N: tl.constexpr, - K: tl.constexpr, - BLOCK_SIZE_K: tl.constexpr, - BLOCK_SIZE_M: tl.constexpr, - BLOCK_SIZE_N: tl.constexpr, + a_ptr, + b_ptr, + c_ptr, + M, + N: tl.constexpr, + K: tl.constexpr, + BLOCK_SIZE_K: tl.constexpr, + BLOCK_SIZE_M: tl.constexpr, + BLOCK_SIZE_N: tl.constexpr, ): pid_m = tl.program_id(axis=0) pid_n = tl.program_id(axis=1) k = tl.cdiv(K, BLOCK_SIZE_K) - offs_m = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) - offs_n = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) + offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) offs_k = tl.arange(0, BLOCK_SIZE_K) a_ptrs = a_ptr + offs_m[:, None] * K + offs_k[None, :] b_ptrs = b_ptr + offs_n[None, :] * K + offs_k[:, None] c = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32) for i in range(k): - a = tl.load(a_ptrs).to(tl.float32) - b = tl.load(b_ptrs).to(tl.float32) - # c += tl.dot(a, b) + a = tl.load(a_ptrs) # .to(tl.float32) + b = tl.load(b_ptrs) # .to(tl.float32) c = tl.dot(a, b, c) a_ptrs += BLOCK_SIZE_K b_ptrs += BLOCK_SIZE_K @@ -59,8 +46,7 @@ def fp32_gemm_kernel( tl.store(c_ptrs, c) - -def triton_fp32_gemm(a: torch.Tensor, b: torch.Tensor): +def triton_fp32_gemm(x: torch.Tensor, w: torch.Tensor): """ return fp32 gemm result with fp16/bf16 inputs, it's mainly used for MoE router GEMM @@ -72,47 +58,53 @@ def triton_fp32_gemm(a: torch.Tensor, b: torch.Tensor): Returns: c: output with fp32 precision """ - assert a.is_contiguous() and b.is_contiguous() - M, K = a.size() - N, K = b.size() - assert N >= 128 - c = torch.empty(M, N, dtype=torch.float32, device=a.device) - grid = lambda META: (triton.cdiv(M, META["BLOCK_SIZE_M"]), - triton.cdiv(N, META["BLOCK_SIZE_N"])) # noqa - BLOCK_SIZE_K = 128 - BLOCK_SIZE_M = 32 - BLOCK_SIZE_N = 128 + assert x.is_contiguous() and w.is_contiguous() + M, K = x.size() + N, K = w.size() + assert M % 32 == 0 and K % 64 == 0 and N % 16 == 0, f"{M=} {K=} {N=}" + c = torch.empty(M, N, dtype=torch.float32, device=x.device) + grid = lambda META: ( + triton.cdiv(M, META["BLOCK_SIZE_M"]), + triton.cdiv(N, META["BLOCK_SIZE_N"]), + ) # noqa + BLOCK_SIZE_K = 64 + BLOCK_SIZE_M = max([x for x in [32, 64] if M % x == 0]) + BLOCK_SIZE_N = max([x for x in [16, 32, 64] if N % x == 0]) num_warps = 4 num_stages = 3 - fp32_gemm_kernel[grid](a, b, c, - M, N, K, - BLOCK_SIZE_K, - BLOCK_SIZE_M, - BLOCK_SIZE_N, - num_warps=num_warps, - num_stages=num_stages - ) + fp32_gemm_kernel[grid]( + x, + w, + c, + M, + N, + K, + BLOCK_SIZE_K, + BLOCK_SIZE_M, + BLOCK_SIZE_N, + num_warps=num_warps, + num_stages=num_stages, + ) return c -# @triton.autotune(configs=fp32_gemm_configs, key=["M", "N", "K"]) @triton.jit def fp32_gemm_for_backward_kernel( - a_ptr, - b_ptr, - c_ptr, - M, - N: tl.constexpr, - K: tl.constexpr, - BLOCK_SIZE_K: tl.constexpr, - BLOCK_SIZE_M: tl.constexpr, - BLOCK_SIZE_N: tl.constexpr, + a_ptr, + b_ptr, + c_ptr, + M, + N: tl.constexpr, + K: tl.constexpr, + BLOCK_SIZE_K: tl.constexpr, + BLOCK_SIZE_M: tl.constexpr, + BLOCK_SIZE_N: tl.constexpr, ): pid_m = tl.program_id(axis=0) pid_n = tl.program_id(axis=1) k = tl.cdiv(K, BLOCK_SIZE_K) - offs_m = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) - offs_n = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) + offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) offs_k = tl.arange(0, BLOCK_SIZE_K) a_ptrs = a_ptr + offs_m[:, None] * K + offs_k[None, :] b_ptrs = b_ptr + offs_n[None, :] + offs_k[:, None] * N @@ -122,7 +114,6 @@ def fp32_gemm_for_backward_kernel( for i in range(k): a = tl.load(a_ptrs) b = tl.load(b_ptrs).to(tl.float32) - # c += tl.dot(a, b) c = tl.dot(a, b, c) a_ptrs += BLOCK_SIZE_K b_ptrs += BLOCK_SIZE_K * N @@ -132,8 +123,7 @@ def fp32_gemm_for_backward_kernel( tl.store(c_ptrs, c) -def triton_fp32_gemm_for_backward(a: torch.Tensor, - b: torch.Tensor): +def triton_fp32_gemm_for_backward(y: torch.Tensor, w: torch.Tensor): """ mix precision gemm for backward, a@b.float() Args: @@ -142,56 +132,60 @@ def triton_fp32_gemm_for_backward(a: torch.Tensor, Returns: c: gradient of activation """ - assert a.is_contiguous() and b.is_contiguous() - M, K = a.size() - K, N = b.size() - c = torch.empty((M, N), dtype=b.dtype, device=b.device) - grid = lambda META: (triton.cdiv(M, META["BLOCK_SIZE_M"]), - triton.cdiv(N, META["BLOCK_SIZE_N"])) # noqa - BLOCK_SIZE_K = 128 - BLOCK_SIZE_M = 32 - BLOCK_SIZE_N = 128 + assert y.is_contiguous() and w.is_contiguous() + M, K = y.size() + K, N = w.size() + c = torch.empty((M, N), dtype=w.dtype, device=w.device) + grid = lambda META: ( + triton.cdiv(M, META["BLOCK_SIZE_M"]), + triton.cdiv(N, META["BLOCK_SIZE_N"]), + ) # noqa + BLOCK_SIZE_K = max([x for x in [16, 32, 64] if K % x == 0]) + BLOCK_SIZE_M = max([x for x in [32, 64] if M % x == 0]) + BLOCK_SIZE_N = 64 num_warps = 4 - num_stages = 2 - fp32_gemm_for_backward_kernel[grid](a, b, c, - M, N, K, - BLOCK_SIZE_K, - BLOCK_SIZE_M, - BLOCK_SIZE_N, - num_warps=num_warps, - num_stages=num_stages - ) + num_stages = 3 + fp32_gemm_for_backward_kernel[grid]( + y, + w, + c, + M, + N, + K, + BLOCK_SIZE_K, + BLOCK_SIZE_M, + BLOCK_SIZE_N, + num_warps=num_warps, + num_stages=num_stages, + ) return c -# @triton.autotune(configs=fp32_gemm_configs, key=["M", "N", "K"]) @triton.jit def fp32_gemm_for_update_kernel( - a_ptr, - b_ptr, - c_ptr, - M, - N: tl.constexpr, - K: tl.constexpr, - BLOCK_SIZE_K: tl.constexpr, - BLOCK_SIZE_M: tl.constexpr, - BLOCK_SIZE_N: tl.constexpr, + a_ptr, + b_ptr, + c_ptr, + M, + N: tl.constexpr, + K: tl.constexpr, + BLOCK_SIZE_K: tl.constexpr, + BLOCK_SIZE_M: tl.constexpr, + BLOCK_SIZE_N: tl.constexpr, ): pid_m = tl.program_id(axis=0) pid_n = tl.program_id(axis=1) k = tl.cdiv(K, BLOCK_SIZE_K) - offs_m = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) - offs_n = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) + offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) offs_k = tl.arange(0, BLOCK_SIZE_K) a_ptrs = a_ptr + offs_m[None, :] + offs_k[:, None] * M b_ptrs = b_ptr + offs_n[None, :] + offs_k[:, None] * N c = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32) - # c = tl.load(c_ptr + offs_m[:, None] * N + offs_n[None, :]).to(tl.float32) for i in range(k): a = tl.trans(tl.load(a_ptrs)).to(tl.float32) b = tl.load(b_ptrs).to(tl.float32) - # c += tl.dot(a, b) c = tl.dot(a, b, c) a_ptrs += BLOCK_SIZE_K * M b_ptrs += BLOCK_SIZE_K * N @@ -202,157 +196,291 @@ def fp32_gemm_for_update_kernel( tl.store(c_ptrs, c) -def triton_fp32_gemm_for_update(a: torch.Tensor, b: torch.Tensor): +def triton_fp32_gemm_for_update(y: torch.Tensor, x: torch.Tensor): """ mix precision gemm for updaing weight Args: - a: gradient of output, fp32 - b: input activation, bf16/fp16 + y: gradient of output, fp32 + x: input activation, bf16/fp16 Returns: c: gradient of weight """ - assert a.is_contiguous() and b.is_contiguous() - K, M = a.size() - K, N = b.size() - c = torch.empty((M, N), dtype=b.dtype, device=b.device) - grid = lambda META: (triton.cdiv(M, META["BLOCK_SIZE_M"]), - triton.cdiv(N, META["BLOCK_SIZE_N"])) # noqa - BLOCK_SIZE_K = 128 + assert y.is_contiguous() and x.is_contiguous() + K, M = y.size() + K, N = x.size() + c = torch.empty((M, N), dtype=torch.float32, device=x.device) + grid = lambda META: ( + triton.cdiv(M, META["BLOCK_SIZE_M"]), + triton.cdiv(N, META["BLOCK_SIZE_N"]), + ) # noqa + BLOCK_SIZE_K = max([x for x in [32, 64, 128] if K % x == 0]) BLOCK_SIZE_M = 32 - BLOCK_SIZE_N = 128 + BLOCK_SIZE_N = 64 num_warps = 4 num_stages = 3 - fp32_gemm_for_update_kernel[grid](a, b, c, - M, N, K, - BLOCK_SIZE_K, - BLOCK_SIZE_M, - BLOCK_SIZE_N, - num_warps=num_warps, - num_stages=num_stages - ) + fp32_gemm_for_update_kernel[grid]( + y, + x, + c, + M, + N, + K, + BLOCK_SIZE_K, + BLOCK_SIZE_M, + BLOCK_SIZE_N, + num_warps=num_warps, + num_stages=num_stages, + ) return c +split_fp32_gemm_configs = [ + Config( + {"BLOCK_SIZE_K": block_k, "BLOCK_SIZE_M": block_m, "BLOCK_SIZE_N": block_n}, + num_stages=num_stages, + num_warps=num_warps, + ) + for block_k in [64, 128] + for block_m in [32, 64] + for block_n in [32, 64] + for num_stages in [2, 3] + for num_warps in [2, 4] +] + +@triton.autotune( + configs=split_fp32_gemm_configs, key=["M", "N", "K"], reset_to_zero=["c_ptr"] +) @triton.jit -def scaled_fp32_gemm_kernel( - a_ptr, - b_ptr, - scale_ptr, - c_ptr, - M, - N: tl.constexpr, - K: tl.constexpr, - BLOCK_SIZE_K: tl.constexpr, - BLOCK_SIZE_M: tl.constexpr, - BLOCK_SIZE_N: tl.constexpr, +def split_fp32_gemm_kernel( + a_ptr, + b_ptr, + c_ptr, + M, + N: tl.constexpr, + K: tl.constexpr, + SPLIT_COUNT: tl.constexpr, + BLOCK_SIZE_K: tl.constexpr, + BLOCK_SIZE_M: tl.constexpr, + BLOCK_SIZE_N: tl.constexpr, ): - pid_m = tl.program_id(axis=0) - pid_n = tl.program_id(axis=1) - k = tl.cdiv(K, BLOCK_SIZE_K) - offs_m = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) - offs_n = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) + pid_k = tl.program_id(axis=0) + pid_m = tl.program_id(axis=1) + pid_n = tl.program_id(axis=2) + + k = tl.cdiv(K, BLOCK_SIZE_K * SPLIT_COUNT) + offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) offs_k = tl.arange(0, BLOCK_SIZE_K) - a_ptrs = a_ptr + offs_m[:, None] * K + offs_k[None, :] - b_ptrs = b_ptr + offs_n[None, :] * K + offs_k[:, None] + a_ptrs = a_ptr + pid_k * K // SPLIT_COUNT + offs_m[:, None] * K + offs_k[None, :] + b_ptrs = b_ptr + pid_k * K // SPLIT_COUNT + offs_n[None, :] * K + offs_k[:, None] c = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32) for i in range(k): - a = tl.load(a_ptrs).to(tl.float32) - b = tl.load(b_ptrs).to(tl.float32) - # c += tl.dot(a, b) + a = tl.load(a_ptrs) # .to(tl.float32) + b = tl.load(b_ptrs) # .to(tl.float32) c = tl.dot(a, b, c) a_ptrs += BLOCK_SIZE_K b_ptrs += BLOCK_SIZE_K - - scale = tl.load( - scale_ptr + pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) - c *= scale[:, None] - offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) c_ptrs = c_ptr + offs_m[:, None] * N + offs_n[None, :] - tl.store(c_ptrs, c) + if SPLIT_COUNT == 1: + tl.store(c_ptrs, c) + else: + tl.atomic_add(c_ptrs, c, sem="relaxed") -def triton_scaled_fp32_gemm(a: torch.Tensor, - b: torch.Tensor, - scale: torch.Tensor): +def triton_split_fp32_gemm(x: torch.Tensor, w: torch.Tensor): """ - c = (a*scale[:,None])*b - this kernel is used to fuse RMSNorm and quantization in MoE layer - native implementation: - y = rms_norm(x), - y_q = quantization(y), - router_logits = y@w - we can not fuse rms_norm and quantization - as we still need bf16 y for moe router gemm - fused implementation: - y_q, rms = quantization(rms_norm(x)) - router_logits = (x/rms)@y - so we need a scaled fp32 gemm kernel + return fp32 gemm result with fp16/bf16 inputs, + it's mainly used for MoE router GEMM + and DO NOT suitable for large size GEMM Args: - a: activation tensor - b: weight tensor - scale: scale for activation tensor, 1/rms + a: left matrix with fp16/bf16 precision + b: right matrix with fp16/bf16 precision Returns: - output tensor + c: output with fp32 precision """ - assert a.is_contiguous() and b.is_contiguous() - M, K = a.size() - N, K = b.size() - c = torch.empty(M, N, dtype=torch.float32, device=a.device) - grid = lambda META: (triton.cdiv(M, META["BLOCK_SIZE_M"]), - triton.cdiv(N, META["BLOCK_SIZE_N"])) # noqa - BLOCK_SIZE_K = 128 - BLOCK_SIZE_M = 32 - BLOCK_SIZE_N = 128 - num_warps = 4 - num_stages = 3 - scaled_fp32_gemm_kernel[grid](a, b, - scale, - c, - M, N, K, - BLOCK_SIZE_K, - BLOCK_SIZE_M, - BLOCK_SIZE_N, - num_warps=num_warps, - num_stages=num_stages - ) + assert x.is_contiguous() and w.is_contiguous() + M, K = x.size() + N, K = w.size() + # BLOCK_SIZE_K = 128 + # BLOCK_SIZE_M = 128 if M % 128 == 0 else 32 + # BLOCK_SIZE_N = max([x for x in [16, 32, 64, 128] if N % x == 0]) + # assert M % BLOCK_SIZE_M == 0 and K % BLOCK_SIZE_K == 0 + # num_warps = 4 + # num_stages = 3 + + if M * N <= 2048 * 256 and K % 4096 == 0: + SPLIT_COUNT = min(K // 2048, 4) + else: + SPLIT_COUNT = 1 + + if SPLIT_COUNT == 1: + c = torch.empty(M, N, dtype=torch.float32, device=x.device) + else: + c = torch.zeros(M, N, dtype=torch.float32, device=x.device) + grid = lambda META: ( + SPLIT_COUNT, + triton.cdiv(M, META["BLOCK_SIZE_M"]), + triton.cdiv(N, META["BLOCK_SIZE_N"]), + ) # noqa + + split_fp32_gemm_kernel[grid]( + x, + w, + c, + M, + N, + K, + SPLIT_COUNT, + # BLOCK_SIZE_K, + # BLOCK_SIZE_M, + # BLOCK_SIZE_N, + # num_warps=num_warps, + # num_stages=num_stages + ) return c - +@triton.autotune( + configs=split_fp32_gemm_configs, key=["M", "N", "K"], reset_to_zero=["c_ptr"] +) @triton.jit -def scaled_fp32_gemm_for_update_kernel( - a_ptr, - b_ptr, - scale_ptr, - c_ptr, +def split_fp32_gemm_for_backward_kernel( + a_ptr, + b_ptr, + c_ptr, + M, + N: tl.constexpr, + K: tl.constexpr, + SPLIT_COUNT: tl.constexpr, + BLOCK_SIZE_K: tl.constexpr, + BLOCK_SIZE_M: tl.constexpr, + BLOCK_SIZE_N: tl.constexpr, +): + pid_k = tl.program_id(axis=0) + pid_m = tl.program_id(axis=1) + pid_n = tl.program_id(axis=2) + + k = tl.cdiv(K, BLOCK_SIZE_K * SPLIT_COUNT) + offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) + offs_k = tl.arange(0, BLOCK_SIZE_K) + a_ptrs = a_ptr + pid_k * K // SPLIT_COUNT + offs_m[:, None] * K + offs_k[None, :] + b_ptrs = ( + b_ptr + pid_k * K // SPLIT_COUNT * N + offs_n[None, :] + offs_k[:, None] * N + ) + + c = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32) + + for i in range(k): + a = tl.load(a_ptrs) + b = tl.load(b_ptrs).to(tl.float32) + c = tl.dot(a, b, c) + a_ptrs += BLOCK_SIZE_K + b_ptrs += BLOCK_SIZE_K * N + offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) + c_ptrs = c_ptr + offs_m[:, None] * N + offs_n[None, :] + if SPLIT_COUNT == 1: + tl.store(c_ptrs, c) + else: + tl.atomic_add(c_ptrs, c, sem="relaxed") + + +def triton_split_fp32_gemm_for_backward(y: torch.Tensor, w: torch.Tensor): + """ + mix precision gemm for backward, a@b.float() + Args: + a: input gradient, fp32 + b: gemm weight, bf16/fp16 + Returns: + c: gradient of activation + """ + assert y.is_contiguous() and w.is_contiguous() + M, K = y.size() + K, N = w.size() + + if M * N <= 256 * 2048 and K % 4096 == 0: + SPLIT_COUNT = min(K // 4096, 4) + else: + SPLIT_COUNT = 1 + + if SPLIT_COUNT == 1: + c = torch.empty((M, N), dtype=w.dtype, device=w.device) + else: + c = torch.zeros((M, N), dtype=torch.float32, device=w.device) + + # BLOCK_SIZE_K = max([x for x in [16, 32, 64, 128] if K % x == 0]) + # BLOCK_SIZE_M = 32 + # BLOCK_SIZE_N = 128 + # assert M % BLOCK_SIZE_M == 0 and N % BLOCK_SIZE_N == 0 + # num_warps = 4 + # num_stages = 2 + + grid = lambda META: ( + SPLIT_COUNT, + triton.cdiv(M, META["BLOCK_SIZE_M"]), + triton.cdiv(N, META["BLOCK_SIZE_N"]), + ) # noqa + + split_fp32_gemm_for_backward_kernel[grid]( + y, + w, + c, M, - N: tl.constexpr, - K: tl.constexpr, - BLOCK_SIZE_K: tl.constexpr, - BLOCK_SIZE_M: tl.constexpr, - BLOCK_SIZE_N: tl.constexpr, + N, + K, + SPLIT_COUNT, + # BLOCK_SIZE_K, + # BLOCK_SIZE_M, + # BLOCK_SIZE_N, + # num_warps=num_warps, + # num_stages=num_stages + ) + if SPLIT_COUNT > 1: + c = c.to(w.dtype) + return c + + +@triton.autotune( + configs=split_fp32_gemm_configs, key=["M", "N", "K"], reset_to_zero=["c_ptr"] +) +@triton.jit +def split_fp32_gemm_for_update_kernel( + a_ptr, + b_ptr, + c_ptr, + K, + M: tl.constexpr, + N: tl.constexpr, + SPLIT_COUNT: tl.constexpr, + BLOCK_SIZE_K: tl.constexpr, + BLOCK_SIZE_M: tl.constexpr, + BLOCK_SIZE_N: tl.constexpr, ): - pid_m = tl.program_id(axis=0) - pid_n = tl.program_id(axis=1) - k = tl.cdiv(K, BLOCK_SIZE_K) - offs_m = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) - offs_n = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) + pid_k = tl.program_id(axis=0) + pid_m = tl.program_id(axis=1) + pid_n = tl.program_id(axis=2) + + k = tl.cdiv(K, BLOCK_SIZE_K * SPLIT_COUNT) + offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) offs_k = tl.arange(0, BLOCK_SIZE_K) - a_ptrs = a_ptr + offs_m[None, :] + offs_k[:, None] * M - b_ptrs = b_ptr + offs_n[None, :] + offs_k[:, None] * N + a_ptrs = ( + a_ptr + pid_k * K // SPLIT_COUNT * M + offs_m[None, :] + offs_k[:, None] * M + ) + b_ptrs = ( + b_ptr + pid_k * K // SPLIT_COUNT * N + offs_n[None, :] + offs_k[:, None] * N + ) c = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32) for i in range(k): - scale = tl.load( - scale_ptr + i * BLOCK_SIZE_K + tl.arange(0, BLOCK_SIZE_K)) - a = tl.trans(tl.load(a_ptrs)).to(tl.float32) * scale[None, :] + a = tl.trans(tl.load(a_ptrs)).to(tl.float32) b = tl.load(b_ptrs).to(tl.float32) - # c += tl.dot(a, b) c = tl.dot(a, b, c) a_ptrs += BLOCK_SIZE_K * M b_ptrs += BLOCK_SIZE_K * N @@ -360,39 +488,59 @@ def scaled_fp32_gemm_for_update_kernel( offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) c_ptrs = c_ptr + offs_m[:, None] * N + offs_n[None, :] - tl.store(c_ptrs, c) + if SPLIT_COUNT == 1: + tl.store(c_ptrs, c) + else: + tl.atomic_add(c_ptrs, c, sem="relaxed") -def triton_scaled_fp32_gemm_for_update(a: torch.Tensor, - b: torch.Tensor, - scale: torch.Tensor): +def triton_split_fp32_gemm_for_update(y: torch.Tensor, x: torch.Tensor): """ - see triton_scaled_fp32_gemm + mix precision gemm for updaing weight Args: - a: y - b: activation before RMS norm - scale: 1/rms - + y: gradient of output, fp32 + x: input activation, bf16/fp16 Returns: - dw + c: gradient of weight """ - assert a.is_contiguous() and b.is_contiguous() - K, M = a.size() - K, N = b.size() - c = torch.empty((M, N), dtype=b.dtype, device=b.device) - grid = lambda META: (triton.cdiv(M, META["BLOCK_SIZE_M"]), - triton.cdiv(N, META["BLOCK_SIZE_N"])) # noqa - BLOCK_SIZE_K = 128 - BLOCK_SIZE_M = 32 - BLOCK_SIZE_N = 128 - num_warps = 4 - num_stages = 3 - scaled_fp32_gemm_for_update_kernel[grid](a, b, scale, c, - M, N, K, - BLOCK_SIZE_K, - BLOCK_SIZE_M, - BLOCK_SIZE_N, - num_warps=num_warps, - num_stages=num_stages - ) + assert y.is_contiguous() and x.is_contiguous() + K, M = y.size() + K, N = x.size() + + if M * N <= 256 * 2048 and K % 4096 == 0: + SPLIT_COUNT = min(K // 4096, 4) + else: + SPLIT_COUNT = 1 + + if SPLIT_COUNT == 1: + c = torch.empty((M, N), dtype=torch.float32, device=x.device) + else: + c = torch.zeros((M, N), dtype=torch.float32, device=x.device) + + # BLOCK_SIZE_K = 128 + # BLOCK_SIZE_M = 64 + # BLOCK_SIZE_N = 32 + # assert M % BLOCK_SIZE_M == 0 and N % BLOCK_SIZE_N == 0 and K % (BLOCK_SIZE_K * SPLIT_COUNT) == 0 + # num_warps = 2 + # num_stages = 3 + + grid = lambda META: ( + SPLIT_COUNT, + triton.cdiv(M, META["BLOCK_SIZE_M"]), + triton.cdiv(N, META["BLOCK_SIZE_N"]), + ) # noqa + split_fp32_gemm_for_update_kernel[grid]( + y, + x, + c, + K, + M, + N, + SPLIT_COUNT, + # BLOCK_SIZE_K, + # BLOCK_SIZE_M, + # BLOCK_SIZE_N, + # num_warps=num_warps, + # num_stages=num_stages + ) return c diff --git a/linghe/gemm/mxfp8_gemm.py b/linghe/gemm/mxfp8_gemm.py new file mode 100644 index 0000000..6364eb1 --- /dev/null +++ b/linghe/gemm/mxfp8_gemm.py @@ -0,0 +1,1066 @@ +# -*- coding: utf-8 -*- +""" +Copyright (c) Ant Financial Service Group and its affiliates. +""" + +from typing import Optional, List + +import torch +import triton +import triton.language as tl + +from linghe.utils.transpose import triton_transpose + + +# os.environ["TRITON_PRINT_AUTOTUNING"] = "1" + + +# fp8_gemm_configs = [ +# Config({"BLOCK_SIZE_M": block_m, "BLOCK_SIZE_N": block_n}, +# num_stages=num_stages, num_warps=8) +# for block_m in [32, 64, 128] +# for block_n in [32, 64, 128] +# for num_stages in [3, 4, 5, 6] +# ] + +# @triton.autotune(configs=fp8_gemm_configs, key=["N", "K"]) + + +@triton.jit +def mxfp8_gemm_kernel( + a_ptr, + b_ptr, + c_ptr, + a_s_ptr, + b_s_ptr, + M, + N: tl.constexpr, + K: tl.constexpr, + BLOCK_SIZE_M: tl.constexpr, + BLOCK_SIZE_N: tl.constexpr, + ACCUM: tl.constexpr, +): + pid_m = tl.program_id(axis=0) + pid_n = tl.program_id(axis=1) + k = K // 32 + offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) + offs_k = tl.arange(0, 32) + a_ptrs = a_ptr + offs_m[:, None] * K + offs_k[None, :] + b_ptrs = b_ptr + offs_n[None, :] * K + offs_k[:, None] + + a_s_ptrs = a_s_ptr + offs_m + b_s_ptrs = b_s_ptr + offs_n + + if ACCUM: + accumulator = tl.load(c_ptr + offs_m[:, None] * N + offs_n[None, :]) + else: + accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32) + for i in range(k): + a = tl.load(a_ptrs) + b = tl.load(b_ptrs) + a_s = tl.exp2(tl.load(a_s_ptrs).to(tl.float32) - 127.0) + b_s = tl.exp2(tl.load(b_s_ptrs).to(tl.float32) - 127.0) + accumulator += tl.dot(a, b) * a_s[:, None] * b_s[None, :] + a_ptrs += 32 + b_ptrs += 32 + + a_s_ptrs += M + b_s_ptrs += N + + c = accumulator.to(c_ptr.dtype.element_ty) + offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) + c_ptrs = c_ptr + offs_m[:, None] * N + offs_n[None, :] + tl.store(c_ptrs, c) + + +def triton_mxfp8_gemm( + a: torch.Tensor, + b: torch.Tensor, + a_s: torch.Tensor, + b_s: torch.Tensor, + out: Optional[torch.tensor] = None, + out_dtype: torch.dtype = torch.bfloat16, + layout: str = "TN", + accumulate: bool = False, +): + """ + triton implementation to simulate mxfp8 grouped gemm + layout is defined as the same in TE: + TN: forward + NN: bakcward + NT: update(wgrad) + layout is used to optimize BLOCK SIZE + """ + assert a.is_contiguous() and b.is_contiguous() + assert a_s.is_contiguous() and b_s.is_contiguous() + + if layout == "TN": + assert not accumulate + a_s = a_s.t().contiguous() + b_s = b_s.t().contiguous() + elif layout == "NN": + assert not accumulate + b = triton_transpose(b) + a_s = a_s.t().contiguous() + else: + a = triton_transpose(a) + b = triton_transpose(b) + + M, K = a.shape + N = b.size(0) + + if out is not None: + assert out.is_contiguous() + else: + out = torch.empty(M, N, dtype=out_dtype, device=a.device) + BLOCK_SIZE_M = 128 + BLOCK_SIZE_N = 128 + grid = (M // BLOCK_SIZE_M, N // BLOCK_SIZE_N) + mxfp8_gemm_kernel[grid]( + a, + b, + out, + a_s, + b_s, + M, + N, + K, + BLOCK_SIZE_M=BLOCK_SIZE_M, + BLOCK_SIZE_N=BLOCK_SIZE_N, + ACCUM=accumulate, + num_warps=4, + num_stages=3, + ) + return out + + +@triton.jit +def mxfp8_gemm_forward_kernel( + a_ptr, + b_ptr, + c_ptr, + a_s_ptr, + b_s_ptr, + M, + N: tl.constexpr, + K: tl.constexpr, + BLOCK_SIZE_M: tl.constexpr, + BLOCK_SIZE_N: tl.constexpr, +): + pid_m = tl.program_id(axis=0) + pid_n = tl.program_id(axis=1) + k = K // 32 + offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) + offs_k = tl.arange(0, 32) + a_ptrs = a_ptr + offs_m[:, None] * K + offs_k[None, :] + b_ptrs = b_ptr + offs_n[None, :] * K + offs_k[:, None] + + a_s_ptrs = a_s_ptr + offs_m + b_s_ptrs = b_s_ptr + offs_n + + accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32) + for i in range(k): + a = tl.load(a_ptrs) + b = tl.load(b_ptrs) + a_s = tl.exp2(tl.load(a_s_ptrs).to(tl.float32) - 127.0) + b_s = tl.exp2(tl.load(b_s_ptrs).to(tl.float32) - 127.0) + accumulator += tl.dot(a, b) * a_s[:, None] * b_s[None, :] + a_ptrs += 32 + b_ptrs += 32 + + a_s_ptrs += M + b_s_ptrs += N + + c = accumulator.to(c_ptr.dtype.element_ty) + offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) + c_ptrs = c_ptr + offs_m[:, None] * N + offs_n[None, :] + tl.store(c_ptrs, c) + + +def triton_mxfp8_gemm_forward( + a: torch.Tensor, + b: torch.Tensor, + a_s: torch.Tensor, + b_s: torch.Tensor, + out: Optional[torch.tensor] = None, + out_dtype: torch.dtype = torch.bfloat16, +): + """ + triton implementation to simulate mxfp8 gemm + """ + assert a.is_contiguous() and b.is_contiguous() + assert a_s.is_contiguous() and b_s.is_contiguous() + + M, K = a.shape + N = b.size(0) + + a_s = a_s.t().contiguous() + b_s = b_s.t().contiguous() + + if out is not None: + assert out.is_contiguous() + else: + out = torch.empty(M, N, dtype=out_dtype, device=a.device) + BLOCK_SIZE_M = 128 + BLOCK_SIZE_N = 128 + grid = (M // BLOCK_SIZE_M, N // BLOCK_SIZE_N) + mxfp8_gemm_forward_kernel[grid]( + a, + b, + out, + a_s, + b_s, + M, + N, + K, + BLOCK_SIZE_M=BLOCK_SIZE_M, + BLOCK_SIZE_N=BLOCK_SIZE_N, + num_warps=4, + num_stages=3, + ) + return out + + +@triton.jit +def mxfp8_gemm_backward_kernel( + a_ptr, + b_ptr, + c_ptr, + a_s_ptr, + b_s_ptr, + M, + N: tl.constexpr, + K: tl.constexpr, + BLOCK_SIZE_M: tl.constexpr, + BLOCK_SIZE_N: tl.constexpr, +): + pid_m = tl.program_id(axis=0) + pid_n = tl.program_id(axis=1) + k = K // 32 + offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) + offs_k = tl.arange(0, 32) + a_ptrs = a_ptr + offs_m[:, None] * K + offs_k[None, :] + b_ptrs = b_ptr + offs_n[None, :] + offs_k[:, None] * N + a_s_ptrs = a_s_ptr + offs_m + b_s_ptrs = b_s_ptr + offs_n + + accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32) + for i in range(k): + a = tl.load(a_ptrs) + b = tl.load(b_ptrs) + a_s = tl.exp2(tl.load(a_s_ptrs).to(tl.float32) - 127.0) + b_s = tl.exp2(tl.load(b_s_ptrs).to(tl.float32) - 127.0) + accumulator += tl.dot(a, b) * a_s[:, None] * b_s[None, :] + a_ptrs += 32 + b_ptrs += 32 * N + a_s_ptrs += M + b_s_ptrs += N + + c = accumulator.to(c_ptr.dtype.element_ty) + offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) + c_ptrs = c_ptr + offs_m[:, None] * N + offs_n[None, :] + tl.store(c_ptrs, c) + + +def triton_mxfp8_gemm_backward( + a: torch.Tensor, + b: torch.Tensor, + a_s: torch.Tensor, + b_s: torch.Tensor, + out: Optional[torch.tensor] = None, + out_dtype: torch.dtype = torch.bfloat16, +): + """ + triton implementation to simulate mxfp8 gemm + """ + assert a.is_contiguous() and b.is_contiguous() + assert a_s.is_contiguous() and b_s.is_contiguous() + + M, K = a.shape + N = b.size(1) + + a_s = a_s.t().contiguous() + + if out is not None: + assert out.is_contiguous() + else: + out = torch.empty(M, N, dtype=out_dtype, device=a.device) + BLOCK_SIZE_M = 128 + BLOCK_SIZE_N = 128 + grid = (M // BLOCK_SIZE_M, N // BLOCK_SIZE_N) + mxfp8_gemm_backward_kernel[grid]( + a, + b, + out, + a_s, + b_s, + M, + N, + K, + BLOCK_SIZE_M=BLOCK_SIZE_M, + BLOCK_SIZE_N=BLOCK_SIZE_N, + num_warps=4, + num_stages=3, + ) + return out + + +@triton.jit +def mxfp8_gemm_update_kernel( + a_ptr, + b_ptr, + c_ptr, + a_s_ptr, + b_s_ptr, + M, + N: tl.constexpr, + K: tl.constexpr, + BLOCK_SIZE_M: tl.constexpr, + BLOCK_SIZE_N: tl.constexpr, + ACCUM: tl.constexpr, +): + pid_m = tl.program_id(axis=0) + pid_n = tl.program_id(axis=1) + k = K // 32 + offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) + offs_k = tl.arange(0, 32) + a_ptrs = a_ptr + offs_m[None, :] + offs_k[:, None] * M + # a_ptrs = a_ptr + offs_m[:, None] + offs_k[None, :] * M + b_ptrs = b_ptr + offs_n[None, :] + offs_k[:, None] * N + a_s_ptrs = a_s_ptr + offs_m + b_s_ptrs = b_s_ptr + offs_n + + if ACCUM: + c_ptrs = c_ptr + offs_m[:, None] * N + offs_n[None, :] + accumulator = tl.load(c_ptrs) + else: + accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32) + for i in range(k): + a = tl.load(a_ptrs) + b = tl.load(b_ptrs) + a_s = tl.exp2(tl.load(a_s_ptrs).to(tl.float32) - 127.0) + b_s = tl.exp2(tl.load(b_s_ptrs).to(tl.float32) - 127.0) + accumulator += tl.dot(tl.trans(a), b) * a_s[:, None] * b_s[None, :] + # accumulator += tl.dot(a, b) * a_s[:, None] * b_s[None, :] + a_ptrs += 32 * M + b_ptrs += 32 * N + a_s_ptrs += M + b_s_ptrs += N + + c = accumulator.to(c_ptr.dtype.element_ty) + # offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + # offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) + c_ptrs = c_ptr + offs_m[:, None] * N + offs_n[None, :] + tl.store(c_ptrs, c) + + +def triton_mxfp8_gemm_update( + a: torch.Tensor, + b: torch.Tensor, + a_s: torch.Tensor, + b_s: torch.Tensor, + out: Optional[torch.tensor] = None, + out_dtype: torch.dtype = torch.float32, + accumulate: bool = False, +): + """ + triton implementation to simulate mxfp8 gemm + """ + assert a.is_contiguous() and b.is_contiguous() + assert a_s.is_contiguous() and b_s.is_contiguous() + + K, M = a.shape + N = b.size(1) + + if out is not None: + assert out.is_contiguous() + else: + assert not accumulate + out = torch.empty(M, N, dtype=out_dtype, device=a.device) + BLOCK_SIZE_M = 128 + BLOCK_SIZE_N = 128 + grid = (M // BLOCK_SIZE_M, N // BLOCK_SIZE_N) # noqa + mxfp8_gemm_update_kernel[grid]( + a, + b, + out, + a_s, + b_s, + M, + N, + K, + BLOCK_SIZE_M=BLOCK_SIZE_M, + BLOCK_SIZE_N=BLOCK_SIZE_N, + ACCUM=accumulate, + num_warps=4, + num_stages=3, + ) + return out + + +@triton.jit +def mxfp8_grouped_gemm_kernel( + a_ptrs, + b_ptrs, + c_ptr, + a_s_ptrs, + b_s_ptrs, + size_ptr, + accum_ptr, + M, + N: tl.constexpr, + K: tl.constexpr, + BLOCK_SIZE_M: tl.constexpr, + BLOCK_SIZE_N: tl.constexpr, + ACCUM: tl.constexpr, + LAYOUT: tl.constexpr, +): + pid_e = tl.program_id(axis=0) + pid_m = tl.program_id(axis=1) + pid_n = tl.program_id(axis=2) + size = tl.load(size_ptr + pid_e) + + a_ptr = tl.load(a_ptrs + pid_e).to(tl.pointer_type(tl.float8e4nv)) + b_ptr = tl.load(b_ptrs + pid_e).to(tl.pointer_type(tl.float8e4nv)) + a_s_ptr = tl.load(a_s_ptrs + pid_e).to(tl.pointer_type(tl.uint8)) + b_s_ptr = tl.load(b_s_ptrs + pid_e).to(tl.pointer_type(tl.uint8)) + + if LAYOUT != "NT": + if pid_m * BLOCK_SIZE_M >= size: + return + + if LAYOUT == "NT": + K = size + k = size // 32 + c_ptr = tl.load(c_ptr + pid_e).to(tl.pointer_type(tl.float32)) + MS = 0 + else: + k = K // 32 + MS = tl.load(accum_ptr + pid_e) - size + M = tl.cdiv(size, 128) * 128 + + offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) + offs_k = tl.arange(0, 32) + a_ps = a_ptr + offs_m[:, None] * K + offs_k[None, :] + b_ps = b_ptr + offs_n[None, :] * K + offs_k[:, None] + + a_s_ps = a_s_ptr + offs_m + b_s_ps = b_s_ptr + offs_n + + if ACCUM: + accumulator = tl.load(c_ptr + offs_m[:, None] * N + offs_n[None, :]) + else: + accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32) + + for i in range(k): + a = tl.load(a_ps) + b = tl.load(b_ps) + a_s = tl.exp2(tl.load(a_s_ps).to(tl.float32) - 127.0) + b_s = tl.exp2(tl.load(b_s_ps).to(tl.float32) - 127.0) + accumulator += tl.dot(a, b) * a_s[:, None] * b_s[None, :] + a_ps += 32 + b_ps += 32 + a_s_ps += M + b_s_ps += N + + c_ps = c_ptr + MS * N + offs_m[:, None] * N + offs_n[None, :] + tl.store(c_ps, accumulator, cache_modifier=".cs") + + +def triton_mxfp8_grouped_gemm( + a: List[torch.Tensor], + b: List[torch.Tensor], + a_s: List[torch.Tensor], + b_s: List[torch.Tensor], + m_splits: List[int], + out: Optional[torch.tensor] = None, + out_dtype: torch.dtype = torch.bfloat16, + layout: str = "TN", + accumulate: bool = False, +): + """ + triton implementation to simulate mxfp8 grouped gemm + """ + device = a[0].device + sizes = torch.tensor(m_splits, dtype=torch.int64).cuda(device, non_blocking=True) + accums = torch.cumsum(sizes, 0) + ms = sum(m_splits) + + if layout == "TN": + assert not accumulate + BLOCK_SIZE_M = 32 + BLOCK_SIZE_N = 128 + M = max(m_splits) + K = a[0].size(1) + N = b[0].size(0) + # not work if in one line + a_s = [x.t().contiguous() for x in a_s] + b_s = [x.t().contiguous() for x in b_s] + if out is None: + out = torch.empty(ms, N, dtype=out_dtype, device=device) + + elif layout == "NN": + assert not accumulate + BLOCK_SIZE_M = 32 + BLOCK_SIZE_N = 128 + M = max(m_splits) + K = a[0].size(1) + N = b[0].size(1) + b = [triton_transpose(x) for x in b] + a_s = [x.t().contiguous() for x in a_s] + if out is None: + out = torch.empty(ms, N, dtype=out_dtype, device=device) + else: + assert out is not None + BLOCK_SIZE_M = 128 + BLOCK_SIZE_N = 128 + M = a[0].size(1) + N = b[0].size(1) + K = 0 + a = [triton_transpose(x) for x in a] + b = [triton_transpose(x) for x in b] + out_ptrs = torch.tensor([x.data_ptr() for x in out], dtype=torch.int64).cuda( + device, non_blocking=True + ) + + as_ptrs = torch.tensor([x.data_ptr() for x in a_s], dtype=torch.int64).cuda( + device, non_blocking=True + ) + bs_ptrs = torch.tensor([x.data_ptr() for x in b_s], dtype=torch.int64).cuda( + device, non_blocking=True + ) + + a_ptrs = torch.tensor([x.data_ptr() for x in a], dtype=torch.int64).cuda( + device, non_blocking=True + ) + b_ptrs = torch.tensor([x.data_ptr() for x in b], dtype=torch.int64).cuda( + device, non_blocking=True + ) + + grid = (len(m_splits), M // BLOCK_SIZE_M, N // BLOCK_SIZE_N) # noqa + mxfp8_grouped_gemm_kernel[grid]( + a_ptrs, + b_ptrs, + out_ptrs if layout == "NT" else out, + as_ptrs, + bs_ptrs, + sizes, + accums, + M, + N, + K, + BLOCK_SIZE_M=BLOCK_SIZE_M, + BLOCK_SIZE_N=BLOCK_SIZE_N, + ACCUM=accumulate, + LAYOUT=layout, + num_warps=4, + num_stages=3, + ) + return out + + +@triton.jit +def mxfp8_grouped_gemm_forward_kernel( + a_ptrs, + b_ptrs, + c_ptr, + a_s_ptrs, + b_s_ptrs, + size_ptr, + accum_ptr, + N: tl.constexpr, + K: tl.constexpr, + BLOCK_SIZE_M: tl.constexpr, + BLOCK_SIZE_N: tl.constexpr, +): + pid_e = tl.program_id(axis=0) + pid_m = tl.program_id(axis=1) + pid_n = tl.program_id(axis=2) + size = tl.load(size_ptr + pid_e) + + a_ptr = tl.load(a_ptrs + pid_e).to(tl.pointer_type(tl.float8e4nv)) + b_ptr = tl.load(b_ptrs + pid_e).to(tl.pointer_type(tl.float8e4nv)) + a_s_ptr = tl.load(a_s_ptrs + pid_e).to(tl.pointer_type(tl.uint8)) + b_s_ptr = tl.load(b_s_ptrs + pid_e).to(tl.pointer_type(tl.uint8)) + + if pid_m * BLOCK_SIZE_M >= size: + return + + offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) + offs_k = tl.arange(0, 32) + a_ps = a_ptr + offs_m[:, None] * K + offs_k[None, :] + b_ps = b_ptr + offs_n[None, :] * K + offs_k[:, None] + + a_s_ps = a_s_ptr + offs_m + b_s_ps = b_s_ptr + offs_n + M = tl.cdiv(size, 128) * 128 + + accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32) + + k = K // 32 + for i in range(k): + a = tl.load(a_ps) + b = tl.load(b_ps) + a_s = tl.exp2(tl.load(a_s_ps).to(tl.float32) - 127.0) + b_s = tl.exp2(tl.load(b_s_ps).to(tl.float32) - 127.0) + accumulator += tl.dot(a, b) * a_s[:, None] * b_s[None, :] + a_ps += 32 + b_ps += 32 + a_s_ps += M + b_s_ps += N + + MS = tl.load(accum_ptr + pid_e) - size + c_ps = c_ptr + MS * N + offs_m[:, None] * N + offs_n[None, :] + tl.store(c_ps, accumulator, cache_modifier=".cs") + + +def triton_mxfp8_grouped_gemm_forward( + a: List[torch.Tensor], + b: List[torch.Tensor], + a_s: List[torch.Tensor], + b_s: List[torch.Tensor], + m_splits: List[int], + out: Optional[torch.tensor] = None, + out_dtype: torch.dtype = torch.bfloat16, +): + """ + triton implementation to simulate mxfp8 grouped gemm + """ + device = a[0].device + sizes = torch.tensor(m_splits, dtype=torch.int64).cuda(device, non_blocking=True) + accums = torch.cumsum(sizes, 0) + a_ptrs = torch.tensor([x.data_ptr() for x in a], dtype=torch.int64).cuda( + device, non_blocking=True + ) + b_ptrs = torch.tensor([x.data_ptr() for x in b], dtype=torch.int64).cuda( + device, non_blocking=True + ) + + # not work if in one line + at_s = [x.t().contiguous() for x in a_s] + as_ptrs = torch.tensor([x.data_ptr() for x in at_s], dtype=torch.int64).cuda( + device, non_blocking=True + ) + bt_s = [x.t().contiguous() for x in b_s] + bs_ptrs = torch.tensor([x.data_ptr() for x in bt_s], dtype=torch.int64).cuda( + device, non_blocking=True + ) + + BLOCK_SIZE_M = 32 + BLOCK_SIZE_N = 128 + N = b[0].size(0) + K = a[0].size(1) + if out is None: + out = torch.empty(sum(m_splits), N, dtype=out_dtype, device=device) + + grid = (len(m_splits), max(m_splits) // BLOCK_SIZE_M, N // BLOCK_SIZE_N) # noqa + mxfp8_grouped_gemm_forward_kernel[grid]( + a_ptrs, + b_ptrs, + out, + as_ptrs, + bs_ptrs, + sizes, + accums, + N, + K, + BLOCK_SIZE_M=BLOCK_SIZE_M, + BLOCK_SIZE_N=BLOCK_SIZE_N, + num_warps=4, + num_stages=3, + ) + return out + + +@triton.jit +def mxfp8_grouped_gemm_backward_kernel( + a_ptrs, + b_ptrs, + c_ptr, + a_s_ptrs, + b_s_ptrs, + size_ptr, + accum_ptr, + N: tl.constexpr, + K: tl.constexpr, + BLOCK_SIZE_M: tl.constexpr, + BLOCK_SIZE_N: tl.constexpr, +): + pid_e = tl.program_id(axis=0) + pid_m = tl.program_id(axis=1) + pid_n = tl.program_id(axis=2) + size = tl.load(size_ptr + pid_e) + + a_ptr = tl.load(a_ptrs + pid_e).to(tl.pointer_type(tl.float8e4nv)) + b_ptr = tl.load(b_ptrs + pid_e).to(tl.pointer_type(tl.float8e4nv)) + a_s_ptr = tl.load(a_s_ptrs + pid_e).to(tl.pointer_type(tl.uint8)) + b_s_ptr = tl.load(b_s_ptrs + pid_e).to(tl.pointer_type(tl.uint8)) + + if pid_m * BLOCK_SIZE_M >= size: + return + + MS = tl.load(accum_ptr + pid_e) - size + k = K // 32 + + offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) + offs_k = tl.arange(0, 32) + a_ps = a_ptr + offs_m[:, None] * K + offs_k[None, :] + b_ps = b_ptr + offs_n[None, :] + offs_k[:, None] * N + + a_s_ps = a_s_ptr + offs_m + b_s_ps = b_s_ptr + offs_n + M = tl.cdiv(size, 128) * 128 + + accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32) + + for i in range(k): + a = tl.load(a_ps) + b = tl.load(b_ps) + a_s = tl.exp2(tl.load(a_s_ps).to(tl.float32) - 127.0) + b_s = tl.exp2(tl.load(b_s_ps).to(tl.float32) - 127.0) + accumulator += tl.dot(a, b) * a_s[:, None] * b_s[None, :] + a_ps += 32 + b_ps += 32 * N + a_s_ps += M + b_s_ps += N + + # offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + # offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) + c_ps = c_ptr + MS * N + offs_m[:, None] * N + offs_n[None, :] + # tl.store(c_ps, c) + tl.store(c_ps, accumulator, cache_modifier=".cs") + + +def triton_mxfp8_grouped_gemm_backward( + a: List[torch.Tensor], + b: List[torch.Tensor], + a_s: List[torch.Tensor], + b_s: List[torch.Tensor], + m_splits: List[int], + out: Optional[torch.tensor] = None, + out_dtype: torch.dtype = torch.bfloat16, +): + """ + triton implementation to simulate mxfp8 grouped gemm + layout is defined as the same in TE: + TN: forward + NN: bakcward + NT: update(wgrad) + """ + device = a[0].device + sizes = torch.tensor(m_splits, dtype=torch.int64).cuda(device, non_blocking=True) + accums = torch.cumsum(sizes, 0) + a_ptrs = torch.tensor([x.data_ptr() for x in a], dtype=torch.int64).cuda( + device, non_blocking=True + ) + b_ptrs = torch.tensor([x.data_ptr() for x in b], dtype=torch.int64).cuda( + device, non_blocking=True + ) + + # not work if in one line + at_s = [x.t().contiguous() for x in a_s] + as_ptrs = torch.tensor([x.data_ptr() for x in at_s], dtype=torch.int64).cuda( + device, non_blocking=True + ) + bs_ptrs = torch.tensor([x.data_ptr() for x in b_s], dtype=torch.int64).cuda( + device, non_blocking=True + ) + + BLOCK_SIZE_M = 32 + BLOCK_SIZE_N = 128 + K = a[0].size(1) + N = b[0].size(1) + if out is None: + out = torch.empty(sum(m_splits), N, dtype=out_dtype, device=device) + + grid = (len(m_splits), max(m_splits) // BLOCK_SIZE_M, N // BLOCK_SIZE_N) # noqa + mxfp8_grouped_gemm_backward_kernel[grid]( + a_ptrs, + b_ptrs, + out, + as_ptrs, + bs_ptrs, + sizes, + accums, + N, + K, + BLOCK_SIZE_M=BLOCK_SIZE_M, + BLOCK_SIZE_N=BLOCK_SIZE_N, + num_warps=4, + num_stages=3, + ) + return out + + +@triton.jit +def mxfp8_grouped_gemm_update_kernel( + a_ptrs, + b_ptrs, + c_ptr, + a_s_ptrs, + b_s_ptrs, + size_ptr, + accum_ptr, + M: tl.constexpr, + N: tl.constexpr, + BLOCK_SIZE_M: tl.constexpr, + BLOCK_SIZE_N: tl.constexpr, + ACCUM: tl.constexpr, +): + pid_e = tl.program_id(axis=0) + pid_m = tl.program_id(axis=1) + pid_n = tl.program_id(axis=2) + size = tl.load(size_ptr + pid_e) + + a_ptr = tl.load(a_ptrs + pid_e).to(tl.pointer_type(tl.float8e4nv)) + b_ptr = tl.load(b_ptrs + pid_e).to(tl.pointer_type(tl.float8e4nv)) + a_s_ptr = tl.load(a_s_ptrs + pid_e).to(tl.pointer_type(tl.uint8)) + b_s_ptr = tl.load(b_s_ptrs + pid_e).to(tl.pointer_type(tl.uint8)) + + c_ptr = tl.load(c_ptr + pid_e).to(tl.pointer_type(tl.float32)) + K = size + k = tl.cdiv(K, 128) * 4 + + offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) + offs_k = tl.arange(0, 32) + a_ps = a_ptr + offs_m[None, :] + offs_k[:, None] * M + b_ps = b_ptr + offs_n[None, :] + offs_k[:, None] * N + a_s_ps = a_s_ptr + offs_m + b_s_ps = b_s_ptr + offs_n + + if ACCUM: + c_ps = c_ptr + offs_m[:, None] * N + offs_n[None, :] + accumulator = tl.load(c_ps) + else: + accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32) + + for i in range(k): + a = tl.load(a_ps) + b = tl.load(b_ps) + a_s = tl.exp2(tl.load(a_s_ps).to(tl.float32) - 127.0) + b_s = tl.exp2(tl.load(b_s_ps).to(tl.float32) - 127.0) + accumulator += tl.dot(tl.trans(a), b) * a_s[:, None] * b_s[None, :] + a_ps += 32 * M + b_ps += 32 * N + a_s_ps += M + b_s_ps += N + + # offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + # offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) + c_ps = c_ptr + offs_m[:, None] * N + offs_n[None, :] + # tl.store(c_ps, c) + tl.store(c_ps, accumulator, cache_modifier=".cs") + + +def triton_mxfp8_grouped_gemm_update( + a: List[torch.Tensor], + b: List[torch.Tensor], + a_s: List[torch.Tensor], + b_s: List[torch.Tensor], + m_splits: List[int], + out: Optional[List[torch.tensor]], + out_dtype: torch.dtype = torch.float32, + accumulate: bool = False, +): + """ + triton implementation to simulate mxfp8 grouped gemm + layout is defined as the same in TE: + TN: forward + NN: bakcward + NT: update(wgrad) + layout is used to optimize BLOCK SIZE + """ + device = a[0].device + sizes = torch.tensor(m_splits, dtype=torch.int64).cuda(device, non_blocking=True) + accums = torch.cumsum(sizes, 0) + a_ptrs = torch.tensor([x.data_ptr() for x in a], dtype=torch.int64).cuda( + device, non_blocking=True + ) + b_ptrs = torch.tensor([x.data_ptr() for x in b], dtype=torch.int64).cuda( + device, non_blocking=True + ) + + as_ptrs = torch.tensor([x.data_ptr() for x in a_s], dtype=torch.int64).cuda( + device, non_blocking=True + ) + bs_ptrs = torch.tensor([x.data_ptr() for x in b_s], dtype=torch.int64).cuda( + device, non_blocking=True + ) + + BLOCK_SIZE_M = 64 + BLOCK_SIZE_N = 64 + M = a[0].size(1) + N = b[0].size(1) + out_ptrs = torch.tensor([x.data_ptr() for x in out], dtype=torch.int64).cuda( + device, non_blocking=True + ) + + grid = (len(m_splits), M // BLOCK_SIZE_M, N // BLOCK_SIZE_N) # noqa + mxfp8_grouped_gemm_update_kernel[grid]( + a_ptrs, + b_ptrs, + out_ptrs, + as_ptrs, + bs_ptrs, + sizes, + accums, + M, + N, + BLOCK_SIZE_M=BLOCK_SIZE_M, + BLOCK_SIZE_N=BLOCK_SIZE_N, + ACCUM=accumulate, + num_warps=4, + num_stages=3, + ) + return out + + +@triton.jit +def native_mxfp8_gemm_kernel( + a_ptr, + b_ptr, + c_ptr, + a_s_ptr, + b_s_ptr, + M, + N: tl.constexpr, + K: tl.constexpr, + BLOCK_SIZE_M: tl.constexpr, + BLOCK_SIZE_N: tl.constexpr, + ACCUM: tl.constexpr, +): + pid_m = tl.program_id(axis=0) + pid_n = tl.program_id(axis=1) + k = K // 32 + PM = tl.cdiv(M, 128) * 128 + PN = tl.cdiv(N, 128) * 128 + + offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) + offs_k = tl.arange(0, 32) + a_ptrs = a_ptr + offs_m[:, None] * K + offs_k[None, :] + b_ptrs = b_ptr + offs_n[None, :] * K + offs_k[:, None] + + a_s_ptrs = a_s_ptr + offs_m + b_s_ptrs = b_s_ptr + offs_n + + if ACCUM: + accumulator = tl.load(c_ptr + offs_m[:, None] * N + offs_n[None, :]) + else: + accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32) + for i in range(k): + a = tl.load(a_ptrs) + b = tl.load(b_ptrs) + a_s = tl.exp2(tl.load(a_s_ptrs).to(tl.float32) - 127.0) + b_s = tl.exp2(tl.load(b_s_ptrs).to(tl.float32) - 127.0) + accumulator += tl.dot(a, b) * a_s[:, None] * b_s[None, :] + a_ptrs += 32 + b_ptrs += 32 + + a_s_ptrs += PM + b_s_ptrs += PN + + c = accumulator.to(c_ptr.dtype.element_ty) + offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) + c_ptrs = c_ptr + offs_m[:, None] * N + offs_n[None, :] + tl.store(c_ptrs, c) + + +def triton_native_mxfp8_gemm( + a: torch.Tensor, + b: torch.Tensor, + a_s: torch.Tensor, + b_s: torch.Tensor, + out: Optional[torch.tensor] = None, + out_dtype: torch.dtype = torch.bfloat16, + layout: str = "TN", + accumulate: bool = False, +): + """ + triton implementation to simulate mxfp8 grouped gemm + layout is defined as the same in TE: + TN: forward + NN: bakcward + NT: update(wgrad) + layout is used to optimize BLOCK SIZE + """ + assert a.is_contiguous() and b.is_contiguous() + assert a_s.is_contiguous() and b_s.is_contiguous() + + if layout == "TN": + assert not accumulate + M, K = a.shape + N = b.size(0) + BLOCK_SIZE_M = max([x for x in [32, 64] if M % x == 0]) + BLOCK_SIZE_N = 128 + a_s = a_s.t().contiguous() + b_s = b_s.t().contiguous() + elif layout == "NN": + assert not accumulate + M, K = a.shape + N = b.size(1) + BLOCK_SIZE_M = max([x for x in [32, 64] if M % x == 0]) + BLOCK_SIZE_N = 128 + b = triton_transpose(b) + a_s = a_s.t().contiguous() + else: + K, M = a.shape + N = b.size(1) + BLOCK_SIZE_M = 128 + BLOCK_SIZE_N = 128 + a = triton_transpose(a) + b = triton_transpose(b) + + if out is not None: + assert out.is_contiguous() + else: + out = torch.empty(M, N, dtype=out_dtype, device=a.device) + grid = (M // BLOCK_SIZE_M, N // BLOCK_SIZE_N) + native_mxfp8_gemm_kernel[grid]( + a, + b, + out, + a_s, + b_s, + M, + N, + K, + BLOCK_SIZE_M=BLOCK_SIZE_M, + BLOCK_SIZE_N=BLOCK_SIZE_N, + ACCUM=accumulate, + num_warps=4, + num_stages=3, + ) + return out + + +def triton_native_mxfp8_grouped_gemm( + a: List[torch.Tensor], + b: List[torch.Tensor], + a_s: List[torch.Tensor], + b_s: List[torch.Tensor], + out: List[torch.Tensor], + m_splits: List[int], + layout: str = "TN", +): + device = a[0].device + dtype = a[0].dtype + if layout != "NT": + outs = torch.split(out, m_splits) + else: + outs = out + for i, m in enumerate(m_splits): + triton_native_mxfp8_gemm(a[i], b[i], a_s[i], b_s[i], out=outs[i], layout=layout) + return out diff --git a/linghe/infer/__init__.py b/linghe/infer/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/linghe/infer/gate.py b/linghe/infer/gate.py new file mode 100644 index 0000000..07c34db --- /dev/null +++ b/linghe/infer/gate.py @@ -0,0 +1,156 @@ +# -*- coding: utf-8 -*- +""" +Copyright (c) Ant Financial Service Group and its affiliates. +""" + +import torch +import triton +import triton.language as tl + + +@triton.jit +def group_rms_norm_gate_kernel( + x_ptr, + gate_ptr, + weight_ptr, + out_ptr, + eps, + DIM: tl.constexpr, + D: tl.constexpr, + GROUP_SIZE: tl.constexpr, +): + pid = tl.program_id(axis=0) + weight = tl.load(weight_ptr + tl.arange(0, DIM)).to(tl.float32) + weight = tl.reshape(weight, [GROUP_SIZE, D]) + x_offs = ( + pid * DIM + tl.arange(0, GROUP_SIZE)[:, None] * D + tl.arange(0, D)[None, :] + ) + x = tl.load(x_ptr + x_offs).to(tl.float32) + offs = pid * DIM + tl.arange(0, GROUP_SIZE)[:, None] * D + tl.arange(0, D)[None, :] + g = tl.load(gate_ptr + offs).to(tl.float32) + rms = tl.rsqrt(tl.sum(x * x, axis=1) / D + eps) + x = x * rms[:, None] * weight * tl.sigmoid(g) + tl.store(out_ptr + offs, x) + + +def triton_group_rms_norm_gate( + x: torch.Tensor, + gate: torch.Tensor, + weight: torch.Tensor, + eps=1e-6, + group_size=4, + dtype=torch.bfloat16, +): + """ + norm and gate in linear attention + Args: + x: output of attn, [tokens, n_heads, head_dim] + gate: gate tensor, [tokens, dim] + weight: rms norm weight, [dim] + eps: epsilon of rms norm + group_size: group size of group rms norm + Returns: + output tensor, [tokens, dim] + """ + # row-wise read, row-wise write + assert x.is_contiguous() and gate.is_contiguous() and weight.is_contiguous() + tokens, dim = gate.shape + assert dim <= 8192 and triton.next_power_of_2(dim) == dim + d = dim // group_size + device = x.device + out = torch.empty((tokens, dim), device=device, dtype=dtype) + grid = (tokens,) + group_rms_norm_gate_kernel[grid]( + x, gate, weight, out, eps, dim, d, group_size, num_stages=3, num_warps=4 + ) + return out + + +@triton.jit +def block_group_rms_norm_gate_kernel( + x_ptr, + gate_ptr, + weight_ptr, + out_ptr, + scale_ptr, + eps, + M, + PM, + D: tl.constexpr, + d: tl.constexpr, + GROUP_SIZE: tl.constexpr, + ROUND: tl.constexpr, +): + pid = tl.program_id(axis=0) + mask = pid < M + b: tl.constexpr = d // 128 + B: tl.constexpr = D // 128 + + weight = tl.load(weight_ptr + tl.arange(0, D)).to(tl.float32) + weight = tl.reshape(weight, [GROUP_SIZE, d]) + offs = pid * D + tl.arange(0, GROUP_SIZE)[:, None] * d + tl.arange(0, d)[None, :] + x = tl.load(x_ptr + offs, mask=mask).to(tl.float32) + g = tl.load(gate_ptr + offs, mask=mask).to(tl.float32) + rms = tl.rsqrt(tl.sum(x * x, axis=1) / d + eps) + x = x * tl.sigmoid(g) * rms[:, None] * weight[None, :] + + x = tl.reshape(x, (GROUP_SIZE, b, 128)) + + scale = tl.maximum(tl.max(tl.abs(x), 2) / 448.0, 1e-30) + if ROUND: + scale = tl.exp2(tl.ceil(tl.log2(scale))) + + tl.store(scale_ptr + pid + tl.arange(0, B) * PM, tl.reshape(scale, (B,))) + + x = x / scale[:, :, None] + x = tl.reshape(x, [GROUP_SIZE, d]) + + tl.store(out_ptr + offs, x, mask=mask) + + +def triton_block_group_rms_norm_gate( + x: torch.Tensor, + gate: torch.Tensor, + weight: torch.Tensor, + eps=1e-6, + group_size=4, + round_scale=False, +): + """ + norm and gate in linear attention + Args: + x: output of attn, [tokens, n_heads * head_dim] + gate: gate tensor, [tokens, dim] + weight: rms norm weight, [dim] + eps: epsilon of rms norm + group_size: group size of group rms norm + Returns: + output tensor, [tokens, dim] + """ + assert x.is_contiguous() and gate.is_contiguous() and weight.is_contiguous() + # print(f'{x.shape=} {gate.shape=} {weight.shape=} {group_size=}') + M, D = gate.shape + assert D <= 8192 and triton.next_power_of_2(D) == D + d = D // group_size + PM = triton.cdiv(M, 4) * 4 + device = x.device + out = torch.empty((M, D), device=device, dtype=torch.float8_e4m3fn) + scale = torch.empty((D // 128, PM), device=device, dtype=torch.float32) + grid = (PM,) + block_group_rms_norm_gate_kernel[grid]( + x, + gate, + weight, + out, + scale, + eps, + M, + PM, + D, + d, + group_size, + round_scale, + num_stages=5, + num_warps=4, + ) + return out, scale[:, :M].t() diff --git a/linghe/infer/gemm.py b/linghe/infer/gemm.py new file mode 100644 index 0000000..a0986ed --- /dev/null +++ b/linghe/infer/gemm.py @@ -0,0 +1,226 @@ +""" +Copyright (c) Ant Financial Service Group and its affiliates. +""" + +import torch +import triton +import triton.language as tl + + +@triton.jit +def split_fp32_gemm_kernel( + a_ptr, + b_ptr, + c_ptr, + M, + N: tl.constexpr, + K: tl.constexpr, + SPLIT_COUNT: tl.constexpr, + BLOCK_SIZE_K: tl.constexpr, + BLOCK_SIZE_M: tl.constexpr, + BLOCK_SIZE_N: tl.constexpr, +): + pid_m = tl.program_id(axis=1) + pid_n = tl.program_id(axis=2) + pid_k = tl.program_id(axis=0) + + k = tl.cdiv(K, BLOCK_SIZE_K * SPLIT_COUNT) + offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) + offs_k = tl.arange(0, BLOCK_SIZE_K) + a_ptrs = a_ptr + pid_k * K // SPLIT_COUNT + offs_m[:, None] * K + offs_k[None, :] + b_ptrs = b_ptr + pid_k * K // SPLIT_COUNT + offs_n[None, :] * K + offs_k[:, None] + + c = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32) + for i in range(k): + a = tl.load(a_ptrs, mask=offs_m[:, None] < M) + b = tl.load(b_ptrs) + c = tl.dot(a, b, c) + a_ptrs += BLOCK_SIZE_K + b_ptrs += BLOCK_SIZE_K + offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) + c_ptrs = c_ptr + offs_m[:, None] * N + offs_n[None, :] + if SPLIT_COUNT == 1: + tl.store(c_ptrs, c, mask=offs_m[:, None] < M) + else: + tl.atomic_add(c_ptrs, c, sem="relaxed", mask=offs_m[:, None] < M) + + +def triton_split_fp32_gemm(x: torch.Tensor, w: torch.Tensor): + """ + return fp32 gemm result with fp16/bf16 inputs, + it's mainly used for MoE router GEMM + Args: + a: left matrix with fp16/bf16 precision + b: right matrix with fp16/bf16 precision + + Returns: + c: output with fp32 precision + """ + assert x.is_contiguous() and w.is_contiguous() + M, K = x.size() + N, K = w.size() + + MP = triton.next_power_of_2(M) + BLOCK_SIZE_M = max(min(MP, 64), 16) + BLOCK_SIZE_N = 64 + BLOCK_SIZE_K = 128 + SPLIT_COUNT = 1 + num_warps = 4 + num_stages = 3 + + if MP <= 64: + if N <= 16 * 128: + SPLIT_COUNT = 4 + BLOCK_SIZE_N = 16 + elif N <= 64 * 128: + BLOCK_SIZE_N = 16 + + assert N % BLOCK_SIZE_N == 0 + assert K % (BLOCK_SIZE_K * SPLIT_COUNT) == 0 + + if SPLIT_COUNT == 1: + c = torch.empty(M, N, dtype=torch.float32, device=x.device) + else: + c = torch.zeros(M, N, dtype=torch.float32, device=x.device) + + grid = ( + SPLIT_COUNT, + triton.cdiv(M, BLOCK_SIZE_M), + triton.cdiv(N, BLOCK_SIZE_N), + ) + split_fp32_gemm_kernel[grid]( + x, + w, + c, + M, + N, + K, + SPLIT_COUNT, + BLOCK_SIZE_K, + BLOCK_SIZE_M, + BLOCK_SIZE_N, + num_warps=num_warps, + num_stages=num_stages, + ) + return c + + +@triton.jit +def split_tile_block_fp8_gemm_kernel( + a_ptr, + b_ptr, + c_ptr, + a_s_ptr, + b_s_ptr, + stride_a_scale, + M, + N: tl.constexpr, + K: tl.constexpr, + BLOCK_SIZE_M: tl.constexpr, + BLOCK_SIZE_N: tl.constexpr, + BLOCK_SIZE_K: tl.constexpr, + SPLIT_COUNT: tl.constexpr, + TRANSPOSE_A_SCALE: tl.constexpr, +): + # a tilewise quantization, b blockwise quantization. + pid_k = tl.program_id(axis=0) + pid_m = tl.program_id(axis=1) + pid_n = tl.program_id(axis=2) + offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) + offs_k = tl.arange(0, BLOCK_SIZE_K) + a_ptrs = a_ptr + offs_m[:, None] * K + offs_k[None, :] + b_ptrs = b_ptr + offs_n[None, :] * K + offs_k[:, None] + if TRANSPOSE_A_SCALE: + a_s_ptrs = a_s_ptr + offs_m + else: + a_s_ptrs = a_s_ptr + offs_m * K // 128 + b_s_ptrs = b_s_ptr + pid_n * BLOCK_SIZE_N // BLOCK_SIZE_K * K // BLOCK_SIZE_K + + accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32) + k = K // BLOCK_SIZE_K // SPLIT_COUNT + for i in range(pid_k * k, pid_k * k + k): + a = tl.load(a_ptrs + i * BLOCK_SIZE_K, mask=offs_m[:, None] < M) + b = tl.load(b_ptrs + i * BLOCK_SIZE_K) + if TRANSPOSE_A_SCALE: + a_s = tl.load(a_s_ptrs + i * stride_a_scale, mask=offs_m < M) + else: + a_s = tl.load(a_s_ptrs + i, mask=offs_m < M) + b_s = tl.load(b_s_ptrs + i) + scale = a_s[:, None] * b_s + accumulator += tl.dot(a, b) * scale + # accumulators = tl.dot(a, b, accumulator) + # accumulator += (accumulators - accumulator) * a_s[:, None] * b_s + offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) + c_ptrs = c_ptr + offs_m[:, None] * N + offs_n[None, :] + # tl.store(c_ptrs, c) + mask = offs_m[:, None] < M + if SPLIT_COUNT > 1: + tl.atomic_add(c_ptrs, accumulator, mask=mask, sem="relaxed") + else: + tl.store(c_ptrs, accumulator, mask=mask) + + +def triton_split_tile_block_fp8_gemm( + a: torch.Tensor, + b: torch.Tensor, + a_s: torch.Tensor, + b_s: torch.Tensor, + out_dtype=torch.bfloat16, +): + assert a.is_contiguous() and b.is_contiguous() + assert b_s.is_contiguous() + + K = a.size(-1) + M = a.numel() // K + N = b.size(0) + stride_a_scale = a_s.stride(1) + # print(f'{a.shape=} {a.stride()=} {b.shape=} {b.stride()=} {a_s.shape=} {a_s.stride()=} {b_s.shape=} {b_s.stride()=}') + TRANSPOSE_A_SCALE = not a_s.is_contiguous() + + MP = triton.next_power_of_2(M) + BLOCK_SIZE_M = max(min(MP, 64), 16) + BLOCK_SIZE_N = 64 + BLOCK_SIZE_K = 128 + SPLIT_COUNT = 1 + num_warps = 4 + num_stages = 3 + + if MP <= 64: + if N <= 16 * 256: + SPLIT_COUNT = min(K // BLOCK_SIZE_K, 4) + BLOCK_SIZE_N = 16 + num_warps = 2 if MP <= 16 else 4 + else: + BLOCK_SIZE_N = 16 + num_warps = 2 if MP <= 16 else 4 + elif M >= 1024 and N >= 4096: + BLOCK_SIZE_N = 128 + + grid = (SPLIT_COUNT, triton.cdiv(M, BLOCK_SIZE_M), triton.cdiv(N, BLOCK_SIZE_N)) + if SPLIT_COUNT > 1: + c = torch.zeros(M, N, dtype=out_dtype, device=a.device) + else: + c = torch.empty(M, N, dtype=out_dtype, device=a.device) + split_tile_block_fp8_gemm_kernel[grid]( + a, + b, + c, + a_s, + b_s, + stride_a_scale, + M, + N, + K, + BLOCK_SIZE_M, + BLOCK_SIZE_N, + BLOCK_SIZE_K, + SPLIT_COUNT, + TRANSPOSE_A_SCALE, + num_stages=num_stages, + num_warps=num_warps, + ) + return c diff --git a/linghe/infer/grouped_gemm.py b/linghe/infer/grouped_gemm.py new file mode 100644 index 0000000..4688a24 --- /dev/null +++ b/linghe/infer/grouped_gemm.py @@ -0,0 +1,146 @@ +""" +Copyright (c) Ant Financial Service Group and its affiliates. +""" + +from typing import Optional +import torch +import triton +import triton.language as tl + + +@triton.jit +def fp8_grouped_gemm_kernel( + a_ptr, + b_ptr, + as_ptr, + bs_ptr, + c_ptr, + token_ids_ptr, + expert_ids_ptr, + token_count_ptr, + padding_value, + topk, + as_stride, + M, + N: tl.constexpr, + K: tl.constexpr, + BLOCK_SIZE_K: tl.constexpr, + BLOCK_SIZE_M: tl.constexpr, + BLOCK_SIZE_N: tl.constexpr, + TRANSPOSE_A_SCALE: tl.constexpr, +): + pid_m = tl.program_id(axis=0) + pid_n = tl.program_id(axis=1) + + token_count = tl.load(token_count_ptr) + eid = tl.load(expert_ids_ptr + pid_m) + if pid_m * BLOCK_SIZE_M >= token_count: + return + + offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) + token_ids = tl.load(token_ids_ptr + offs_m) + tids = token_ids // topk + mask = token_ids < padding_value + + k = K // BLOCK_SIZE_K + offs_k = tl.arange(0, BLOCK_SIZE_K) + a_ptrs = a_ptr + tids[:, None] * K + offs_k[None, :] + b_ptrs = b_ptr + eid * N * K + offs_n[None, :] * K + offs_k[:, None] + + sk = K // 128 + if TRANSPOSE_A_SCALE: + as_ptrs = as_ptr + tids + else: + as_ptrs = as_ptr + tids * sk + bs_ptrs = bs_ptr + eid * N // 128 * sk + pid_n * BLOCK_SIZE_N // 128 * sk + + accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32) + for i in range(k): + a = tl.load(a_ptrs + i * BLOCK_SIZE_K, mask=mask[:, None]) + b = tl.load(b_ptrs + i * BLOCK_SIZE_K) + if TRANSPOSE_A_SCALE: + a_s = tl.load(as_ptrs + i * BLOCK_SIZE_K // 128 * as_stride, mask=mask) + else: + a_s = tl.load(as_ptrs + i * BLOCK_SIZE_K // 128, mask=mask) + b_s = tl.load(bs_ptrs + i * BLOCK_SIZE_K // 128) + scale = a_s[:, None] * b_s + accumulator += tl.dot(a, b) * scale + offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) + c_ptrs = c_ptr + token_ids[:, None] * N + offs_n[None, :] + tl.store(c_ptrs, accumulator, mask=mask[:, None]) + + +def triton_fp8_grouped_gemm( + xq: torch.Tensor, + wq: torch.Tensor, + xs: torch.Tensor, + ws: torch.Tensor, + token_ids: torch.Tensor, + expert_ids: torch.Tensor, + token_count: torch.Tensor, + c: Optional[torch.Tensor] = None, + block_size_m: int = 16, + block_size_n: int = 128, + padding_value: int = 9, + topk: int = 9, +): + """ + grouped gemm with fp8 + Returns: + c: output with fp32 precision + """ + assert xq.is_contiguous() + assert wq.is_contiguous() and ws.is_contiguous() + + M = expert_ids.numel() + N, K = wq.shape[1:] + device = xq.device + if c is None: + c = torch.empty(M, N, dtype=torch.bfloat16, device=device) + else: + assert c.is_contiguous() + + TRANSPOSE_A_SCALE = not xs.is_contiguous() + xs_stride = xs.stride(1) + BLOCK_SIZE_K = 128 # only support BLOCK_SIZE_K <= 128 + num_warps = 4 + num_stages = 5 + + if block_size_m == 16: + if M <= 16 * 64 and N <= 512: + block_size_n = 16 + num_warps = 2 + num_stages = 3 + else: + block_size_n = 64 + num_warps = 4 + num_stages = 3 + elif block_size_m == 32: + block_size_n = 64 + + grid = (triton.cdiv(M, block_size_m), triton.cdiv(N, block_size_n)) + fp8_grouped_gemm_kernel[grid]( + xq, + wq, + xs, + ws, + c, + token_ids, + expert_ids, + token_count, + padding_value, + topk, + xs_stride, + M, + N, + K, + BLOCK_SIZE_K, + block_size_m, + block_size_n, + TRANSPOSE_A_SCALE, + num_warps=num_warps, + num_stages=num_stages, + ) + return c diff --git a/linghe/infer/norm.py b/linghe/infer/norm.py new file mode 100644 index 0000000..ab0cd52 --- /dev/null +++ b/linghe/infer/norm.py @@ -0,0 +1,220 @@ +# -*- coding: utf-8 -*- +""" +Copyright (c) Ant Financial Service Group and its affiliates. +""" + +from typing import Optional + +import torch +import triton +import triton.language as tl + + +@triton.jit +def rms_norm_and_block_quant_kernel( + x_ptr, + weight_ptr, + residual_ptr, + out_ptr, + scale_ptr, + eps, + M, + PM, + N: tl.constexpr, + nb: tl.constexpr, + W: tl.constexpr, + ROUND: tl.constexpr, +): + pid = tl.program_id(axis=0) + + # row-wise read, row-wise write + weight = tl.load(weight_ptr + tl.arange(0, N)).to(tl.float32)[None, :] + offs = pid * W * N + tl.arange(0, W)[:, None] * N + tl.arange(0, N)[None, :] + # PM = (M + 3) // 4 * 4 + indices = pid * W + tl.arange(0, W) + x = tl.load(x_ptr + offs, mask=indices[:, None] < M).to(tl.float32) + + if residual_ptr is not None: + r = tl.load(residual_ptr + offs, mask=indices[:, None] < M).to(tl.float32) + x = x + r + tl.debug_barrier() + tl.store(residual_ptr + offs, x, mask=indices[:, None] < M) + + rms = tl.rsqrt(tl.sum(x * x, axis=1) / N + eps) + x = x * rms[:, None] * weight + tl.store(x_ptr + offs, x, mask=indices[:, None] < M) + + x = tl.reshape(x, [W, nb, 128], can_reorder=False) + + scale = tl.max(tl.abs(x), 2) / 448.0 + scale = tl.where(scale == 0.0, 1.0, scale) + if ROUND: + scale = tl.exp2(tl.ceil(tl.log2(scale))) + tl.store( + scale_ptr + tl.arange(0, nb)[:, None] * PM + indices[None, :], + tl.trans(scale), + mask=indices[None, :] < M, + ) + + x = x / scale[:, :, None] + x = tl.reshape(x, [W, N], can_reorder=False) + + tl.store(out_ptr + offs, x, mask=indices[:, None] < M) + + +def triton_rms_norm_and_block_quant( + x: torch.Tensor, + weight: torch.Tensor, + residual: Optional[torch.Tensor] = None, + eps: float = 1e-6, + round_scale: bool = False, +): + """ + Fused RMSNorm forward and block quantization. + Args: + x: Input tensor, shape [M, N] + weight: RMSNorm weight, shape [N] + residual: Residual tensor, shape [M, N] + eps: epsilon value for L2 normalization. + out: output of quantization data + scale: output of quantization scale. + rms: output of rms + round_scale: Set whether to force power of 2 scales. + Returns: + - x: rmsnorm data. + - out: quantization data. + - scale: quantization scale. + - residual: residual tensor. + """ + assert x.is_contiguous() and weight.is_contiguous() + if residual is not None: + assert residual.is_contiguous() + M, N = x.shape + assert N <= 8192 and 8192 % N == 0 and N >= 2048 + device = x.device + + out = torch.empty((M, N), device=device, dtype=torch.float8_e4m3fn) + + PM = (M + 3) // 4 * 4 + scale = torch.empty((N // 128, PM), device=device, dtype=torch.float32) + + if M > 256: + W = 8192 // N + else: + W = 1 + grid = (triton.cdiv(triton.cdiv(M, 4) * 4, W),) + + rms_norm_and_block_quant_kernel[grid]( + x, + weight, + residual, + out, + scale, + eps, + M, + PM, + N, + N // 128, + W, + round_scale, + num_stages=3, + num_warps=4, + ) + + return out, scale[:, :M].t(), residual + + +@triton.jit +def rms_norm_and_token_quant_kernel( + x_ptr, + weight_ptr, + residual_ptr, + out_ptr, + scale_ptr, + eps, + n, + N: tl.constexpr, + ROUND: tl.constexpr, +): + pid = tl.program_id(axis=0) + + weight = tl.load(weight_ptr + tl.arange(0, N)).to(tl.float32) + offs = pid * n + tl.arange(0, N) + mask = tl.arange(0, N) < n + x = tl.load(x_ptr + offs, mask=mask).to(tl.float32) + + if residual_ptr is not None: + r = tl.load(residual_ptr + offs, mask=mask).to(tl.float32) + x = x + r + tl.debug_barrier() + tl.store(residual_ptr + offs, x, mask=mask) + + rms = tl.rsqrt(tl.sum(x * x, axis=0) / n + eps) + x = x * rms * weight + + scale = tl.max(tl.abs(x), 0) / 448.0 + scale = tl.where(scale == 0.0, 1.0, scale) + if ROUND: + scale = tl.exp2(tl.ceil(tl.log2(scale))) + tl.store(scale_ptr + pid, scale, mask=mask) + + x = x / scale + + tl.store(out_ptr + offs, x.to(out_ptr.dtype.element_ty), mask=mask) + + +def triton_rms_norm_and_token_quant( + x: torch.Tensor, + weight: torch.Tensor, + residual: Optional[torch.Tensor] = None, + eps: float = 1e-6, + out: Optional[torch.Tensor] = None, + scale: Optional[torch.Tensor] = None, + round_scale: bool = False, +): + """ + Fused RMSNorm forward and block quantization. + Args: + x: Input tensor, shape [M, N] + weight: RMSNorm weight, shape [N] + residual: Residual tensor, shape [M, N] + eps: epsilon value for L2 normalization. + out: output of quantization data + scale: output of quantization scale. + rms: output of rms + round_scale: Set whether to force power of 2 scales. + Returns: + - out: quantization data. + - scale: quantization scale. + - residual: residual tensor. + """ + assert x.is_contiguous() and weight.is_contiguous() + if residual is not None: + assert residual.is_contiguous() + M, n = x.shape + N = triton.next_power_of_2(n) + device = x.device + + if out is None: + out = torch.empty((M, n), device=device, dtype=torch.float8_e4m3fn) + + if scale is None: + scale = torch.empty((M, 1), device=device, dtype=torch.float32) + + grid = (M,) + + rms_norm_and_token_quant_kernel[grid]( + x, + weight, + residual, + out, + scale, + eps, + n, + N, + round_scale, + num_stages=3, + num_warps=2, + ) + + return out, scale, residual diff --git a/linghe/infer/quant.py b/linghe/infer/quant.py new file mode 100644 index 0000000..516583e --- /dev/null +++ b/linghe/infer/quant.py @@ -0,0 +1,62 @@ +# -*- coding: utf-8 -*- +""" +Copyright (c) Ant Financial Service Group and its affiliates. +""" + +import torch +import triton +import triton.language as tl + + +@triton.jit +def group_quant_kernel( + x_ptr, + y_ptr, + s_ptr, + N, + BLOCK_SIZE: tl.constexpr, + n: tl.constexpr, + ROUND: tl.constexpr, +): + pid = tl.program_id(axis=0) + bid = tl.program_id(axis=1) + offs = pid * N + bid * n + tl.arange(0, n) + SB: tl.constexpr = n // BLOCK_SIZE + soffs = pid * (N // BLOCK_SIZE) + bid * SB + tl.arange(0, SB) + x = tl.load(x_ptr + offs).to(tl.float32) + x = tl.reshape(x, (SB, BLOCK_SIZE), can_reorder=False) + s = tl.maximum(tl.max(tl.abs(x), 1) / 448.0, 1e-30) + if ROUND: + s = tl.exp2(tl.ceil(tl.log2(s))) + y = x / s[:, None] + y = y.to(y_ptr.dtype.element_ty) + y = tl.reshape(y, (n,), can_reorder=False) + tl.store(y_ptr + offs, y) + tl.store(s_ptr + soffs, s) + + +def triton_group_quant(x, dtype=torch.float8_e4m3fn, group_size=128, round_scale=False): + """ + groupwise quantize x, group is in under rowwise format + Args: + x: input tensor + group_size: group wise + round_scale: whether round scale to power of 2 + + Returns: + - y: quantized tensor, float8_e4m3fn + - s: quantization scale, float32 + """ + M, N = x.shape + assert N % group_size == 0 + assert x.is_contiguous() + + y = torch.empty((M, N), device=x.device, dtype=dtype) + s = torch.empty(M, N // group_size, device=x.device, dtype=torch.float32) + B = 4 if M <= 256 else 1 + n = N // B + grid = (M, B) + group_quant_kernel[grid]( + x, y, s, N, group_size, n, round_scale, num_stages=5, num_warps=4 + ) + return y, s diff --git a/linghe/infer/rope.py b/linghe/infer/rope.py new file mode 100644 index 0000000..d90feef --- /dev/null +++ b/linghe/infer/rope.py @@ -0,0 +1,574 @@ +# -*- coding: utf-8 -*- +import torch +import triton +import triton.language as tl + + +@triton.jit +def varlen_qk_norm_and_half_rope_kernel( + qkv_ptr, + q_norm_weight_ptr, + k_norm_weight_ptr, + freqs_ptr, + position_ids, + qo_ptr, + ko_ptr, + vo_ptr, + stride, + eps, + linear_scale_value, + H: tl.constexpr, + h: tl.constexpr, + PH: tl.constexpr, + ph: tl.constexpr, + B: tl.constexpr, + D: tl.constexpr, + d: tl.constexpr, + INTERLEAVED: tl.constexpr, + SILU: tl.constexpr, + SCALE: tl.constexpr, +): + pid = tl.program_id(0) + hid = tl.program_id(1) + + pos = tl.load(position_ids + pid) + + DD = D * 2 + + cos = tl.load(freqs_ptr + pos * D + tl.arange(0, D) % d).to(tl.float32) + + sin = tl.load(freqs_ptr + pos * D + d + tl.arange(0, D) % d).to(tl.float32) + + signs = tl.arange(0, 2).to(tl.float32) * 2 - 1 + + q_weight_0 = tl.load(q_norm_weight_ptr + tl.arange(0, D)).to(tl.float32) + q_weight_1 = tl.load(q_norm_weight_ptr + D + tl.arange(0, D)).to(tl.float32) + q_ptr = qkv_ptr + G: tl.constexpr = H // h + + if INTERLEAVED: + if B == 1: + # query per token + row_offs = tl.arange(0, PH) + tl.arange(0, PH) // G * 2 + row_mask = tl.arange(0, PH)[:, None] < H + else: + # query per kv head + row_offs = hid * (G + 2) + tl.arange(0, G) + row_mask = None + else: + if B == 1: + # query per token + row_offs = tl.arange(0, PH) + row_mask = row_offs[:, None] < H + else: + # query per kv head + row_offs = hid * G + tl.arange(0, G) + row_mask = None + + q0 = tl.load( + q_ptr + pid * stride + DD * row_offs[:, None] + tl.arange(0, D)[None, :], + mask=row_mask, + ).to(tl.float32) + q1 = tl.load( + q_ptr + pid * stride + D + DD * row_offs[:, None] + tl.arange(0, D)[None, :], + mask=row_mask, + ).to(tl.float32) + + if SILU: + q0 = q0 * tl.sigmoid(q0) + q1 = q1 * tl.sigmoid(q1) + rms = tl.rsqrt((tl.sum(q0 * q0, 1) + tl.sum(q1 * q1, 1)) / DD + eps) + q1 *= rms[:, None] + q1 *= q_weight_1 + + if SCALE: + q1 *= linear_scale_value + + if B == 1: + q_mask = tl.arange(0, PH)[:, None] < H + tl.store( + qo_ptr + + pid * H * DD + + D + + DD * tl.arange(0, PH)[:, None] + + tl.arange(0, D)[None, :], + q1, + mask=q_mask, + ) + else: + tl.store( + qo_ptr + + pid * H * DD + + D + + hid * G * DD + + DD * tl.arange(0, G)[:, None] + + tl.arange(0, D)[None, :], + q1, + mask=None, + ) + + q0 *= rms[:, None] + q0 *= q_weight_0 + + if SCALE: + q0 *= linear_scale_value + + if B == 1: + qr = tl.reshape( + tl.permute( + tl.flip(tl.permute(tl.reshape(q0, (PH, 2, d)), (0, 2, 1)), dim=2) + * signs, + (0, 2, 1), + ), + (PH, D), + ) + q0 = q0 * cos + qr * sin + tl.store( + qo_ptr + + pid * H * DD + + DD * tl.arange(0, PH)[:, None] + + tl.arange(0, D)[None, :], + q0, + mask=q_mask, + ) + else: + qr = tl.reshape( + tl.permute( + tl.flip(tl.permute(tl.reshape(q0, (G, 2, d)), (0, 2, 1)), dim=2) + * signs, + (0, 2, 1), + ), + (G, D), + ) + q0 = q0 * cos + qr * sin + tl.store( + qo_ptr + + pid * H * DD + + hid * G * DD + + DD * tl.arange(0, G)[:, None] + + tl.arange(0, D)[None, :], + q0, + mask=None, + ) + + k_weight_0 = tl.load(k_norm_weight_ptr + tl.arange(0, D)).to(tl.float32) + k_weight_1 = tl.load(k_norm_weight_ptr + D + tl.arange(0, D)).to(tl.float32) + + if INTERLEAVED: + if B == 1: + k_ptr = qkv_ptr + DD * G + row_offs = tl.arange(0, ph) * (G + 2) + row_mask = row_offs[:, None] < (h * (G + 2)) + else: + k_ptr = qkv_ptr + DD * G + row_offs = hid * (G + 2) + tl.arange(0, 1) * (G + 2) + row_mask = None + else: + if B == 1: + row_offs = tl.arange(0, ph) + k_ptr = qkv_ptr + DD * H + row_mask = tl.arange(0, ph)[:, None] < h + else: + row_offs = hid + tl.arange(0, 1) + k_ptr = qkv_ptr + DD * H + row_mask = None + + k0 = tl.load( + k_ptr + pid * stride + DD * row_offs[:, None] + tl.arange(0, D)[None, :], + mask=row_mask, + ).to(tl.float32) + k1 = tl.load( + k_ptr + pid * stride + D + DD * row_offs[:, None] + tl.arange(0, D)[None, :], + mask=row_mask, + ).to(tl.float32) + + if SILU: + k0 = k0 * tl.sigmoid(k0) + k1 = k1 * tl.sigmoid(k1) + rms = tl.rsqrt((tl.sum(k0 * k0, 1) + tl.sum(k1 * k1, 1)) / DD + eps) + k1 *= rms[:, None] + k1 *= k_weight_1 + + if B == 1: + k_mask = tl.arange(0, ph)[:, None] < h + tl.store( + ko_ptr + + pid * h * DD + + D + + DD * tl.arange(0, ph)[:, None] + + tl.arange(0, D)[None, :], + k1, + mask=k_mask, + ) + else: + tl.store( + ko_ptr + + pid * h * DD + + D + + hid * DD + + DD * tl.arange(0, 1)[:, None] + + tl.arange(0, D)[None, :], + k1, + mask=None, + ) + + k0 *= rms[:, None] + k0 *= k_weight_0 + if B == 1: + kr = tl.reshape( + tl.permute( + tl.flip(tl.permute(tl.reshape(k0, (ph, 2, d)), (0, 2, 1)), dim=2) + * signs, + (0, 2, 1), + ), + (ph, D), + ) + k0 = k0 * cos + kr * sin + tl.store( + ko_ptr + + pid * h * DD + + DD * tl.arange(0, ph)[:, None] + + tl.arange(0, D)[None, :], + k0, + mask=k_mask, + ) + else: + kr = tl.reshape( + tl.permute( + tl.flip(tl.permute(tl.reshape(k0, (1, 2, d)), (0, 2, 1)), dim=2) + * signs, + (0, 2, 1), + ), + (1, D), + ) + k0 = k0 * cos + kr * sin + tl.store( + ko_ptr + + pid * h * DD + + hid * DD + + DD * tl.arange(0, 1)[:, None] + + tl.arange(0, D)[None, :], + k0, + mask=None, + ) + + if INTERLEAVED: + if B == 1: + v_ptr = qkv_ptr + DD * G + DD + row_offs = hid * h * (G + 2) + tl.arange(0, ph) * (G + 2) + row_mask = row_offs[:, None] < (h * (G + 2)) + else: + v_ptr = qkv_ptr + DD * G + DD + row_offs = hid * (G + 2) + tl.arange(0, 1) * (G + 2) + row_mask = None + else: + if B == 1: + v_ptr = qkv_ptr + DD * H + DD * h + row_offs = hid * h + tl.arange(0, ph) + row_mask = tl.arange(0, ph)[:, None] < h + else: + v_ptr = qkv_ptr + DD * H + DD * h + row_offs = hid + tl.arange(0, 1) + row_mask = None + + v0 = tl.load( + v_ptr + pid * stride + DD * row_offs[:, None] + tl.arange(0, D)[None, :], + mask=row_mask, + ).to(tl.float32) + v1 = tl.load( + v_ptr + pid * stride + D + DD * row_offs[:, None] + tl.arange(0, D)[None, :], + mask=row_mask, + ).to(tl.float32) + + if SILU: + v0 = v0 * tl.sigmoid(v0) + v1 = v1 * tl.sigmoid(v1) + + v_mask = tl.arange(0, ph)[:, None] < h + if B == 1: + tl.store( + vo_ptr + + pid * h * DD + + DD * tl.arange(0, ph)[:, None] + + tl.arange(0, D)[None, :], + v0, + mask=v_mask, + ) + tl.store( + vo_ptr + + pid * h * DD + + D + + DD * tl.arange(0, ph)[:, None] + + tl.arange(0, D)[None, :], + v1, + mask=v_mask, + ) + else: + tl.store( + vo_ptr + + pid * h * DD + + hid * DD + + DD * tl.arange(0, 1)[:, None] + + tl.arange(0, D)[None, :], + v0, + mask=None, + ) + tl.store( + vo_ptr + + pid * h * DD + + hid * DD + + D + + DD * tl.arange(0, 1)[:, None] + + tl.arange(0, D)[None, :], + v1, + mask=None, + ) + + +def triton_varlen_qk_norm_and_half_rope( + qkv, + q_norm_weight, + k_norm_weight, + freqs, + position_ids, + H=32, + h=4, + eps=1e-6, + scaling=1.0, + interleaved=False, + silu=False, + linear_scale=False, + output_dtype=None, +): + """ + split qkv to q/k/v, apply qk norm and half rope to q/k + TOOD: support arbitrary rotary percent rather than 0.5 + Args: + qkv: QKV tensor with size of [S, dim] + q_norm_weight: rms norm weight for query + k_norm_weight: rms norm weight for key + freqs: Freqs tensor based on half dim. + H: Number of attention heads. + h: Number of key/value heads. + eps: epsilon value for L2 normalization. + interleaved: whether head of qkv is interleaved, + interleaved: [q...qkvq...qkv] + non-interleaved: [q...qk...kv...v] + silu: apply silu on qkv before qk norm and rope + output_dtype: dtype of output tensors + Returns: + - qo: shape [S, H, head_dim] + - ko: shape [S, h, head_dim] + - vo: shape [S, h, head_dim] + """ + assert qkv.is_contiguous() and freqs.is_contiguous() + assert q_norm_weight.is_contiguous() and k_norm_weight.is_contiguous() + T, Dim = qkv.shape + stride = qkv.stride(0) # qkv may be a slice of a tensor + D = Dim // (H + 2 * h) + dtype = qkv.dtype if output_dtype is None else output_dtype + device = qkv.device + qo = torch.empty((T, H, D), dtype=dtype, device=device) + ko = torch.empty((T, h, D), dtype=dtype, device=device) + vo = torch.empty((T, h, D), dtype=dtype, device=device) + + num_stages = 3 + num_warps = 2 + + PH = triton.next_power_of_2(H) + ph = triton.next_power_of_2(h) + + if h >= 2 and T < 128: + B = h + else: + B = 1 + grid = (T, B) + + varlen_qk_norm_and_half_rope_kernel[grid]( + qkv, + q_norm_weight, + k_norm_weight, + freqs, + position_ids, + qo, + ko, + vo, + stride, + eps, + scaling, + H, + h, + PH, + ph, + B, + D // 2, + D // 4, + interleaved, + silu, + linear_scale, + num_stages=num_stages, + num_warps=num_warps, + ) + return qo, ko, vo + + +@triton.jit +def mla_rope_kernel( + q_ptr, + k_ptr, + freqs_ptr, + position_ids_ptr, + q_stride_0, + q_stride_1, + k_stride_0, + k_stride_1, + H: tl.constexpr, + h: tl.constexpr, + SINGLE: tl.constexpr, + D: tl.constexpr, + d: tl.constexpr, + INTERLEAVE: tl.constexpr, +): + pid = tl.program_id(0) + hid = tl.program_id(1) + + pos = tl.load(position_ids_ptr + pid) + + cos = tl.load(freqs_ptr + pos * D + tl.arange(0, D) % d).to(tl.float32) + sin = tl.load(freqs_ptr + pos * D + d + tl.arange(0, D) % d).to(tl.float32) + if INTERLEAVE: + cos = tl.reshape(tl.trans(tl.reshape(cos, (2, d))), (D,)) + sin = tl.reshape(tl.trans(tl.reshape(sin, (2, d))), (D,)) + + signs = tl.arange(0, 2).to(tl.float32) * 2 - 1 + + q = tl.load( + q_ptr + + pid * q_stride_0 + + hid * h * q_stride_1 + + tl.arange(0, h)[:, None] * q_stride_1 + + tl.arange(0, D)[None, :] + ).to(tl.float32) + + if INTERLEAVE: + qr = tl.reshape(tl.flip(tl.reshape(q, (h, d, 2)), dim=2) * signs, (h, D)) + else: + qr = tl.reshape( + tl.permute( + tl.flip(tl.permute(tl.reshape(q, (h, 2, d)), (0, 2, 1)), dim=2) * signs, + (0, 2, 1), + ), + (h, D), + ) + q = q * cos + qr * sin + tl.store( + q_ptr + + pid * q_stride_0 + + hid * h * q_stride_1 + + tl.arange(0, h)[:, None] * q_stride_1 + + tl.arange(0, D)[None, :], + q, + ) + + if SINGLE: + if hid == 0: + k = tl.load(k_ptr + pid * k_stride_0 + tl.arange(0, D)).to(tl.float32) + + if INTERLEAVE: + kr = tl.reshape(tl.flip(tl.reshape(k, (d, 2)), dim=1) * signs, (D,)) + else: + kr = tl.reshape( + tl.permute( + tl.flip(tl.permute(tl.reshape(k, (2, d)), (1, 0)), dim=1) + * signs, + (1, 0), + ), + (D,), + ) + k = k * cos + kr * sin + tl.store(k_ptr + pid * k_stride_0 + tl.arange(0, D), k) + else: + + k = tl.load( + k_ptr + + pid * k_stride_0 + + hid * h * k_stride_1 + + tl.arange(0, h)[:, None] * k_stride_1 + + tl.arange(0, D)[None, :] + ).to(tl.float32) + + if INTERLEAVE: + kr = tl.reshape(tl.flip(tl.reshape(k, (h, d, 2)), dim=2) * signs, (h, D)) + else: + kr = tl.reshape( + tl.permute( + tl.flip(tl.permute(tl.reshape(k, (h, 2, d)), (0, 2, 1)), dim=2) + * signs, + (0, 2, 1), + ), + (h, D), + ) + k = k * cos + kr * sin + tl.store( + k_ptr + + pid * k_stride_0 + + hid * h * k_stride_1 + + tl.arange(0, h)[:, None] * k_stride_1 + + tl.arange(0, D)[None, :], + k, + ) + + +def triton_mla_rope(q, k, freqs, position_ids, interleave=True): + """ + apply MLA-type rope + Args: + q: query tensor, [t, n_heads, 64] + k: key tensor, [t, 1, 64] + freqs: rope freqs, [len, 64] + position_ids: position_ids for rope + interleave: whether q/k is interleaved + Returns: + + """ + + assert freqs.is_contiguous() + + N, H, D = q.shape + hk = k.shape[1] + assert hk == H or hk == 1 + SINGLE = k.shape[1] == 1 + q_stride_0 = q.stride(0) + q_stride_1 = q.stride(1) + k_stride_0 = k.stride(0) + k_stride_1 = k.stride(1) + + num_stages = 3 + num_warps = 4 + if N <= 64 and H % 4 == 0: + B = 4 + h = H // B + else: + B = 1 + h = H + + grid = (N, B) + mla_rope_kernel[grid]( + q, + k, + freqs, + position_ids, + q_stride_0, + q_stride_1, + k_stride_0, + k_stride_1, + H, + h, + SINGLE, + D, + D // 2, + interleave, + num_stages=num_stages, + num_warps=num_warps, + ) + return q, k diff --git a/linghe/infer/silu.py b/linghe/infer/silu.py new file mode 100644 index 0000000..7560817 --- /dev/null +++ b/linghe/infer/silu.py @@ -0,0 +1,223 @@ +# -*- coding: utf-8 -*- +""" +Copyright (c) Ant Financial Service Group and its affiliates. +""" + +import torch +import triton +import triton.language as tl + + +@triton.jit +def silu_and_block_quant_kernel( + x_ptr, + weight_ptr, + out_ptr, + scale_ptr, + routing_scale, + M, + n: tl.constexpr, + K: tl.constexpr, + ROUND: tl.constexpr, + TRANSPOSE: tl.constexpr, +): + rid = tl.program_id(axis=0) + cid = tl.program_id(axis=1) + + offs = ( + rid * n * 2 + + cid * K * 128 + + tl.arange(0, K)[:, None] * 128 + + tl.arange(0, 128)[None, :] + ) + + x1 = tl.load(x_ptr + offs).to(tl.float32) + x2 = tl.load(x_ptr + n + offs).to(tl.float32) + x = x1 * tl.sigmoid(x1) * x2 * routing_scale + if weight_ptr is not None: + weight = tl.load(weight_ptr + rid).to(tl.float32) + x = x * weight + + scale = tl.maximum(tl.max(x.abs(), 1) / 448, 1e-30) + if ROUND: + scale = tl.exp2(tl.ceil(tl.log2(scale))) + + if TRANSPOSE: + PM = tl.cdiv(M, 4) * 4 + tl.store(scale_ptr + cid * K * PM + rid + tl.arange(0, K) * PM, scale) + else: + tl.store(scale_ptr + rid * n // 128 + cid * K + tl.arange(0, K), scale) + + xq = x / scale[:, None] + tl.store( + out_ptr + + rid * n + + cid * K * 128 + + tl.arange(0, K)[:, None] * 128 + + tl.arange(0, 128)[None, :], + xq, + ) + + +def triton_silu_and_block_quant( + x, + weight=None, + out=None, + scale=None, + routing_scale=1.0, + transpose_scale=False, + round_scale=False, +): + """ + fused silu and blockwise quantization in mlp/moe kernel + scale is not transposed, which is different from deepgemm kernel + Args: + x: input tensor + weight: router weight + round_scale: whether round scale to power of 2 + Returns: + - out: quantized tensor + - scale: quantization scale + """ + assert x.is_contiguous() + if weight is not None: + # print(f'{x.shape=} {weight.shape=} {weight.stride()=}') + assert weight.is_contiguous() + + M, N = x.shape + n = N // 2 + assert n % 128 == 0 + + if weight is not None: + M = weight.numel() + + device = x.device + if out is None: + out = torch.empty((M, n), device=device, dtype=torch.float8_e4m3fn) + else: + assert out.is_contiguous() + + if scale is None: + if transpose_scale: + scale = torch.empty( + (n // 128, (M + 3) // 4 * 4), device=device, dtype=torch.float32 + ) + else: + scale = torch.empty((M, n // 128), device=device, dtype=torch.float32) + else: + assert scale.is_contiguous() + + num_warps = 4 + num_stages = 3 + B = n // 128 + if M >= 256: + if B % 4 == 0: + B = B // 4 + elif B % 2 == 0: + B = B // 2 + elif M >= 64: + if B % 2 == 0: + B = B // 2 + else: + num_warps = 2 + + K = n // B // 128 + grid = (M, B) + silu_and_block_quant_kernel[grid]( + x, + weight, + out, + scale, + routing_scale, + M, + n, + K, + round_scale, + transpose_scale, + num_stages=num_stages, + num_warps=num_warps, + ) + + if transpose_scale: + scale = scale[:, :M].t() + return out, scale + + +@triton.jit +def silu_and_token_quant_kernel( + x_ptr, + out_ptr, + scale_ptr, + M, + n, + pn: tl.constexpr, + ROUND: tl.constexpr, +): + pid = tl.program_id(axis=0) + + offs = pid * n * 2 + tl.arange(0, pn) + mask = tl.arange(0, pn) < n + + x1 = tl.load(x_ptr + offs, mask=mask).to(tl.float32) + x2 = tl.load(x_ptr + n + offs, mask=mask).to(tl.float32) + x = x1 * tl.sigmoid(x1) * x2 + + scale = tl.maximum(tl.max(x.abs(), 0) / 448, 1e-30) + if ROUND: + scale = tl.exp2(tl.ceil(tl.log2(scale))) + + tl.store(scale_ptr + pid, scale) + + xq = (x / scale).to(out_ptr.dtype.element_ty) + tl.store( + out_ptr + pid * n + tl.arange(0, pn), + xq, + mask=mask, + ) + + +def triton_silu_and_token_quant( + x, weight=None, out=None, scale=None, round_scale=False +): + """ + fused silu and tokenwise quantization + Args: + x: input tensor + round_scale: whether round scale to power of 2 + Returns: + - out: quantized tensor + - scale: quantization scale + """ + + assert x.is_contiguous() + if weight is not None: + assert weight.is_contiguous() + + M, N = x.shape + n = N // 2 + pn = triton.next_power_of_2(n) + device = x.device + if out is None: + out = torch.empty((M, n), device=device, dtype=torch.float8_e4m3fn) + else: + assert out.is_contiguous() + + if scale is None: + scale = torch.empty((M, 1), device=device, dtype=torch.float32) + else: + assert scale.is_contiguous() + + grid = (M,) + silu_and_token_quant_kernel[grid]( + x, + out, + scale, + M, + n, + pn, + round_scale, + num_stages=2, + num_warps=2, + ) + + return out, scale diff --git a/linghe/infer/topk.py b/linghe/infer/topk.py new file mode 100644 index 0000000..d5f974a --- /dev/null +++ b/linghe/infer/topk.py @@ -0,0 +1,121 @@ +# -*- coding: utf-8 -*- +""" +Copyright (c) Ant Financial Service Group and its affiliates. +""" + +from typing import Optional +import torch +import triton +import triton.language as tl + + +@triton.jit +def group_topk_score_kernel( + input_ptr, + bias_ptr, + topk_weight_ptr, + topk_ids_ptr, + scale, + eps, + N: tl.constexpr, + K: tl.constexpr, + G: tl.constexpr, + GK: tl.constexpr, + SHARE_EXPERTS: tl.constexpr, +): + pid = tl.program_id(axis=0) + GS: tl.constexpr = N // G + k: tl.constexpr = K // GK + + logit = tl.load(input_ptr + pid * N + tl.arange(0, N)).to(tl.float32) + x = tl.sigmoid(logit) + b = tl.load(bias_ptr + tl.arange(0, N)).to(tl.float32) + m = tl.reshape(x + b, (G, GS)) + + gt = tl.topk(m, k, dim=1) + + gts = tl.sum(gt, 1) + gtst = tl.topk(gts, GK, dim=0) + sum_min_value = tl.min(gtst) + + group_filling = tl.where((gts[:, None] >= sum_min_value), m, -1.0) + group_filling = tl.reshape(group_filling, [N]) + t = tl.min(tl.topk(group_filling, K, dim=0)) + mask = group_filling >= t + + hit_count = tl.sum(tl.where(mask, 1, 0)) + + if hit_count > K: + group_fillings = ( + group_filling.to(tl.float64) - tl.arange(0, N).to(tl.float64) * 1e-12 + ) + ts = tl.min(tl.topk(group_fillings, K, dim=0)) + masks = group_fillings >= ts + else: + masks = mask + + hitmap = tl.where(masks, 1, 0) + filling = tl.where(masks, x, 0.0) + score = filling / (tl.sum(filling) + eps) + acc = tl.cumsum(hitmap, 0) - 1 + + if SHARE_EXPERTS == 1: + tl.store(topk_weight_ptr + pid * (K + 1) + acc, score, mask=masks) + tl.store(topk_ids_ptr + pid * (K + 1) + acc, tl.arange(0, N), mask=masks) + tl.store(topk_weight_ptr + pid * (K + 1) + K, 1.0 / scale) + tl.store(topk_ids_ptr + pid * (K + 1) + K, N) + else: + tl.store(topk_weight_ptr + pid * K + acc, score, mask=masks) + tl.store(topk_ids_ptr + pid * K + acc, tl.arange(0, N), mask=masks) + + +def triton_group_topk_score( + x: torch.Tensor, + k: int, + expert_bias: torch.Tensor, + num_groups=8, + group_topk=4, + scaling_factor=1.0, + score_function="sigmoid", + num_shared_experts=0, + eps=1e-20, +): + """ + calculate topk. + Args: + x: input tensor. + expert_bias: expert bias + k: topk + Returns: + topk_weights: + topk_ids: + """ + device = x.device + M, N = x.shape + assert x.is_contiguous() and score_function == "sigmoid" + assert num_shared_experts in (0, 1) + topk_weights = torch.empty( + (M, k + num_shared_experts), device=device, dtype=torch.float32 + ) + topk_ids = torch.empty( + (M, k + num_shared_experts), device=device, dtype=torch.int32 + ) + + grid = (M,) + group_topk_score_kernel[grid]( + x, + expert_bias, + topk_weights, + topk_ids, + scaling_factor, + eps, + N, + k, + num_groups, + group_topk, + num_shared_experts, + num_stages=1, + num_warps=1, + ) + + return topk_weights, topk_ids diff --git a/linghe/quant/block.py b/linghe/quant/block.py index e7d093f..594126d 100644 --- a/linghe/quant/block.py +++ b/linghe/quant/block.py @@ -9,8 +9,9 @@ @triton.jit -def block_quant_kernel(x_ptr, y_ptr, s_ptr, M, N, BLOCK_SIZE: tl.constexpr, - ROUND: tl.constexpr): +def block_quant_kernel( + x_ptr, y_ptr, s_ptr, M, N, BLOCK_SIZE: tl.constexpr, ROUND: tl.constexpr +): pid_m = tl.program_id(axis=0) pid_n = tl.program_id(axis=1) n = tl.cdiv(N, BLOCK_SIZE) @@ -23,16 +24,13 @@ def block_quant_kernel(x_ptr, y_ptr, s_ptr, M, N, BLOCK_SIZE: tl.constexpr, if ROUND: s = tl.exp2(tl.ceil(tl.log2(s))) y = x / s - y = y.to(y_ptr.dtype.element_ty) tl.store(y_ptr + offs, y, mask=mask) tl.store(s_ptr + pid_m * n + pid_n, s) -def triton_block_quant(x, - block_size=128, - round_scale=False): +def triton_block_quant(x, out=None, scale=None, block_size=128, round_scale=False): """ - blockwise quantize x + blockwise quantize x, used for blockwise recipe for weight in megatron Args: x: input tensor block_size: block wise @@ -42,18 +40,262 @@ def triton_block_quant(x, - y: quantized tensor, float8_e4m3fn - s: quantization scale, float32 """ + assert x.is_contiguous() M, N = x.size() - y = torch.empty((M, N), dtype=torch.float8_e4m3fn, device=x.device) - s = x.new_empty(x.size(-2) // block_size, x.size(-1) // block_size, - dtype=torch.float32) + if out is None: + out = torch.empty((M, N), dtype=torch.float8_e4m3fn, device=x.device) + if scale is None: + scale = torch.empty( + triton.cdiv(M, block_size), + triton.cdiv(N, block_size), + dtype=torch.float32, + device=x.device, + ) grid = (triton.cdiv(M, block_size), triton.cdiv(N, block_size)) - block_quant_kernel[grid](x, - y, - s, - M, - N, - BLOCK_SIZE=block_size, - ROUND=round_scale, - num_stages=6, - num_warps=8) - return y, s + block_quant_kernel[grid]( + x, + out, + scale, + M, + N, + BLOCK_SIZE=block_size, + ROUND=round_scale, + num_stages=3, + num_warps=4, + ) + return out, scale + + +@triton.jit +def blockwise_quant_kernel( + x_ptr, + x_q_ptr, + x_scale_ptr, + xt_q_ptr, + xt_scale_ptr, + M, + N: tl.constexpr, + ROUND: tl.constexpr, + OUTPUT_MODE: tl.constexpr, +): + rid = tl.program_id(axis=0) + cid = tl.program_id(axis=1) + + offs = ( + rid * 128 * N + + cid * 128 + + tl.arange(0, 128)[:, None] * N + + tl.arange(0, 128)[None, :] + ) + indices = rid * 128 + tl.arange(0, 128) + mask = indices[:, None] < M + + x = tl.load(x_ptr + offs, mask=mask).to(tl.float32) + + if OUTPUT_MODE % 2 == 0: + scale = tl.maximum(tl.max(x.abs(), 1) / 448, 1e-30) + if ROUND: + scale = tl.exp2(tl.ceil(tl.log2(scale))) + + tl.store( + x_scale_ptr + rid * 128 + cid * M + tl.arange(0, 128), + scale, + mask=indices < M, + ) + xq = (x / scale[:, None]).to(x_q_ptr.dtype.element_ty) + tl.store( + x_q_ptr + + rid * 128 * N + + cid * 128 + + tl.arange(0, 128)[:, None] * N + + tl.arange(0, 128)[None, :], + xq, + mask=mask, + ) + + if OUTPUT_MODE > 0: + scale = tl.maximum(tl.max(x.abs(), 0) / 448, 1e-30) + if ROUND: + scale = tl.exp2(tl.ceil(tl.log2(scale))) + tl.store(xt_scale_ptr + rid * N + cid * 128 + tl.arange(0, 128), scale) + xq = (x / scale).to(xt_q_ptr.dtype.element_ty) + tl.store( + xt_q_ptr + + rid * 128 + + cid * 128 * M + + tl.arange(0, 128)[:, None] * M + + tl.arange(0, 128)[None, :], + tl.trans(xq), + mask=indices[None, :] < M, + ) + + +def triton_blockwise_quant(x, round_scale=False, output_mode=2): + """ + blockwise quantization, used in blockwise recipt in megatron + Args: + x: input tensor + round_scale: whether round scale to power of 2 + output_mode: one of {0, 1, 2} + 0: only output non-transposed quantized tensor + 1: only output transposed quantized tensor + 2: output both + + Returns: + x_q: + x_scale: + xt_q: + xt_scale: + """ + M, N = x.shape + assert M % 16 == 0 and x.is_contiguous() + device = x.device + x_q = torch.empty((M, N), device=device, dtype=torch.float8_e4m3fn) + x_scale = torch.empty((N // 128, M), device=device, dtype=torch.float32) + + xt_q = torch.empty((N, M), device=device, dtype=torch.float8_e4m3fn) + xt_scale = torch.empty((triton.cdiv(M, 128), N), device=device, dtype=torch.float32) + + grid = (triton.cdiv(M, 128), N // 128) + blockwise_quant_kernel[grid]( + x, + x_q, + x_scale, + xt_q, + xt_scale, + M, + N, + round_scale, + output_mode, + num_stages=2, + num_warps=2, + ) + + return x_q, x_scale, xt_q, xt_scale + + +@triton.jit +def batch_blockwise_quant_kernel( + x_ptr, + count_ptr, + xq_ptr, + xs_ptr, + xtq_ptr, + xts_ptr, + N: tl.constexpr, + E: tl.constexpr, + ROUND: tl.constexpr, +): + eid = tl.program_id(axis=0) + rid = tl.program_id(axis=1) + cid = tl.program_id(axis=2) + + count = tl.load(count_ptr + eid) + counts = tl.load(count_ptr + tl.arange(0, E)) + + if rid >= tl.cdiv(count, 128): + return + + nb = N // 128 + + m_block = tl.sum(tl.where(tl.arange(0, E) < eid, tl.cdiv(counts, 128), 0)) + si = tl.sum(tl.where(tl.arange(0, E) < eid, counts, 0)) + + rids = rid * 128 + tl.arange(0, 128) + + x = tl.load( + x_ptr + + si * N + + rid * 128 * N + + cid * 128 + + tl.arange(0, 128)[:, None] * N + + tl.arange(0, 128)[None, :], + mask=rids[:, None] < count, + ).to(tl.float32) + + scale = tl.maximum(tl.max(tl.abs(x), 1) / 448.0, 1e-30) + if ROUND: + scale = tl.exp2(tl.ceil(tl.log2(scale))) + xq = x / scale[:, None] + + tl.store( + xs_ptr + si * nb + cid * count + rid * 128 + tl.arange(0, 128), + scale, + mask=rids < count, + ) + tl.store( + xq_ptr + + si * N + + rid * 128 * N + + cid * 128 + + tl.arange(0, 128)[:, None] * N + + tl.arange(0, 128)[None, :], + xq, + mask=rids[:, None] < count, + ) + + scale = tl.maximum(tl.max(tl.abs(x), 0) / 448.0, 1e-30) + if ROUND: + scale = tl.exp2(tl.ceil(tl.log2(scale))) + xq = x / scale[None, :] + + tl.store(xts_ptr + m_block * N + rid * N + cid * 128 + tl.arange(0, 128), scale) + tl.store( + xtq_ptr + + si * N + + rid * 128 + + cid * 128 * count + + tl.arange(0, 128)[:, None] * count + + tl.arange(0, 128)[None, :], + tl.trans(xq), + mask=rids[None, :] < count, + ) + + +def triton_batch_blockwise_quant(xs, token_count_per_expert, splits, round_scale=False): + """ + select and quant, used in megatron 0.12 flex moe + Args: + xs: [bs, dim] + token_count_per_expert: [n_experts] + splits: python int list of token_count_per_expert + round_scale: whether round scale to power of 2 + + Returns: + - x_q: + - x_scale: + - xt_q: + - xt_scale: + + """ + assert xs.is_contiguous() + M, N = xs.shape + n_experts = token_count_per_expert.size(0) + device = xs.device + x_q = torch.empty((M, N), device=device, dtype=torch.float8_e4m3fn) + # intra layout and inner layput are not consist, + # tensors will be viewed after splitting + x_scale = torch.empty((M * N // 128,), device=device, dtype=torch.float32) + xt_q = torch.empty((M * N,), device=device, dtype=torch.float8_e4m3fn) + blocks = sum([(x + 127) // 128 for x in splits]) + xt_scale = torch.empty((blocks * N,), device=device, dtype=torch.float32) + + if M == 0: + return x_q, x_scale, xt_q, xt_scale + + grid = (n_experts, triton.cdiv(max(splits), 128), N // 128) + batch_blockwise_quant_kernel[grid]( + xs, + token_count_per_expert, + x_q, + x_scale, + xt_q, + xt_scale, + N, + n_experts, + round_scale, + num_stages=2, + num_warps=4, + ) + + return x_q, x_scale, xt_q, xt_scale diff --git a/linghe/quant/channel.py b/linghe/quant/channel.py index 462b6cb..7394370 100644 --- a/linghe/quant/channel.py +++ b/linghe/quant/channel.py @@ -4,14 +4,16 @@ """ from typing import Optional + import torch import triton import triton.language as tl @triton.jit -def row_quant_kernel(x_ptr, q_ptr, s_ptr, M, N, BLOCK_SIZE: tl.constexpr, - ROUND: tl.constexpr): +def row_quant_kernel( + x_ptr, q_ptr, s_ptr, M, N, BLOCK_SIZE: tl.constexpr, ROUND: tl.constexpr +): pid = tl.program_id(0) n_block = tl.cdiv(N, BLOCK_SIZE) indices = tl.arange(0, BLOCK_SIZE) @@ -20,8 +22,9 @@ def row_quant_kernel(x_ptr, q_ptr, s_ptr, M, N, BLOCK_SIZE: tl.constexpr, for j in range(n_block): offs = pid * N + j * BLOCK_SIZE + indices - x = tl.load(x_ptr + offs, mask=j * BLOCK_SIZE + indices < N, - other=0).to(tl.float32) + x = tl.load(x_ptr + offs, mask=j * BLOCK_SIZE + indices < N, other=0).to( + tl.float32 + ) max_val = tl.maximum(tl.max(tl.abs(x)), max_val) scale = tl.maximum(max_val / 448.0, 1e-30) if ROUND: @@ -53,20 +56,15 @@ def triton_row_quant(x, round_scale=False): x_scale = torch.empty((M,), dtype=torch.float32, device=x.device) grid = (M,) row_quant_kernel[grid]( - x, x_q, x_scale, - M, N, - BLOCK_SIZE, - round_scale, - num_stages=5, - num_warps=4 + x, x_q, x_scale, M, N, BLOCK_SIZE, round_scale, num_stages=5, num_warps=4 ) return x_q, x_scale @triton.jit -def deprecated_tokenwise_row_quant_kernel(x_ptr, out_ptr, scale_ptr, M, - T: tl.constexpr, N: tl.constexpr, - ROUND: tl.constexpr): +def deprecated_tokenwise_row_quant_kernel( + x_ptr, out_ptr, scale_ptr, M, T: tl.constexpr, N: tl.constexpr, ROUND: tl.constexpr +): pid = tl.program_id(axis=0) offs = pid * T * N + tl.arange(0, N) for i in range(T): @@ -84,10 +82,12 @@ def deprecated_tokenwise_row_quant_kernel(x_ptr, out_ptr, scale_ptr, M, offs += N -def triton_deprecated_tokenwise_row_quant(x: torch.Tensor, - out: Optional[torch.Tensor] = None, - scale: Optional[torch.Tensor] = None, - round_scale: bool = False): +def triton_deprecated_tokenwise_row_quant( + x: torch.Tensor, + out: Optional[torch.Tensor] = None, + scale: Optional[torch.Tensor] = None, + round_scale: bool = False, +): M, N = x.shape device = x.device if out is None: @@ -98,20 +98,15 @@ def triton_deprecated_tokenwise_row_quant(x: torch.Tensor, T = triton.cdiv(M, sm) grid = (sm,) deprecated_tokenwise_row_quant_kernel[grid]( - x, - out, - scale, - M, T, N, - round_scale, - num_stages=3, - num_warps=16 + x, out, scale, M, T, N, round_scale, num_stages=3, num_warps=16 ) return out, scale @triton.jit -def tokenwise_row_quant_kernel(x_ptr, out_ptr, scale_ptr, N: tl.constexpr, - ROUND: tl.constexpr): +def tokenwise_row_quant_kernel( + x_ptr, out_ptr, scale_ptr, N: tl.constexpr, ROUND: tl.constexpr +): pid = tl.program_id(axis=0) x = tl.load(x_ptr + pid * N + tl.arange(0, N)).to(tl.float32) x_max = tl.max(tl.abs(x)) @@ -144,13 +139,7 @@ def triton_tokenwise_row_quant(x, out=None, scale=None, round_scale=False): scale = torch.empty((M,), dtype=torch.float32, device=device) grid = (M,) tokenwise_row_quant_kernel[grid]( - x, - out, - scale, - N, - round_scale, - num_stages=3, - num_warps=16 + x, out, scale, N, round_scale, num_stages=3, num_warps=16 ) return out, scale @@ -159,8 +148,9 @@ def triton_tokenwise_row_quant(x, out=None, scale=None, round_scale=False): # dx = y @ wT # dwT = yT @ x @triton.jit -def transpose_row_quant_kernel(x_ptr, q_ptr, s_ptr, M, N, H: tl.constexpr, - W: tl.constexpr, ROUND: tl.constexpr): +def transpose_row_quant_kernel( + x_ptr, q_ptr, s_ptr, M, N, H: tl.constexpr, W: tl.constexpr, ROUND: tl.constexpr +): pid = tl.program_id(axis=0) # col-wise read, row-wise write # read block: [BLOCK_SIZE, B] @@ -181,8 +171,7 @@ def transpose_row_quant_kernel(x_ptr, q_ptr, s_ptr, M, N, H: tl.constexpr, tl.store(s_ptr + pid * W + tl.arange(0, W), scale) offs = pid * W + tl.arange(0, H)[:, None] * N + tl.arange(0, W)[None, :] - toffs = pid * W * M + tl.arange(0, W)[:, None] * M + tl.arange(0, H)[None, - :] + toffs = pid * W * M + tl.arange(0, W)[:, None] * M + tl.arange(0, H)[None, :] for i in range(m): x = tl.trans(tl.load(x_ptr + offs, mask=i * H + indices[:, None] < M)) x = (x * s).to(q_ptr.dtype.element_ty) @@ -210,12 +199,7 @@ def triton_transpose_row_quant(x, round_scale=False): x_scale = torch.empty((N, 1), dtype=torch.float32, device=x.device) grid = (N // W,) transpose_row_quant_kernel[grid]( - x, x_q, x_scale, - M, N, - H, W, - round_scale, - num_stages=6, - num_warps=4 + x, x_q, x_scale, M, N, H, W, round_scale, num_stages=6, num_warps=4 ) return x_q, x_scale @@ -240,33 +224,38 @@ def triton_channel_quant_tn(y, x): def channel_quant_forward(x, w): x_q, x_scale, w_q, w_scale = triton_channel_quant_nt(x, w) - output = torch._scaled_mm(x_q, - w_q.t(), - scale_a=x_scale, - scale_b=w_scale.view(1, -1), - out_dtype=torch.bfloat16, - use_fast_accum=True) + output = torch._scaled_mm( + x_q, + w_q.t(), + scale_a=x_scale, + scale_b=w_scale.view(1, -1), + out_dtype=torch.bfloat16, + use_fast_accum=True, + ) return output, x_q, w_q, x_scale, w_scale def channel_quant_backward(y, w): y_q, y_scale, w_q, w_scale = triton_channel_quant_nn(y, w) - output = torch._scaled_mm(y_q, - w_q.t(), - scale_a=y_scale, - scale_b=w_scale.view(1, -1), - out_dtype=torch.bfloat16, - use_fast_accum=True) + output = torch._scaled_mm( + y_q, + w_q.t(), + scale_a=y_scale, + scale_b=w_scale.view(1, -1), + out_dtype=torch.bfloat16, + use_fast_accum=True, + ) return output, y_q, w_q, y_scale, w_scale def channel_quant_update(y, x): y_q, y_scale, x_q, x_scale = triton_channel_quant_tn(y, x) - output = torch._scaled_mm(y_q, - x_q.t(), - scale_a=y_scale, - scale_b=x_scale.view(1, -1), - out_dtype=torch.bfloat16, - use_fast_accum=True) + output = torch._scaled_mm( + y_q, + x_q.t(), + scale_a=y_scale, + scale_b=x_scale.view(1, -1), + out_dtype=torch.bfloat16, + use_fast_accum=True, + ) return output, y_q, x_q, y_scale, x_scale - diff --git a/linghe/quant/group.py b/linghe/quant/group.py index b46ce47..1c4d9ce 100644 --- a/linghe/quant/group.py +++ b/linghe/quant/group.py @@ -9,13 +9,20 @@ @triton.jit -def group_quant_kernel(x_ptr, y_ptr, s_ptr, N, BLOCK_SIZE: tl.constexpr, - K: tl.constexpr, ROUND: tl.constexpr): +def group_quant_kernel( + x_ptr, + y_ptr, + s_ptr, + N, + BLOCK_SIZE: tl.constexpr, + K: tl.constexpr, + ROUND: tl.constexpr, +): pid = tl.program_id(axis=0) offs = pid * N + tl.arange(0, K * BLOCK_SIZE) n = tl.cdiv(N, K * BLOCK_SIZE) - soffs = pid * n * K + tl.arange(0, K) - for i in range(n): + soffs = pid * (N // BLOCK_SIZE) + tl.arange(0, K) + for i in tl.range(n): x = tl.load(x_ptr + offs).to(tl.float32) x = tl.reshape(x, (K, BLOCK_SIZE), can_reorder=False) s = tl.maximum(tl.max(tl.abs(x), 1) / 448.0, 1e-30) @@ -30,10 +37,7 @@ def group_quant_kernel(x_ptr, y_ptr, s_ptr, N, BLOCK_SIZE: tl.constexpr, soffs += K -def triton_group_quant(x, - dtype=torch.float8_e4m3fn, - group_size=128, - round_scale=False): +def triton_group_quant(x, dtype=torch.float8_e4m3fn, group_size=128, round_scale=False): """ groupwise quantize x, group is in under rowwise format Args: @@ -46,21 +50,14 @@ def triton_group_quant(x, - s: quantization scale, float32 """ M, N = x.shape - K = 16 - assert N % group_size == 0 and N % (group_size * K) == 0 + K = 16 if N > 1024 else 8 + assert N % group_size == 0 assert x.is_contiguous() y = torch.empty((M, N), device=x.device, dtype=dtype) s = torch.empty(M, N // group_size, device=x.device, dtype=torch.float32) - grid = (M,) # noqa - group_quant_kernel[grid](x, - y, - s, - N, - group_size, - K, - round_scale, - num_stages=5, - num_warps=4) + grid = (M,) + group_quant_kernel[grid]( + x, y, s, N, group_size, K, round_scale, num_stages=5, num_warps=4 + ) return y, s - diff --git a/linghe/quant/hadamard.py b/linghe/quant/hadamard.py index 0c59f96..b2f7d0b 100644 --- a/linghe/quant/hadamard.py +++ b/linghe/quant/hadamard.py @@ -10,14 +10,14 @@ @triton.jit def hadamard_quant_row_kernel( - x_ptr, - hm_ptr, - x_q_ptr, - x_scale_ptr, - M, - N, - BLOCK_SIZE: tl.constexpr, - R: tl.constexpr, + x_ptr, + hm_ptr, + x_q_ptr, + x_scale_ptr, + M, + N, + BLOCK_SIZE: tl.constexpr, + R: tl.constexpr, ): pid = tl.program_id(0) row_start = pid * R * BLOCK_SIZE @@ -25,9 +25,10 @@ def hadamard_quant_row_kernel( mask_rows = rows < M hm = tl.load( - hm_ptr + tl.arange(0, BLOCK_SIZE)[:, None] * BLOCK_SIZE + tl.arange(0, - BLOCK_SIZE)[ - None, :]) + hm_ptr + + tl.arange(0, BLOCK_SIZE)[:, None] * BLOCK_SIZE + + tl.arange(0, BLOCK_SIZE)[None, :] + ) max_val = tl.zeros((R * BLOCK_SIZE,), dtype=tl.float32) + 1.17e-38 @@ -38,8 +39,9 @@ def hadamard_quant_row_kernel( mask_cols = cols < N offs = rows[:, None] * N + cols[None, :] - x = tl.load(x_ptr + offs, mask=mask_rows[:, None] & mask_cols[None, :], - other=0.0) + x = tl.load( + x_ptr + offs, mask=mask_rows[:, None] & mask_cols[None, :], other=0.0 + ) x_transformed = tl.dot(x, hm) current_max = tl.max(tl.abs(x_transformed), axis=1) max_val = tl.maximum(max_val, current_max) @@ -54,24 +56,26 @@ def hadamard_quant_row_kernel( mask_cols = cols < N offs = rows[:, None] * N + cols[None, :] - x = tl.load(x_ptr + offs, mask=mask_rows[:, None] & mask_cols[None, :], - other=0.0) + x = tl.load( + x_ptr + offs, mask=mask_rows[:, None] & mask_cols[None, :], other=0.0 + ) x_transformed = tl.dot(x, hm) quantized = (x_transformed * s[:, None]).to(x_q_ptr.dtype.element_ty) - tl.store(x_q_ptr + offs, quantized, - mask=mask_rows[:, None] & mask_cols[None, :]) + tl.store( + x_q_ptr + offs, quantized, mask=mask_rows[:, None] & mask_cols[None, :] + ) @triton.jit def hadamard_quant_col_kernel( - x_ptr, - hm_ptr, - xt_q_ptr, - xt_scale_ptr, - M, - N, - BLOCK_SIZE: tl.constexpr, - R: tl.constexpr, + x_ptr, + hm_ptr, + xt_q_ptr, + xt_scale_ptr, + M, + N, + BLOCK_SIZE: tl.constexpr, + R: tl.constexpr, ): pid = tl.program_id(0) col_start = pid * R * BLOCK_SIZE @@ -79,9 +83,10 @@ def hadamard_quant_col_kernel( mask_cols = cols < N hm = tl.load( - hm_ptr + tl.arange(0, BLOCK_SIZE)[:, None] * BLOCK_SIZE + tl.arange(0, - BLOCK_SIZE)[ - None, :]) + hm_ptr + + tl.arange(0, BLOCK_SIZE)[:, None] * BLOCK_SIZE + + tl.arange(0, BLOCK_SIZE)[None, :] + ) max_val = tl.zeros((R * BLOCK_SIZE,), dtype=tl.float32) + 1.17e-38 @@ -92,8 +97,9 @@ def hadamard_quant_col_kernel( mask_rows = rows < M offs = rows[:, None] * N + cols[None, :] - x = tl.load(x_ptr + offs, mask=mask_rows[:, None] & mask_cols[None, :], - other=0.0) + x = tl.load( + x_ptr + offs, mask=mask_rows[:, None] & mask_cols[None, :], other=0.0 + ) x_transformed = tl.dot(hm, x) current_max = tl.max(tl.abs(x_transformed), axis=0) max_val = tl.maximum(max_val, current_max) @@ -108,14 +114,18 @@ def hadamard_quant_col_kernel( mask_rows = rows < M offs = rows[:, None] * N + cols[None, :] - x = tl.load(x_ptr + offs, mask=mask_rows[:, None] & mask_cols[None, :], - other=0.0) + x = tl.load( + x_ptr + offs, mask=mask_rows[:, None] & mask_cols[None, :], other=0.0 + ) x_transformed = tl.dot(hm, x) quantized = (x_transformed * s[None, :]).to(xt_q_ptr.dtype.element_ty) quantized_t = tl.trans(quantized) store_offs = cols[:, None] * M + rows[None, :] - tl.store(xt_q_ptr + store_offs, quantized_t, - mask=mask_cols[:, None] & mask_rows[None, :]) + tl.store( + xt_q_ptr + store_offs, + quantized_t, + mask=mask_cols[:, None] & mask_rows[None, :], + ) def triton_hadamard_quant(x, hm): @@ -136,35 +146,17 @@ def triton_hadamard_quant(x, hm): R = 1 x_q = torch.empty((M, N), dtype=torch.float8_e4m3fn, device=device) xt_q = torch.empty((N, M), dtype=torch.float8_e4m3fn, device=device) - x_scale = torch.empty((M, ), dtype=torch.float32, device=device) - xt_scale = torch.empty((N, ), dtype=torch.float32, device=device) + x_scale = torch.empty((M,), dtype=torch.float32, device=device) + xt_scale = torch.empty((N,), dtype=torch.float32, device=device) grid_row = (triton.cdiv(M, R * BLOCK_SIZE),) hadamard_quant_row_kernel[grid_row]( - x, - hm, - x_q, - x_scale, - M, - N, - BLOCK_SIZE, - R, - num_stages=6, - num_warps=4 + x, hm, x_q, x_scale, M, N, BLOCK_SIZE, R, num_stages=6, num_warps=4 ) grid_col = (triton.cdiv(N, R * BLOCK_SIZE),) hadamard_quant_col_kernel[grid_col]( - x, - hm, - xt_q, - xt_scale, - M, - N, - BLOCK_SIZE, - R, - num_stages=6, - num_warps=4 + x, hm, xt_q, xt_scale, M, N, BLOCK_SIZE, R, num_stages=6, num_warps=4 ) - return x_q, x_scale,xt_q, xt_scale + return x_q, x_scale, xt_q, xt_scale diff --git a/linghe/quant/mxfp8.py b/linghe/quant/mxfp8.py new file mode 100644 index 0000000..1a56e87 --- /dev/null +++ b/linghe/quant/mxfp8.py @@ -0,0 +1,233 @@ +# -*- coding: utf-8 -*- +""" +Copyright (c) Ant Financial Service Group and its affiliates. +""" + +import torch +import triton +import triton.language as tl + + +@triton.jit +def mxfp8_quant_kernel( + x_ptr, + x_q_ptr, + x_s_ptr, + xt_q_ptr, + xt_s_ptr, + m, + N: tl.constexpr, + B: tl.constexpr, + OUTPUT_MODE: tl.constexpr, +): + rid = tl.program_id(axis=0) + cid = tl.program_id(axis=1) + offs = ( + rid * 32 * N + + cid * B + + tl.arange(0, 32)[:, None] * N + + tl.arange(0, B)[None, :] + ) + indices = rid * 32 + tl.arange(0, 32) + mask = indices[:, None] < m + b = N // 32 + sb: tl.constexpr = B // 32 + + x = tl.load(x_ptr + offs, mask=mask).to(tl.float32) + + if OUTPUT_MODE % 2 == 0: + xr = tl.reshape(x, [32, sb, 32]) + scale = tl.maximum(tl.max(xr.abs(), 2) / 448, 1e-30) + log_scale = tl.ceil(tl.log2(scale)) + scale = tl.exp2(log_scale) + tl.store( + x_s_ptr + + rid * 32 * b + + cid * B // 32 + + tl.arange(0, 32)[:, None] * b + + tl.arange(0, sb), + log_scale + 127, + ) + xq = tl.reshape(xr / scale[:, :, None], (32, B)).to(x_q_ptr.dtype.element_ty) + tl.store( + x_q_ptr + + rid * 32 * N + + cid * B + + tl.arange(0, 32)[:, None] * N + + tl.arange(0, B)[None, :], + xq, + mask=mask, + ) + + if OUTPUT_MODE > 0: + scale = tl.maximum(tl.max(x.abs(), 0) / 448, 1e-30) + log_scale = tl.ceil(tl.log2(scale)) + scale = tl.exp2(log_scale) + tl.store(xt_s_ptr + rid * N + cid * B + tl.arange(0, B), log_scale + 127) + xq = (x / scale).to(xt_q_ptr.dtype.element_ty) + tl.store( + xt_q_ptr + + rid * 32 * N + + cid * B + + tl.arange(0, 32)[:, None] * N + + tl.arange(0, B)[None, :], + xq, + mask=mask, + ) + + +def triton_mxfp8_quant(x, output_mode=2): + """ + fused silu and mxfp8 quantization, used in shared expert + Args: + x: input tensor + output_mode: one of {0, 1, 2} + 0: only output non-transposed quantized tensor + 1: only output transposed quantized tensor + 2: output both + + Returns: + - x_q: quantized tensor + - x_scale: quantization scale + - xt_q: quantized tensor of transposed output + - xt_scale: quantization scale of transposed output + """ + m, N = x.shape + M = (m + 127) // 128 * 128 + assert N % 128 == 0 # transposed scaled should be multiplier of 128 + assert x.is_contiguous() + device = x.device + x_q = torch.empty((m, N), device=device, dtype=torch.float8_e4m3fn) + x_scale = torch.empty((M, N // 32), device=device, dtype=torch.uint8) + + xt_q = torch.empty((m, N), device=device, dtype=torch.float8_e4m3fn) + xt_scale = torch.empty((M // 32, N), device=device, dtype=torch.uint8) + B = 128 + grid = (M // 32, N // B) + mxfp8_quant_kernel[grid]( + x, x_q, x_scale, xt_q, xt_scale, m, N, B, output_mode, num_stages=3, num_warps=2 + ) + + return x_q, x_scale, xt_q, xt_scale + + +@triton.jit +def batch_mxfp8_quant_kernel( + x_ptr, + count_ptr, + xq_ptr, + xs_ptr, + xtq_ptr, + xts_ptr, + N: tl.constexpr, + B: tl.constexpr, + E: tl.constexpr, + OUTPUT_MODE: tl.constexpr, +): + eid = tl.program_id(axis=0) + rid = tl.program_id(axis=1) + cid = tl.program_id(axis=2) + + count = tl.load(count_ptr + eid) + counts = tl.load(count_ptr + tl.arange(0, E)) + + if rid >= tl.cdiv(count, 128) * 4: + return + + m_block = tl.sum(tl.where(tl.arange(0, E) < eid, tl.cdiv(counts, 128), 0)) * 4 + si = tl.sum(tl.where(tl.arange(0, E) < eid, counts, 0)) + + offs = ( + si * N + + rid * 32 * N + + cid * B + + tl.arange(0, 32)[:, None] * N + + tl.arange(0, B)[None, :] + ) + indices = rid * 32 + tl.arange(0, 32) + mask = indices[:, None] < count + b = N // 32 + sb: tl.constexpr = B // 32 + + x = tl.load(x_ptr + offs, mask=mask).to(tl.float32) + + if OUTPUT_MODE % 2 == 0: + xr = tl.reshape(x, [32, sb, 32]) + scale = tl.maximum(tl.max(xr.abs(), 2) / 448, 1e-30) + log_scale = tl.ceil(tl.log2(scale)) + scale = tl.exp2(log_scale) + tl.store( + xs_ptr + + m_block * N + + rid * 32 * b + + cid * B // 32 + + tl.arange(0, 32)[:, None] * b + + tl.arange(0, sb), + log_scale + 127, + ) + xq = tl.reshape(xr / scale[:, :, None], (32, B)).to(xq_ptr.dtype.element_ty) + tl.store(xq_ptr + offs, xq, mask=mask) + + if OUTPUT_MODE > 0: + scale = tl.maximum(tl.max(x.abs(), 0) / 448, 1e-30) + log_scale = tl.ceil(tl.log2(scale)) + scale = tl.exp2(log_scale) + tl.store( + xts_ptr + m_block * N + rid * N + cid * B + tl.arange(0, B), log_scale + 127 + ) + xq = (x / scale).to(xtq_ptr.dtype.element_ty) + tl.store(xtq_ptr + offs, xq, mask=mask) + + +def triton_batch_mxfp8_quant(xs, token_count_per_expert, splits, output_mode=2): + """ + select and quant, used in megatron 0.12 flex moe + Args: + xs: [bs, dim] + token_count_per_expert: [n_experts] + splits: python int list of token_count_per_expert + output_mode: one of {0, 1, 2} + 0: only output non-transposed quantized tensor + 1: only output transposed quantized tensor + 2: output both + + Returns: + - x_q: + - x_scale: + - xt_q: + - xt_scale: + + """ + assert xs.is_contiguous() + m, N = xs.shape + assert N % 128 == 0 + n_experts = token_count_per_expert.size(0) + device = xs.device + M = sum([(x + 127) // 128 for x in splits]) * 128 + + x_q = torch.empty((m, N), device=device, dtype=torch.float8_e4m3fn) + x_scale = torch.empty((M, N // 32), device=device, dtype=torch.uint8) + xt_q = torch.empty((m, N), device=device, dtype=torch.float8_e4m3fn) + xt_scale = torch.empty((M // 32, N), device=device, dtype=torch.uint8) + + if m == 0: + return x_q, x_scale, xt_q, xt_scale + + B = 256 if N % 256 == 0 else 128 + grid = (n_experts, triton.cdiv(max(splits), 128) * 4, N // B) + batch_mxfp8_quant_kernel[grid]( + xs, + token_count_per_expert, + x_q, + x_scale, + xt_q, + xt_scale, + N, + B, + n_experts, + output_mode, + num_stages=3, + num_warps=4, + ) + + return x_q, x_scale, xt_q, xt_scale diff --git a/linghe/quant/smooth.py b/linghe/quant/smooth.py index 93ac054..092655c 100644 --- a/linghe/quant/smooth.py +++ b/linghe/quant/smooth.py @@ -7,87 +7,95 @@ import triton import triton.language as tl -from linghe.tools.util import round_up -from linghe.utils.transpose import triton_transpose_and_pad - @triton.jit -def tokenwise_smooth_quant_kernel(x_ptr, q_ptr, ss_ptr, qs_ptr, max_ptr, - M, T, - N: tl.constexpr, - W: tl.constexpr, - EVEN: tl.constexpr, - REVERSE: tl.constexpr, - ROUND: tl.constexpr, - CALIBRATE: tl.constexpr): +def tokenwise_smooth_quant_kernel( + x_ptr, + q_ptr, + ss_ptr, + qs_ptr, + M, + T, + N: tl.constexpr, + W: tl.constexpr, + EVEN: tl.constexpr, + REVERSE: tl.constexpr, + ROUND: tl.constexpr, +): pid = tl.program_id(axis=0) - # row-wise read, row-wise write smooth_scale = tl.load(ss_ptr + tl.arange(0, N))[None, :] if not REVERSE: smooth_scale = 1.0 / smooth_scale - if CALIBRATE: - output_maxs = tl.zeros((W, N), dtype=tl.float32) for i in range(T): indices = pid * W * T + i * W + tl.arange(0, W) if EVEN: - x = tl.load(x_ptr + pid * W * T * N + i * N * W + tl.arange(0, W)[:, - None] * N + tl.arange( - 0, N)[None, :]).to( - tl.float32) + x = tl.load( + x_ptr + + pid * W * T * N + + i * N * W + + tl.arange(0, W)[:, None] * N + + tl.arange(0, N)[None, :] + ).to(tl.float32) else: - x = tl.load(x_ptr + pid * W * T * N + i * N * W + tl.arange(0, W)[:, - None] * N + tl.arange( - 0, N)[None, :], - mask=indices[:, None] < M).to( - tl.float32) - if CALIBRATE: - output_maxs = tl.maximum(tl.abs(x), output_maxs) + x = tl.load( + x_ptr + + pid * W * T * N + + i * N * W + + tl.arange(0, W)[:, None] * N + + tl.arange(0, N)[None, :], + mask=indices[:, None] < M, + ).to(tl.float32) x *= smooth_scale x_max = tl.max(tl.abs(x), axis=1) scale = tl.maximum(x_max / 448.0, 1e-30) if ROUND: scale = tl.exp2(tl.ceil(tl.log2(scale))) if EVEN: - tl.store(qs_ptr + pid * W * T + i * W + tl.arange(0, W), scale, ) + tl.store(qs_ptr + pid * W * T + i * W + tl.arange(0, W), scale) else: - tl.store(qs_ptr + pid * W * T + i * W + tl.arange(0, W), scale, - mask=indices < M) + tl.store( + qs_ptr + pid * W * T + i * W + tl.arange(0, W), scale, mask=indices < M + ) x /= scale[:, None] xq = x.to(q_ptr.dtype.element_ty) if EVEN: - tl.store(q_ptr + pid * W * T * N + i * N * W + tl.arange(0, W)[:, - None] * N + tl.arange( - 0, - N)[ - None, :], - xq) + tl.store( + q_ptr + + pid * W * T * N + + i * N * W + + tl.arange(0, W)[:, None] * N + + tl.arange(0, N)[None, :], + xq, + ) else: - tl.store(q_ptr + pid * W * T * N + i * N * W + tl.arange(0, W)[:, - None] * N + tl.arange( - 0, - N)[ - None, :], - xq, - mask=indices[:, None] < M) - if CALIBRATE: - output_maxs = tl.max(output_maxs, 0) - tl.store(max_ptr + pid * N + tl.arange(0, N), output_maxs) + tl.store( + q_ptr + + pid * W * T * N + + i * N * W + + tl.arange(0, W)[:, None] * N + + tl.arange(0, N)[None, :], + xq, + mask=indices[:, None] < M, + ) @triton.jit -def blockwise_smooth_quant_kernel(x_ptr, q_ptr, ss_ptr, qs_ptr, max_ptr, - M, - N, - H: tl.constexpr, - W: tl.constexpr, - EVEN: tl.constexpr, - REVERSE: tl.constexpr, - ROUND: tl.constexpr, - CALIBRATE: tl.constexpr): +def blockwise_smooth_quant_kernel( + x_ptr, + q_ptr, + ss_ptr, + qs_ptr, + M, + N, + H: tl.constexpr, + W: tl.constexpr, + EVEN: tl.constexpr, + REVERSE: tl.constexpr, + ROUND: tl.constexpr, +): pid = tl.program_id(axis=0) - # row-wise read, row-wise write offs = pid * W * N + tl.arange(0, W)[:, None] * N + tl.arange(0, H)[None, :] soffs = tl.arange(0, H) x_max = tl.zeros((W,), dtype=tl.float32) @@ -97,12 +105,9 @@ def blockwise_smooth_quant_kernel(x_ptr, q_ptr, ss_ptr, qs_ptr, max_ptr, if EVEN: x = tl.load(x_ptr + offs).to(tl.float32) else: - x = tl.load(x_ptr + offs, - mask=pid * W + tl.arange(0, W)[:, None] < M).to( - tl.float32) - if CALIBRATE: - output_maxs = tl.max(x.abs(), 0) - tl.store(max_ptr + pid * N + i * H + tl.arange(0, H), output_maxs) + x = tl.load(x_ptr + offs, mask=pid * W + tl.arange(0, W)[:, None] < M).to( + tl.float32 + ) if REVERSE: x = x * smooth_scale else: @@ -115,8 +120,9 @@ def blockwise_smooth_quant_kernel(x_ptr, q_ptr, ss_ptr, qs_ptr, max_ptr, if ROUND: scale = tl.exp2(tl.ceil(tl.log2(scale))) - tl.store(qs_ptr + pid * W + tl.arange(0, W), scale, - mask=pid * W + tl.arange(0, W) < M) + tl.store( + qs_ptr + pid * W + tl.arange(0, W), scale, mask=pid * W + tl.arange(0, W) < M + ) s = (1.0 / scale)[:, None] @@ -127,30 +133,26 @@ def blockwise_smooth_quant_kernel(x_ptr, q_ptr, ss_ptr, qs_ptr, max_ptr, if EVEN: x = tl.load(x_ptr + offs) else: - x = tl.load(x_ptr + offs, - mask=pid * W + tl.arange(0, W)[:, None] < M) + x = tl.load(x_ptr + offs, mask=pid * W + tl.arange(0, W)[:, None] < M) if REVERSE: - xq = (x.to(tl.float32) * smooth_scale * s).to( - q_ptr.dtype.element_ty) + xq = (x.to(tl.float32) * smooth_scale * s).to(q_ptr.dtype.element_ty) else: - xq = (x.to(tl.float32) / smooth_scale * s).to( - q_ptr.dtype.element_ty) + xq = (x.to(tl.float32) / smooth_scale * s).to(q_ptr.dtype.element_ty) if EVEN: tl.store(q_ptr + offs, xq) else: # tl.store(q_ptr+offs, xq, mask=(i*H+tl.arange(0, H)[None,:] 8192 else 4 EVEN = M % W == 0 T = triton.cdiv(M, W) - if calibrate: - x_maxs = torch.empty((T, N), device=device, dtype=torch.bfloat16) - else: - x_maxs = None grid = (T,) blockwise_smooth_quant_kernel[grid]( x, x_q, smooth_scale, x_scale, - x_maxs, M, N, H, @@ -210,215 +198,157 @@ def triton_smooth_quant(x, smooth_scale, x_q=None, x_scale=None, EVEN, reverse, round_scale, - calibrate, num_stages=3, - num_warps=4 + num_warps=4, ) - if calibrate: - x_maxs = x_maxs.amax(0).float() - return x_q, x_scale, x_maxs + return x_q, x_scale @triton.jit -def subrow_smooth_quant_kernel(x_ptr, q_ptr, ss_ptr, qs_ptr, - subrow_scales_ptr, - tail_ri, - tail_si, - head_ri, - head_ei, - size, - N, - W: tl.constexpr, - TAIL: tl.constexpr, - HEAD: tl.constexpr, - REVERSE: tl.constexpr, - ROUND: tl.constexpr): - if TAIL: - # scale is saved as max/448 - scale = tl.maximum(tl.load(subrow_scales_ptr), 1e-30) - if ROUND: - scale = tl.exp2(tl.ceil(tl.log2(scale))) - # scale only stores in subrow with leading values - - T = tl.cdiv(N - tail_si, W) - for i in range(T): - mask = tail_si + i * W + tl.arange(0, W) < N - if REVERSE: - smooth_scale = tl.load( - ss_ptr + tail_si + i * W + tl.arange(0, W), mask=mask) - else: - smooth_scale = tl.load( - ss_ptr + tail_si + i * W + tl.arange(0, W), other=1e30, - mask=mask) - smooth_scale = 1.0 / smooth_scale - x = tl.load(x_ptr + i * W + tl.arange(0, W), mask=mask).to( - tl.float32) - x *= smooth_scale - x /= scale - xq = tl.minimum(tl.maximum(x, -448), 448) - tl.store(q_ptr + tail_ri * N + tail_si + i * W + tl.arange(0, W), - xq.to(q_ptr.dtype.element_ty), mask=mask) - - if HEAD: - # scale is saved as max/448 - scale = tl.maximum(tl.load(subrow_scales_ptr + 1), 1e-30) - if ROUND: - scale = tl.exp2(tl.ceil(tl.log2(scale))) - tl.store(qs_ptr + head_ri, scale) - - T = tl.cdiv(head_ei, W) - for i in range(T): - mask = i * W + tl.arange(0, W) < head_ei - if REVERSE: - smooth_scale = tl.load(ss_ptr + i * W + tl.arange(0, W), - mask=mask) - else: - smooth_scale = tl.load(ss_ptr + i * W + tl.arange(0, W), - other=1e30, mask=mask) - smooth_scale = 1.0 / smooth_scale - x = tl.load(x_ptr + size - head_ei + i * W + tl.arange(0, W), - mask=mask).to(tl.float32) - x *= smooth_scale - x /= scale - xq = tl.minimum(tl.maximum(x, -448), 448) - tl.store(q_ptr + head_ri * N + i * W + tl.arange(0, W), - xq.to(q_ptr.dtype.element_ty), mask=mask) - +def transpose_smooth_quant_kernel( + x_ptr, + q_ptr, + ss_ptr, + qs_ptr, + M, + N, + P, + H: tl.constexpr, + W: tl.constexpr, + EVEN: tl.constexpr, + REVERSE: tl.constexpr, + ROUND: tl.constexpr, +): + pid = tl.program_id(axis=0) + offs = pid * W + tl.arange(0, H)[:, None] * N + tl.arange(0, W)[None, :] + soffs = tl.arange(0, H) + x_max = tl.zeros((W,), dtype=tl.float32) + m = tl.cdiv(P, H) + for i in range(m): + if EVEN: + x = tl.load(x_ptr + offs) + smooth_scale = tl.load(ss_ptr + soffs)[:, None] + else: + mask = (i * H + tl.arange(0, H)[:, None] < M) & ( + pid * W + tl.arange(0, W)[None, :] < N + ) + x = tl.load(x_ptr + offs, mask=mask) + other = 0.0 if REVERSE else 1e30 + smooth_scale = tl.load(ss_ptr + soffs, mask=soffs < M, other=other)[:, None] + if REVERSE: + x = x * smooth_scale + else: + x = x / smooth_scale + x_max = tl.maximum(tl.max(tl.abs(x), axis=0), x_max) + offs += H * N + soffs += H -def triton_subrow_smooth_quant(x, smooth_scale, x_q, x_scale, - subrow_scales, offset, size, - reverse=False, round_scale=False): - """""" - M, N = x_q.shape - W = 128 - if offset % N == 0: - tail_ri = 0 - tail_si = 0 - TAIL = False - else: - tail_ri = offset // N - tail_si = offset % N - TAIL = True + scale = tl.maximum(x_max / 448.0, 1e-30) + if ROUND: + scale = tl.exp2(tl.ceil(tl.log2(scale))) - if (offset + size) % N == 0: - head_ri = 0 - head_ei = 0 # head_size = head_ei - HEAD = False + if EVEN: + tl.store(qs_ptr + pid * W + tl.arange(0, W), scale) else: - head_ri = (offset + size) // N - head_ei = (offset + size) % N - HEAD = True - - grid = (1,) - subrow_smooth_quant_kernel[grid]( - x, - x_q, - smooth_scale, - x_scale, - subrow_scales, - tail_ri, - tail_si, - head_ri, - head_ei, - size, - N, - W, - TAIL, - HEAD, - reverse, - round_scale, - num_stages=3, - num_warps=1 - ) - - -@triton.jit -def depracated_tokenwise_smooth_quant_kernel(x_ptr, q_ptr, ss_ptr, - qs_ptr, M, W, - N: tl.constexpr, - REVERSE: tl.constexpr, - ROUND: tl.constexpr): - pid = tl.program_id(axis=0) - # row-wise read, row-wise write - smooth_scale = tl.load(ss_ptr + tl.arange(0, N)) - if not REVERSE: - smooth_scale = 1.0 / smooth_scale - - for i in range(W): - x = tl.load(x_ptr + pid * W * N + i * N + tl.arange(0, N), - mask=pid * W + i < M).to(tl.float32) - x *= smooth_scale - x_max = tl.maximum(tl.max(tl.abs(x)), 1e-30) + tl.store( + qs_ptr + pid * W + tl.arange(0, W), + scale, + mask=pid * W + tl.arange(0, W) < N, + ) - scale = x_max / 448.0 - if ROUND: - scale = tl.exp2(tl.ceil(tl.log2(scale))) - tl.store(qs_ptr + pid * W + i, scale, mask=pid * W + i < M) + s = (1.0 / scale)[None, :] + offs = pid * W + tl.arange(0, H)[:, None] * N + tl.arange(0, W)[None, :] + soffs = tl.arange(0, H) + toffs = pid * W * P + tl.arange(0, W)[:, None] * P + tl.arange(0, H)[None, :] + for i in range(m): + if EVEN: + x = tl.load(x_ptr + offs).to(tl.float32) + smooth_scale = tl.load(ss_ptr + soffs)[:, None] + else: + x = tl.load(x_ptr + offs, mask=(i * H + tl.arange(0, H)[:, None] < M)).to( + tl.float32 + ) + other = 0.0 if REVERSE else 1e30 + smooth_scale = tl.load(ss_ptr + soffs, mask=soffs < M, other=other)[:, None] - x /= scale - xq = x.to(q_ptr.dtype.element_ty) - tl.store(q_ptr + pid * W * N + i * N + tl.arange(0, N), xq, - mask=pid * W + i < M) + if REVERSE: + x = (x * smooth_scale * s).to(q_ptr.dtype.element_ty) + else: + x = (x / smooth_scale * s).to(q_ptr.dtype.element_ty) + if EVEN: + tl.store(q_ptr + toffs, tl.trans(x)) + else: + # mask with P instead of M + tl.store( + q_ptr + toffs, tl.trans(x), mask=(i * H + tl.arange(0, H)[None, :] < P) + ) + offs += H * N + toffs += H + soffs += H -def triton_depracated_tokenwise_smooth_quant(x, smooth_scale, x_q=None, - x_scale=None, reverse=False, - round_scale=False): - """""" - # row-wise read, row-wise write +def triton_transpose_smooth_quant( + x, smooth_scale, reverse=False, pad=True, round_scale=False +): + # M should be padded to mutiple of 32 if pad is True M, N = x.shape device = x.device - if x_q is None: - x_q = torch.empty((M, N), device=device, dtype=torch.float8_e4m3fn) - if x_scale is None: - x_scale = torch.empty((M,), device=device, dtype=torch.float32) - sm = torch.cuda.get_device_properties(device).multi_processor_count - W = triton.cdiv(M, sm) - grid = (sm,) - depracated_tokenwise_smooth_quant_kernel[grid]( + P = (M + 31) // 32 * 32 if pad else M + x_q = torch.empty((N, P), device=device, dtype=torch.float8_e4m3fn) + x_scale = torch.empty((N,), device=device, dtype=torch.float32) + H = 1024 + W = 16 # if N >= 4096 else 16 + assert N % W == 0 + EVEN = P % H == 0 and M == P + + grid = (triton.cdiv(N, W),) + transpose_smooth_quant_kernel[grid]( x, x_q, smooth_scale, x_scale, M, - W, N, + P, + H, + W, + EVEN, reverse, round_scale, num_stages=3, - num_warps=8 + num_warps=4 if N >= 8192 else 4, ) return x_q, x_scale @triton.jit -def batch_smooth_quant_kernel(x_ptr, q_ptr, ss_ptr, qs_ptr, xm_ptr, count_ptr, - accum_ptr, T, N: tl.constexpr, - REVERSE: tl.constexpr, ROUND: tl.constexpr, - CALIBRATE: tl.constexpr): - pid = tl.program_id(axis=0) - - i_expert = pid // T - i_batch = pid % T +def batch_smooth_quant_kernel( + x_ptr, + q_ptr, + ss_ptr, + qs_ptr, + count_ptr, + accum_ptr, + T, + N: tl.constexpr, + REVERSE: tl.constexpr, + ROUND: tl.constexpr, +): + eid = tl.program_id(axis=0) + tid = tl.program_id(axis=1) - # row-wise read, row-wise write - smooth_scale = tl.load(ss_ptr + i_expert * N + tl.arange(0, N)) + smooth_scale = tl.load(ss_ptr + eid * N + tl.arange(0, N)) if not REVERSE: smooth_scale = 1.0 / smooth_scale - if CALIBRATE: - x_maxs = tl.zeros((N,), dtype=tl.float32) - - count = tl.load(count_ptr + i_expert) - ei = tl.load(accum_ptr + i_expert) + count = tl.load(count_ptr + eid) + ei = tl.load(accum_ptr + eid) si = ei - count - n = tl.cdiv(count, T) # samples for each task - for i in range(i_batch * n, min((i_batch + 1) * n, count)): + n = tl.cdiv(count, T) # tokens per block + for i in range(tid * n, min((tid + 1) * n, count)): x = tl.load(x_ptr + si * N + i * N + tl.arange(0, N)).to(tl.float32) - if CALIBRATE: - x_maxs = tl.maximum(x_maxs, x.abs()) x *= smooth_scale scale = tl.maximum(tl.max(tl.abs(x)) / 448.0, 1e-30) if ROUND: @@ -431,99 +361,89 @@ def batch_smooth_quant_kernel(x_ptr, q_ptr, ss_ptr, qs_ptr, xm_ptr, count_ptr, xq = x.to(q_ptr.dtype.element_ty) tl.store(q_ptr + si * N + i * N + tl.arange(0, N), xq) - if CALIBRATE: - tl.store(xm_ptr + pid * N + tl.arange(0, N), x_maxs) - -""" -select and smooth and quant -x: [bs, dim] -smooth_scales: [n_experts, dim] -token_count_per_expert: [n_experts] -x_q: [bs, dim] -x_scale: [bs] -""" - - -def triton_batch_smooth_quant(x, smooth_scales, token_count_per_expert, - x_q=None, x_scale=None, x_maxs=None, - reverse=False, round_scale=False, - calibrate=False): - """""" +def triton_batch_smooth_quant( + x, smooth_scales, token_count_per_expert, reverse=False, round_scale=False +): + """ + smooth quant + x: [sum(tokens), dim] + smooth_scales: [n_experts, dim] + token_count_per_expert: [n_experts] + reverse: x * smooth_scale if reverse else x / smooth_scale + x_scale: [bs] + """ M, N = x.shape device = x.device n_expert = token_count_per_expert.shape[0] - assert 128 % n_expert == 0 - if x_q is None: - x_q = torch.empty((M, N), device=device, dtype=torch.float8_e4m3fn) - if x_scale is None: - x_scale = torch.empty((M,), device=device, dtype=torch.float32) + x_q = torch.empty((M, N), device=device, dtype=torch.float8_e4m3fn) + x_scale = torch.empty((M,), device=device, dtype=torch.float32) accum_token_count = torch.cumsum(token_count_per_expert, 0) - T = 128 // n_expert - if calibrate and x_maxs is None: - x_maxs = torch.empty((128, N), device=device, dtype=torch.float32) + T = 128 - grid = (128,) + grid = (n_expert, T) batch_smooth_quant_kernel[grid]( x, x_q, smooth_scales, x_scale, - x_maxs, token_count_per_expert, accum_token_count, - T, N, + T, + N, reverse, round_scale, - calibrate, num_stages=3, - num_warps=8 + num_warps=4, ) - if calibrate: - x_maxs = x_maxs.view(n_expert, T, N).amax(1) - return x_q, x_scale, x_maxs + return x_q, x_scale @triton.jit -def batch_pad_transpose_smooth_quant_kernel(x_ptr, q_ptr, ss_ptr, qs_ptr, - count_ptr, - accum_ptr, - N, - H: tl.constexpr, - W: tl.constexpr, - E: tl.constexpr, - REVERSE: tl.constexpr, - ROUND: tl.constexpr): +def batch_transpose_smooth_quant_kernel( + x_ptr, + q_ptr, + ss_ptr, + qs_ptr, + count_ptr, + N, + H: tl.constexpr, + W: tl.constexpr, + E: tl.constexpr, + REVERSE: tl.constexpr, + ROUND: tl.constexpr, +): eid = tl.program_id(axis=0) bid = tl.program_id(axis=1) count = tl.load(count_ptr + eid) - ei = tl.load(accum_ptr + eid) - si = ei - count round_count = tl.cdiv(count, 32) * 32 counts = tl.load(count_ptr + tl.arange(0, E)) - n_blocks = tl.cdiv(counts, 128) - bias = tl.sum(tl.where(tl.arange(0, E) < eid, n_blocks, 0)) + si = tl.sum(tl.where(tl.arange(0, E) < eid, counts, 0)) + + round_si = tl.sum(tl.where(tl.arange(0, E) < eid, tl.cdiv(counts, 32), 0)) * 32 n = tl.cdiv(count, H) maxs = tl.zeros((H, W), dtype=tl.float32) for i in range(n): - # col-wise read, row-wise write indices = i * H + tl.arange(0, H) - smooth_scale = tl.load(ss_ptr + indices, mask=indices < count) + smooth_scale = tl.load(ss_ptr + si + indices, mask=indices < count) if not REVERSE: smooth_scale = 1.0 / smooth_scale - x = tl.load(x_ptr + si * N + i * H * N + bid * W + tl.arange(0, H)[:, - None] + tl.arange(0, - W)[ - None, :], - mask=indices[:, None] < count).to(tl.float32) + x = tl.load( + x_ptr + + si * N + + i * H * N + + bid * W + + tl.arange(0, H)[:, None] * N + + tl.arange(0, W)[None, :], + mask=indices[:, None] < count, + ).to(tl.float32) x *= smooth_scale[:, None] maxs = tl.maximum(maxs, tl.abs(x)) - maxs = tl.max(maxs, 0) scale = tl.maximum(tl.max(maxs, 0) / 448.0, 1e-30) if ROUND: scale = tl.exp2(tl.ceil(tl.log2(scale))) @@ -531,25 +451,33 @@ def batch_pad_transpose_smooth_quant_kernel(x_ptr, q_ptr, ss_ptr, qs_ptr, s = 1.0 / scale for i in range(n): - # col-wise read, row-wise write indices = i * H + tl.arange(0, H) - smooth_scale = tl.load(ss_ptr + indices, mask=indices < count) + smooth_scale = tl.load(ss_ptr + si + indices, mask=indices < count) if not REVERSE: smooth_scale = 1.0 / smooth_scale - x = tl.load(x_ptr + si * N + i * H * N + bid * W + tl.arange(0, H)[:, - None] + tl.arange(0, - W)[ - None, :], - mask=indices[:, None] < count).to(tl.float32) + x = tl.load( + x_ptr + + si * N + + i * H * N + + bid * W + + tl.arange(0, H)[:, None] * N + + tl.arange(0, W)[None, :], + mask=indices[:, None] < count, + ).to(tl.float32) x *= smooth_scale[:, None] x *= s xq = tl.trans(x.to(q_ptr.dtype.element_ty)) tl.store( - q_ptr + bias * N + bid * W * round_count + i * H + tl.arange(0, W)[ - :, - None] + tl.arange( - 0, H)[None, :], xq, mask=indices[None, :] < round_count) + q_ptr + + round_si * N + + bid * W * round_count + + i * H + + tl.arange(0, W)[:, None] * round_count + + tl.arange(0, H)[None, :], + xq, + mask=indices[None, :] < round_count, + ) """ @@ -564,34 +492,32 @@ def batch_pad_transpose_smooth_quant_kernel(x_ptr, q_ptr, ss_ptr, qs_ptr, """ -def triton_batch_pad_transpose_smooth_quant(x, - smooth_scales, - token_count_per_expert, - splits, - x_q=None, x_scale=None, x_maxs=None, - reverse=False, round_scale=False): +def triton_batch_transpose_smooth_quant( + x, + smooth_scales, + token_count_per_expert, + splits, + pad=True, + reverse=False, + round_scale=False, +): """""" + assert pad and reverse M, N = x.shape device = x.device n_expert = token_count_per_expert.shape[0] round_splits = [(x + 31) // 32 * 32 for x in splits] - round_size = sum(round_splits) - if x_q is None: - x_q = torch.empty((round_size, N), device=device, - dtype=torch.float8_e4m3fn) - if x_scale is None: - x_scale = torch.empty((n_expert, N), device=device, dtype=torch.float32) - accum_token_count = torch.cumsum(token_count_per_expert, 0) + x_q = torch.empty((sum(round_splits), N), device=device, dtype=torch.float8_e4m3fn) + x_scale = torch.empty((n_expert, N), device=device, dtype=torch.float32) H = 128 W = 32 grid = (n_expert, N // W) - batch_pad_transpose_smooth_quant_kernel[grid]( + batch_transpose_smooth_quant_kernel[grid]( x, x_q, smooth_scales, x_scale, token_count_per_expert, - accum_token_count, N, H, W, @@ -599,155 +525,53 @@ def triton_batch_pad_transpose_smooth_quant(x, reverse, round_scale, num_stages=3, - num_warps=8 - ) - return x_q, x_scale - - -@triton.jit -def transpose_smooth_quant_kernel(x_ptr, q_ptr, ss_ptr, qs_ptr, M, N, P, - H: tl.constexpr, W: tl.constexpr, - EVEN: tl.constexpr, - REVERSE: tl.constexpr, - ROUND: tl.constexpr): - pid = tl.program_id(axis=0) - # col-wise read, row-wise write - offs = pid * W + tl.arange(0, H)[:, None] * N + tl.arange(0, W)[None, :] - soffs = tl.arange(0, H) - x_max = tl.zeros((W,), dtype=tl.float32) - m = tl.cdiv(P, H) - for i in range(m): - if EVEN: - x = tl.load(x_ptr + offs) - smooth_scale = tl.load(ss_ptr + soffs)[:, None] - else: - x = tl.load(x_ptr + offs, - mask=(i * H + tl.arange(0, H)[:, None] < M) & ( - pid * W + tl.arange(0, W)[None, :] < N)) - other = 0.0 if REVERSE else 1e30 - smooth_scale = tl.load(ss_ptr + soffs, mask=soffs < M, other=other)[ - :, None] - if REVERSE: - x = x * smooth_scale - else: - x = x / smooth_scale - x_max = tl.maximum(tl.max(tl.abs(x), axis=0), x_max) - offs += H * N - soffs += H - - scale = tl.maximum(x_max / 448.0, 1e-30) - if ROUND: - scale = tl.exp2(tl.ceil(tl.log2(scale))) - - if EVEN: - tl.store(qs_ptr + pid * W + tl.arange(0, W), scale) - else: - tl.store(qs_ptr + pid * W + tl.arange(0, W), scale, - mask=pid * W + tl.arange(0, W) < N) - - s = (1.0 / scale)[None, :] - offs = pid * W + tl.arange(0, H)[:, None] * N + tl.arange(0, W)[None, :] - soffs = tl.arange(0, H) - toffs = pid * W * P + tl.arange(0, W)[:, None] * P + tl.arange(0, H)[None, - :] - for i in range(m): - if EVEN: - x = tl.load(x_ptr + offs).to(tl.float32) - smooth_scale = tl.load(ss_ptr + soffs)[:, None] - else: - x = tl.load(x_ptr + offs, - mask=(i * H + tl.arange(0, H)[:, None] < M)).to( - tl.float32) - other = 0.0 if REVERSE else 1e30 - smooth_scale = tl.load(ss_ptr + soffs, mask=soffs < M, other=other)[ - :, None] - - if REVERSE: - x = (x * smooth_scale * s).to(q_ptr.dtype.element_ty) - else: - x = (x / smooth_scale * s).to(q_ptr.dtype.element_ty) - if EVEN: - tl.store(q_ptr + toffs, tl.trans(x)) - else: - # mask with P instead of M - tl.store(q_ptr + toffs, tl.trans(x), - mask=(i * H + tl.arange(0, H)[None, :] < P)) - offs += H * N - toffs += H - soffs += H - - -def triton_transpose_smooth_quant(x, - smooth_scale, - reverse=False, - pad=False, - round_scale=False): - # col-wise read, row-wise write - # M should be padded if M % 32 != 0 - """""" - M, N = x.shape - device = x.device - P = (M + 31) // 32 * 32 if pad else M - x_q = torch.empty((N, P), device=device, dtype=torch.float8_e4m3fn) - x_scale = torch.empty((N,), device=device, dtype=torch.float32) - H = 1024 - W = 16 # if N >= 4096 else 16 - assert N % W == 0 - EVEN = P % H == 0 and M == P - - grid = (triton.cdiv(N, W),) - transpose_smooth_quant_kernel[grid]( - x, - x_q, - smooth_scale, - x_scale, - M, - N, - P, - H, - W, - EVEN, - reverse, - round_scale, - num_stages=3, - num_warps=4 if N >= 8192 else 4 + num_warps=8, ) return x_q, x_scale @triton.jit -def transpose_rescale_smooth_quant_kernel(x_ptr, q_ptr, - org_smooth_scale_ptr, - org_quant_scale_ptr, - transpose_smooth_scale_ptr, - transpose_quant_scale_ptr, M, - N, P, H: tl.constexpr, - W: tl.constexpr, - EVEN: tl.constexpr, - ROUND: tl.constexpr): +def transpose_rescale_smooth_quant_kernel( + x_ptr, + q_ptr, + org_smooth_scale_ptr, + org_quant_scale_ptr, + transpose_smooth_scale_ptr, + transpose_quant_scale_ptr, + M, + N, + P, + H: tl.constexpr, + W: tl.constexpr, + EVEN: tl.constexpr, + ROUND: tl.constexpr, +): pid = tl.program_id(axis=0) - # col-wise read, row-wise write offs = pid * W + tl.arange(0, H)[:, None] * N + tl.arange(0, W)[None, :] soffs = tl.arange(0, H) x_max = tl.zeros((W,), dtype=tl.float32) - org_smooth_scale = tl.load( - org_smooth_scale_ptr + pid * W + tl.arange(0, W))[None, :] + org_smooth_scale = tl.load(org_smooth_scale_ptr + pid * W + tl.arange(0, W))[ + None, : + ] m = tl.cdiv(P, H) for i in range(m): if EVEN: x = tl.load(x_ptr + offs).to(tl.float32) org_quant_scale = tl.load(org_quant_scale_ptr + soffs)[:, None] - transpose_smooth_scale = tl.load( - transpose_smooth_scale_ptr + soffs)[:, None] + transpose_smooth_scale = tl.load(transpose_smooth_scale_ptr + soffs)[ + :, None + ] else: - x = tl.load(x_ptr + offs, - mask=(i * H + tl.arange(0, H)[:, None] < M)).to( - tl.float32) - org_quant_scale = tl.load(org_quant_scale_ptr + soffs, - mask=soffs < M, other=0.0)[:, None] - transpose_smooth_scale = tl.load(transpose_smooth_scale_ptr + soffs, - mask=soffs < M, other=0.0)[:, None] + x = tl.load(x_ptr + offs, mask=(i * H + tl.arange(0, H)[:, None] < M)).to( + tl.float32 + ) + org_quant_scale = tl.load( + org_quant_scale_ptr + soffs, mask=soffs < M, other=0.0 + )[:, None] + transpose_smooth_scale = tl.load( + transpose_smooth_scale_ptr + soffs, mask=soffs < M, other=0.0 + )[:, None] x = x / org_smooth_scale * (org_quant_scale * transpose_smooth_scale) x_max = tl.maximum(tl.max(tl.abs(x), axis=0), x_max) @@ -764,33 +588,34 @@ def transpose_rescale_smooth_quant_kernel(x_ptr, q_ptr, offs = pid * W + tl.arange(0, H)[:, None] * N + tl.arange(0, W)[None, :] soffs = tl.arange(0, H) - toffs = pid * W * P + tl.arange(0, W)[:, None] * P + tl.arange(0, H)[None, - :] + toffs = pid * W * P + tl.arange(0, W)[:, None] * P + tl.arange(0, H)[None, :] for i in range(m): if EVEN: x = tl.load(x_ptr + offs).to(tl.float32) org_quant_scale = tl.load(org_quant_scale_ptr + soffs)[:, None] - transpose_smooth_scale = tl.load( - transpose_smooth_scale_ptr + soffs)[:, None] + transpose_smooth_scale = tl.load(transpose_smooth_scale_ptr + soffs)[ + :, None + ] else: - x = tl.load(x_ptr + offs, - mask=(i * H + tl.arange(0, H)[:, None] < M) & ( - pid * W + tl.arange(0, W)[None, :] < N)).to( - tl.float32) - org_quant_scale = tl.load(org_quant_scale_ptr + soffs, - mask=soffs < M, other=0.0)[:, None] - transpose_smooth_scale = tl.load(transpose_smooth_scale_ptr + soffs, - mask=soffs < M, other=0.0)[:, None] - - x = x * s / org_smooth_scale * ( - org_quant_scale * transpose_smooth_scale) + x = tl.load( + x_ptr + offs, + mask=(i * H + tl.arange(0, H)[:, None] < M) + & (pid * W + tl.arange(0, W)[None, :] < N), + ).to(tl.float32) + org_quant_scale = tl.load( + org_quant_scale_ptr + soffs, mask=soffs < M, other=0.0 + )[:, None] + transpose_smooth_scale = tl.load( + transpose_smooth_scale_ptr + soffs, mask=soffs < M, other=0.0 + )[:, None] + + x = x * s / org_smooth_scale * (org_quant_scale * transpose_smooth_scale) x = tl.trans(x.to(q_ptr.dtype.element_ty)) if EVEN: tl.store(q_ptr + toffs, x) else: - tl.store(q_ptr + toffs, x, - mask=(i * H + tl.arange(0, H)[None, :] < P)) + tl.store(q_ptr + toffs, x, mask=(i * H + tl.arange(0, H)[None, :] < P)) offs += H * N toffs += H soffs += H @@ -805,17 +630,20 @@ def transpose_rescale_smooth_quant_kernel(x_ptr, q_ptr, """ -def triton_transpose_rescale_smooth_quant(x_q, org_smooth_scale, - org_quant_scale, - transpose_smooth_scale, - reverse=True, - pad=False, - round_scale=False): +def triton_transpose_rescale_smooth_quant( + x_q, + org_smooth_scale, + org_quant_scale, + transpose_smooth_scale, + reverse=True, + pad=False, + round_scale=False, +): """""" assert reverse M, N = x_q.shape device = x_q.device - P = round_up(M, b=32) if pad else M + P = (M + 31) // 32 * 32 if pad else M xt_q = torch.empty((N, P), device=device, dtype=torch.float8_e4m3fn) x_scale = torch.empty((N,), device=device, dtype=torch.float32) H = 256 @@ -831,138 +659,148 @@ def triton_transpose_rescale_smooth_quant(x_q, org_smooth_scale, org_quant_scale, transpose_smooth_scale, x_scale, - M, N, P, - H, W, + M, + N, + P, + H, + W, EVEN, round_scale, num_stages=4, - num_warps=8 + num_warps=8, ) return xt_q, x_scale +@triton.jit +def subrow_smooth_quant_kernel( + x_ptr, + q_ptr, + ss_ptr, + qs_ptr, + subrow_scales_ptr, + tail_ri, + tail_si, + head_ri, + head_ei, + size, + N, + W: tl.constexpr, + TAIL: tl.constexpr, + HEAD: tl.constexpr, + REVERSE: tl.constexpr, + ROUND: tl.constexpr, +): + if TAIL: + # scale is saved as max/448 + scale = tl.maximum(tl.load(subrow_scales_ptr), 1e-30) + if ROUND: + scale = tl.exp2(tl.ceil(tl.log2(scale))) + # scale only stores in subrow with leading values -""" -megatron fp8 training steps: -step 0: init w smooth scale w_smooth -step 1: smooth and quant w after w is updated by optimizer -step 2: in forward step, columnwise smooth x and rowwise quant x, calc y=x@w; - meanwhile, record the columnwise max of x, it is used to update w_smooth -step 3: in dgrad step, columnwise smooth y and rowwise quant y, transpose x, calc dx=y@wT -step 4: in wgrad step, dequant then smooth an then quant y_q to get yt_q, calc dw=yT@x - -alternative (it's not suitable for fp8 combine): -step 4: in wgrad step, rowwise smooth y and columnwise quant y and transpose to get yt_q, calc dw=yT@x - -""" - -""" -divide x by smooth_scale and row-wise quantization -smooth scale is updated by square root of x's column-wise maxs, and set in weight's x_maxs attr - -transpose: transpose quantized x for wgrad -pad: # pad M to be multiplier of 32, including quant scales and transposed x - -""" - + T = tl.cdiv(N - tail_si, W) + for i in range(T): + mask = tail_si + i * W + tl.arange(0, W) < N + if REVERSE: + smooth_scale = tl.load( + ss_ptr + tail_si + i * W + tl.arange(0, W), mask=mask + ) + else: + smooth_scale = tl.load( + ss_ptr + tail_si + i * W + tl.arange(0, W), other=1e30, mask=mask + ) + smooth_scale = 1.0 / smooth_scale + x = tl.load(x_ptr + i * W + tl.arange(0, W), mask=mask).to(tl.float32) + x *= smooth_scale + x /= scale + xq = tl.minimum(tl.maximum(x, -448), 448) + tl.store( + q_ptr + tail_ri * N + tail_si + i * W + tl.arange(0, W), + xq.to(q_ptr.dtype.element_ty), + mask=mask, + ) -# y = x @ w -# dx = y @ wT -# dwT = yT @ x -def triton_smooth_quant_input(x, smooth_scale, x_q=None, x_scale=None, xt_q=None, - transpose=True, pad=True, round_scale=False): - """""" - x_q, x_scale, x_maxs = triton_smooth_quant(x, smooth_scale, x_q=x_q, - x_scale=x_scale, reverse=False, - round_scale=round_scale) + if HEAD: + # scale is saved as max/448 + scale = tl.maximum(tl.load(subrow_scales_ptr + 1), 1e-30) + if ROUND: + scale = tl.exp2(tl.ceil(tl.log2(scale))) + tl.store(qs_ptr + head_ri, scale) - if transpose: - xt_q = triton_transpose_and_pad(x_q, out=xt_q, pad=pad) - else: - xt_q = None - xt_scale = smooth_scale - - return x_q, xt_q, x_scale, xt_scale - - -# y = x @ w -# dx = y @ wT -# dwT = yT @ x -def triton_smooth_quant_gradient(y, - smooth_scale, - transpose_smooth_scale, - reverse=True, - transpose=True, - pad=True, - round_scale=False): + T = tl.cdiv(head_ei, W) + for i in range(T): + mask = i * W + tl.arange(0, W) < head_ei + if REVERSE: + smooth_scale = tl.load(ss_ptr + i * W + tl.arange(0, W), mask=mask) + else: + smooth_scale = tl.load( + ss_ptr + i * W + tl.arange(0, W), other=1e30, mask=mask + ) + smooth_scale = 1.0 / smooth_scale + x = tl.load(x_ptr + size - head_ei + i * W + tl.arange(0, W), mask=mask).to( + tl.float32 + ) + x *= smooth_scale + x /= scale + xq = tl.minimum(tl.maximum(x, -448), 448) + tl.store( + q_ptr + head_ri * N + i * W + tl.arange(0, W), + xq.to(q_ptr.dtype.element_ty), + mask=mask, + ) + + +def triton_subrow_smooth_quant( + x, + smooth_scale, + x_q, + x_scale, + subrow_scales, + offset, + size, + reverse=False, + round_scale=False, +): """""" - assert reverse, ("args `smooth_scale` and/or `transpose_smooth_scale` " - "must be in reciprocal format in triton_smooth_quant_grad") - y_q, y_scale, _ = triton_smooth_quant(y, smooth_scale, reverse=True, - round_scale=round_scale) - if transpose: - yt_q, yt_scale = triton_transpose_smooth_quant(y, - transpose_smooth_scale, - reverse=True, - pad=pad, - round_scale=round_scale) + M, N = x_q.shape + W = 128 + if offset % N == 0: + tail_ri = 0 + tail_si = 0 + TAIL = False else: - yt_q, yt_scale = None, None - - return y_q, yt_q, y_scale, yt_scale - + tail_ri = offset // N + tail_si = offset % N + TAIL = True -def triton_smooth_quant_weight(w, - smooth_scale, - w_q, - quant_scale, - subrow_scales, offset=0, - round_scale=False): - """""" - assert w.ndim == 1 - assert w_q.size(1) == smooth_scale.size(0) - - size = w.numel() - M, N = w_q.shape - - if size == M * N: - triton_smooth_quant(w.view(M, N), smooth_scale, x_q=w_q, - x_scale=quant_scale, - round_scale=round_scale) - elif offset % N == 0 and size % N == 0: - n_row = size // N - row_id = offset // N - w_q_slice = w_q[row_id:row_id + n_row] - quant_scale_slice = quant_scale[row_id:row_id + n_row] - triton_smooth_quant(w.view(n_row,N), smooth_scale, x_q=w_q_slice, - x_scale=quant_scale_slice, - round_scale=round_scale) + if (offset + size) % N == 0: + head_ri = 0 + head_ei = 0 # head_size = head_ei + HEAD = False else: - row_si = (offset - 1)//N + 1 - row_ei = (offset + size) // N - col_si = offset % N - col_ei = (offset + size ) % N - n_row = row_ei - row_si - mw_offset = 0 if col_si == 0 else N - col_si - w_q_slice = w_q[row_si:row_ei] - quant_scale_slice = quant_scale[row_si:row_ei] - w_slice = w[mw_offset:mw_offset+n_row*N].view(n_row,N) - triton_smooth_quant(w_slice, - smooth_scale, - x_q=w_q_slice, - x_scale=quant_scale_slice, - round_scale=round_scale) - - # subrow scale is writed by the row with leading master weights - if col_si > 0 or col_ei > 0: - triton_subrow_smooth_quant(w, - smooth_scale, - w_q, - quant_scale, - subrow_scales, - offset, - size, - reverse=False, - round_scale=round_scale) + head_ri = (offset + size) // N + head_ei = (offset + size) % N + HEAD = True + grid = (1,) + subrow_smooth_quant_kernel[grid]( + x, + x_q, + smooth_scale, + x_scale, + subrow_scales, + tail_ri, + tail_si, + head_ri, + head_ei, + size, + N, + W, + TAIL, + HEAD, + reverse, + round_scale, + num_stages=3, + num_warps=1, + ) diff --git a/linghe/tools/__init__.py b/linghe/tools/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/linghe/tools/benchmark.py b/linghe/tools/benchmark.py index 3fa7737..e44fe10 100644 --- a/linghe/tools/benchmark.py +++ b/linghe/tools/benchmark.py @@ -9,18 +9,26 @@ from torch.profiler import profile, ProfilerActivity -def benchmark_func(fn, *args, n_warmup=10, n_repeat=100, ref_flops=None, - ref_bytes=None, ref_time=None, - n_profile=0, trace_dir=None, - name='', **kwargs): - func_name = getattr(fn, '__name__', None) - func_name = name if func_name == 'apply' or func_name is None else func_name +def benchmark_func( + fn, + *args, + n_warmup=10, + n_repeat=100, + ref_flops=None, + ref_bytes=None, + ref_time=None, + n_profile=0, + trace_dir=None, + name="", + **kwargs, +): + func_name = getattr(fn, "__name__", None) + func_name = name if func_name == "apply" or func_name is None else func_name for i in range(n_warmup): fn(*args, **kwargs) - start_events = [torch.cuda.Event(enable_timing=True) for _ in - range(n_repeat)] + start_events = [torch.cuda.Event(enable_timing=True) for _ in range(n_repeat)] end_events = [torch.cuda.Event(enable_timing=True) for _ in range(n_repeat)] ts = time.time() @@ -32,36 +40,46 @@ def benchmark_func(fn, *args, n_warmup=10, n_repeat=100, ref_flops=None, torch.cuda.synchronize() if n_profile > 0: - with profile(activities=[ProfilerActivity.CPU, - ProfilerActivity.CUDA, - ProfilerActivity.XPU]) as prof: + with profile( + activities=[ + ProfilerActivity.CPU, + ProfilerActivity.CUDA, + ProfilerActivity.XPU, + ], + with_stack=True, + ) as prof: for i in range(n_profile): fn(*args, **kwargs) - print(prof.key_averages().table(sort_by="cuda_time_total", - top_level_events_only=True, - row_limit=100)) + print( + prof.key_averages().table( + sort_by="cuda_time_total", top_level_events_only=True, row_limit=100 + ) + ) if trace_dir is not None: - assert trace_dir.endswith('.json') + assert trace_dir.endswith(".json") prof.export_chrome_trace(trace_dir) times = [s.elapsed_time(e) for s, e in zip(start_events, end_events)] times = sorted(times) clip = max(1, n_repeat // 100) - times = sum(times[clip:-clip]) + if 2 * clip < n_repeat: + times = sum(times[clip:-clip]) + n_repeat = n_repeat - 2 * clip + else: + times = sum(times) - average_event_time = times * 1000 / (n_repeat - 2 * clip) + average_event_time = times * 1000 / n_repeat - fs = '' + fs = "" if ref_flops is not None: linghe = ref_flops / 1e12 / (average_event_time / 1e6) - fs = f'FLOPS:{linghe:.2f}T' - bs = '' + fs = f"FLOPS:{linghe:.2f}T" + bs = "" if ref_bytes is not None: - bs = f'bandwidth:{ref_bytes / average_event_time / 1e3:.1f}G/S' - ss = '' + bs = f"bandwidth:{ref_bytes / average_event_time / 1e3:.1f}G/S" + ss = "" if ref_time is not None: - ss = f'speedup:{ref_time / average_event_time:.3f}' + ss = f"speedup:{ref_time / average_event_time:.3f}" - print( - f'{func_name:<30} {name} time:{average_event_time:.1f} us {fs} {bs} {ss}') + print(f"{func_name:<30} {name} time:{average_event_time:.1f} us {fs} {bs} {ss}") return average_event_time diff --git a/linghe/tools/check.py b/linghe/tools/check.py new file mode 100644 index 0000000..7755590 --- /dev/null +++ b/linghe/tools/check.py @@ -0,0 +1,166 @@ +# -*- coding: utf-8 -*- +""" +Copyright (c) Ant Financial Service Group and its affiliates. +""" + +import math + +import torch + + +def output_check( + org_out, opt_out, name="", rtol=None, atol=None, itol=0, amp=1.0, digest=4 +): + org_out = org_out.detach() + opt_out = opt_out.detach() + assert org_out.dtype == opt_out.dtype, f"ref:{org_out.dtype} != out:{opt_out.dtype}" + assert org_out.shape == opt_out.shape, f"ref:{org_out.shape} != out:{opt_out.shape}" + if org_out.numel() == 0: + return + org_dtype = org_out.dtype + opt_dtype = opt_out.dtype + + if org_dtype in ( + torch.bfloat16, + torch.float16, + torch.float8_e4m3fn, + torch.float8_e5m2, + ): + org_out = org_out.float() + elif org_dtype in (torch.bool, torch.uint8, torch.int8, torch.uint16, torch.int16): + org_out = org_out.int() + + if opt_dtype in ( + torch.bfloat16, + torch.float16, + torch.float8_e4m3fn, + torch.float8_e5m2, + ): + opt_out = opt_out.float() + elif org_dtype in (torch.bool, torch.uint8, torch.int8, torch.uint16, torch.int16): + opt_out = opt_out.int() + + if rtol is None: + if org_dtype == torch.float64: + rtol = 1e-8 + elif org_dtype == torch.float32: + rtol = 1e-4 + elif org_dtype == torch.float16: + rtol = 4e-3 + elif org_dtype == torch.bfloat16: + rtol = 2e-2 + elif org_dtype == torch.float8_e4m3fn: + rtol = 0.125 + elif org_dtype == torch.float8_e5m2: + rtol = 0.25 + + if atol is None: + if org_dtype == torch.float64: + atol = 1e-12 + elif org_dtype == torch.float32: + atol = 1e-6 + elif org_dtype == torch.float16: + atol = 1e-4 + elif org_dtype == torch.bfloat16: + atol = 1e-3 + elif org_dtype == torch.float8_e4m3fn: + atol = 0.125 + elif org_dtype == torch.float8_e5m2: + atol = 0.25 + + if org_out.dtype in (torch.float32, torch.float64): + rtol = rtol * amp + atol = atol * amp + diff = (opt_out - org_out).abs() + abs_error = diff.mean().item() + rel_error = abs_error / max(org_out.abs().mean().item(), 1e-30) + if rel_error >= 0.005: + rel_err_str = f"\033[91m {rel_error:.6f}\033[00m" + else: + rel_err_str = f"{rel_error:.6f}" + org_max = org_out.abs().max() + org_mean = org_out.abs().mean() + opt_max = opt_out.abs().max() + opt_mean = opt_out.abs().mean() + print( + f"\n{name:<16} rel:{rel_err_str} abs:{abs_error:.6f} " + f"org:{org_max:.3f}/{org_mean:.3f} " + f"opt:{opt_max:.3f}/{opt_mean:.3f} " + ) + if rtol >= 0 and atol >= 0: + # torch.testing.assert_close(opt_out, org_out, rtol=rtol, atol=atol) + mistake_mask = diff >= (rtol * org_out.abs() + atol) + if mistake_mask.float().sum().item() > 0: + mistake_indices = torch.where(mistake_mask)[0] + org_val = org_out[mistake_mask] + opt_val = opt_out[mistake_mask] + mismatch_count = org_val.numel() + tot_cnt = org_out.numel() + itv = max(mismatch_count // digest, 1) + org_val = org_val[::itv].tolist() + opt_val = opt_val[::itv].tolist() + mistake_indices = ", ".join( + [f"{x}" for x in mistake_indices[::itv].tolist()] + ) + if org_dtype == torch.float64: + org_str = ", ".join([f"{x:.8g}" for x in org_val]) + opt_str = ", ".join([f"{x:.8g}" for x in opt_val]) + elif org_dtype == torch.float32: + org_str = ", ".join([f"{x:.5g}" for x in org_val]) + opt_str = ", ".join([f"{x:.5g}" for x in opt_val]) + else: + org_str = ", ".join([f"{x:.3g}" for x in org_val]) + opt_str = ", ".join([f"{x:.3g}" for x in opt_val]) + info = ( + f"Mismatched elements: {mismatch_count} / {tot_cnt} ({mismatch_count / tot_cnt * 100:.1f}%) " + f"with {rtol} rtol and {atol} atol \n org: {org_str} \n " + f"opt: {opt_str} \n idx: {mistake_indices} \n" + ) + assert mismatch_count == 0, info + return rel_error + else: + # int dtype + diff = (opt_out - org_out).abs() + mismatch_count = (diff > itol).sum().item() + if mismatch_count > 0: + diff_err_str = f"\033[91m {mismatch_count}\033[00m" + else: + diff_err_str = f"{mismatch_count}" + max_error = diff.max() + print(f"\n{name:<16} diff:{diff_err_str} max:{max_error}") + assert ( + mismatch_count == 0 + ), f"Mismatched elements: {mismatch_count} with {itol} itol" + return mismatch_count + + +def quant_check(org_out, xq, wq, opt_out, mode): + abs_error = (opt_out.float() - org_out.float()).abs().mean().item() + rel_error = abs_error / org_out.float().abs().mean().item() + x_underflow = (xq == 0.0).sum().item() / xq.numel() + w_underflow = (wq == 0.0).sum().item() / wq.numel() + x_overflow = (torch.isnan(xq)).sum().item() + w_overflow = (torch.isnan(wq)).sum().item() + print( + f"\n{mode} rel:{rel_error:.3f} abs:{abs_error:.3f} " + f"org:{org_out.abs().max():.3f}/{org_out.abs().mean():.3f} " + f"opt:{opt_out.abs().max():.3f}/{opt_out.abs().mean():.3f} " + f"x_underflow:{x_underflow:.5f} w_underflow:{w_underflow:.5f} " + f"x_overflow:{x_overflow} w_overflow:{w_overflow}" + ) + + +def inf_or_nan(xs, name=""): + if not isinstance(xs, (list, tuple)): + xs = [xs] + hit = False + for x in xs: + value = x.abs().max().item() + if math.isnan(value) or math.isinf(value): + hit = True + break + if hit: + for x in xs: + print( + f"{name=} {x.shape=} {x.argmax()=} {x.max()=} {x.argmin()=} {x.min()=} {x=}" + ) diff --git a/linghe/tools/util.py b/linghe/tools/util.py index a97768e..cb368d0 100644 --- a/linghe/tools/util.py +++ b/linghe/tools/util.py @@ -4,7 +4,10 @@ """ import math + import torch +import triton +import triton.language as tl def round_up(x, b=16): @@ -21,10 +24,12 @@ def torch_tensor_quant(x, dtype=torch.float8_e4m3fn, round_scale=False): def torch_row_quant(x, dtype=torch.float8_e4m3fn, round_scale=False): + x = x.float() fmax = torch.finfo(dtype).max scale = torch.abs(x).amax(1) / fmax - scale = torch.maximum(scale, 1e-30 * torch.ones((1,), dtype=torch.float32, - device=x.device)) + scale = torch.maximum( + scale, 1e-30 * torch.ones((1,), dtype=torch.float32, device=x.device) + ) if round_scale: scale = torch.exp2(torch.ceil(torch.log2(scale))) x_q = (x / scale[:, None]).to(dtype) @@ -52,8 +57,9 @@ def torch_group_quant(x, B=128, dtype=torch.float8_e4m3fn, round_scale=False): xp = torch.reshape(x.contiguous(), (M, P // B, B)) scale = torch.amax(torch.abs(xp).float(), dim=2) / fmax - scaoe = torch.maximum(scale, 1e-30 * torch.ones((1,), dtype=torch.float32, - device=x.device)) + scale = torch.maximum( + scale, 1e-30 * torch.ones((1,), dtype=torch.float32, device=x.device) + ) if round_scale: scale = torch.exp2(torch.ceil(torch.log2(scale))) xq = (xp / scale[:, :, None]).to(dtype) @@ -62,46 +68,139 @@ def torch_group_quant(x, B=128, dtype=torch.float8_e4m3fn, round_scale=False): return xq, scale +def torch_group_dequant(x_q, x_s, B=128): + m, n = x_s.shape + x_dq = x_q.float() * x_s.repeat_interleave(128, 1) + return x_dq + + +def torch_blockwise_quant(x, round_scale=True, padding=False): + m, N = x.shape + + if padding: + padding_size = (m + 15) // 16 * 16 - m + if padding_size > 0: + x = torch.nn.functional.pad(x, (0, 0, 0, padding_size)) + + x = x.float() + + y_q, y_scale = torch_group_quant(x, round_scale=round_scale) + yt_q, yt_scale = torch_group_quant(x.t(), round_scale=round_scale) + + return y_q, y_scale.t().contiguous(), yt_q, yt_scale.t().contiguous() + + +def torch_blockwise_dequant(w_q, w_s, B=128): + w_s = w_s.repeat_interleave(B, 1) + w_s = w_s.repeat_interleave(B, 0) + x_dq = w_q.float() * w_s + return x_dq + + def torch_block_quant(w, B=128, dtype=torch.float8_e4m3fn, round_scale=False): fmax = torch.finfo(dtype).max w = w.clone() N, K = w.shape - wp = torch.reshape(w.t().contiguous(), (K // B, B, N // B, B)).permute(0, 2, - 1, 3) + wp = torch.reshape(w, (N // B, B, K // B, B)).permute(0, 2, 1, 3) scale = torch.amax(torch.amax(torch.abs(wp).float(), dim=2), dim=2) / fmax if round_scale: scale = torch.exp2(torch.ceil(torch.log2(scale))) wq = (wp / scale[:, :, None, None]).to(dtype) wq = wq.permute(0, 2, 1, 3) - wq = torch.reshape(wq, (K, N)).t().contiguous() + wq = torch.reshape(wq, (N, K)).contiguous() return wq, scale +def torch_mxfp8_quant(x, padding=False, zero=False): + m_ori, N = x.shape + if padding: + padding_size = (m_ori + 31) // 32 * 32 - m_ori + if padding_size > 0: + x = torch.nn.functional.pad(x, (0, 0, 0, padding_size)) + + x = x.float() + m, N = x.shape # current m is multiple of 32 + assert N % 128 == 0 + if m % 128 != 0: + M = (m + 127) // 128 * 128 + x = torch.cat([x, torch.zeros((M - m, N), dtype=x.dtype, device=x.device)], 0) + else: + M = m + xs = x.view(M, N // 32, 32) + xm = xs.abs().amax(2) + scale = torch.maximum(xm / 448, 1e-30 * torch.ones_like(xm)) + scale = torch.exp2(torch.ceil(torch.log2(scale))) + x_q = (xs / scale[:, :, None]).to(torch.float8_e4m3fn).view(M, N)[:m] # 取得前m行 + if zero: + scale[m_ori:, :] = 0 + x_scale = scale.to(torch.float8_e8m0fnu).view(torch.uint8) + + xs = x.view(M // 32, 32, N) + xm = xs.abs().amax(1) + scale = torch.maximum(xm / 448, 1e-30 * torch.ones_like(xm)) + scale = torch.exp2(torch.ceil(torch.log2(scale))) + xt_q = (xs / scale[:, None, :]).to(torch.float8_e4m3fn).view(M, N)[:m] + if zero: + scale[(m_ori + 31) // 32 :, :] = 0 + xt_scale = scale.to(torch.float8_e8m0fnu).view(torch.uint8) + + return x_q, x_scale, xt_q, xt_scale + + +def torch_batch_mxfp8_quant(x, token_count_per_expert_list): + M, DIM = x.shape + q_refs = [] + s_refs = [] + qt_refs = [] + st_refs = [] + s = 0 + for i, c in enumerate(token_count_per_expert_list): + c = token_count_per_expert_list[i] + if c == 0: + continue + y = x[s : s + c] + y = y.float() + + y_q, y_scale, yt_q, yt_scale = torch_mxfp8_quant(y) + q_refs.append(y_q) + s_refs.append(y_scale) + qt_refs.append(yt_q) + st_refs.append(yt_scale) + s += c + q_ref = torch.cat(q_refs, 0) + s_ref = torch.cat(s_refs, 0) + qt_ref = torch.cat(qt_refs, 0) + st_ref = torch.cat(st_refs, 0) + return q_ref, s_ref, qt_ref, st_ref + + def torch_smooth_quant(x, smooth_scale, reverse=False, round_scale=False): x = x.float() - x_maxs = x.abs().amax(0) + # x_maxs = x.abs().amax(0) if reverse: x_smooth = x * smooth_scale else: - x_smooth = x / torch.maximum(smooth_scale, - 1e-30 * torch.ones_like(smooth_scale)) + x_smooth = x / torch.maximum( + smooth_scale, 1e-30 * torch.ones_like(smooth_scale) + ) scale = x_smooth.abs().amax(1) / 448 scale = torch.maximum(scale, 1e-30 * torch.ones_like(scale)) if round_scale: scale = torch.exp2(torch.ceil(torch.log2(scale))) x_q = (x_smooth / scale[:, None]).to(torch.float8_e4m3fn) - return x_q, scale, x_maxs + return x_q, scale -def torch_batch_smooth_quant(xs, smooth_scales, indices, token_count_per_expert, - reverse=False, round_scale=False): +def torch_batch_smooth_quant( + xs, smooth_scales, indices, token_count_per_expert, reverse=False, round_scale=False +): q_refs = [] scale_refs = [] s = 0 for i, c in enumerate(token_count_per_expert): - idx = indices[s:s + c] + idx = indices[s : s + c] y_slice = xs[idx] if reverse: y_smooth = y_slice * smooth_scales[i] @@ -121,9 +220,9 @@ def torch_batch_smooth_quant(xs, smooth_scales, indices, token_count_per_expert, def torch_make_indices(logits, topk=8, bias=-0.01): M, n_experts = logits.shape device = logits.device - logits = logits.to(torch.float64) + 1e-10 * torch.arange(n_experts, - device=device).to( - torch.float32) + logits = logits.to(torch.float64) + 1e-10 * torch.arange( + n_experts, device=device + ).to(torch.float32) topk_values, topk_indices = torch.topk(logits, topk, dim=-1, sorted=True) logits[logits < topk_values[:, -1:] + bias] = -1000000 probs = torch.nn.Softmax(dim=1)(logits) @@ -136,11 +235,15 @@ def torch_make_indices(logits, topk=8, bias=-0.01): torch.arange(M, device=logits.device).unsqueeze(0).expand(n_experts, -1) ) indices = token_indices.masked_select(route_map.T.contiguous()) - row_id_map = torch.reshape( - torch.cumsum(route_map.T.contiguous().view(-1), 0), (n_experts, M)) - 1 + row_id_map = ( + torch.reshape( + torch.cumsum(route_map.T.contiguous().view(-1), 0), (n_experts, M) + ) + - 1 + ) row_id_map[torch.logical_not(route_map.T)] = -1 row_id_map = row_id_map.T.contiguous() - return probs, route_map, token_count_per_expert, indices, row_id_map + return probs.float(), route_map, token_count_per_expert, indices, row_id_map # quant with scaling to 448 @@ -201,26 +304,25 @@ def torch_outlier_quant(x, w, dtype): return xq, wq, x_scale, w_scale, max_idx[:4], x_outlier -def make_hadamard_matrix(n, device='cuda:0', dtype=torch.bfloat16, norm=False): +def make_hadamard_matrix(n, device="cuda:0", dtype=torch.bfloat16, norm=False): assert 2 ** int(math.log2(n)) == n - m2 = torch.tensor([[1, 1], [1, -1]], device='cpu', dtype=torch.float32) + m2 = torch.tensor([[1, 1], [1, -1]], device="cpu", dtype=torch.float32) m = m2 for i in range(int(math.log2(n)) - 1): m = torch.kron(m, m2) if norm: - m = m / n ** 0.5 + m = m / n**0.5 return m.to(dtype=dtype, device=device) -def torch_hadamard_transform(x, hm, side='right'): - assert side in ('right', 'left') +def torch_hadamard_transform(x, hm, side="right"): + assert side in ("right", "left") x = x.clone() hm = hm.clone() M, K = x.shape B = hm.size(0) - xp = torch.reshape(x, (M // B, B, K // B, B)).permute(0, 2, 1, - 3).contiguous() - if side == 'right': + xp = torch.reshape(x, (M // B, B, K // B, B)).permute(0, 2, 1, 3).contiguous() + if side == "right": xp = xp @ hm else: xp = hm @ xp @@ -238,12 +340,14 @@ def torch_channel_quant_f_and_b(x, w, y): w_scale = w.abs().float().amax(dim=1, keepdim=True) / 448.0 # [N,1] xq = (x / x_scale).to(torch.float8_e4m3fn) wq = (w / w_scale).to(torch.float8_e4m3fn) - o = torch._scaled_mm(xq, - wq.t(), - scale_a=x_scale.view(-1, 1), - scale_b=w_scale.view(1, -1), - out_dtype=torch.bfloat16, - use_fast_accum=True) + o = torch._scaled_mm( + xq, + wq.t(), + scale_a=x_scale.view(-1, 1), + scale_b=w_scale.view(1, -1), + out_dtype=torch.bfloat16, + use_fast_accum=True, + ) # dx = y @ wT # absort w quant scale to y @@ -251,29 +355,33 @@ def torch_channel_quant_f_and_b(x, w, y): y_scale = ys.abs().float().amax(dim=1, keepdim=True) / 448.0 + 1e-9 yq = (ys / y_scale).to(torch.float8_e4m3fn) w_dummy_scale = torch.ones((1, K), dtype=torch.float32, device=x.device) - dx = torch._scaled_mm(yq, - wq.t().contiguous().t(), - scale_a=y_scale, - scale_b=w_dummy_scale, - out_dtype=torch.bfloat16, - use_fast_accum=True) + dx = torch._scaled_mm( + yq, + wq.t().contiguous().t(), + scale_a=y_scale, + scale_b=w_dummy_scale, + out_dtype=torch.bfloat16, + use_fast_accum=True, + ) # dw = yT@x yt = y.t().contiguous() yts = yt * x_scale.view(1, M) yt_scale = yts.abs().float().amax(dim=1, keepdim=True) / 448.0 + 1e-9 ytq = (yts / yt_scale).to(torch.float8_e4m3fn) - dw = torch._scaled_mm(ytq, - xq.t().contiguous().t(), - scale_a=yt_scale.view(-1, 1), - scale_b=w_dummy_scale, - out_dtype=torch.bfloat16, - use_fast_accum=True) + dw = torch._scaled_mm( + ytq, + xq.t().contiguous().t(), + scale_a=yt_scale.view(-1, 1), + scale_b=w_dummy_scale, + out_dtype=torch.bfloat16, + use_fast_accum=True, + ) return xq, wq, yq, ytq, o, dx, dw # smooth and token-wise/channel-wise -def torch_reuse_smooth_quant_f_and_b(x, w, y): +def torch_smooth_quant_f_and_b(x, w, y): x = x.clone() w = w.clone() y = y.clone() @@ -296,12 +404,14 @@ def torch_reuse_smooth_quant_f_and_b(x, w, y): xq = (x_smooth / x_quant_scale).to(torch.float8_e4m3fn) wq = (w_smooth / w_quant_scale).to(torch.float8_e4m3fn) - o = torch._scaled_mm(xq, - wq.t(), - scale_a=x_quant_scale.view(-1, 1), - scale_b=w_quant_scale.view(1, -1), - out_dtype=torch.bfloat16, - use_fast_accum=True) + o = torch._scaled_mm( + xq, + wq.t(), + scale_a=x_quant_scale.view(-1, 1), + scale_b=w_quant_scale.view(1, -1), + out_dtype=torch.bfloat16, + use_fast_accum=True, + ) # print(f'{x_smooth_scale=} {x_quant_scale[:,0]=} {w_quant_scale=}') @@ -310,24 +420,28 @@ def torch_reuse_smooth_quant_f_and_b(x, w, y): ys = y * w_quant_scale.view(1, N) y_scale = ys.abs().float().amax(dim=1, keepdim=True) / 448.0 + 1e-9 yq = (ys / y_scale).to(torch.float8_e4m3fn) - dx = torch._scaled_mm(yq, - wq.t().contiguous().t(), - scale_a=y_scale, - scale_b=w_smooth_scale.view(1, -1), - out_dtype=torch.bfloat16, - use_fast_accum=True) + dx = torch._scaled_mm( + yq, + wq.t().contiguous().t(), + scale_a=y_scale, + scale_b=w_smooth_scale.view(1, -1), + out_dtype=torch.bfloat16, + use_fast_accum=True, + ) # dw = yT@x yt = y.t().contiguous() # [N, M] yts = yt * x_quant_scale.view(1, M) yt_scale = yts.abs().amax(dim=1, keepdim=True) / 448.0 + 1e-9 ytq = (yts / yt_scale).to(torch.float8_e4m3fn) - dw = torch._scaled_mm(ytq, - xq.t().contiguous().t(), - scale_a=yt_scale.view(-1, 1), - scale_b=x_smooth_scale.view(1, -1), - out_dtype=torch.bfloat16, - use_fast_accum=True) + dw = torch._scaled_mm( + ytq, + xq.t().contiguous().t(), + scale_a=yt_scale.view(-1, 1), + scale_b=x_smooth_scale.view(1, -1), + out_dtype=torch.bfloat16, + use_fast_accum=True, + ) return xq, wq, yq, ytq, o, dx, dw @@ -359,60 +473,17 @@ def fp16_f_and_b(x, w, y): return o, dx, dw -def output_check(org_out, opt_out, mode='', rtol=None, atol=None): - assert org_out.shape == opt_out.shape, f"ref:{org_out.shape} != out:{opt_out.shape}" - dtype = org_out.dtype - assert opt_out.dtype == dtype or dtype == torch.float32, f"ref:{dtype} != out:{opt_out.dtype}" - if org_out.numel() == 0: - return - - if dtype != torch.float32: - org_out = org_out.float() - opt_out = opt_out.float() - if dtype == torch.float8_e4m3fn: - rtol = 0.1 - abs_error = (opt_out - org_out).abs().mean().item() - rel_error = abs_error / max(org_out.abs().mean().item(), 1e-38) - if rel_error >= 0.005: - rel_err_str = f"\033[91m {rel_error:.6f}\033[00m" - else: - rel_err_str = f"{rel_error:.6f}" - org_max = org_out.abs().max() - org_mean = org_out.abs().mean() - opt_max = opt_out.abs().max() - opt_mean = opt_out.abs().mean() - print(f'\n{mode:<16} rel:{rel_err_str} abs:{abs_error:.6f} ' \ - f'org:{org_max:.3f}/{org_mean:.3f} ' \ - f'opt:{opt_max:.3f}/{opt_mean:.3f} ') - if rtol is not None and atol is not None: - torch.testing.assert_close(opt_out, org_out, rtol=rtol, atol=atol) - - -def quant_check(org_out, xq, wq, opt_out, mode): - abs_error = (opt_out.float() - org_out.float()).abs().mean().item() - rel_error = abs_error / org_out.float().abs().mean().item() - x_underflow = (xq == 0.0).sum().item() / xq.numel() - w_underflow = (wq == 0.0).sum().item() / wq.numel() - x_overflow = (torch.isnan(xq)).sum().item() - w_overflow = (torch.isnan(wq)).sum().item() - print(f'\n{mode} rel:{rel_error:.3f} abs:{abs_error:.3f} ' \ - f'org:{org_out.abs().max():.3f}/{org_out.abs().mean():.3f} ' \ - f'opt:{opt_out.abs().max():.3f}/{opt_out.abs().mean():.3f} ' \ - f'x_underflow:{x_underflow:.5f} w_underflow:{w_underflow:.5f} ' \ - f'x_overflow:{x_overflow} w_overflow:{w_overflow}') - - def read_and_tile(filename, tile=True): - device = 'cuda:0' + device = "cuda:0" dtype = torch.bfloat16 d = torch.load(filename, weights_only=True) # x = d['x'][0].to(dtype).to(device) # w = d['w'].to(dtype).to(device) # y = d['y'][0].to(dtype).to(device) - x = d['x'] - y = d['y'] + x = d["x"] + y = d["y"] x = x.view(-1, x.size(2)).to(dtype).to(device) - w = d['w'].to(dtype).to(device) + w = d["w"].to(dtype).to(device) y = y.view(-1, y.size(2)).to(dtype).to(device) if tile: @@ -438,51 +509,122 @@ def read_and_tile(filename, tile=True): batch_size, in_dim = x.shape out_dim, in_dim = w.shape - print(f'\ndataset: {batch_size=} {in_dim=} {out_dim=} ' \ - f'x.max={x.abs().max().item():.3f} x.mean={x.abs().mean().item():.3f} ' \ - f'w.max={w.abs().max().item():.3f} w.mean={w.abs().mean().item():.3f} ' \ - f'y.max={y.abs().max().item():.3f} y.mean={y.abs().mean().item():.3f}') + print( + f"\ndataset: {batch_size=} {in_dim=} {out_dim=} " + f"x.max={x.abs().max().item():.3f} x.mean={x.abs().mean().item():.3f} " + f"w.max={w.abs().max().item():.3f} w.mean={w.abs().mean().item():.3f} " + f"y.max={y.abs().max().item():.3f} y.mean={y.abs().mean().item():.3f}" + ) return x, w, y def torch_fp16_vector_scaled_mm(x, weight, x_scale, weight_scale): - output = torch._scaled_mm(x, - weight, - scale_a=x_scale, - scale_b=weight_scale, - out_dtype=torch.bfloat16, - use_fast_accum=True) + output = torch._scaled_mm( + x, + weight, + scale_a=x_scale, + scale_b=weight_scale, + out_dtype=torch.bfloat16, + use_fast_accum=True, + ) return output -def torch_fp32_vector_scaled_mm(x, weight, x_scale, weight_scale, ones, - out=None): - output = torch._scaled_mm(x, - weight, - scale_a=ones, - scale_b=ones, - out_dtype=torch.float32, - use_fast_accum=True, - out=out) +def torch_fp32_vector_scaled_mm(x, weight, x_scale, weight_scale, ones, out=None): + output = torch._scaled_mm( + x, + weight, + scale_a=ones, + scale_b=ones, + out_dtype=torch.float32, + use_fast_accum=True, + out=out, + ) return output * x_scale * weight_scale def torch_fp16_scaler_scaled_mm(x, weight, x_scale, weight_scale): - output = torch._scaled_mm(x, - weight, - scale_a=x_scale, - scale_b=weight_scale, - out_dtype=torch.bfloat16, - use_fast_accum=True) + output = torch._scaled_mm( + x, + weight, + scale_a=x_scale, + scale_b=weight_scale, + out_dtype=torch.bfloat16, + use_fast_accum=True, + ) return output def torch_fp32_scaler_scaled_mm(x, weight, x_scale, weight_scale): - output = torch._scaled_mm(x, - weight, - scale_a=x_scale, - scale_b=weight_scale, - out_dtype=torch.float32, - use_fast_accum=True) + output = torch._scaled_mm( + x, + weight, + scale_a=x_scale, + scale_b=weight_scale, + out_dtype=torch.float32, + use_fast_accum=True, + ) return output + + +def torch_make_chunk_sort_map(num_global_tokens_per_local_expert: torch.Tensor): + device = num_global_tokens_per_local_expert.device + num_ranks, num_local_experts = num_global_tokens_per_local_expert.shape + + flat_sizes = num_global_tokens_per_local_expert.flatten() + src_chunk_starts = torch.cumsum(flat_sizes, dim=0) - flat_sizes + total_tokens = flat_sizes.sum().item() + + tokens_per_expert = num_global_tokens_per_local_expert.sum(dim=0) + padded_tokens_per_expert = (tokens_per_expert + 31) // 32 * 32 + expert_dst_starts = ( + torch.cumsum(padded_tokens_per_expert, dim=0) - padded_tokens_per_expert + ) + total_padded_size = padded_tokens_per_expert.sum().item() + + row_id_map = torch.zeros((total_padded_size,), dtype=torch.int32, device=device) + row_id_map_inverse = torch.empty((total_tokens,), dtype=torch.int32, device=device) + + for e_idx in range(num_local_experts): + current_dst_offset = expert_dst_starts[e_idx].item() + + for r_idx in range(num_ranks): + count = num_global_tokens_per_local_expert[r_idx, e_idx].item() + if count == 0: + continue + + src_idx_in_flat = r_idx * num_local_experts + e_idx + src_start = src_chunk_starts[src_idx_in_flat].item() + + src_range = torch.arange( + src_start, src_start + count, device=device, dtype=torch.int32 + ) + dst_range = torch.arange( + current_dst_offset, + current_dst_offset + count, + device=device, + dtype=torch.int32, + ) + + row_id_map[dst_range] = src_range + row_id_map_inverse[src_range] = dst_range + + current_dst_offset += count + + return row_id_map, row_id_map_inverse + + +@triton.jit +def print_kernel(x_ptr): + pid = tl.program_id(axis=0) + x = tl.load(x_ptr + tl.arange(0, 16)).to(tl.float32) + if pid == 0: + tl.device_print("x", x) + + +# the kernel is used to debug cuda graph tensor +def triton_print(x): + grid = lambda META: (1,) + print_kernel[grid](x, num_stages=1, num_warps=1) + return x diff --git a/linghe/utils/add.py b/linghe/utils/add.py index 868e1f0..6d774c9 100644 --- a/linghe/utils/add.py +++ b/linghe/utils/add.py @@ -3,49 +3,40 @@ Copyright (c) Ant Financial Service Group and its affiliates. """ +from typing import List + import torch -from typing import Iterable, Optional, Tuple import triton import triton.language as tl @triton.jit -def inplace_add_kernel(x_ptr, y_ptr, M, N, H: tl.constexpr, W: tl.constexpr, - EVEN: tl.constexpr, ACCUM: tl.constexpr): - rid = tl.program_id(axis=0) - cid = tl.program_id(axis=1) - offs = rid * H * N + cid * W + tl.arange(0, H)[:, None] * N + tl.arange(0, - W)[ - None, :] +def inplace_add_kernel( + x_ptr, y_ptr, N, B: tl.constexpr, EVEN: tl.constexpr, ACCUM: tl.constexpr +): + pid = tl.program_id(axis=0) + offs = pid * B + tl.arange(0, B) if ACCUM: if EVEN: - x = tl.load(x_ptr + offs) + x = tl.load(x_ptr + offs).to(tl.float32) y = tl.load(y_ptr + offs).to(tl.float32) tl.store(x_ptr + offs, x + y) else: - x = tl.load(x_ptr + offs, - mask=(cid * W + tl.arange(0, W)[None, :] < N) & ( - rid * H + tl.arange(0, H)[:, None] < M)) - y = tl.load(y_ptr + offs, - mask=(cid * W + tl.arange(0, W)[None, :] < N) & ( - rid * H + tl.arange(0, H)[:, None] < M)) - tl.store(x_ptr + offs, x + y, - mask=(cid * W + tl.arange(0, W)[:, None] < N) & ( - rid * H + tl.arange(0, H)[None, :] < M)) + mask = offs < N + x = tl.load(x_ptr + offs, mask=mask).to(tl.float32) + y = tl.load(y_ptr + offs, mask=mask).to(tl.float32) + tl.store(x_ptr + offs, x + y, mask=mask) else: if EVEN: - y = tl.load(y_ptr + offs).to(tl.float32) + y = tl.load(y_ptr + offs) tl.store(x_ptr + offs, y) else: - y = tl.load(y_ptr + offs, - mask=(cid * W + tl.arange(0, W)[None, :] < N) & ( - rid * H + tl.arange(0, H)[:, None] < M)) - tl.store(x_ptr + offs, y, - mask=(cid * W + tl.arange(0, W)[:, None] < N) & ( - rid * H + tl.arange(0, H)[None, :] < M)) + mask = offs < N + y = tl.load(y_ptr + offs, mask=mask) + tl.store(x_ptr + offs, y, mask=mask) -def triton_inplace_add(x: torch.Tensor, y: torch.Tensor, accum : bool = True): +def triton_inplace_add(x: torch.Tensor, y: torch.Tensor, accum: bool = True): """ inplace add y to x Args: @@ -56,23 +47,124 @@ def triton_inplace_add(x: torch.Tensor, y: torch.Tensor, accum : bool = True): Returns: updated x """ - N = x.shape[-1] - M = x.numel() // N - # M, N = x.shape - H = 128 - W = 128 - EVEN = M % H == 0 and N % W == 0 + assert x.is_contiguous() and y.is_contiguous() + N = x.numel() + B = 512 + EVEN = N % B == 0 num_stages = 2 - num_warps = 8 + num_warps = 4 - grid = (triton.cdiv(M, H), triton.cdiv(N, W)) + grid = (triton.cdiv(N, B),) inplace_add_kernel[grid]( - x, y, - M, N, - H, W, - EVEN, + x, y, N, B, EVEN, accum, num_stages=num_stages, num_warps=num_warps + ) + return x + + +@triton.jit +def batch_inplace_add_kernel( + x_ptrs, + y_ptrs, + size_ptr, + T, + B: tl.constexpr, + ACCUM: tl.constexpr, + XT: tl.constexpr, + YT: tl.constexpr, +): + tid = tl.program_id(axis=0) + bid = tl.program_id(axis=1) + size = tl.load(size_ptr + tid) + + if XT == 0: + x_ptr = tl.load(x_ptrs + tid).to(tl.pointer_type(tl.float32)) + elif XT == 1: + x_ptr = tl.load(x_ptrs + tid).to(tl.pointer_type(tl.bfloat16)) + else: + x_ptr = tl.load(x_ptrs + tid).to(tl.pointer_type(tl.float16)) + + if YT == 0: + y_ptr = tl.load(y_ptrs + tid).to(tl.pointer_type(tl.float32)) + elif YT == 1: + y_ptr = tl.load(y_ptrs + tid).to(tl.pointer_type(tl.bfloat16)) + else: + y_ptr = tl.load(y_ptrs + tid).to(tl.pointer_type(tl.float16)) + + t = tl.cdiv(size, B * T) + offs = bid * t * B + tl.arange(0, B) + + if ACCUM: + for i in range(t): + x = tl.load(x_ptr + offs, mask=offs < size).to(tl.float32) + y = tl.load(y_ptr + offs, mask=offs < size).to(tl.float32) + tl.store(x_ptr + offs, x + y, mask=offs < size) + offs += B + else: + for i in range(t): + y = tl.load(y_ptr + offs, mask=offs < size) + tl.store(x_ptr + offs, y, mask=offs < size) + offs += B + + +def triton_batch_inplace_add( + xs: List[torch.Tensor], ys: List[torch.Tensor], accum: bool = True +): + """ + inplace add y to x + Args: + xs: a list of Tensor + ys: a list of Tensor + accum: x += y if accum=True else x.copy_(y) + + Returns: + updated xs + """ + # assert all([x.is_contiguous() for x in xs]) + # assert all([y.is_contiguous() for y in ys]) + + device = xs[0].device + sizes = torch.tensor([x.numel() for x in xs], dtype=torch.int64).cuda( + device, non_blocking=True + ) + x_ptrs = torch.tensor([x.data_ptr() for x in xs], dtype=torch.int64).cuda( + device, non_blocking=True + ) + y_ptrs = torch.tensor([y.data_ptr() for y in ys], dtype=torch.int64).cuda( + device, non_blocking=True + ) + x_dtype = xs[0].dtype + assert x_dtype in (torch.bfloat16, torch.float32, torch.float16) + if x_dtype == torch.float32: + XT = 0 + elif x_dtype == torch.bfloat16: + XT = 1 + else: + XT = 2 + y_dtype = ys[0].dtype + assert y_dtype in (torch.bfloat16, torch.float32, torch.float16) + if y_dtype == torch.float32: + YT = 0 + elif y_dtype == torch.bfloat16: + YT = 1 + else: + YT = 2 + + T = 512 + B = 1024 + num_stages = 3 + num_warps = 4 + + grid = (len(xs), T) + batch_inplace_add_kernel[grid]( + x_ptrs, + y_ptrs, + sizes, + T, + B, accum, + XT, + YT, num_stages=num_stages, - num_warps=num_warps + num_warps=num_warps, ) - return x + return xs diff --git a/linghe/utils/dot.py b/linghe/utils/dot.py deleted file mode 100644 index f57acab..0000000 --- a/linghe/utils/dot.py +++ /dev/null @@ -1,57 +0,0 @@ -# -*- coding: utf-8 -*- -""" -Copyright (c) Ant Financial Service Group and its affiliates. -""" - -import torch -import triton -import triton.language as tl - - -@triton.jit -def dot_kernel(x_ptr, y_ptr, sum_ptr, M, N, H: tl.constexpr, W: tl.constexpr): - # rowwise read, rowwise write - pid = tl.program_id(axis=0) - offs = pid * W * N + tl.arange(0, W)[:, None] * N + tl.arange(0, H)[None, :] - - n = tl.cdiv(N, H) - sums = tl.zeros((W,), dtype=tl.float32) - for i in range(n): - x = tl.load(x_ptr + offs).to(tl.float32) - y = tl.load(y_ptr + offs).to(tl.float32) - sums += tl.sum(x * y, axis=1) - offs += H - - tl.store(sum_ptr + pid * W + tl.arange(0, W), sums) - - -def triton_dot(x, y): - """ - vector dot multiply, output = sum(x*y, 1), - it is used to calculate gradient of router weight - Args: - x: - y: - - Returns: - output of sum(x*y, 1) - """ - M, N = x.shape - H = 128 - W = 16 - assert M % W == 0 - - num_stages = 5 - num_warps = 8 - device = x.device - s = torch.empty((M,), device=device, dtype=x.dtype) - grid = (triton.cdiv(M, W),) - dot_kernel[grid]( - x, y, s, - M, N, - H, W, - num_stages=num_stages, - num_warps=num_warps - ) - return s - diff --git a/linghe/utils/emb.py b/linghe/utils/emb.py new file mode 100644 index 0000000..756ddcd --- /dev/null +++ b/linghe/utils/emb.py @@ -0,0 +1,479 @@ +# -*- coding: utf-8 -*- +""" +Copyright (c) Ant Financial Service Group and its affiliates. +""" + +import torch +import triton +import triton.language as tl +import math + + +@triton.jit +def embedding_forward_kernel(x_ptr, y_ptr, w_ptr, dim, DIM: tl.constexpr): + pid = tl.program_id(axis=0).to(tl.int64) + index = tl.load(x_ptr + pid) + weight_ptr = w_ptr.to(tl.pointer_type(tl.bfloat16)) + + w = tl.load( + weight_ptr + index * dim + tl.arange(0, DIM), mask=tl.arange(0, DIM) < dim + ) + tl.store(y_ptr + pid * dim + tl.arange(0, DIM), w, mask=tl.arange(0, DIM) < dim) + + +def triton_embedding_forward(x, w_ptr, dim=4096, dtype=torch.bfloat16): + """ + inplace add y to x + Args: + x: input ids Tensor + w_ptr: data_ptr of embedding weight + Returns: + embedding output + """ + assert x.is_contiguous() + assert dtype == torch.bfloat16 + M = x.numel() + y = torch.empty((x.shape + (dim,)), device=x.device, dtype=dtype) + DIM = triton.next_power_of_2(dim) + num_stages = 2 + num_warps = 8 + + grid = (M,) + embedding_forward_kernel[grid]( + x, y, w_ptr, dim, DIM, num_stages=num_stages, num_warps=num_warps + ) + return y + + +@triton.jit +def atomic_embedding_backward_kernel( + y_ptr, x_ptr, g_ptr, stride_0, stride_1, dim, DIM: tl.constexpr, T: tl.constexpr +): + bid = tl.program_id(axis=0).to(tl.int64) + lid = tl.program_id(axis=1) + B = tl.num_programs(0) + L = tl.num_programs(1) + index = tl.load(x_ptr + bid * L + lid) + + if T == 0: + grad_ptr = g_ptr.to(tl.pointer_type(tl.float32)) + elif T == 1: + grad_ptr = g_ptr.to(tl.pointer_type(tl.bfloat16)) + else: + grad_ptr = g_ptr.to(tl.pointer_type(tl.float16)) + + y = tl.load( + y_ptr + bid * stride_0 + lid * stride_1 + tl.arange(0, DIM), + mask=tl.arange(0, DIM) < dim, + ) + tl.atomic_add( + grad_ptr + index * dim + tl.arange(0, DIM), y, mask=tl.arange(0, DIM) < dim + ) + + +def triton_atomic_embedding_backward(y, x, g_ptr, dtype=torch.bfloat16): + """ + inplace update embedding weight gradient + Args: + y: gradient of output + x: input ids Tensor + g_ptr: data_ptr of embedding weight gradient + Returns: + None + """ + assert dtype in (torch.float32, torch.bfloat16, torch.float16) + shape = x.shape + assert len(shape) == 2 + if dtype == torch.float32: + T = 0 + elif dtype == torch.bfloat16: + T = 1 + else: + T = 2 + B, L, dim = y.shape + stride_0 = y.stride(0) + stride_1 = y.stride(1) + + DIM = triton.next_power_of_2(dim) + num_stages = 2 + num_warps = 16 + + grid = (B, L) + atomic_embedding_backward_kernel[grid]( + y, + x, + g_ptr, + stride_0, + stride_1, + dim, + DIM, + T, + num_stages=num_stages, + num_warps=num_warps, + ) + + +@triton.jit +def sync_embedding_backward_kernel( + grad_output_ptr, + unique_ids_ptr, + sorted_indices_ptr, + accum_counts_ptr, + g_ptr, + stride_0, + stride_1, + dim, + B, + L, + DIM: tl.constexpr, + T: tl.constexpr, +): + pid = tl.program_id(axis=0).to(tl.int64) + + if pid == 0: + c0 = 0 + c0 = c0.to(tl.int64) + c1 = tl.load(accum_counts_ptr) + else: + c01 = tl.load(accum_counts_ptr + pid - 1 + tl.arange(0, 2)) + c0, c1 = tl.split(c01) + count = c1 - c0 + input_id = tl.load(unique_ids_ptr + pid).to(tl.int64) + + if T == 0: + grad_ptr = g_ptr.to(tl.pointer_type(tl.float32)) + elif T == 1: + grad_ptr = g_ptr.to(tl.pointer_type(tl.bfloat16)) + else: + grad_ptr = g_ptr.to(tl.pointer_type(tl.float16)) + + outputs = tl.zeros((DIM,), dtype=tl.float32) + + for i in range(count): + pos = tl.load(sorted_indices_ptr + c0 + i) + bid = pos // L + lid = pos % L + g = tl.load( + grad_output_ptr + bid * stride_0 + lid * stride_1 + tl.arange(0, DIM), + mask=tl.arange(0, DIM) < dim, + ).to(tl.float32) + outputs += g + tl.store( + grad_ptr + input_id * dim + tl.arange(0, DIM), + outputs, + mask=tl.arange(0, DIM) < dim, + ) + + +def triton_sync_embedding_backward(grad_output, x, g_ptr, dtype=torch.bfloat16): + """ + inplace update embedding weight gradient + Args: + y: gradient of output + x: input ids Tensor + g_ptr: data_ptr of embedding weight gradient + Returns: + None + """ + assert dtype in (torch.float32, torch.bfloat16, torch.float16) + if dtype == torch.float32: + T = 0 + elif dtype == torch.bfloat16: + T = 1 + else: + T = 2 + shape = x.shape + assert len(shape) == 2 + B, L, dim = grad_output.shape + stride_0 = grad_output.stride(0) + stride_1 = grad_output.stride(1) + + sorted_ids, sorted_indices = torch.sort(x.view(-1), stable=False) + unique_ids, unique_counts = torch.unique_consecutive(sorted_ids, return_counts=True) + accum_counts = torch.cumsum(unique_counts, 0) + DIM = triton.next_power_of_2(dim) + num_stages = 3 + num_warps = 4 + + grid = (unique_ids.size(0),) + sync_embedding_backward_kernel[grid]( + grad_output, + unique_ids, + sorted_indices, + accum_counts, + g_ptr, + stride_0, + stride_1, + dim, + B, + L, + DIM, + T, + num_stages=num_stages, + num_warps=num_warps, + ) + + +@triton.jit +def scan_and_count_split_kernel( + id_ptr, counts_ptr, unique_id_ptr, unique_count_ptr, L, B: tl.constexpr +): + bid = tl.program_id(axis=0) + sid = tl.program_id(axis=1) + ns = tl.num_programs(1) + + offsets = sid * B + tl.arange(0, B) + + ids = tl.load(id_ptr + bid * L + offsets, mask=offsets < L, other=2**30) + + write_index = bid * L + sid * B + unique_count = 0 + stop = False + while not stop: + min_id = tl.min(ids) + if min_id == 2**30: + stop = True + else: + count = tl.sum(tl.where(ids == min_id, 1, 0)) + ids = tl.where(ids <= min_id, 2**30, ids) + tl.store(counts_ptr + write_index, count) + tl.store(unique_id_ptr + write_index, min_id) + write_index += 1 + unique_count += 1 + tl.store(unique_count_ptr + bid * ns + sid, unique_count) + + +@triton.jit +def scan_and_count_merge_kernel( + counts_ptr, + unique_id_ptr, + unique_count_ptr, + accum_counts_ptr, + L, + B: tl.constexpr, + T: tl.constexpr, +): + bid = tl.program_id(axis=0) + write_index = bid * (L + 1) + 1 + tl.store(accum_counts_ptr + bid * (L + 1), 0) + pre_id = -1 + for i in range(T): + uc = tl.load(unique_count_ptr + bid * T + i) + counts = tl.load( + counts_ptr + bid * L + i * B + tl.arange(0, B), mask=tl.arange(0, B) < uc + ) + uids = tl.load( + unique_id_ptr + bid * L + i * B + tl.arange(0, B), + mask=tl.arange(0, B) < uc, + other=2**30, + ) + min_id = tl.min(uids) + offset = tl.where(min_id == pre_id, -1, 0) + pre_id = tl.max(tl.where(tl.arange(0, B) < uc, uids, -1)) + tl.atomic_add( + accum_counts_ptr + write_index + offset + tl.arange(0, B), + counts, + mask=tl.arange(0, B) < uc, + ) + write_index += uc + offset + + +def triton_scan_and_count(ids): + assert ids.is_contiguous() + shape = ids.shape + device = ids.device + assert len(shape) in (1, 2) + + BLOCK = 128 + if len(shape) == 2: + B, L = ids.shape + else: + B, L = 1, shape[0] + + T = math.ceil(L / BLOCK) + counts = torch.empty((B, L), dtype=torch.int32, device=device) + unique_ids = torch.empty((B, L), dtype=torch.int32, device=device) + unique_counts = torch.empty((B, T), dtype=torch.int32, device=device) + accum_counts = torch.zeros((B, L + 1), dtype=torch.int32, device=device) + + num_stages = 3 + num_warps = 4 + grid = (B, T) + scan_and_count_split_kernel[grid]( + ids, + counts, + unique_ids, + unique_counts, + L, + BLOCK, + num_stages=num_stages, + num_warps=num_warps, + ) + + num_stages = 3 + num_warps = 4 + grid = (B,) + scan_and_count_merge_kernel[grid]( + counts, + unique_ids, + unique_counts, + accum_counts, + L, + BLOCK, + T, + num_stages=num_stages, + num_warps=num_warps, + ) + accum_counts = torch.cumsum(accum_counts, -1) + + if len(shape) == 1: + accum_counts = accum_counts.squeeze(0) + + return accum_counts + + +@triton.jit +def deprecated_scan_and_count_kernel( + id_ptr, accum_counts_ptr, B: tl.constexpr, T: tl.constexpr +): + accum = 0 + write_index = 0 + last_min_id = -1 + for i in range(T): + ids = tl.load(id_ptr + i * B + tl.arange(0, B)) + stop = False + while not stop: + min_id = tl.min(ids) + if min_id == 2**30: + stop = True + else: + if min_id != last_min_id: + tl.store(accum_counts_ptr + write_index, accum) + last_min_id = min_id + write_index += 1 + count = tl.sum(tl.where(ids == min_id, 1, 0)) + ids = tl.where(ids <= min_id, 2**30, ids) + accum += count + tl.store(accum_counts_ptr + write_index, accum) + + +def triton_deprecated_scan_and_count(ids): + M = ids.numel() + accum_counts = -torch.ones((M + 1,), dtype=torch.int32, device=ids.device) + num_stages = 3 + num_warps = 1 + + B = 64 + assert M % B == 0 + T = M // B + grid = (1,) + deprecated_scan_and_count_kernel[grid]( + ids, accum_counts, B, T, num_stages=num_stages, num_warps=num_warps + ) + return accum_counts + + +@triton.jit +def embedding_backward_kernel( + grad_output_ptr, + sorted_ids_ptr, + sorted_indices_ptr, + accum_counts_ptr, + g_ptr, + stride_0, + stride_1, + dim, + B, + L, + BLOCK: tl.constexpr, + T: tl.constexpr, +): + pid = tl.program_id(axis=0).to(tl.int64) + cid = tl.program_id(axis=1) + c01 = tl.load(accum_counts_ptr + pid + tl.arange(0, 2)) + c0, c1 = tl.split(c01) + if c0 == c1: + return + + count = c1 - c0 + input_id = tl.load(sorted_ids_ptr + c0).to(tl.int64) + + if T == 0: + grad_ptr = g_ptr.to(tl.pointer_type(tl.float32)) + elif T == 1: + grad_ptr = g_ptr.to(tl.pointer_type(tl.bfloat16)) + else: + grad_ptr = g_ptr.to(tl.pointer_type(tl.float16)) + + mask = cid * BLOCK + tl.arange(0, BLOCK) < dim + outputs = tl.load( + grad_ptr + input_id * dim + cid * BLOCK + tl.arange(0, BLOCK), mask=mask + ).to(tl.float32) + + for i in range(count): + pos = tl.load(sorted_indices_ptr + c0 + i) + bid = pos // L + lid = pos % L + g = tl.load( + grad_output_ptr + + bid * stride_0 + + lid * stride_1 + + cid * BLOCK + + tl.arange(0, BLOCK), + mask=mask, + ).to(tl.float32) + outputs += g + tl.store( + grad_ptr + input_id * dim + cid * BLOCK + tl.arange(0, BLOCK), + outputs, + mask=mask, + ) + + +def triton_embedding_backward(grad_output, x, g_ptr, dtype=torch.bfloat16): + """ + inplace update embedding weight gradient + Args: + y: gradient of output + x: input ids Tensor + g_ptr: data_ptr of embedding weight gradient + Returns: + None + """ + assert dtype in (torch.bfloat16, torch.float32, torch.float16) + if dtype == torch.float32: + T = 0 + elif dtype == torch.bfloat16: + T = 1 + else: + T = 2 + shape = x.shape + assert len(shape) == 2 + B, L, dim = grad_output.shape + stride_0 = grad_output.stride(0) + stride_1 = grad_output.stride(1) + + sorted_ids, sorted_indices = torch.sort(x.view(-1), stable=False) + accum_counts = triton_scan_and_count(sorted_ids) + BLOCK = 512 + assert dim % BLOCK == 0 + num_stages = 3 + num_warps = 2 + + grid = (B * L, dim // BLOCK) + embedding_backward_kernel[grid]( + grad_output, + sorted_ids, + sorted_indices, + accum_counts, + g_ptr, + stride_0, + stride_1, + dim, + B, + L, + BLOCK, + T, + num_stages=num_stages, + num_warps=num_warps, + ) diff --git a/linghe/utils/gate.py b/linghe/utils/gate.py new file mode 100644 index 0000000..aaada50 --- /dev/null +++ b/linghe/utils/gate.py @@ -0,0 +1,639 @@ +import torch +import triton +import triton.language as tl +import math + + +@triton.jit +def group_rms_norm_gate_forward_kernel( + x_ptr, + gate_ptr, + weight_ptr, + out_ptr, + stride_g, + eps, + bs, + length, + DIM: tl.constexpr, + d: tl.constexpr, + D: tl.constexpr, + GROUP_SIZE: tl.constexpr, + SHARE: tl.constexpr, + NATIVE: tl.constexpr, +): + pid = tl.program_id(axis=0) + bid = pid // length + sid = pid % length + + if SHARE: + weight = tl.load(weight_ptr + tl.arange(0, D), mask=tl.arange(0, D) < d)[ + None, : + ] + else: + weight = tl.load( + weight_ptr + tl.arange(0, GROUP_SIZE)[:, None] * d + tl.arange(0, D), + mask=tl.arange(0, D)[None, :] < d, + ) + + if NATIVE: + x_offs = ( + pid * DIM + tl.arange(0, GROUP_SIZE)[:, None] * d + tl.arange(0, D)[None, :] + ) + else: + x_offs = ( + sid * bs * DIM + + bid * DIM + + tl.arange(0, GROUP_SIZE)[:, None] * d + + tl.arange(0, D)[None, :] + ) + + x_offs_mask = tl.arange(0, D)[None, :] < d + x = tl.load(x_ptr + x_offs, mask=x_offs_mask).to(tl.float32) + + g_offs = ( + sid * bs * stride_g + + bid * stride_g + + tl.arange(0, GROUP_SIZE)[:, None] * d + + tl.arange(0, D)[None, :] + ) + + g = tl.load(gate_ptr + g_offs, mask=tl.arange(0, D)[None, :] < d).to(tl.float32) + + rms = tl.rsqrt(tl.sum(x * x, axis=1) / d + eps) + + x = (x * rms[:, None]) * weight * tl.sigmoid(g) + + g_offs = ( + sid * bs * DIM + + bid * DIM + + tl.arange(0, GROUP_SIZE)[:, None] * d + + tl.arange(0, D)[None, :] + ) + tl.store(out_ptr + g_offs, x, mask=tl.arange(0, D)[None, :] < d) + + +def triton_group_rms_norm_gate_forward( + x: torch.Tensor, gate: torch.Tensor, weight: torch.Tensor, eps=1e-6, group_size=4 +): + """ + norm and gate in linear attention + Args: + x: output of attn, [bs, length, n_heads, head_dim] + gate: gate tensor, [length, bs, dim] if transpose=True else [bs, length, dim] + weight: rms norm weight, [dim] + eps: epsilon of rms norm + group_size: group size of group rms norm + layout: layout of x, should in {'bsd', 'sbd} + + Returns: + output tensor, [length, bs, dim] + """ + length, bs, dim = gate.shape + + assert dim <= 8192 and triton.next_power_of_2(group_size) == group_size + assert x.is_contiguous() and weight.is_contiguous() + assert gate.stride(2) == 1 and gate.stride(0) == gate.stride(1) * bs + assert length != bs + wd = weight.shape[0] + SHARE = wd != dim # all groups share the same weight + NATIVE = x.size(0) == bs + d = dim // group_size + device = x.device + + D = triton.next_power_of_2(d) + + out = torch.empty((length, bs, dim), device=device, dtype=x.dtype) + + grid = (bs * length,) + group_rms_norm_gate_forward_kernel[grid]( + x, + gate, + weight, + out, + gate.stride(1), + eps, + bs, + length, + dim, + d, + D, + group_size, + SHARE, + NATIVE, + num_stages=3, + num_warps=4, + ) + return out + + +@triton.jit +def group_rms_norm_gate_backward_kernel( + grad_output_ptr, + x_ptr, + gate_ptr, + w_ptr, + dx_ptr, + dg_ptr, + dw_ptr, + eps, + bs, + length, + stride_g, + DIM: tl.constexpr, + d: tl.constexpr, + D: tl.constexpr, + GROUP_SIZE: tl.constexpr, + T: tl.constexpr, + SHARE: tl.constexpr, + NATIVE: tl.constexpr, +): + pid = tl.program_id(0) + bid = pid * T // length + sid = pid * T % length + + if SHARE: + w = tl.load(w_ptr + tl.arange(0, d), mask=tl.arange(0, D) < D)[None, :] + else: + w = tl.load( + w_ptr + tl.arange(0, GROUP_SIZE)[:, None] * d + tl.arange(0, D), + mask=tl.arange(0, D)[None, :] < d, + ) + + if NATIVE: + x_offs = ( + pid * DIM * T + + tl.arange(0, GROUP_SIZE)[:, None] * d + + tl.arange(0, D)[None, :] + ) + else: + x_offs = ( + sid * bs * DIM + + bid * DIM + + tl.arange(0, GROUP_SIZE)[:, None] * d + + tl.arange(0, D)[None, :] + ) + + x_offs_mask = tl.arange(0, D)[None, :] < d + offs = ( + sid * bs * DIM + + bid * DIM + + tl.arange(0, GROUP_SIZE)[:, None] * d + + tl.arange(0, D)[None, :] + ) + offs_mask = tl.arange(0, D)[None, :] < d + g_offs = ( + sid * bs * stride_g + + bid * stride_g + + tl.arange(0, GROUP_SIZE)[:, None] * d + + tl.arange(0, D)[None, :] + ) + + dw = tl.zeros((GROUP_SIZE, D), dtype=tl.float32) + for i in range(T): + x = tl.load(x_ptr + x_offs, mask=x_offs_mask).to(tl.float32) + g = tl.load(grad_output_ptr + offs, mask=offs_mask).to(tl.float32) + gate = tl.load(gate_ptr + g_offs, mask=offs_mask).to(tl.float32) + gate = tl.sigmoid(gate) + r = tl.rsqrt(tl.sum(x * x, 1) / d + eps)[:, None] + w_grad = x * g * r * gate + dw += w_grad + + dx = ( + r * g * w * gate + - r * r * r * x * tl.sum(x * g * w * gate, 1, keep_dims=True) / d + ) + + tl.store(dx_ptr + x_offs, dx, mask=x_offs_mask) + + dg = x * r * w * g * gate * (1 - gate) + tl.store(dg_ptr + offs, dg, mask=offs_mask) + + if NATIVE: + x_offs += DIM + else: + x_offs += DIM * bs + offs += DIM * bs + g_offs += bs * stride_g + + if SHARE: + dw = tl.sum(dw, 0) + tl.store(dw_ptr + pid * d + tl.arange(0, d), dw, mask=tl.arange(0, D) < d) + else: + tl.store( + dw_ptr + + pid * DIM + + tl.arange(0, GROUP_SIZE)[:, None] * d + + tl.arange(0, D)[None, :], + dw, + mask=tl.arange(0, D)[None, :] < d, + ) + + +def triton_group_rms_norm_gate_backward( + grad_output, x, gate, weight, eps=1e-6, group_size=4 +): + length, bs, dim = gate.shape + assert dim <= 8192 and triton.next_power_of_2(group_size) == group_size + assert grad_output.is_contiguous() + assert length != bs + d = dim // group_size + wd = weight.shape[0] + SHARE = wd != dim # all groups share the same weight + NATIVE = x.size(0) == bs + + device = x.device + dx = torch.empty_like(x) + dg = torch.empty_like(gate) + + T = 8 + g = (bs * length) // T + if SHARE: + tmp_dw = torch.empty(g, d, dtype=torch.float32, device=device) + else: + tmp_dw = torch.empty(g, dim, dtype=torch.float32, device=device) + + D = triton.next_power_of_2(d) + grid = (g,) + group_rms_norm_gate_backward_kernel[grid]( + grad_output, + x, + gate, + weight, + dx, + dg, + tmp_dw, + eps, + bs, + length, + gate.stride(1), + dim, + d, + D, + group_size, + T, + SHARE, + NATIVE, + num_stages=3, + num_warps=8, + ) + dw = tmp_dw.sum(dim=0) + return dx, dg, dw + + +@triton.jit +def group_rms_norm_gate_and_mxfp8_quant_forward_kernel( + x_ptr, + gate_ptr, + weight_ptr, + x_q_ptr, + x_s_ptr, + xt_q_ptr, + xt_s_ptr, + stride_g, + eps, + bs, + length, + m, # bs * length + DIM: tl.constexpr, + d: tl.constexpr, # TODO: add support for d not power of 2d, now d == D + D: tl.constexpr, + GROUP_SIZE: tl.constexpr, + SB: tl.constexpr, + SHARE: tl.constexpr, + NATIVE: tl.constexpr, # True → x layout is [bs, length, DIM] or x layout is [length, bs, DIM] + OUTPUT_MODE: tl.constexpr, +): + rid = tl.program_id(axis=0) + gid = tl.program_id(axis=1) + + out_rows = rid * 32 + tl.arange(0, 32) + row_mask = out_rows < m + col_mask = tl.arange(0, D)[None, :] < d + + if SHARE: + w_off = tl.arange(0, D) + else: + w_off = gid * d + tl.arange(0, D) + weight = tl.load(weight_ptr + w_off, mask=tl.arange(0, D) < d) + + if NATIVE: + sids = out_rows // bs + bids = out_rows % bs + x_row_base = bids * length * DIM + sids * DIM + else: + x_row_base = out_rows * DIM + + x_offs = x_row_base[:, None] + gid * d + tl.arange(0, D)[None, :] # [32, D] + x = tl.load(x_ptr + x_offs, mask=row_mask[:, None] & col_mask).to(tl.float32) + + g_offs = out_rows[:, None] * stride_g + gid * d + tl.arange(0, D)[None, :] + g = tl.load(gate_ptr + g_offs, mask=row_mask[:, None] & col_mask).to(tl.float32) + + rms = tl.rsqrt(tl.sum(x * x, axis=1) / d + eps) # [32] current group norm + x = x * rms[:, None] * weight[None, :] * tl.sigmoid(g) # [32, D] + + if OUTPUT_MODE % 2 == 0: + x_blocks = tl.reshape(x, [32, SB, 32]) + scale_row = tl.maximum(tl.max(tl.abs(x_blocks), 2) / 448.0, 1e-30) # [32, SB] + log_scale_row = tl.ceil(tl.log2(scale_row)) + x_q = tl.reshape(x_blocks / tl.exp2(log_scale_row)[:, :, None], [32, D]) + + tl.store( + x_q_ptr + out_rows[:, None] * DIM + gid * d + tl.arange(0, D)[None, :], + x_q.to(x_q_ptr.dtype.element_ty), + mask=row_mask[:, None] & col_mask, + ) + + tl.store( + x_s_ptr + + out_rows[:, None] * (GROUP_SIZE * SB) + + gid * SB + + tl.arange(0, SB)[None, :], + (log_scale_row + 127).to(tl.uint8), + mask=row_mask[:, None], + ) + + if OUTPUT_MODE > 0: + scale_col = tl.maximum(tl.max(tl.abs(x), 0) / 448.0, 1e-30) + log_scale_col = tl.ceil(tl.log2(scale_col)) + xt_q = x / tl.exp2(log_scale_col)[None, :] + + tl.store( + xt_q_ptr + out_rows[:, None] * DIM + gid * d + tl.arange(0, D)[None, :], + xt_q.to(xt_q_ptr.dtype.element_ty), + mask=row_mask[:, None] & col_mask, + ) + + tl.store( + xt_s_ptr + rid * DIM + gid * d + tl.arange(0, D), + (log_scale_col + 127).to(tl.uint8), + mask=tl.arange(0, D) < d, + ) + + +def triton_group_rms_norm_gate_and_mxfp8_quant_forward( + x: torch.Tensor, + gate: torch.Tensor, + weight: torch.Tensor, + eps: float = 1e-6, + group_size: int = 4, + output_mode: int = 2, +): + """ + Fused group RMSnorm + sigmoid-gate + MXFP8 quantization. + + """ + length, bs, dim = gate.shape + m = length * bs + M = (m + 127) // 128 * 128 + + d = dim // group_size + D = triton.next_power_of_2(d) + assert D == d, f"d = dim // group_size = {d} only support D==d now" + assert d >= 32 and d % 32 == 0, f"d = {d} must be >= 32 and divisible by 32" + assert dim <= 8192 and triton.next_power_of_2(group_size) == group_size + assert x.is_contiguous() and weight.is_contiguous() + assert gate.stride(2) == 1 and gate.stride(0) == gate.stride(1) * bs + + SB = d // 32 # D // 32 + SHARE = weight.shape[0] != dim + NATIVE = x.size(0) == bs + device = x.device + + x_q = torch.empty((m, dim), device=device, dtype=torch.float8_e4m3fn) + x_s = torch.empty((M, dim // 32), device=device, dtype=torch.uint8) + xt_q = torch.empty((m, dim), device=device, dtype=torch.float8_e4m3fn) + xt_s = torch.empty((M // 32, dim), device=device, dtype=torch.uint8) + + grid = (M // 32, group_size) + group_rms_norm_gate_and_mxfp8_quant_forward_kernel[grid]( + x, + gate, + weight, + x_q, + x_s, + xt_q, + xt_s, + gate.stride(1), + eps, + bs, + length, + m, + dim, + d, + D, + group_size, + SB, + SHARE, + NATIVE, + output_mode, + num_stages=3, + num_warps=4, + ) + return x_q, x_s, xt_q, xt_s + + +@triton.jit +def group_rms_norm_gate_and_mxfp8_quant_backward_kernel( + grad_output_ptr, + x_ptr, + gate_ptr, # [length, bs, DIM] + w_ptr, + dx_ptr, + dg_q_ptr, + dg_s_ptr, + dgt_q_ptr, + dgt_s_ptr, + tmp_dw_ptr, # [num_row_blocks * GROUP_SIZE, D] fp32 + stride_g, # gate.stride(1) + eps, + bs, + length, + m, # = bs * length + DIM: tl.constexpr, + d: tl.constexpr, # d = DIM // GROUP_SIZE + D: tl.constexpr, # next_power_of_2(d) + GROUP_SIZE: tl.constexpr, + SB: tl.constexpr, # d // 32 + T: tl.constexpr, + SHARE: tl.constexpr, + NATIVE: tl.constexpr, # True → x layout is [bs, length, DIM] or x layout is [length, bs, DIM] + OUTPUT_MODE: tl.constexpr, +): + rid_base = tl.program_id(axis=0) # eache block covers T*32 rows + gid = tl.program_id(axis=1) + + if SHARE: + w = tl.load(w_ptr + tl.arange(0, D), mask=tl.arange(0, D) < d) + else: + w = tl.load(w_ptr + gid * d + tl.arange(0, D), mask=tl.arange(0, D) < d) + + dw = tl.zeros([D], dtype=tl.float32) + + for t in tl.static_range(T): + chunk_rid = rid_base * T + t + out_rows = chunk_rid * 32 + tl.arange(0, 32) + row_mask = out_rows < m + col_mask = tl.arange(0, D)[None, :] < d + + if NATIVE: + sids = out_rows // bs + bids = out_rows % bs + x_row_base = bids * length * DIM + sids * DIM + else: + x_row_base = out_rows * DIM + x_offs = x_row_base[:, None] + gid * d + tl.arange(0, D)[None, :] + x = tl.load(x_ptr + x_offs, mask=row_mask[:, None] & col_mask).to(tl.float32) + + base = out_rows[:, None] * DIM + gid * d + tl.arange(0, D)[None, :] + go = tl.load(grad_output_ptr + base, mask=row_mask[:, None] & col_mask).to( + tl.float32 + ) + gv = tl.load( + gate_ptr + + out_rows[:, None] * stride_g + + gid * d + + tl.arange(0, D)[None, :], + mask=row_mask[:, None] & col_mask, + ).to(tl.float32) + gs = tl.sigmoid(gv) # [32, D] + + r = tl.rsqrt(tl.sum(x * x, 1) / d + eps) + r3 = r * r * r + + dw += tl.sum(x * go * r[:, None] * gs, 0) # [D] + + xgwgs = x * go * w[None, :] * gs + dx = ( + r[:, None] * go * w[None, :] * gs + - r3[:, None] * x * tl.sum(xgwgs, 1, keep_dims=True) / d + ) + + tl.store(dx_ptr + x_offs, dx.to(tl.bfloat16), mask=row_mask[:, None] & col_mask) + + dg = x * r[:, None] * w[None, :] * go * gs * (1.0 - gs) + dg = tl.where(row_mask[:, None] & col_mask, dg, 0.0) + + if OUTPUT_MODE % 2 == 0: + dg_b = tl.reshape(dg, [32, SB, 32]) + s_row = tl.maximum(tl.max(tl.abs(dg_b), 2) / 448.0, 1e-30) + ls_r = tl.ceil(tl.log2(s_row)) + dg_qr = tl.reshape(dg_b / tl.exp2(ls_r)[:, :, None], [32, D]) + + tl.store( + dg_q_ptr + out_rows[:, None] * DIM + gid * d + tl.arange(0, D)[None, :], + dg_qr.to(dg_q_ptr.dtype.element_ty), + mask=row_mask[:, None] & col_mask, + ) + tl.store( + dg_s_ptr + + out_rows[:, None] * (GROUP_SIZE * SB) + + gid * SB + + tl.arange(0, SB)[None, :], + (ls_r + 127).to(tl.uint8), + mask=row_mask[:, None], + ) + + if OUTPUT_MODE > 0: + s_col = tl.maximum(tl.max(tl.abs(dg), 0) / 448.0, 1e-30) + ls_c = tl.ceil(tl.log2(s_col)) + dg_qc = dg / tl.exp2(ls_c)[None, :] + + tl.store( + dgt_q_ptr + + out_rows[:, None] * DIM + + gid * d + + tl.arange(0, D)[None, :], + dg_qc.to(dgt_q_ptr.dtype.element_ty), + mask=row_mask[:, None] & col_mask, + ) + tl.store( + dgt_s_ptr + chunk_rid * DIM + gid * d + tl.arange(0, D), + (ls_c + 127).to(tl.uint8), + mask=tl.arange(0, D) < d, + ) + + tmp_row = rid_base * GROUP_SIZE + gid + tl.store(tmp_dw_ptr + tmp_row * D + tl.arange(0, D), dw, mask=tl.arange(0, D) < d) + + +def triton_group_rms_norm_gate_and_mxfp8_quant_backward( + grad_output: torch.Tensor, + x: torch.Tensor, + gate: torch.Tensor, + weight: torch.Tensor, + eps: float = 1e-6, + group_size: int = 4, + output_mode: int = 2, +): + """ + Fused backward of group RMSnorm + sigmoid-gate + MXFP8 quantization. + """ + length, bs, dim = gate.shape + m = length * bs + M = (m + 127) // 128 * 128 + + lbs = length * bs + T = 1 if lbs <= 8192 else 2 ** (math.ceil(math.log2(lbs / 8192))) + + d = dim // group_size + D = triton.next_power_of_2(d) + assert D == d, f"d = dim // group_size = {d} must be a power of 2" + assert d >= 32 and d % 32 == 0, f"d = {d} must be >= 32 and divisible by 32" + assert (M // 32) % T == 0, f"M//32 = {M // 32} must be divisible by T = {T}" + assert dim <= 8192 and triton.next_power_of_2(group_size) == group_size + assert grad_output.is_contiguous() + assert gate.stride(2) == 1 and gate.stride(0) == gate.stride(1) * bs + assert length != bs + + SB = d // 32 + SHARE = weight.shape[0] != dim + NATIVE = x.size(0) == bs + device = x.device + + num_row_blocks = M // 32 // T + + dx = torch.empty_like(x, dtype=torch.bfloat16) + dg_q = torch.empty((m, dim), device=device, dtype=torch.float8_e4m3fn) + dg_s = torch.empty((M, dim // 32), device=device, dtype=torch.uint8) + dgt_q = torch.empty((m, dim), device=device, dtype=torch.float8_e4m3fn) + dgt_s = torch.empty((M // 32, dim), device=device, dtype=torch.uint8) + tmp_dw = torch.empty( + (num_row_blocks * group_size, D), dtype=torch.float32, device=device + ) + + grid = (num_row_blocks, group_size) + group_rms_norm_gate_and_mxfp8_quant_backward_kernel[grid]( + grad_output, + x, + gate, + weight, + dx, + dg_q, + dg_s, + dgt_q, + dgt_s, + tmp_dw, + gate.stride(1), + eps, + bs, + length, + m, + dim, + d, + D, + group_size, + SB, + T, + SHARE, + NATIVE, + output_mode, + num_stages=3, + num_warps=4, + ) + + if SHARE: + dw = tmp_dw.sum(0).to(x.dtype) + else: + dw = tmp_dw.view(num_row_blocks, group_size, D).sum(0).reshape(dim).to(x.dtype) + return dx, dg_q, dg_s, dgt_q, dgt_s, dw diff --git a/linghe/utils/gather.py b/linghe/utils/gather.py index 642db26..87a01cb 100644 --- a/linghe/utils/gather.py +++ b/linghe/utils/gather.py @@ -11,16 +11,110 @@ @triton.jit -def block_count_kernel(map_ptr, count_ptr, M, B, T: tl.constexpr, - b: tl.constexpr, E: tl.constexpr): +def _chunk_sort_map_kernel( + split_sizes_ptr, + row_id_map_ptr, + row_id_map_inv_ptr, + num_ranks: tl.constexpr, + num_local_experts: tl.constexpr, + BLOCK_SIZE_SPLITS: tl.constexpr, + BLOCK_SIZE_ROW: tl.constexpr, +): + chunk_idx = tl.program_id(0) + num_splits = num_ranks * num_local_experts + if chunk_idx >= num_splits: + return + + r_idx = chunk_idx // num_local_experts + e_idx = chunk_idx % num_local_experts + + split_idx = tl.arange(0, BLOCK_SIZE_SPLITS) + mask = split_idx < num_splits + all_sizes = tl.load(split_sizes_ptr + split_idx, mask=mask, other=0) + + src_start = tl.sum(tl.where(split_idx < chunk_idx, all_sizes, 0)).to(tl.int32) + count = tl.sum(tl.where(split_idx == chunk_idx, all_sizes, 0)).to(tl.int32) + + rank_mask = (split_idx % num_local_experts == e_idx) & (split_idx < chunk_idx) + offset_within_expert = tl.sum(tl.where(rank_mask, all_sizes, 0)).to(tl.int32) + + expert_base_offset = 0 + for e in range(num_local_experts): + e_mask = split_idx % num_local_experts == e + e_total = tl.sum(tl.where(e_mask, all_sizes, 0)).to(tl.int32) + e_padded = ((e_total + 31) // 32) * 32 + + if e < e_idx: + expert_base_offset += e_padded + + dst_start = expert_base_offset + offset_within_expert + + for i in range(0, count, BLOCK_SIZE_ROW): + offsets = i + tl.arange(0, BLOCK_SIZE_ROW) + row_mask = offsets < count + + curr_src = src_start + offsets + curr_dst = dst_start + offsets + + tl.store(row_id_map_ptr + curr_dst, curr_src, mask=row_mask) + tl.store(row_id_map_inv_ptr + curr_src, curr_dst, mask=row_mask) + + if r_idx == num_ranks - 1: + e_total_for_expert = offset_within_expert + count + e_padded_for_expert = ((e_total_for_expert + 31) // 32) * 32 + pad_offsets = tl.arange(0, BLOCK_SIZE_ROW) + pad_mask = pad_offsets < (e_padded_for_expert - e_total_for_expert) + pad_dst = expert_base_offset + e_total_for_expert + pad_offsets + tl.store( + row_id_map_ptr + pad_dst, + tl.full([BLOCK_SIZE_ROW], 0, dtype=tl.int32), + mask=pad_mask, + ) + + +def triton_make_chunk_sort_map( + num_global_tokens_per_local_expert: torch.Tensor, token_per_expert_cpu +): + device = num_global_tokens_per_local_expert.device + num_ranks, num_local_experts = num_global_tokens_per_local_expert.shape + num_splits = num_ranks * num_local_experts + + total_padded_size = sum([(x + 31) // 32 * 32 for x in token_per_expert_cpu]) + total_tokens = sum(token_per_expert_cpu) + + # row_id_map = torch.zeros((total_padded_size,), dtype=torch.int32, device=device) + # row_id_map = torch.ones((total_padded_size,), dtype=torch.int32, device=device) * -1 + row_id_map = torch.empty((total_padded_size,), dtype=torch.int32, device=device) + row_id_map_inverse = torch.empty((total_tokens,), dtype=torch.int32, device=device) + + grid = (num_splits,) + block_splits = triton.next_power_of_2(num_splits) + + _chunk_sort_map_kernel[grid]( + num_global_tokens_per_local_expert, + row_id_map, + row_id_map_inverse, + num_ranks, + num_local_experts, + BLOCK_SIZE_SPLITS=block_splits, + BLOCK_SIZE_ROW=256, + num_warps=4, + ) + + return row_id_map, row_id_map_inverse + + +@triton.jit +def block_count_kernel( + map_ptr, count_ptr, M, B, T: tl.constexpr, b: tl.constexpr, E: tl.constexpr +): pid = tl.program_id(axis=0) counts = tl.zeros((E,), dtype=tl.int32) offs = pid * B * E + tl.arange(0, b)[:, None] * E + tl.arange(0, E)[None, :] t = tl.cdiv(B, b) for i in range(t): - mask = pid * B + i * b + tl.arange(0, b)[:, None] < tl.minimum(M, - pid * B + B) + mask = pid * B + i * b + tl.arange(0, b)[:, None] < tl.minimum(M, pid * B + B) values = tl.load(map_ptr + offs, mask=mask).to(tl.int32) counts += tl.sum(values, 0) offs += b * E @@ -29,8 +123,17 @@ def block_count_kernel(map_ptr, count_ptr, M, B, T: tl.constexpr, @triton.jit -def make_row_id_map_kernel(map_ptr, count_ptr, output_ptr, M, B, P, - T: tl.constexpr, b: tl.constexpr, E: tl.constexpr): +def make_row_id_map_kernel( + map_ptr, + count_ptr, + output_ptr, + M, + B, + P, + T: tl.constexpr, + b: tl.constexpr, + E: tl.constexpr, +): pid = tl.program_id(axis=0) indices = tl.arange(0, T)[:, None] * E + tl.arange(0, E)[None, :] @@ -45,8 +148,7 @@ def make_row_id_map_kernel(map_ptr, count_ptr, output_ptr, M, B, P, offs = pid * B * E + tl.arange(0, b)[:, None] * E + tl.arange(0, E)[None, :] t = tl.cdiv(B, b) for i in range(t): - mask = pid * B + i * b + tl.arange(0, b)[:, None] < tl.minimum(M, - pid * B + B) + mask = pid * B + i * b + tl.arange(0, b)[:, None] < tl.minimum(M, pid * B + B) values = tl.load(map_ptr + offs, mask=mask).to(tl.int32) acc = count_offset + tl.cumsum(values, 0) count_offset = tl.max(acc, 0) @@ -55,11 +157,7 @@ def make_row_id_map_kernel(map_ptr, count_ptr, output_ptr, M, B, P, offs += b * E - -def triton_make_row_id_map( - routing_map: torch.Tensor, - multiple_of: int = 1 -): +def triton_make_row_id_map(routing_map: torch.Tensor, multiple_of: int = 1): """ make row id map, values in the tensor are the row indices Args: @@ -69,12 +167,15 @@ def triton_make_row_id_map( Returns: row id map with shape [n_tokens, n_experts] """ + assert routing_map.is_contiguous() n_tokens, n_experts = routing_map.shape T = 128 - block_counts = torch.empty((T, n_experts), dtype=torch.int32, - device=routing_map.device) - output = torch.empty((n_tokens, n_experts), dtype=torch.int32, - device=routing_map.device) + block_counts = torch.empty( + (T, n_experts), dtype=torch.int32, device=routing_map.device + ) + output = torch.empty( + (n_tokens, n_experts), dtype=torch.int32, device=routing_map.device + ) B = triton.cdiv(n_tokens, T) b = 16 @@ -88,7 +189,7 @@ def triton_make_row_id_map( b, n_experts, num_stages=3, - num_warps=8 + num_warps=8, ) make_row_id_map_kernel[grid]( @@ -102,17 +203,25 @@ def triton_make_row_id_map( b, n_experts, num_stages=3, - num_warps=8 + num_warps=8, ) return output @triton.jit -def make_row_id_map_and_indices_kernel(map_ptr, count_ptr, row_map_ptr, - row_indices_ptr, M, B, P, - T: tl.constexpr, b: tl.constexpr, - E: tl.constexpr): +def make_row_id_map_and_index_kernel( + map_ptr, + count_ptr, + row_map_ptr, + row_indices_ptr, + M, + B, + P, + T: tl.constexpr, + b: tl.constexpr, + E: tl.constexpr, +): pid = tl.program_id(axis=0) indices = tl.arange(0, T)[:, None] * E + tl.arange(0, E)[None, :] @@ -127,27 +236,26 @@ def make_row_id_map_and_indices_kernel(map_ptr, count_ptr, row_map_ptr, offs = pid * B * E + tl.arange(0, b)[:, None] * E + tl.arange(0, E)[None, :] t = tl.cdiv(B, b) for i in range(t): - mask = pid * B + i * b + tl.arange(0, b)[:, None] < tl.minimum(M, - pid * B + B) + mask = pid * B + i * b + tl.arange(0, b)[:, None] < tl.minimum(M, pid * B + B) values = tl.load(map_ptr + offs, mask=mask).to(tl.int32) acc = count_offset + tl.cumsum(values, 0) count_offset = tl.max(acc, 0) output_acc = tl.where(values == 0, -1, acc) tl.store(row_map_ptr + offs, output_acc, mask=mask) - tl.store(row_indices_ptr + acc, - pid * B + i * b + tl.arange(0, b)[:, None] + (0 * tl.arange(0, - E))[ - None, :], - mask=mask & values != 0) + tl.store( + row_indices_ptr + acc, + pid * B + i * b + tl.arange(0, b)[:, None] + (0 * tl.arange(0, E))[None, :], + mask=mask & values != 0, + ) offs += b * E -def triton_make_row_id_map_and_indices( - routing_map: torch.Tensor, - num_out_tokens: int, - multiple_of: int = 1, +def triton_make_row_id_map_and_index( + routing_map: torch.Tensor, + num_out_tokens: int, + multiple_of: int = 1, ): """ similar with triton_make_row_id_map, but output an indices tensor as well @@ -159,14 +267,18 @@ def triton_make_row_id_map_and_indices( row_in_map: [n_tokens, n_experts] row_indices: [num_out_tokens] """ + assert routing_map.is_contiguous() n_tokens, n_experts = routing_map.shape T = 128 - block_counts = torch.empty((T, n_experts), dtype=torch.int32, - device=routing_map.device) - row_id_map = torch.empty((n_tokens, n_experts), dtype=torch.int32, - device=routing_map.device) - row_id_indices = torch.empty((num_out_tokens,), dtype=torch.int32, - device=routing_map.device) + block_counts = torch.empty( + (T, n_experts), dtype=torch.int32, device=routing_map.device + ) + row_id_map = torch.empty( + (n_tokens, n_experts), dtype=torch.int32, device=routing_map.device + ) + row_id_indices = torch.zeros( + (num_out_tokens,), dtype=torch.int32, device=routing_map.device + ) B = triton.cdiv(n_tokens, T) b = 16 @@ -180,10 +292,10 @@ def triton_make_row_id_map_and_indices( b, n_experts, num_stages=3, - num_warps=8 + num_warps=8, ) - make_row_id_map_and_indices_kernel[grid]( + make_row_id_map_and_index_kernel[grid]( routing_map, block_counts, row_id_map, @@ -195,16 +307,24 @@ def triton_make_row_id_map_and_indices( b, n_experts, num_stages=3, - num_warps=8 + num_warps=8, ) return row_id_map, row_id_indices @triton.jit -def index_select_kernel(x_ptr, out_ptr, scale_ptr, scale_out_ptr, index_ptr, M, - T, N: tl.constexpr, SCALE: tl.constexpr): +def permute_with_indices_kernel( + x_ptr, + out_ptr, + scale_ptr, + scale_out_ptr, + index_ptr, + M, + T, + N: tl.constexpr, + SCALE: tl.constexpr, +): pid = tl.program_id(axis=0) - # row-wise read, row-wise write for i in range(T): dst_idx = pid * T + i src_idx = tl.load(index_ptr + dst_idx, mask=dst_idx < M) @@ -216,9 +336,9 @@ def index_select_kernel(x_ptr, out_ptr, scale_ptr, scale_out_ptr, index_ptr, M, tl.store(scale_out_ptr + dst_idx, scale, mask=dst_idx < M) -def triton_index_select(x, indices, scale=None, out=None, scale_out=None): +def triton_permute_with_indices(x, indices, scale=None, out=None, scale_out=None): """ - index select for quantized tensor + index select, mainly used for fp8 dispatch Args: x: [bs, dim] indices: [K] @@ -227,43 +347,40 @@ def triton_index_select(x, indices, scale=None, out=None, scale_out=None): out: output of selected x scale_out: scale of selected scale """ - # row-wise read, row-wise write + assert x.is_contiguous() M, N = x.shape + assert triton.next_power_of_2(N) == N E = indices.shape[0] device = x.device if out is None: out = torch.empty((E, N), device=device, dtype=x.dtype) if scale is not None and scale_out is None: scale_out = torch.empty((E,), device=device, dtype=scale.dtype) - sm = torch.cuda.get_device_properties(device).multi_processor_count + sm = 2048 T = triton.cdiv(E, sm) SCALE = scale is not None grid = (sm,) - index_select_kernel[grid]( - x, - out, - scale, - scale_out, - indices, - E, T, N, - SCALE, - num_stages=3, - num_warps=8 + permute_with_indices_kernel[grid]( + x, out, scale, scale_out, indices, E, T, N, SCALE, num_stages=3, num_warps=4 ) return out, scale_out @triton.jit -def permute_with_mask_map_kernel(data_ptr, scale_ptr, probs_ptr, - mask_map_ptr, - output_data_ptr, - output_scale_ptr, - output_probs_ptr, - num_experts: tl.constexpr, - N: tl.constexpr, - hs: tl.constexpr, - SCALE: tl.constexpr, - PROB: tl.constexpr): +def permute_with_mask_map_kernel( + data_ptr, + scale_ptr, + probs_ptr, + mask_map_ptr, + output_data_ptr, + output_scale_ptr, + output_probs_ptr, + num_experts: tl.constexpr, + N: tl.constexpr, + hs: tl.constexpr, + SCALE: tl.constexpr, + PROB: tl.constexpr, +): pid = tl.program_id(axis=0) x = tl.load(data_ptr + pid * N + tl.arange(0, N)) if SCALE == 1: @@ -271,10 +388,9 @@ def permute_with_mask_map_kernel(data_ptr, scale_ptr, probs_ptr, elif SCALE == 2: scale = tl.load(scale_ptr + pid * hs + tl.arange(0, hs)) - indices = tl.load( - mask_map_ptr + pid * num_experts + tl.arange(0, num_experts)) + indices = tl.load(mask_map_ptr + pid * num_experts + tl.arange(0, num_experts)) count = tl.sum(tl.where(indices >= 0, 1, 0)) - mask_indices = tl.where(indices < 0, 2 ** 20, indices) + mask_indices = tl.where(indices < 0, 2**20, indices) idx = tl.argmin(mask_indices, 0) index = tl.min(mask_indices) for i in range(count): @@ -290,19 +406,23 @@ def permute_with_mask_map_kernel(data_ptr, scale_ptr, probs_ptr, prob = tl.load(probs_ptr + pid * num_experts + idx) tl.store(output_probs_ptr + index, prob) - mask_indices = tl.where(indices <= index, 2 ** 20, indices) + mask_indices = tl.where(indices <= index, 2**20, indices) idx = tl.argmin(mask_indices, 0) index = tl.min(mask_indices) @triton.jit -def fill_padded_token_with_zero_kernel(data_ptr, scale_ptr, probs_ptr, - max_indices_ptr, - token_per_expert_ptr, - N: tl.constexpr, - hs: tl.constexpr, - SCALE: tl.constexpr, - PROB: tl.constexpr): +def fill_padded_token_with_zero_kernel( + data_ptr, + scale_ptr, + probs_ptr, + max_indices_ptr, + token_per_expert_ptr, + N: tl.constexpr, + hs: tl.constexpr, + SCALE: tl.constexpr, + PROB: tl.constexpr, +): pid = tl.program_id(axis=0) x = tl.zeros((N,), dtype=tl.float32) si = tl.load(max_indices_ptr + pid) @@ -322,20 +442,20 @@ def fill_padded_token_with_zero_kernel(data_ptr, scale_ptr, probs_ptr, def triton_permute_with_mask_map( - inp: torch.Tensor, - scale: torch.Tensor, - probs: torch.Tensor, - row_id_map: torch.Tensor, - num_out_tokens: int, - contiguous: bool = True, - tokens_per_expert: Optional[torch.Tensor] = None + inp: torch.Tensor, + scale: torch.Tensor, + probs: torch.Tensor, + row_id_map: torch.Tensor, + num_out_tokens: int, + contiguous: bool = True, + tokens_per_expert: Optional[torch.Tensor] = None, ): """ gather quantized tensor with row id map Args: inp: [num_tokens, hidden_size], rowwise quantized tensor - scale: [num_tokens], quantization scale - probs: router prob, used as weight + scale: optional, [num_tokens], quantization scale + probs: optional, router prob, used as weight row_id_map: [n_experts, num_tokens] index >= 0: row index of output tensor index == -1: ignore @@ -352,8 +472,10 @@ def triton_permute_with_mask_map( permuted_probs: permuted router prob """ + assert inp.is_contiguous() num_tokens, hidden_size = inp.shape num_tokens_, num_experts = row_id_map.shape # not transposed + assert triton.next_power_of_2(hidden_size) == hidden_size assert num_tokens == num_tokens_ SCALE = 0 # NO SCALE hs = 0 @@ -365,24 +487,20 @@ def triton_permute_with_mask_map( ZERO = not contiguous and tokens_per_expert is None if ZERO: - output = torch.zeros((num_out_tokens, hidden_size), dtype=inp.dtype, - device="cuda") + output = torch.zeros( + (num_out_tokens, hidden_size), dtype=inp.dtype, device="cuda" + ) else: - output = torch.empty((num_out_tokens, hidden_size), dtype=inp.dtype, - device="cuda") + output = torch.empty( + (num_out_tokens, hidden_size), dtype=inp.dtype, device="cuda" + ) if SCALE > 0: shape = (num_out_tokens, hs) if SCALE == 2 else (num_out_tokens,) if ZERO: - permuted_scale = torch.zeros( - shape, - dtype=scale.dtype, device="cuda" - ) + permuted_scale = torch.zeros(shape, dtype=scale.dtype, device="cuda") else: - permuted_scale = torch.empty( - shape, - dtype=scale.dtype, device="cuda" - ) + permuted_scale = torch.empty(shape, dtype=scale.dtype, device="cuda") else: permuted_scale = None @@ -417,64 +535,188 @@ def triton_permute_with_mask_map( SCALE, PROB, num_stages=3, - num_warps=8 + num_warps=8, ) if not contiguous and tokens_per_expert is not None: max_indices = row_id_map.amax(0) - fill_padded_token_with_zero_kernel[(num_experts,)](output, - permuted_scale, - permuted_probs, - max_indices, - tokens_per_expert, - hidden_size, - hs, - SCALE, - PROB) + fill_padded_token_with_zero_kernel[(num_experts,)]( + output, + permuted_scale, + permuted_probs, + max_indices, + tokens_per_expert, + hidden_size, + hs, + SCALE, + PROB, + ) return output, permuted_scale, permuted_probs @triton.jit -def batch_smooth_transpose_smooth_permute_kernel(x_ptr, scale_ptr, oss_ptr, - ss_ptr, index_ptr, count_ptr, - accum_ptr, q_ptr, qs_ptr, - N: tl.constexpr, - E: tl.constexpr, - H: tl.constexpr, - W: tl.constexpr, - SMOOTHED: tl.constexpr, - ROUND: tl.constexpr): +def batch_smooth_permute_with_indices_kernel( + x_ptr, + ss_ptr, + prob_ptr, + q_ptr, + qs_ptr, + prob_out_ptr, + count_ptr, + accum_ptr, + index_ptr, + T, + N: tl.constexpr, + REVERSE: tl.constexpr, + ROUND: tl.constexpr, +): + eid = tl.program_id(axis=0) + NE = tl.num_programs(0) + tid = tl.program_id(axis=1) + + smooth_scale = tl.load(ss_ptr + eid * N + tl.arange(0, N)) + if not REVERSE: + smooth_scale = 1.0 / smooth_scale + count = tl.load(count_ptr + eid) + ei = tl.load(accum_ptr + eid) + si = ei - count + c = tl.cdiv(count, T) + for i in range(si + tid * c, tl.minimum(si + tid * c + c, ei)): + index = tl.load(index_ptr + i) + x = tl.load(x_ptr + index * N + tl.arange(0, N)).to(tl.float32) + + x *= smooth_scale + x_max = tl.max(tl.abs(x)) + + scale = tl.maximum(x_max / 448.0, 1e-30) + if ROUND: + scale = tl.exp2(tl.ceil(tl.log2(scale))) + + tl.store(qs_ptr + i, scale) + + s = 1.0 / scale + x *= s + xq = x.to(q_ptr.dtype.element_ty) + tl.store(q_ptr + i * N + tl.arange(0, N), xq) + + if prob_ptr is not None: + prob = tl.load(prob_ptr + index * NE + eid) + tl.store(prob_out_ptr + i, prob) + + +def triton_batch_smooth_permute_with_indices( + x, + smooth_scales, + token_count_per_expert, + indices, + probs=None, + x_q=None, + x_scale=None, + reverse=False, + round_scale=False, +): + """ + TODO: opt perfermance + used for permutation with megatron flex backend + step: select, smooth, quant + Args: + x: [bs, dim] + smooth_scales: [n_experts, dim] + token_count_per_expert: [n_experts] + indices: [n_experts*topk] + x_q: [bs*topk, dim] + x_scale: [bs*topk] + reverse: + round_scale: + + Returns: + + """ + assert x.is_contiguous() + M, N = x.shape + n_expert, n = smooth_scales.shape + assert N == n + assert triton.next_power_of_2(N) == N + + E = indices.size(0) + device = x.device + if x_q is None: + x_q = torch.empty((E, N), device=device, dtype=torch.float8_e4m3fn) + if x_scale is None: + x_scale = torch.empty((E,), device=device, dtype=torch.float32) + + PROB = probs is not None + if PROB: + prob_output = torch.empty((E,), device=device, dtype=probs.dtype) + + else: + prob_output = None + + if E == 0: + return x_q, x_scale, prob_output + + accum_token_count = torch.cumsum(token_count_per_expert, 0) + T = 128 + grid = (n_expert, T) + batch_smooth_permute_with_indices_kernel[grid]( + x, + smooth_scales, + probs, + x_q, + x_scale, + prob_output, + token_count_per_expert, + accum_token_count, + indices, + T, + N, + reverse, + round_scale, + num_stages=3, + num_warps=4, + ) + return x_q, x_scale, prob_output + + +@triton.jit +def batch_transpose_smooth_permute_kernel( + x_ptr, + ss_ptr, + index_ptr, + count_ptr, + accum_ptr, + q_ptr, + qs_ptr, + N: tl.constexpr, + E: tl.constexpr, + H: tl.constexpr, + W: tl.constexpr, + ROUND: tl.constexpr, +): eid = tl.program_id(axis=0) cid = tl.program_id(axis=1) count = tl.load(count_ptr + eid) counts = tl.load(count_ptr + tl.arange(0, E)) - si = tl.load(accum_ptr + eid) - count + si = (tl.load(accum_ptr + eid) - count).to(tl.int64) pad = tl.cdiv(count, 32) * 32 loop = tl.cdiv(pad, H) - bias = tl.sum( - tl.where(tl.arange(0, E) < eid, tl.cdiv(counts, 32), 0)) * 32 * N + bias = tl.sum(tl.where(tl.arange(0, E) < eid, tl.cdiv(counts, 32), 0)) * 32 * N - # col-wise read, row-wise write - if SMOOTHED: - org_smooth_scale = tl.load(oss_ptr + cid * W + tl.arange(0, W)) x_max = tl.zeros((H, W), dtype=tl.float32) for i in range(loop): idx = i * H + tl.arange(0, H) - indices = tl.load(index_ptr + si + i * H + tl.arange(0, H), - mask=idx < count) + indices = tl.load(index_ptr + si + i * H + tl.arange(0, H), mask=idx < count) x = tl.load( x_ptr + cid * W + indices[:, None] * N + tl.arange(0, W)[None, :], - mask=idx[:, None] < count).to(tl.float32) - smooth_scale = tl.load(ss_ptr + si + i * H + tl.arange(0, H), - mask=idx < count)[:, None] - if SMOOTHED: - s = tl.load(scale_ptr + indices, mask=idx < count)[:, None] - x = x * org_smooth_scale * (s * smooth_scale) - else: - x = x * smooth_scale + mask=idx[:, None] < count, + ).to(tl.float32) + smooth_scale = tl.load(ss_ptr + si + i * H + tl.arange(0, H), mask=idx < count)[ + :, None + ] + x = x * smooth_scale x_max = tl.maximum(tl.abs(x), x_max) scale = tl.maximum(tl.max(x_max, 0) / 448.0, 1e-30) @@ -484,45 +726,40 @@ def batch_smooth_transpose_smooth_permute_kernel(x_ptr, scale_ptr, oss_ptr, tl.store(qs_ptr + eid * N + cid * W + tl.arange(0, W), scale) scale = 1.0 / scale - toffs = bias + cid * pad * W + tl.arange(0, W)[:, None] * pad + tl.arange(0, - H) + toffs = bias + cid * pad * W + tl.arange(0, W)[:, None] * pad + tl.arange(0, H) for i in range(loop): idx = i * H + tl.arange(0, H) - indices = tl.load(index_ptr + si + i * H + tl.arange(0, H), - mask=idx < count) + indices = tl.load(index_ptr + si + i * H + tl.arange(0, H), mask=idx < count) x = tl.load( x_ptr + cid * W + indices[:, None] * N + tl.arange(0, W)[None, :], - mask=idx[:, None] < count).to(tl.float32) - smooth_scale = tl.load(ss_ptr + si + i * H + tl.arange(0, H), - mask=idx < count)[:, None] - if SMOOTHED: - s = tl.load(scale_ptr + indices, mask=idx < count)[:, None] - x = x * (org_smooth_scale * scale) * (s * smooth_scale) - else: - x = x * scale * smooth_scale + mask=idx[:, None] < count, + ).to(tl.float32) + smooth_scale = tl.load(ss_ptr + si + i * H + tl.arange(0, H), mask=idx < count)[ + :, None + ] + + x = x * scale * smooth_scale xq = tl.trans(x.to(q_ptr.dtype.element_ty)) tl.store(q_ptr + toffs, xq, mask=idx[None, :] < pad) toffs += H -def triton_batch_transpose_smooth_permute_with_indices(x, - scale, - org_smooth_scale, - smooth_scales, - indices, - token_count_per_expert, - splits, - x_q=None, - x_scale=None, - round_scale=False): +def triton_batch_transpose_smooth_permute_with_indices( + x, + smooth_scales, + indices, + token_count_per_expert, + splits, + x_q=None, + x_scale=None, + round_scale=False, +): """ - used for smooth quantization backward in megatron 0.12, - x is gathered, requantized, padded to multiple of 32 and tranposed + used for unpermutation with megatron flex backend + x is gathered, padded to multiple of 32, tranposed and smooth quantized Args: - x: dy, [bs, dim], it is smooth quantized - scale: [bs], quantized scale - org_smooth_scale: [dim] - smooth_scales: [n_experts, dim] + x: dy, bf16, [bs, dim] + smooth_scales: [sum(tokens_per_experts)] indices: [sum(tokens_per_experts)] token_count_per_expert: [n_experts], tensor of token count per expert splits: [n_experts], list of token_count_per_expert @@ -532,7 +769,7 @@ def triton_batch_transpose_smooth_permute_with_indices(x, x_q: [sum(roundup(tokens_per_experts)) * dim] x_scale: [sum(roundup(tokens_per_experts))] """ - # row-wise read, row-wise write + assert x.is_contiguous() M, N = x.shape n_expert = len(splits) out_tokens = sum([(x + 31) // 32 for x in splits]) * 32 @@ -544,20 +781,19 @@ def triton_batch_transpose_smooth_permute_with_indices(x, W = 32 device = x.device accum_token_count = torch.cumsum(token_count_per_expert, 0) - smoothed = scale is not None if x_q is None: - # TODO(nanxiao): opt performance - x_q = torch.empty((out_tokens * N,), device=device, - dtype=torch.float8_e4m3fn) + x_q = torch.empty((out_tokens, N), device=device, dtype=torch.float8_e4m3fn) if x_scale is None: x_scale = torch.empty((n_expert, N), device=device, dtype=torch.float32) + if out_tokens == 0: + return x_q, x_scale + # import pydevd # pydevd.settrace(suspend=False, trace_only_current_thread=True) + assert N % W == 0 grid = (n_expert, N // W) - batch_smooth_transpose_smooth_permute_kernel[grid]( + batch_transpose_smooth_permute_kernel[grid]( x, - scale, - org_smooth_scale, smooth_scales, indices, token_count_per_expert, @@ -568,126 +804,31 @@ def triton_batch_transpose_smooth_permute_with_indices(x, n_expert, H, W, - smoothed, round_scale, num_stages=3, - num_warps=8 + num_warps=8, ) return x_q, x_scale @triton.jit -def smooth_weighted_permute_with_indices_kernel(grads_ptr, - tokens_ptr, q_ptr, - ss_ptr, qs_ptr, - count_ptr, accum_ptr, - index_ptr, sum_ptr, - M, N: tl.constexpr, - REVERSE: tl.constexpr, - ROUND: tl.constexpr): - pid = tl.program_id(axis=0) - # row-wise read, row-wise write - smooth_scale = tl.load(ss_ptr + pid * N + tl.arange(0, N)) - if not REVERSE: - smooth_scale = 1.0 / smooth_scale - count = tl.load(count_ptr + pid) - ei = tl.load(accum_ptr + pid) - si = ei - count - for i in range(count): - index = tl.load(index_ptr + si + i) - x = tl.load(grads_ptr + index * N + tl.arange(0, N)).to(tl.float32) - t = tl.load(tokens_ptr + si * N + i * N + tl.arange(0, N)).to( - tl.float32) - sums = tl.sum(x * t) - tl.store(sum_ptr + si + i, sums) - - x *= smooth_scale - x_max = tl.max(tl.abs(x)) - scale = tl.maximum(x_max / 448.0, 1e-30) - if ROUND: - scale = tl.exp2(tl.ceil(tl.log2(scale))) - - tl.store(qs_ptr + si + i, scale) - - s = 1.0 / scale - x *= s - xq = x.to(q_ptr.dtype.element_ty) - tl.store(q_ptr + si * N + i * N + tl.arange(0, N), xq) - - -def triton_smooth_weighted_permute_with_indices(grads, - tokens, - smooth_scales, - token_count_per_expert, - indices, - x_q=None, - x_scale=None, - x_sum=None, - reverse=False, - round_scale=False): - """ - select and smooth and quant, used in megatron 0.11 all2all moe - Args: - grads: [bs, dim] - tokens: [bs, dim] - smooth_scales: [n_experts, dim] - token_count_per_expert: [n_experts] - indices: [n_experts*topk] - reverse: whether scale is 1/scale - round_scale: whether round scale to power of 2 - - Returns: - x_q: [bs*topk, dim] - x_scale: [bs*topk] - x_sum: [bs*topk] - """ - M, N = grads.shape - n_expert, n = smooth_scales.shape - assert N == n, f'{N=} {n=}' - E = indices.shape[0] - device = grads.device - if x_q is None: - x_q = torch.empty((E, N), device=device, dtype=torch.float8_e4m3fn) - if x_scale is None: - x_scale = torch.empty((E,), device=device, dtype=torch.float32) - if x_sum is None: - x_sum = torch.empty((E,), device=device, dtype=grads.dtype) - accum_token_count = torch.cumsum(token_count_per_expert, 0) - grid = (n_expert,) - smooth_weighted_permute_with_indices_kernel[grid]( - grads, - tokens, - x_q, - smooth_scales, - x_scale, - token_count_per_expert, - accum_token_count, - indices, - x_sum, - M, N, - reverse, - round_scale, - num_stages=3, - num_warps=8 - ) - return x_q, x_scale, x_sum - - -@triton.jit -def smooth_permute_with_indices_kernel(grads_data_ptr, - grads_scale_ptr, q_ptr, - ss_ptr, qs_ptr, count_ptr, - accum_ptr, index_ptr, - N: tl.constexpr, - hs: tl.constexpr, - REVERSE: tl.constexpr, - ROUND: tl.constexpr, - GROUP: tl.constexpr): +def batch_smooth_fused_permute_with_indices_kernel( + grads_data_ptr, + grads_scale_ptr, + q_ptr, + ss_ptr, + qs_ptr, + count_ptr, + accum_ptr, + index_ptr, + N: tl.constexpr, + REVERSE: tl.constexpr, + ROUND: tl.constexpr, +): eid = tl.program_id(axis=0) wid = tl.program_id(axis=1) T = tl.num_programs(axis=1) - # row-wise read, row-wise write smooth_scale = tl.load(ss_ptr + eid * N + tl.arange(0, N)) if not REVERSE: smooth_scale = 1.0 / smooth_scale @@ -698,12 +839,9 @@ def smooth_permute_with_indices_kernel(grads_data_ptr, for i in range(si + wid * c, tl.minimum(si + wid * c + c, ei)): index = tl.load(index_ptr + i) x = tl.load(grads_data_ptr + index * N + tl.arange(0, N)).to(tl.float32) - if GROUP: - gs = tl.load(grads_scale_ptr + index * hs + tl.arange(0, hs)) - x = tl.reshape(tl.reshape(x, (hs, N // hs)) * gs[:, None], (N,)) - else: - gs = tl.load(grads_scale_ptr + index) - x *= gs + + gs = tl.load(grads_scale_ptr + index) + x *= gs x *= smooth_scale x_max = tl.max(tl.abs(x)) @@ -720,21 +858,25 @@ def smooth_permute_with_indices_kernel(grads_data_ptr, tl.store(q_ptr + i * N + tl.arange(0, N), xq) -def triton_smooth_permute_with_indices(grad_data, - grad_scale, - smooth_scales, - token_count_per_expert, - indices, - x_q=None, - x_scale=None, - reverse=False, - round_scale=False): +def triton_batch_smooth_fused_permute_with_indices( + grad_data, + grad_scale, + org_smooth_scales, + smooth_scales, + token_count_per_expert, + indices, + x_q=None, + x_scale=None, + reverse=False, + round_scale=False, +): """ + used in unpermutation backward with megatron flex backend and fp8 combine, select and smooth and quant Args: grad_data: [bs, dim] grad_scale: [bs] - smooth_scales: [n_experts, dim] + smooth_scales: [n_experts, dim], it is composite of two smooth scales token_count_per_expert: [n_experts] indices: [n_experts*topk] x_q: [bs*topk, dim] @@ -745,14 +887,14 @@ def triton_smooth_permute_with_indices(grad_data, Returns: """ - # row-wise read, row-wise write + assert grad_data.is_contiguous() M, N = grad_data.shape n_expert, n = smooth_scales.shape assert 128 % n_expert == 0 assert N == n + assert triton.next_power_of_2(N) == N - group = grad_scale.ndim > 1 - hs = grad_scale.shape[1] if group else 1 + smooth_scales = smooth_scales * org_smooth_scales E = indices.size(0) device = grad_data.device @@ -762,8 +904,9 @@ def triton_smooth_permute_with_indices(grad_data, x_scale = torch.empty((E,), device=device, dtype=torch.float32) accum_token_count = torch.cumsum(token_count_per_expert, 0) W = 128 // n_expert + # TODO: opt perf grid = (n_expert, W) - smooth_permute_with_indices_kernel[grid]( + batch_smooth_fused_permute_with_indices_kernel[grid]( grad_data, grad_scale, x_q, @@ -773,120 +916,507 @@ def triton_smooth_permute_with_indices(grad_data, accum_token_count, indices, N, - hs, reverse, round_scale, - group, num_stages=3, - num_warps=16 + num_warps=16, ) return x_q, x_scale @triton.jit -def smooth_permute_with_mask_map_kernel(grads_data_ptr, quant_data_ptr, - mask_map_ptr, grads_scale_ptr, - smooth_scale_ptr, quant_scale_ptr, - M, T, - N: tl.constexpr, - hs: tl.constexpr, - REVERSE: tl.constexpr, - ROUND: tl.constexpr): +def batch_transpose_smooth_fused_permute_kernel( + x_ptr, + scale_ptr, + oss_ptr, + ss_ptr, + index_ptr, + count_ptr, + accum_ptr, + q_ptr, + qs_ptr, + N: tl.constexpr, + E: tl.constexpr, + H: tl.constexpr, + W: tl.constexpr, + SMOOTHED: tl.constexpr, + ROUND: tl.constexpr, +): eid = tl.program_id(axis=0) - bid = tl.program_id(axis=1) - n_experts = tl.num_programs(axis=0) + cid = tl.program_id(axis=1) - # smooth_scale_ptr = tl.load(smooth_scale_ptrs + eid).to(tl.pointer_type(tl.float32)) - smooth_scale = tl.load(smooth_scale_ptr + eid * N + tl.arange(0, N)) - if not REVERSE: - smooth_scale = 1.0 / smooth_scale - for i in range(bid * T, tl.minimum(bid * T + T, M)): - index = tl.load(mask_map_ptr + i * n_experts + eid) - mask = index >= 0 - if index >= 0: - x = tl.load(grads_data_ptr + i * N + tl.arange(0, N), mask=mask).to( - tl.float32) - - if hs > 1: - gs = tl.load(grads_scale_ptr + i * hs + tl.arange(0, hs), - mask=mask) - x = tl.reshape(tl.reshape(x, (hs, N // hs)) * gs[:, None], (N,)) - elif hs == 1: - gs = tl.load(grads_scale_ptr + i, mask=mask) - x *= gs - - x *= smooth_scale - x_max = tl.max(tl.abs(x)) - - scale = tl.maximum(x_max / 448.0, 1e-30) - if ROUND: - scale = tl.exp2(tl.ceil(tl.log2(scale))) - - tl.store(quant_scale_ptr + index, scale, mask=mask) - - x /= scale - xq = x.to(quant_data_ptr.dtype.element_ty) - tl.store(quant_data_ptr + index * N + tl.arange(0, N), xq, - mask=mask) - - -def triton_smooth_permute_with_mask_map( - inp: torch.Tensor, - row_id_map: torch.Tensor, - scale: torch.Tensor, - num_tokens: int, - num_experts: int, - num_out_tokens: int, - hidden_size: int, - smooth_scales: torch.Tensor, - reverse=True, - round_scale=False + count = tl.load(count_ptr + eid) + counts = tl.load(count_ptr + tl.arange(0, E)) + si = (tl.load(accum_ptr + eid) - count).to(tl.int64) + + pad = tl.cdiv(count, 32) * 32 + loop = tl.cdiv(pad, H) + bias = tl.sum(tl.where(tl.arange(0, E) < eid, tl.cdiv(counts, 32), 0)) * 32 * N + + if SMOOTHED: + org_smooth_scale = tl.load(oss_ptr + cid * W + tl.arange(0, W)) + x_max = tl.zeros((H, W), dtype=tl.float32) + for i in range(loop): + idx = i * H + tl.arange(0, H) + indices = tl.load(index_ptr + si + i * H + tl.arange(0, H), mask=idx < count) + x = tl.load( + x_ptr + cid * W + indices[:, None] * N + tl.arange(0, W)[None, :], + mask=idx[:, None] < count, + ).to(tl.float32) + smooth_scale = tl.load(ss_ptr + si + i * H + tl.arange(0, H), mask=idx < count)[ + :, None + ] + if SMOOTHED: + s = tl.load(scale_ptr + indices, mask=idx < count)[:, None] + x = x * org_smooth_scale * (s * smooth_scale) + else: + x = x * smooth_scale + x_max = tl.maximum(tl.abs(x), x_max) + + scale = tl.maximum(tl.max(x_max, 0) / 448.0, 1e-30) + if ROUND: + scale = tl.exp2(tl.ceil(tl.log2(scale))) + + tl.store(qs_ptr + eid * N + cid * W + tl.arange(0, W), scale) + + scale = 1.0 / scale + toffs = bias + cid * pad * W + tl.arange(0, W)[:, None] * pad + tl.arange(0, H) + for i in range(loop): + idx = i * H + tl.arange(0, H) + indices = tl.load(index_ptr + si + i * H + tl.arange(0, H), mask=idx < count) + x = tl.load( + x_ptr + cid * W + indices[:, None] * N + tl.arange(0, W)[None, :], + mask=idx[:, None] < count, + ).to(tl.float32) + smooth_scale = tl.load(ss_ptr + si + i * H + tl.arange(0, H), mask=idx < count)[ + :, None + ] + if SMOOTHED: + s = tl.load(scale_ptr + indices, mask=idx < count)[:, None] + x = x * (org_smooth_scale * scale) * (s * smooth_scale) + else: + x = x * scale * smooth_scale + xq = tl.trans(x.to(q_ptr.dtype.element_ty)) + tl.store(q_ptr + toffs, xq, mask=idx[None, :] < pad) + toffs += H + + +def triton_batch_transpose_smooth_fused_permute_with_indices( + x, + scale, + org_smooth_scale, + smooth_scales, + indices, + token_count_per_expert, + splits, + x_q=None, + x_scale=None, + round_scale=False, ): """ - gather ( and optional dequant) and smooth quant + used for calculating fc2 wgrad with megatron flex backend and fp8 combine, + x is gathered, requantized, padded to multiple of 32 and tranposed Args: - inp: [num_tokens, hidden_size], rowwise quantized tensor - row_id_map: [n_experts, num_tokens], indices - scale: [num_tokens, hs], rowwise_scale_inv, optional - num_tokens: [n_experts] - num_experts: - num_out_tokens: - hidden_size: - smooth_scales: [n_experts, hidden_size] - reverse: - round_scale: + x: dy, [bs, dim], it may be smooth quantized + scale: [bs], quantized scale + org_smooth_scale: [dim] + smooth_scales: [n_experts, dim] + indices: [sum(tokens_per_experts)] + token_count_per_expert: [n_experts], tensor of token count per expert + splits: [n_experts], list of token_count_per_expert + round_scale: round quantization scale to power of 2 Returns: - - output: output tensor - - permuted_scale: permuted scale if scale is not None + x_q: [sum(roundup(tokens_per_experts)) * dim] + x_scale: [sum(roundup(tokens_per_experts))] """ - assert row_id_map.shape[1] == num_experts - output = torch.empty((num_out_tokens, hidden_size), - dtype=torch.float8_e4m3fn, - device=row_id_map.device) - if scale is None: - hs = 0 + assert x.is_contiguous() + M, N = x.shape + n_expert = len(splits) + out_tokens = sum([(x + 31) // 32 for x in splits]) * 32 + if N >= 4096: + H = 64 + W = 64 else: - hs = scale.shape[1] if scale.ndim == 2 else 1 - permuted_scale = torch.empty( - (num_out_tokens,), dtype=torch.float32, device=inp.device + H = 128 + W = 32 + device = x.device + accum_token_count = torch.cumsum(token_count_per_expert, 0) + smoothed = scale is not None + if x_q is None: + # TODO(nanxiao): opt performance + x_q = torch.empty((out_tokens * N,), device=device, dtype=torch.float8_e4m3fn) + if x_scale is None: + x_scale = torch.empty((n_expert, N), device=device, dtype=torch.float32) + # import pydevd + # pydevd.settrace(suspend=False, trace_only_current_thread=True) + assert N % W == 0 + grid = (n_expert, N // W) + batch_transpose_smooth_fused_permute_kernel[grid]( + x, + scale, + org_smooth_scale, + smooth_scales, + indices, + token_count_per_expert, + accum_token_count, + x_q, + x_scale, + N, + n_expert, + H, + W, + smoothed, + round_scale, + num_stages=3, + num_warps=8, ) + return x_q, x_scale - sm = torch.cuda.get_device_properties(inp.device).multi_processor_count - T = triton.cdiv(num_tokens, sm) - grid = (num_experts, sm) - smooth_permute_with_mask_map_kernel[grid]( - inp, - output, - row_id_map, + +@triton.jit +def batch_block_pad_permute_with_indices_kernel( + x_ptr, + prob_ptr, + indices_ptr, + count_ptr, + xq_ptr, + xs_ptr, + xtq_ptr, + xts_ptr, + output_prob_ptr, + N, + E: tl.constexpr, + ROUND: tl.constexpr, + PROB: tl.constexpr, +): + eid = tl.program_id(axis=0) + rid = tl.program_id(axis=1) + cid = tl.program_id(axis=2) + + count = tl.load(count_ptr + eid) + counts = tl.load(count_ptr + tl.arange(0, E)) + + if rid >= tl.cdiv(count, 128): + return + + N = N.to(tl.int64) + nb = N // 128 + + padding_count = tl.cdiv(count, 16) * 16 + m_block = tl.sum(tl.where(tl.arange(0, E) < eid, tl.cdiv(counts, 128), 0)) + psi = tl.sum(tl.where(tl.arange(0, E) < eid, tl.cdiv(counts, 16) * 16, 0)) + + rids = rid * 128 + tl.arange(0, 128) + + indices = tl.load(indices_ptr + psi + rids, mask=rids < count) + + x = tl.load( + x_ptr + cid * 128 + indices[:, None] * N + tl.arange(0, 128)[None, :], + mask=rids[:, None] < count, + ).to(tl.float32) + + scale = tl.maximum(tl.max(tl.abs(x), 1) / 448.0, 1e-30) + if ROUND: + scale = tl.exp2(tl.ceil(tl.log2(scale))) + xq = x / scale[:, None] + + tl.store( + xs_ptr + psi * nb + cid * padding_count + rid * 128 + tl.arange(0, 128), scale, - smooth_scales, - permuted_scale, - num_tokens, - T, - hidden_size, - hs, - reverse, - round_scale + mask=rids < padding_count, + ) + tl.store( + xq_ptr + + psi * N + + rid * 128 * N + + cid * 128 + + tl.arange(0, 128)[:, None] * N + + tl.arange(0, 128)[None, :], + xq, + mask=rids[:, None] < padding_count, + ) + + scale = tl.maximum(tl.max(tl.abs(x), 0) / 448.0, 1e-30) + if ROUND: + scale = tl.exp2(tl.ceil(tl.log2(scale))) + xq = x / scale[None, :] + + tl.store(xts_ptr + m_block * N + rid * N + cid * 128 + tl.arange(0, 128), scale) + tl.store( + xtq_ptr + + psi * N + + rid * 128 + + cid * 128 * padding_count + + tl.arange(0, 128)[:, None] * padding_count + + tl.arange(0, 128)[None, :], + tl.trans(xq), + mask=rids[None, :] < padding_count, + ) + + if PROB: + prob = tl.load(prob_ptr + eid + indices * E, mask=rids < count) + tl.store( + output_prob_ptr + psi + rid * 128 + tl.arange(0, 128), + prob, + mask=rids < padding_count, + ) + + +def triton_batch_block_pad_permute_with_indices( + xs, token_count_per_expert, indices, splits, probs=None, round_scale=False +): + """ + select and quant, used in megatron 0.12 megatron flex backend + Args: + xs: [bs, dim] + token_count_per_expert: [n_experts] + indices: [n_experts*topk] + splits: python int list of token_count_per_expert + probs: route weights, [bs, n_experts] + round_scale: whether round scale to power of 2 + + Returns: + x_q: + x_scale: + xt_q: + xt_scale: + prob_output: + + """ + assert xs.is_contiguous() + m, N = xs.shape + assert N % 128 == 0 + n_experts = token_count_per_expert.size(0) + M = indices.shape[0] + device = xs.device + x_q = torch.empty((M, N), device=device, dtype=torch.float8_e4m3fn) + # intra layout is [N//128, m] + x_scale = torch.empty((M, N // 128), device=device, dtype=torch.float32) + blocks = sum([(x + 127) // 128 for x in splits]) + # intra layout is [N, m] + xt_q = torch.empty((M, N), device=device, dtype=torch.float8_e4m3fn) + # intra layout is [ceil(m/128), N] + xt_scale = torch.empty((blocks, N), device=device, dtype=torch.float32) + PROB = probs is not None + if PROB: + prob_output = torch.empty((M,), device=device, dtype=probs.dtype) + + else: + prob_output = None + + if m == 0: + return x_q, x_scale, xt_q, xt_scale, prob_output + + grid = (n_experts, triton.cdiv(max(splits), 128), N // 128) + batch_block_pad_permute_with_indices_kernel[grid]( + xs, + probs, + indices, + token_count_per_expert, + x_q, + x_scale, + xt_q, + xt_scale, + prob_output, + N, + n_experts, + round_scale, + PROB, + num_stages=2, + num_warps=4, + ) + + return x_q, x_scale, xt_q, xt_scale, prob_output + + +@triton.jit +def batch_mxfp8_permute_with_indices_kernel( + x_ptr, + prob_ptr, + indices_ptr, + count_ptr, + xq_ptr, + xs_ptr, + xtq_ptr, + xts_ptr, + output_prob_ptr, + N, + E: tl.constexpr, + B: tl.constexpr, + PROB: tl.constexpr, + OUTPUT_MODE: tl.constexpr, + ARRAY_PROB: tl.constexpr, +): + eid = tl.program_id(axis=0) + rid = tl.program_id(axis=1) + cid = tl.program_id(axis=2) + + count = tl.load(count_ptr + eid) + counts = tl.load(count_ptr + tl.arange(0, E)) + + if rid >= tl.cdiv(count, 128) * 4: + return + + N = N.to(tl.int64) + + m_block = tl.sum(tl.where(tl.arange(0, E) < eid, tl.cdiv(counts, 128), 0)) * 4 + si = tl.sum(tl.where(tl.arange(0, E) < eid, tl.cdiv(counts, 32) * 32, 0)) + + rids = rid * 32 + tl.arange(0, 32) + indices = tl.load(indices_ptr + si + rids, mask=rids < count) + + offs = ( + si * N + + rid * 32 * N + + cid * B + + tl.arange(0, 32)[:, None] * N + + tl.arange(0, B)[None, :] + ) + + pad_count = tl.cdiv(count, 32) * 32 + b = N // 32 + sb: tl.constexpr = B // 32 + + x = tl.load( + x_ptr + cid * B + indices[:, None] * N + tl.arange(0, B)[None, :], + mask=rids[:, None] < count, + other=0.0, + ).to(tl.float32) + + if OUTPUT_MODE % 2 == 0: + xr = tl.reshape(x, [32, sb, 32]) + valid_mask = (rids < count)[:, None] + scale = tl.maximum(tl.max(xr.abs(), 2) / 448, 1e-30) + log_scale = tl.ceil(tl.log2(scale)) + scale = tl.exp2(log_scale) + stored_log_scale = tl.where(valid_mask, log_scale + 127, 0).to(tl.uint8) + tl.store( + xs_ptr + + m_block * N + + rid * 32 * b + + cid * B // 32 + + tl.arange(0, 32)[:, None] * b + + tl.arange(0, sb), + stored_log_scale, + ) + + xq = tl.reshape(xr / scale[:, :, None], (32, B)).to(xq_ptr.dtype.element_ty) + tl.store(xq_ptr + offs, xq, mask=rids[:, None] < pad_count) + + if OUTPUT_MODE > 0: + valid_mask_t = rid * 32 < count + scale_t = tl.maximum(tl.max(x.abs(), 0) / 448, 1e-30) + log_scale_t = tl.ceil(tl.log2(scale_t)) + scale_t = tl.exp2(log_scale_t) + stored_log_scale_t = tl.where(valid_mask_t, log_scale_t + 127, 0).to(tl.uint8) + tl.store( + xts_ptr + m_block * N + rid * N + cid * B + tl.arange(0, B), + stored_log_scale_t, + ) + + xq = (x / scale_t).to(xtq_ptr.dtype.element_ty) + tl.store(xtq_ptr + offs, xq, mask=rids[:, None] < pad_count) + + if PROB: + if cid == 0: + if ARRAY_PROB: # for all to all + prob = tl.load( + prob_ptr + indices, mask=rid * 32 + tl.arange(0, 32) < count + ) + tl.store( + output_prob_ptr + si + rid * 32 + tl.arange(0, 32), + prob, + mask=rid * 32 + tl.arange(0, 32) < pad_count, + ) + else: + prob = tl.load( + prob_ptr + eid + indices * E, + mask=rid * 32 + tl.arange(0, 32) < count, + ) + tl.store( + output_prob_ptr + si + rid * 32 + tl.arange(0, 32), + prob, + mask=rid * 32 + tl.arange(0, 32) < pad_count, + ) + + +def triton_batch_mxfp8_permute_with_indices( + xs, + token_count_per_expert, + indices, + splits, + probs=None, + output_mode=2, + dispatch_type="alltoall", +): + """ + select and quant, use in dispatch type in deepep or alltoall. + Args: + xs: [bs, dim] + token_count_per_expert: [n_experts] + indices: [n_experts*topk] + splits: python int list of token_count_per_expert + probs: route weights, [bs, n_experts] + output_mode: one of {0, 1, 2} + 0: only output non-transposed quantized tensor + 1: only output transposed quantized tensor + 2: output both + dispatch_type: ("alltoall", "deepep") + + Returns: + x_q: + x_scale: + xt_q: + xt_scale: + prob_output: + + """ + assert xs.is_contiguous() + bs, N = xs.shape + n_experts = token_count_per_expert.size(0) + m = indices.shape[0] + device = xs.device + + assert N % 128 == 0 + M = sum([(x + 127) // 128 for x in splits]) * 128 + + x_q = torch.empty((m, N), device=device, dtype=torch.float8_e4m3fn) + x_scale = torch.empty((M, N // 32), device=device, dtype=torch.uint8) + xt_q = torch.empty((m, N), device=device, dtype=torch.float8_e4m3fn) + xt_scale = torch.empty((M // 32, N), device=device, dtype=torch.uint8) + + PROB = probs is not None + if PROB: + prob_output = torch.empty((m,), device=device, dtype=probs.dtype) + else: + prob_output = None + + ARRAY_PROB = dispatch_type == "alltoall" + + if bs == 0: + return x_q, x_scale, xt_q, xt_scale, prob_output + + B = 128 + grid = (n_experts, triton.cdiv(max(splits), 128) * 4, N // B) + batch_mxfp8_permute_with_indices_kernel[grid]( + xs, + probs, + indices, + token_count_per_expert, + x_q, + x_scale, + xt_q, + xt_scale, + prob_output, + N, + n_experts, + B, + PROB, + output_mode, + ARRAY_PROB, + num_stages=2, + num_warps=2, ) - return output, permuted_scale + + return x_q, x_scale, xt_q, xt_scale, prob_output diff --git a/linghe/utils/loss.py b/linghe/utils/loss.py index 1eae2fa..0c3cb6c 100644 --- a/linghe/utils/loss.py +++ b/linghe/utils/loss.py @@ -7,43 +7,51 @@ import triton import triton.language as tl + @triton.jit -def softmax_cross_entropy_forward_kernel(logit_ptr, label_ptr, loss_ptr, - sum_exp_ptr, max_logit_ptr, N, - B: tl.constexpr): +def softmax_cross_entropy_forward_kernel( + logit_ptr, + label_ptr, + loss_ptr, + sum_exp_ptr, + max_logit_ptr, + N, + ignore_index, + B: tl.constexpr, +): pid = tl.program_id(axis=0).to(tl.int64) label = tl.load(label_ptr + pid) + if label == ignore_index: + tl.store(sum_exp_ptr + pid, 0.0) + tl.store(max_logit_ptr + pid, 0.0) + tl.store(loss_ptr + pid, 0.0) + return + sum_exp = 0.0 + sum_exp = sum_exp.to(tl.float64) T = tl.cdiv(N, B) - max_logit = -1e30 + max_logit = -1e9 for i in range(T): - logit = tl.load(logit_ptr + pid * N + i * B + tl.arange(0, B), - mask=i * B + tl.arange(0, B) < N, other=-1e30).to( - tl.float32) - max_logit = tl.maximum(max_logit, tl.max(logit)) - sum_exp += tl.sum(tl.exp(logit)) - - retry = sum_exp > 3.389e38 - max_logit = tl.where(retry, max_logit, 0.0) - retry_sum_exp = 0.0 - if retry: - for i in range(T): - logit = tl.load(logit_ptr + pid * N + i * B + tl.arange(0, B), - mask=i * B + tl.arange(0, B) < N, other=-1e30).to( - tl.float32) - retry_sum_exp += tl.sum(tl.exp(logit - max_logit)) - sum_exp = tl.where(retry, retry_sum_exp, sum_exp) + logit = tl.load( + logit_ptr + pid * N + i * B + tl.arange(0, B), + mask=i * B + tl.arange(0, B) < N, + other=-1e10, + ).to(tl.float32) + latest_max_logit = tl.maximum(max_logit, tl.max(logit)) + + sum_exp = sum_exp * tl.exp(max_logit - latest_max_logit) + tl.sum( + tl.exp(logit - latest_max_logit) + ) + max_logit = latest_max_logit + tl.store(sum_exp_ptr + pid, sum_exp) + tl.store(max_logit_ptr + pid, max_logit) target_logit = tl.load(logit_ptr + pid * N + label) loss = tl.log(sum_exp) - (target_logit - max_logit) tl.store(loss_ptr + pid, loss) - tl.store(max_logit_ptr + pid, max_logit) -""" -TODO: support distributed loss with pytorch ongoing nvshmem feature -""" -def triton_softmax_cross_entropy_forward(logits, labels): +def triton_softmax_cross_entropy_forward(logits, labels, ignore_index=-100): """ compute token-wise softmax cross entropy loss Args: @@ -55,10 +63,11 @@ def triton_softmax_cross_entropy_forward(logits, labels): """ M, N = logits.shape device = logits.device + assert logits.is_contiguous() and labels.is_contiguous() loss = torch.empty((M,), device=device, dtype=torch.float32) sum_exp = torch.empty((M,), device=device, dtype=torch.float32) max_logit = torch.empty((M,), device=device, dtype=torch.float32) - B = 4096 + B = 2048 grid = (M,) softmax_cross_entropy_forward_kernel[grid]( logits, @@ -67,41 +76,83 @@ def triton_softmax_cross_entropy_forward(logits, labels): sum_exp, max_logit, N, + ignore_index, B, num_stages=3, - num_warps=8 + num_warps=2, ) return loss, sum_exp, max_logit @triton.jit -def softmax_cross_entropy_backward_kernel(logit_ptr, label_ptr, sum_exp_ptr, - max_logit_ptr, - input_grad_ptr, output_grad_ptr, - N, B: tl.constexpr): +def softmax_cross_entropy_backward_kernel( + logit_ptr, + label_ptr, + sum_exp_ptr, + max_logit_ptr, + output_grad_ptr, + input_grad_ptr, + N, + ignore_index, + B: tl.constexpr, + INPLACE: tl.constexpr, +): pid = tl.program_id(axis=0).to(tl.int64) + T = tl.cdiv(N, B) label = tl.load(label_ptr + pid) - input_grad = tl.load(input_grad_ptr + pid).to(tl.float32) + if label == ignore_index: + for i in range(T): + grad = tl.zeros((B,), dtype=tl.float32) + if INPLACE: + tl.store( + logit_ptr + pid * N + i * B + tl.arange(0, B), + grad, + mask=i * B + tl.arange(0, B) < N, + ) + else: + tl.store( + input_grad_ptr + pid * N + i * B + tl.arange(0, B), + grad, + mask=i * B + tl.arange(0, B) < N, + ) + return + + output_grad = tl.load(output_grad_ptr + pid).to(tl.float32) sum_exp = tl.load(sum_exp_ptr + pid) max_logit = tl.load(max_logit_ptr + pid) - coef = input_grad / sum_exp - T = tl.cdiv(N, B) + coef = output_grad / sum_exp + target_logit = tl.load(logit_ptr + pid * N + label).to(tl.float32) + target_grad = (tl.exp(target_logit - max_logit) / sum_exp - 1) * output_grad + tl.debug_barrier() # must add barrier here, or it may read stored values for i in range(T): - logit = tl.load(logit_ptr + pid * N + i * B + tl.arange(0, B), - mask=i * B + tl.arange(0, B) < N, other=-1e30).to( - tl.float32) + logit = tl.load( + logit_ptr + pid * N + i * B + tl.arange(0, B), + mask=i * B + tl.arange(0, B) < N, + other=-1e10, + ).to(tl.float32) grad = tl.exp(logit - max_logit) * coef - tl.store(output_grad_ptr + pid * N + i * B + tl.arange(0, B), grad, - mask=i * B + tl.arange(0, B) < N) - tl.debug_barrier() - target_grad = tl.load(output_grad_ptr + pid * N + label) - target_grad -= input_grad - tl.store(output_grad_ptr + pid * N + label, target_grad) + if INPLACE: + tl.store( + logit_ptr + pid * N + i * B + tl.arange(0, B), + grad, + mask=i * B + tl.arange(0, B) < N, + ) + else: + tl.store( + input_grad_ptr + pid * N + i * B + tl.arange(0, B), + grad, + mask=i * B + tl.arange(0, B) < N, + ) + tl.debug_barrier() # must add barrier here, or it may execute before loop + if INPLACE: + tl.store(logit_ptr + pid * N + label, target_grad) + else: + tl.store(input_grad_ptr + pid * N + label, target_grad) -def triton_softmax_cross_entropy_backward(logits, labels, sum_exp, max_logit, - input_grad, - output_grad=None): +def triton_softmax_cross_entropy_backward( + logits, labels, sum_exp, max_logit, output_grad, ignore_index=-100, inplace=False +): """ backward of softmax cross entropy loss Args: @@ -109,27 +160,414 @@ def triton_softmax_cross_entropy_backward(logits, labels, sum_exp, max_logit, labels: label tensor, [bs] sum_exp: [bs] max_logit: [bs] - input_grad: gradient, [bs, dim] + output_grad: gradient, [bs, dim] + inplace: whether to reuse logits as gradient Returns: - output_grad: [bs, dim] + grad of input: [bs, dim] """ + assert output_grad.is_contiguous() M, N = logits.shape device = logits.device - if output_grad is None: - output_grad = torch.empty((M, N), device=device, dtype=logits.dtype) - B = 4096 + if not inplace: + dx = torch.empty((M, N), device=device, dtype=logits.dtype) + else: + dx = None + B = 2048 grid = (M,) softmax_cross_entropy_backward_kernel[grid]( logits, labels, sum_exp, max_logit, - input_grad, output_grad, + dx, N, + ignore_index, B, + inplace, num_stages=3, - num_warps=8 + num_warps=8, + ) + if inplace: + dx = logits + return dx + + +@triton.jit +def parallel_logit_stat_kernel( + logit_ptr, + label_ptr, + sum_exp_ptr, + max_logit_ptr, + target_logit_ptr, + N, + ignore_index, + group_rank, + group_size, + B: tl.constexpr, +): + pid = tl.program_id(axis=0).to(tl.int64) + label = tl.load(label_ptr + pid) + + if label == ignore_index: + tl.store(sum_exp_ptr + pid, 0.0) + tl.store(max_logit_ptr + pid, 0.0) + tl.store(target_logit_ptr + pid, 0.0) + return + + sum_exp = 0.0 + sum_exp = sum_exp.to(tl.float64) + T = tl.cdiv(N, B) + max_logit = -1e9 + for i in range(T): + logit = tl.load( + logit_ptr + pid * N + i * B + tl.arange(0, B), + mask=i * B + tl.arange(0, B) < N, + other=-1e10, + ).to(tl.float32) + latest_max_logit = tl.maximum(max_logit, tl.max(logit)) + + sum_exp = sum_exp * tl.exp(max_logit - latest_max_logit) + tl.sum( + tl.exp(logit - latest_max_logit) + ) + max_logit = latest_max_logit + + tl.store(sum_exp_ptr + pid, sum_exp) + tl.store(max_logit_ptr + pid, max_logit) + + if label // N == group_rank: + target_logit = tl.load(logit_ptr + pid * N + label % N).to(tl.float32) + else: + target_logit = float("-inf") + tl.store(target_logit_ptr + pid, target_logit) + + +@triton.jit +def parallel_calc_loss_kernel( + label_ptr, + stats, + sum_exp_ptr, + max_logit_ptr, + loss_ptr, + M, + N, + ignore_index, + group_size, +): + pid = tl.program_id(axis=0).to(tl.int64) + label = tl.load(label_ptr + pid) + if label == ignore_index: + tl.store(loss_ptr + pid, 0.0) + tl.store(sum_exp_ptr + pid, 0.0) + tl.store(max_logit_ptr + pid, 0.0) + return + + sum_exp = 0.0 + sum_exp = sum_exp.to(tl.float64) + max_logit = -1e9 + tg = float("-inf") # target logit + for i in range(group_size): + se = tl.load(stats + i * M * 3 + pid) + ml = tl.load(stats + i * M * 3 + M + pid) + tg = tl.maximum(tl.load(stats + i * M * 3 + 2 * M + pid), tg) + latest_max_logit = tl.maximum(max_logit, ml) + sum_exp = sum_exp * tl.exp(max_logit - latest_max_logit) + se * tl.exp( + ml - latest_max_logit + ) + max_logit = latest_max_logit + + loss = tl.log(sum_exp) - (tg - max_logit) + tl.store(loss_ptr + pid, loss) + tl.store(sum_exp_ptr + pid, sum_exp) + tl.store(max_logit_ptr + pid, max_logit) + + +""" +TODO1: support distributed loss with pytorch ongoing nvshmem feature +TODO2: optimize performance when vocab size is not multiple of 16 +""" + + +def triton_parallel_softmax_cross_entropy_forward( + logits, labels, group, ignore_index=-100 +): + """ + compute token-wise softmax cross entropy loss + Args: + logits: logits tensor + labels: labels tensor + + Returns: + loss of each token + """ + M, N = logits.shape + device = logits.device + assert logits.is_contiguous() and labels.is_contiguous() + loss = torch.empty((M,), device=device, dtype=torch.float32) + + group_size = group.size() + group_rank = group.rank() + stats = torch.empty((3, M), device=device, dtype=torch.float32) + sum_exp = stats[0] + max_logit = stats[1] + target_logit = stats[2] + statistic = torch.empty((3 * group_size, M), device=device, dtype=torch.float32) + B = 2048 + grid = (M,) + parallel_logit_stat_kernel[grid]( + logits, + labels, + sum_exp, + max_logit, + target_logit, + N, + ignore_index, + group_rank, + group_size, + B, + num_stages=3, + num_warps=2, + ) + torch.distributed.all_gather_into_tensor(statistic, stats, group=group) + parallel_calc_loss_kernel[grid]( + labels, + statistic, + sum_exp, + max_logit, + loss, + M, + N, + ignore_index, + group_size, + num_stages=3, + num_warps=2, + ) + return loss, sum_exp, max_logit + + +@triton.jit +def parallel_softmax_cross_entropy_backward_kernel( + logit_ptr, + label_ptr, + sum_exp_ptr, + max_logit_ptr, + output_grad_ptr, + input_grad_ptr, + N, + ignore_index, + group_rank, + group_size, + B: tl.constexpr, + INPLACE: tl.constexpr, +): + pid = tl.program_id(axis=0).to(tl.int64) + label = tl.load(label_ptr + pid) + T = tl.cdiv(N, B) + + if label == ignore_index: + for i in range(T): + grad = tl.zeros((B,), dtype=tl.float32) + if INPLACE: + tl.store( + logit_ptr + pid * N + i * B + tl.arange(0, B), + grad, + mask=i * B + tl.arange(0, B) < N, + ) + else: + tl.store( + input_grad_ptr + pid * N + i * B + tl.arange(0, B), + grad, + mask=i * B + tl.arange(0, B) < N, + ) + return + + output_grad = tl.load(output_grad_ptr + pid).to(tl.float32) + sum_exp = tl.load(sum_exp_ptr + pid) + max_logit = tl.load(max_logit_ptr + pid) + coef = output_grad / sum_exp + tl.debug_barrier() # must add barrier here, or it may read stored values + for i in range(T): + logit = tl.load( + logit_ptr + pid * N + i * B + tl.arange(0, B), + mask=i * B + tl.arange(0, B) < N, + other=-1e10, + ).to(tl.float32) + grad = tl.exp(logit - max_logit) * coef + if INPLACE: + tl.store( + logit_ptr + pid * N + i * B + tl.arange(0, B), + grad, + mask=i * B + tl.arange(0, B) < N, + ) + else: + tl.store( + input_grad_ptr + pid * N + i * B + tl.arange(0, B), + grad, + mask=i * B + tl.arange(0, B) < N, + ) + tl.debug_barrier() # must add barrier here, or it may execute before loop + + if label // N == group_rank: + target_logit = tl.load(logit_ptr + pid * N + label % N).to(tl.float32) + target_grad = (tl.exp(target_logit - max_logit) / sum_exp - 1) * output_grad + if INPLACE: + tl.store(logit_ptr + pid * N + label % N, target_grad) + else: + tl.store(input_grad_ptr + pid * N + label % N, target_grad) + + +def triton_parallel_softmax_cross_entropy_backward( + logits, + labels, + sum_exp, + max_logit, + output_grad, + group, + ignore_index=-100, + inplace=False, +): + """ + backward of softmax cross entropy loss + Args: + logits: logit tensor, [bs, dim] + labels: label tensor, [bs] + sum_exp: [bs] + max_logit: [bs] + output_grad: gradient, [bs, dim] + inplace: whether to reuse logits as gradient + + Returns: + grad of input: [bs, dim] + """ + assert output_grad.is_contiguous() + M, N = logits.shape + device = logits.device + if not inplace: + dx = torch.empty((M, N), device=device, dtype=logits.dtype) + else: + dx = None + + group_size = group.size() + group_rank = group.rank() + + B = 2048 + grid = (M,) + parallel_softmax_cross_entropy_backward_kernel[grid]( + logits, + labels, + sum_exp, + max_logit, + output_grad, + dx, + N, + ignore_index, + group_rank, + group_size, + B, + inplace, + num_stages=3, + num_warps=8, + ) + if inplace: + dx = logits + return dx + + +@triton.jit +def moe_z_loss_forward_kernel( + logit_ptr, loss_ptr, coef, T: tl.constexpr, D: tl.constexpr +): + pid = tl.program_id(axis=0) + + logit = tl.load( + logit_ptr + pid * T * D + tl.arange(0, T)[:, None] * D + tl.arange(0, D) + ).to(tl.float32) + max_logit = tl.max(logit, 1) + lse = tl.log(tl.sum(tl.exp(logit - max_logit[:, None]), 1)) + max_logit + loss = coef / T * tl.sum(lse * lse) + + tl.store(loss_ptr + pid, loss) + + +def triton_moe_z_loss_forward(logits, coef=1e-6): + """ + compute moe z loss, + z_loss = torch.mean(torch.square(torch.logsumexp(logits, dim=-1))) * coef + Args: + logits: logits tensor + coef: z loss coef + Returns: + z loss + """ + assert logits.is_contiguous() + shape = logits.shape + if len(shape) == 3: + L, B, D = logits.shape + M = L * B + else: + M, D = logits.shape + device = logits.device + T = 4 + assert M % T == 0 + loss = torch.empty((M // T,), device=device, dtype=torch.float32) + grid = (M // T,) + moe_z_loss_forward_kernel[grid](logits, loss, coef, T, D, num_stages=3, num_warps=1) + return loss.mean() + + +@triton.jit +def moe_z_loss_backward_kernel( + input_grad_ptr, logit_ptr, output_grad_ptr, coef, T: tl.constexpr, D: tl.constexpr +): + pid = tl.program_id(axis=0) + n_tokens = tl.num_programs(axis=0) * T + grad = tl.load(input_grad_ptr).to(tl.float32) + + logit = tl.load( + logit_ptr + + pid * T * D + + tl.arange(0, T)[:, None] * D + + tl.arange(0, D)[None, :] + ).to(tl.float32) + max_logit = tl.max(logit, 1, keep_dims=True) + e = tl.exp(logit - max_logit) + se = tl.sum(e, 1, keep_dims=True) + lse = tl.log(se) + max_logit + + grads = 2 * coef / n_tokens * grad * lse * e / se + + tl.store( + output_grad_ptr + pid * T * D + tl.arange(0, T)[:, None] * D + tl.arange(0, D), + grads, + ) + + +def triton_moe_z_loss_backward(grads, logits, coef=1e-6): + """ + backward of moe z loss + Args: + grads: grad scalar tensor + logits: logit tensor, [L, B, dim] + coef: python scalar + Returns: + output_grad: [L, B, dim] + """ + assert grads.is_contiguous() + device = logits.device + shape = logits.shape + if len(shape) == 3: + L, B, D = logits.shape + M = L * B + output_grad = torch.empty((L, B, D), device=device, dtype=logits.dtype) + else: + M, D = logits.shape + output_grad = torch.empty((M, D), device=device, dtype=logits.dtype) + + T = 4 + assert M % T == 0 + grid = (M // T,) + moe_z_loss_backward_kernel[grid]( + grads, logits, output_grad, coef, T, D, num_stages=3, num_warps=1 ) return output_grad diff --git a/linghe/utils/mul.py b/linghe/utils/mul.py new file mode 100644 index 0000000..f324001 --- /dev/null +++ b/linghe/utils/mul.py @@ -0,0 +1,142 @@ +# -*- coding: utf-8 -*- +""" +Copyright (c) Ant Financial Service Group and its affiliates. +""" + +import torch +import triton +import triton.language as tl + + +@triton.jit +def dot_kernel(x_ptr, y_ptr, sum_ptr, M, N, H: tl.constexpr, W: tl.constexpr): + # rowwise read, rowwise write + pid = tl.program_id(axis=0) + offs = pid * W * N + tl.arange(0, W)[:, None] * N + tl.arange(0, H)[None, :] + + n = tl.cdiv(N, H) + sums = tl.zeros((W,), dtype=tl.float32) + for i in range(n): + x = tl.load(x_ptr + offs).to(tl.float32) + y = tl.load(y_ptr + offs).to(tl.float32) + sums += tl.sum(x * y, axis=1) + offs += H + + tl.store(sum_ptr + pid * W + tl.arange(0, W), sums) + + +def triton_dot(x, y): + """ + vector dot multiply, output = sum(x*y, 1), + it is used to calculate gradient of router weight + Args: + x: + y: + + Returns: + output of sum(x*y, 1) + """ + assert x.is_contiguous() and y.is_contiguous() + M, N = x.shape + H = 128 + W = 16 + assert M % W == 0 + + num_stages = 5 + num_warps = 8 + device = x.device + s = torch.empty((M,), device=device, dtype=x.dtype) + grid = (triton.cdiv(M, W),) + dot_kernel[grid](x, y, s, M, N, H, W, num_stages=num_stages, num_warps=num_warps) + return s + + +@triton.jit +def inplace_scale_kernel(input_ptr, scale, m, B: tl.constexpr): + pid = tl.program_id(axis=0).to(tl.int64) + + offs = pid * B + tl.arange(0, B) + x = tl.load(input_ptr + offs, mask=offs < m) + x = x * scale + tl.store(input_ptr + offs, x, mask=offs < m) + + +def triton_inplace_scale(x, scale): + """ + inplace scale a tensor. + Args: + x: Tensor. + scale: a python float scale + Returns: + x + """ + assert x.is_contiguous() + B = 512 + m = x.numel() + grid = (triton.cdiv(m, B),) + inplace_scale_kernel[grid](x, scale, m, B, num_stages=2, num_warps=2) + return x + + +@triton.jit +def batch_scale_kernel( + input_ptrs, + size_ptr, + scale, + DT: tl.constexpr, + B: tl.constexpr, + ZERO: tl.constexpr, +): + tid = tl.program_id(axis=0) + bid = tl.program_id(axis=1) + T = tl.num_programs(axis=1) + + size = tl.load(size_ptr + tid) + if DT == 0: + input_ptr = tl.load(input_ptrs + tid).to(tl.pointer_type(tl.float32)) + else: + input_ptr = tl.load(input_ptrs + tid).to(tl.pointer_type(tl.bfloat16)) + t = tl.cdiv(size, B * T) + offs = bid.to(tl.int64) * t * B + tl.arange(0, B) + for i in range(t): + x = tl.load(input_ptr + offs, mask=offs < size, other=0).to(tl.float32) + if ZERO: + x = tl.where(tl.abs(x) == float("inf"), 1.0, 0.0) + else: + x = x * scale + tl.store(input_ptr + offs, x, mask=offs < size) + offs += B + + +def triton_batch_scale(xs, scale): + """ + return [x*scale for x in xs], + used to scale gradient. + Args: + xs: Tensor lists. + scale: a python float scale + Returns: + xs + """ + if len(xs) == 0: + return + dtype = xs[0].dtype + assert dtype in (torch.float32, torch.bfloat16) + assert all([x.is_contiguous() and x.dtype == dtype for x in xs]) + + device = xs[0].device + sizes = torch.tensor([x.numel() for x in xs], dtype=torch.int64).cuda( + device, non_blocking=True + ) + ptrs = torch.tensor([x.data_ptr() for x in xs], dtype=torch.int64).cuda( + device, non_blocking=True + ) + + DT = 0 if dtype == torch.float32 else 1 + T = 256 + tensor_count = len(xs) + B = 512 + ZERO = scale == 0.0 + grid = (tensor_count, T) + batch_scale_kernel[grid](ptrs, sizes, scale, DT, B, ZERO, num_stages=2, num_warps=2) + return xs diff --git a/linghe/utils/norm.py b/linghe/utils/norm.py index 3388a5b..e7b393e 100644 --- a/linghe/utils/norm.py +++ b/linghe/utils/norm.py @@ -1,278 +1,355 @@ -from re import S +# -*- coding: utf-8 -*- +""" +Copyright (c) Ant Financial Service Group and its affiliates. +""" + +from typing import Optional + import torch import triton import triton.language as tl -from typing import Optional @triton.jit -def rms_norm_forward_kernel(x_ptr, weight_ptr, out_ptr, eps, M, T, - N: tl.constexpr, W: tl.constexpr): +def rms_norm_forward_kernel( + x_ptr, + weight_ptr, + out_ptr, + rms_ptr, + eps, + M, + T, + n, + N: tl.constexpr, + W: tl.constexpr, + REUSE: tl.constexpr, +): pid = tl.program_id(axis=0) - weight = tl.load(weight_ptr + tl.arange(0, N)).to(tl.float32)[None, :] + weight = tl.load(weight_ptr + tl.arange(0, N), mask=tl.arange(0, N) < n).to( + tl.float32 + )[None, :] - offs = pid * W * T * N + tl.arange(0, W)[:, None] * N + tl.arange(0, N)[ - None, :] + offs = pid * W * T * n + tl.arange(0, W)[:, None] * n + tl.arange(0, N)[None, :] for i in range(T): - x = tl.load(x_ptr + offs, - mask=pid * W * T + i * W + tl.arange(0, W)[:, None] < M).to( - tl.float32) - rms = tl.sqrt(tl.sum(x * x, axis=1) / N + eps) + mask = (pid * W * T + i * W + tl.arange(0, W)[:, None] < M) & ( + tl.arange(0, N) < n + ) + x = tl.load(x_ptr + offs, mask=mask).to(tl.float32) + if REUSE: + rms = tl.load( + rms_ptr + pid * W * T + i * W + tl.arange(0, W), + mask=pid * W * T + i * W + tl.arange(0, W) < M, + other=1.0, + ) + else: + rms = tl.rsqrt(tl.sum(x * x, axis=1) / n + eps) + tl.store( + rms_ptr + pid * W * T + i * W + tl.arange(0, W), + rms, + mask=pid * W * T + i * W + tl.arange(0, W) < M, + ) - x = (x / rms[:, None]) * weight + x = (x * rms[:, None]) * weight - tl.store(out_ptr + offs, x, - mask=pid * W * T + i * W + tl.arange(0, W)[:, None] < M) - offs += N * W + tl.store(out_ptr + offs, x, mask=mask) + offs += n * W -def triton_rms_norm_forward(x, weight, eps=1e-6, out=None): +def triton_rms_norm_forward(x, weight, eps=1e-6, out=None, rms=None): """ rms norm Args: x: input tensor weight: weight of rms norm eps: epsilon of rms norm + rms: use x*rms to calculate output if rms is not None, + it will accelerate recompute of rms norm Returns: out: output tensor + rms: 1/rms of input tensor """ - # row-wise read, row-wise write - M, N = x.shape + assert x.is_contiguous() and weight.is_contiguous() + shape = x.shape + assert len(shape) in (2, 3) + if len(shape) == 3: + M, n = shape[0] * shape[1], shape[2] + else: + M, n = x.shape + N = triton.next_power_of_2(n) W = 8192 // N T = 4 - assert N <= 8192 and M % (W*T) == 0 + assert N <= 8192 device = x.device if out is None: - out = torch.empty((M, N), device=device, dtype=x.dtype) + out = torch.empty_like(x) + REUSE = rms is not None + if not REUSE: + rms = torch.empty((M,), device=device, dtype=torch.float32) - grid = (M//(T*W),) + grid = (triton.cdiv(M, T * W),) rms_norm_forward_kernel[grid]( - x, - weight, - out, - eps, - M, - T, - N, - W, - num_stages=3, - num_warps=4 + x, weight, out, rms, eps, M, T, n, N, W, REUSE, num_stages=3, num_warps=4 ) - return out + return out, rms @triton.jit def rms_norm_backward_kernel( - grad_output_ptr, - x_ptr, - w_ptr, - dx_ptr, - dw_ptr, - eps, - M, - T, - N: tl.constexpr, - W: tl.constexpr + grad_output_ptr, + x_ptr, + w_ptr, + rms_ptr, + dx_ptr, + dw_ptr, + eps, + M, + T, + n, + N: tl.constexpr, + W: tl.constexpr, + REUSE: tl.constexpr, ): pid = tl.program_id(0) - w = tl.load(w_ptr + tl.arange(0, N)).to(tl.float32) + w = tl.load(w_ptr + tl.arange(0, N), mask=tl.arange(0, N) < n).to(tl.float32) - offs = pid * W * T * N + tl.arange(0, W)[:, None] * N + tl.arange(0, N)[ - None, :] + offs = pid * W * T * n + tl.arange(0, W)[:, None] * n + tl.arange(0, N)[None, :] w_grads = tl.zeros((N,), dtype=tl.float32) for i in range(T): - mask = pid * T + i < M + mask = (pid * W * T + i * W + tl.arange(0, W)[:, None] < M) & ( + tl.arange(0, N) < n + ) + x = tl.load(x_ptr + offs, mask=mask).to(tl.float32) g = tl.load(grad_output_ptr + offs, mask=mask).to(tl.float32) - rms = tl.sqrt(tl.sum(x * x, 1) / N + eps) - r = 1.0 / rms[:, None] + if REUSE: + r = tl.load( + rms_ptr + pid * W * T + i * W + tl.arange(0, W), + mask=pid * W * T + i * W + tl.arange(0, W) < M, + )[:, None] + else: + r = tl.rsqrt(tl.sum(x * x, 1) / n + eps)[:, None] w_grad = x * g * r w_grads += tl.sum(w_grad, 0) - dx = r * g * w - r * r * r * x * tl.sum(x * g * w, 1, keep_dims=True) / N + dx = r * g * w - r * r * r * x * tl.sum(x * g * w, 1, keep_dims=True) / n tl.store(dx_ptr + offs, dx, mask=mask) - offs += N * W + offs += n * W - tl.store(dw_ptr + pid * N + tl.arange(0, N), w_grads) + tl.store(dw_ptr + pid * n + tl.arange(0, N), w_grads, mask=tl.arange(0, N) < n) -def triton_rms_norm_backward(grad_output, x, w, eps=1e-6): - M, N = x.shape - dx = torch.empty(M, N, dtype=x.dtype, device=x.device) +def triton_rms_norm_backward(grad_output, x, w, eps=1e-6, rms=None): + assert grad_output.is_contiguous() + shape = x.shape + if len(shape) == 3: + M, n = shape[0] * shape[1], shape[2] + else: + M, n = x.shape + N = triton.next_power_of_2(n) + assert N <= 8192 + + dx = torch.empty_like(x) + REUSE = rms is not None W = 8192 // N T = 16 - assert 8192 % N ==0 and M % (T*W) == 0 - g = M//(T*W) - tmp_dw = torch.empty(g, N, dtype=torch.float32, device=w.device) + g = triton.cdiv(M, T * W) + tmp_dw = torch.empty(g, n, dtype=torch.float32, device=w.device) grid = (g,) rms_norm_backward_kernel[grid]( grad_output, x, w, + rms, dx, tmp_dw, eps, M, T, + n, N, W, + REUSE, num_stages=3, - num_warps=4 + num_warps=4, ) - return dx, tmp_dw.sum(dim=0).to(w.dtype) + return dx, tmp_dw.sum(dim=0) # output non-transposed and transposed together -# should used with batchsize >= 16384 +# performance is bad with batchsize < 16384 @triton.jit -def rms_norm_and_block_quant_forward_kernel(x_ptr, - weight_ptr, - out_ptr, - scale_ptr, - transpose_output_ptr, - transpose_scale_ptr, - rms_ptr, - eps, - M, - T: tl.constexpr, - N: tl.constexpr, - nb: tl.constexpr, - W: tl.constexpr, - H : tl.constexpr, - ROUND: tl.constexpr): +def rms_norm_and_block_quant_forward_kernel( + x_ptr, + weight_ptr, + out_ptr, + scale_ptr, + transpose_output_ptr, + transpose_scale_ptr, + rms_ptr, + eps, + M, + n, + N: tl.constexpr, + T: tl.constexpr, + W: tl.constexpr, + H: tl.constexpr, + ROUND: tl.constexpr, +): pid = tl.program_id(axis=0) - # row-wise read, row-wise write - weight = tl.load(weight_ptr + tl.arange(0, N)).to(tl.float32)[None, :] - offs = pid * W * T * N + tl.arange(0, W)[:, None] * N + tl.arange(0, N)[ - None, :] + NB: tl.constexpr = N // 128 + nb = n // 128 + + mask = tl.arange(0, N) < n + weight = tl.load(weight_ptr + tl.arange(0, N), mask=mask).to(tl.float32)[None, :] + offs = pid * W * T * n + tl.arange(0, W)[:, None] * n + tl.arange(0, N)[None, :] for i in range(T): indices = pid * W * T + i * W + tl.arange(0, W) - x = tl.load(x_ptr + offs, mask=indices[:, None] < M).to(tl.float32) - rms = 1/tl.sqrt(tl.sum(x * x, axis=1) / N + eps) + masks = (indices[:, None] < M) & (tl.arange(0, N) < n) + x = tl.load(x_ptr + offs, mask=masks).to(tl.float32) + rms = tl.rsqrt(tl.sum(x * x, axis=1) / n + eps) tl.store(rms_ptr + indices, rms, mask=indices < M) x = (x * rms[:, None]) * weight - x = tl.reshape(x, [W, nb, 128]) + x = tl.reshape(x, [W, NB, 128]) scale = tl.maximum(tl.max(tl.abs(x), 2) / 448.0, 1e-30) if ROUND: scale = tl.exp2(tl.ceil(tl.log2(scale))) - x = (x / scale[:,:, None]).to(out_ptr.dtype.element_ty) + x = (x / scale[:, :, None]).to(out_ptr.dtype.element_ty) x = tl.reshape(x, [W, N]) - tl.store(scale_ptr + indices[:, None] * nb + tl.arange(0, nb)[None, :], scale, mask=indices[:, None] < M) - tl.store(out_ptr + offs, x, mask=indices[:, None] < M) - offs += N * W - + tl.store( + scale_ptr + tl.arange(0, NB)[:, None] * M + indices[None, :], + tl.trans(scale), + mask=(indices[None, :] < M) & (tl.arange(0, NB)[:, None] < nb), + ) + tl.store(out_ptr + offs, x, mask=masks) + offs += n * W - offs = pid * W * T * N + tl.arange(0, 128)[:, None] * N + tl.arange(0, H)[ - None, :] - toffs = pid * 128 + tl.arange(0, H)[:, None] * M + tl.arange(0, 128)[ - None, :] + offs = pid * W * T * n + tl.arange(0, 128)[:, None] * n + tl.arange(0, H)[None, :] + toffs = pid * 128 + tl.arange(0, H)[:, None] * M + tl.arange(0, 128)[None, :] indices = pid * W * T + tl.arange(0, 128) tl.debug_barrier() rms = tl.load(rms_ptr + indices, mask=indices < M)[:, None] - for i in range(N//H): - x = tl.load(x_ptr + offs, mask=indices[:, None] < M).to(tl.float32) - wgt = tl.load(weight_ptr + i * H + tl.arange(0, H)).to(tl.float32) + for i in range(n // H): + x = tl.load(x_ptr + offs, mask=indices[:, None] < M).to(tl.float32) + wgt = tl.load(weight_ptr + i * H + tl.arange(0, H)).to(tl.float32) x = x * rms * wgt scale = tl.maximum(tl.max(x.abs(), 0) / 448.0, 1e-30) if ROUND: scale = tl.exp2(tl.ceil(tl.log2(scale))) - tl.store(transpose_scale_ptr + pid * N + i * H + tl.arange(0, H), scale) - x = (x/scale).to(transpose_output_ptr.dtype.element_ty) + tl.store(transpose_scale_ptr + pid * n + i * H + tl.arange(0, H), scale) + + x = (x / scale).to(transpose_output_ptr.dtype.element_ty) tl.store(transpose_output_ptr + toffs, tl.trans(x), mask=indices[None, :] < M) offs += H toffs += M * H - # output non-transposed tensor only @triton.jit -def rms_norm_and_block_quant_forward_n_kernel(x_ptr, - weight_ptr, - out_ptr, - scale_ptr, - rms_ptr, - eps, - M: tl.constexpr, - T: tl.constexpr, - N: tl.constexpr, - nb: tl.constexpr, - W: tl.constexpr, - ROUND: tl.constexpr): +def rms_norm_and_block_quant_forward_n_kernel( + x_ptr, + weight_ptr, + out_ptr, + scale_ptr, + rms_ptr, + eps, + M, + n, + N: tl.constexpr, + T: tl.constexpr, + W: tl.constexpr, + ROUND: tl.constexpr, +): pid = tl.program_id(axis=0) - # row-wise read, row-wise write - weight = tl.load(weight_ptr + tl.arange(0, N)).to(tl.float32)[None, :] - offs = pid * W * T * N + tl.arange(0, W)[:, None] * N + tl.arange(0, N)[ - None, :] + NB: tl.constexpr = N // 128 + + mask = tl.arange(0, N) < n + weight = tl.load(weight_ptr + tl.arange(0, N), mask=mask).to(tl.float32)[None, :] + offs = pid * W * T * n + tl.arange(0, W)[:, None] * n + tl.arange(0, N)[None, :] for i in range(T): indices = pid * W * T + i * W + tl.arange(0, W) - x = tl.load(x_ptr + offs, mask=indices[:, None] < M).to(tl.float32) - rms = tl.rsqrt(tl.sum(x * x, axis=1) / N + eps) + masks = (indices[:, None] < M) & (tl.arange(0, N) < n) + x = tl.load(x_ptr + offs, mask=masks).to(tl.float32) + rms = tl.rsqrt(tl.sum(x * x, axis=1) / n + eps) tl.store(rms_ptr + indices, rms, mask=indices < M) x = x * rms[:, None] * weight - x = tl.reshape(x, [W, nb, 128]) + x = tl.reshape(x, [W, NB, 128]) scale = tl.maximum(tl.max(tl.abs(x), 2) / 448.0, 1e-30) if ROUND: scale = tl.exp2(tl.ceil(tl.log2(scale))) # x = (x / scale[:,:, None]).to(out_ptr.dtype.element_ty) # x = tl.reshape(x, [W, N]) - - x = x / scale[:,:, None] + x = x / scale[:, :, None] x = tl.reshape(x, [W, N]) - tl.store(scale_ptr + indices[:, None] * nb + tl.arange(0, nb)[None, :], scale, mask=indices[:, None] < M) - tl.store(out_ptr + offs, x, mask=indices[:, None] < M) - offs += N * W + tl.store( + scale_ptr + tl.arange(0, NB)[:, None] * M + indices[None, :], + tl.trans(scale), + mask=(indices[None, :] < M) & (tl.arange(0, NB)[:, None] < n // 128), + ) + tl.store(out_ptr + offs, x, mask=masks) + offs += n * W # output transposed tensor only @triton.jit -def rms_norm_and_block_quant_forward_t_kernel(x_ptr, - weight_ptr, - transpose_output_ptr, - transpose_scale_ptr, - rms_ptr, - M, - N, - W: tl.constexpr, - ROUND: tl.constexpr): +def rms_norm_and_block_quant_forward_t_kernel( + x_ptr, + weight_ptr, + transpose_output_ptr, + transpose_scale_ptr, + rms_ptr, + M, + N, + W: tl.constexpr, + ROUND: tl.constexpr, +): rid = tl.program_id(axis=0) cid = tl.program_id(axis=1) - offs = rid * 128 * N + cid * W + tl.arange(0, 128)[:, None] * N + tl.arange(0, W)[ - None, :] - toffs = rid * 128 + cid * M * W + tl.arange(0, W)[:, None] * M + tl.arange(0, 128)[ - None, :] + offs = ( + rid * 128 * N + + cid * W + + tl.arange(0, 128)[:, None] * N + + tl.arange(0, W)[None, :] + ) + toffs = ( + rid * 128 + + cid * M * W + + tl.arange(0, W)[:, None] * M + + tl.arange(0, 128)[None, :] + ) weight = tl.load(weight_ptr + cid * W + tl.arange(0, W)).to(tl.float32) indices = rid * 128 + tl.arange(0, 128) rms = tl.load(rms_ptr + indices, mask=indices < M)[:, None] - x = tl.load(x_ptr + offs, mask=indices[:, None] < M).to(tl.float32) + x = tl.load(x_ptr + offs, mask=indices[:, None] < M).to(tl.float32) x = x * rms * weight scale = tl.maximum(tl.max(x.abs(), 0) / 448.0, 1e-30) if ROUND: scale = tl.exp2(tl.ceil(tl.log2(scale))) tl.store(transpose_scale_ptr + rid * N + cid * W + tl.arange(0, W), scale) - x = (tl.trans(x/scale)).to(transpose_output_ptr.dtype.element_ty) + x = (tl.trans(x / scale)).to(transpose_output_ptr.dtype.element_ty) tl.store(transpose_output_ptr + toffs, x, mask=indices[None, :] < M) - -def triton_rms_norm_and_block_quant_forward(x: torch.Tensor, - weight: torch.Tensor, - eps: float = 1e-6, - out: Optional[torch.Tensor] = None, - scale: Optional[torch.Tensor] = None, - rms: Optional[torch.Tensor] = None, - round_scale: bool = False, - output_mode: int = 2): +def triton_rms_norm_and_block_quant_forward( + x: torch.Tensor, + weight: torch.Tensor, + eps: float = 1e-6, + out: Optional[torch.Tensor] = None, + scale: Optional[torch.Tensor] = None, + rms: Optional[torch.Tensor] = None, + round_scale: bool = False, + output_mode: int = 2, +): """ Fused RMSNorm forward and block quantization. Args: @@ -296,23 +373,28 @@ def triton_rms_norm_and_block_quant_forward(x: torch.Tensor, - transpose_scale: quantization scale of transposed gradient. """ # row-wise read, row-wise write - M, N = x.shape - assert N <= 8192 and 8192 % N == 0 + assert x.is_contiguous() and weight.is_contiguous() + M, n = x.shape + N = triton.next_power_of_2(n) + assert N <= 8192 device = x.device - if out is None and output_mode in (0, 2): - out = torch.empty((M, N), device=device, dtype=torch.float8_e4m3fn) + if out is None and output_mode in (0, 2): + out = torch.empty((M, n), device=device, dtype=torch.float8_e4m3fn) if scale is None and output_mode in (0, 2): - scale = torch.empty((M, N//128), device=device, dtype=torch.float32) - if rms is None: - rms = torch.empty((M,), dtype=torch.float32, device=device) + scale = torch.empty((n // 128, M), device=device, dtype=torch.float32) + # transpose_output should be initialized, or else can not make splitted tensors - transpose_output = torch.empty((N, M), device=device, dtype=torch.float8_e4m3fn) - transpose_scale = torch.empty(((M+127)//128, N), device=device, dtype=torch.float32) - if output_mode == 0: # only output non-transpose tensor + transpose_output = torch.empty((n, M), device=device, dtype=torch.float8_e4m3fn) + transpose_scale = torch.empty( + ((M + 127) // 128, n), device=device, dtype=torch.float32 + ) + if output_mode == 0: # only output non-transpose tensor + assert rms is None + rms = torch.empty((M,), dtype=torch.float32, device=device) W = 8192 // N - T = 16 // W + T = 16 // W grid = (triton.cdiv(M, 16),) rms_norm_and_block_quant_forward_n_kernel[grid]( x, @@ -322,261 +404,133 @@ def triton_rms_norm_and_block_quant_forward(x: torch.Tensor, rms, eps, M, - T, + n, N, - N//128, + T, W, round_scale, num_stages=3, - num_warps=4 + num_warps=4, ) - scale = scale.t().contiguous() elif output_mode == 1: # only output transposed tensor - # W = N//512 - # grid = (512,) - W = 32 - grid = (triton.cdiv(M, 128), N//W) - rms_norm_and_block_quant_forward_t_kernel[grid](x, - weight, - transpose_output, - transpose_scale, - rms, - M, - N, - W, - round_scale, - num_stages=3, - num_warps=4) - - elif output_mode == 2: # output non-transposed and transposed tensor together - W = 8192 // N - T = 128 // W # BLOCK SIZE - H = 64 - grid = (triton.cdiv(M, 128),) - rms_norm_and_block_quant_forward_kernel[grid]( + assert rms is not None + W = 32 + assert n % W == 0 + grid = (triton.cdiv(M, 128), n // W) + rms_norm_and_block_quant_forward_t_kernel[grid]( x, weight, - out, - scale, transpose_output, transpose_scale, rms, - eps, M, - T, - N, - N//128, + n, W, - H, round_scale, num_stages=3, - num_warps=16 + num_warps=4, ) - scale = scale.t().contiguous() - - return out, scale, rms, transpose_output, transpose_scale + elif output_mode == 2: # output non-transposed and transposed tensor together + # we force set output_mode=2 when recompute qkv, but it has rms + # assert rms is None + rms = torch.empty((M,), dtype=torch.float32, device=device) + if M >= 1048576: # not used + W = 4096 // N + T = 128 // W + H = 32 + assert n % 128 == 0 + grid = (triton.cdiv(M, 128),) + rms_norm_and_block_quant_forward_kernel[grid]( + x, + weight, + out, + scale, + transpose_output, + transpose_scale, + rms, + eps, + M, + n, + N, + T, + W, + H, + round_scale, + num_stages=2, + num_warps=4, + ) + else: + W = 8192 // N + T = 16 // W + grid = (triton.cdiv(M, 16),) + rms_norm_and_block_quant_forward_n_kernel[grid]( + x, + weight, + out, + scale, + rms, + eps, + M, + n, + N, + T, + W, + round_scale, + num_stages=3, + num_warps=4, + ) + + W = 32 + assert n % W == 0, f" {n=} {W=}" + grid = (triton.cdiv(M, 128), n // W) + rms_norm_and_block_quant_forward_t_kernel[grid]( + x, + weight, + transpose_output, + transpose_scale, + rms, + M, + n, + W, + round_scale, + num_stages=3, + num_warps=4, + ) -# TOOD(nanxiao): opt performance -@triton.jit -def group_rms_norm_gate_forward_kernel(x_ptr, gate_ptr, weight_ptr, out_ptr, eps, bs, length, - DIM: tl.constexpr, - D: tl.constexpr, - GROUP_SIZE: tl.constexpr): - pid = tl.program_id(axis=0) - bid = pid // length - sid = pid % length - - weight = tl.load(weight_ptr + tl.arange(0, DIM)) - weight = tl.reshape(weight, [GROUP_SIZE, D]) - - x_offs = pid * DIM + tl.arange(0, GROUP_SIZE)[:, None] * D + tl.arange(0, D)[ - None, :] - x = tl.load(x_ptr + x_offs).to(tl.float32) - offs = sid * bs * DIM + bid * DIM + tl.arange(0, GROUP_SIZE)[:, None] * D + tl.arange(0, D)[ - None, :] - g = tl.load(gate_ptr + offs).to(tl.float32) - rms = tl.sqrt(tl.sum(x * x, axis=1) / D + eps) - - x = (x / rms[:, None]) * weight * tl.sigmoid(g) - - tl.store(out_ptr + offs, x) - - -def triton_group_rms_norm_gate_forward(x: torch.Tensor, - gate: torch.Tensor, - weight: torch.Tensor, - eps=1e-6, - group_size=4, - transpose=True): - """ - norm and gate in linear attention - Args: - x: output of attn, [bs, length, n_heads, head_dim] - gate: gate tensor, [length, bs, dim] if transpose=True else [bs, length, dim] - weight: rms norm weight, [dim] - eps: epsilon of rms norm - group_size: group size of group rms norm - transpose: whether gate is transposed and output will be transposed - - Returns: - output tensor, [length, bs, dim] if transpose=True else [bs, length, dim] - """ - # row-wise read, row-wise write - if transpose: - length, bs, dim = gate.shape - else: - bs, length, dim = gate.shape - assert dim <= 8192 and triton.next_power_of_2(dim) == dim and triton.next_power_of_2(group_size) == group_size - d = dim // group_size - device = x.device - if transpose: - out = torch.empty((length, bs, dim), device=device, dtype=x.dtype) - else: - out = torch.empty((bs, length, dim), device=device, dtype=x.dtype) - - grid = (bs*length,) - group_rms_norm_gate_forward_kernel[grid]( - x, - gate, - weight.data, - out, - eps, - bs, - length, - dim, - d, - group_size, - num_stages=3, - num_warps=4 - ) - return out + return out, scale, rms, transpose_output, transpose_scale @triton.jit -def group_rms_gate_backward_kernel( - grad_output_ptr, - x_ptr, - gate_ptr, - w_ptr, - dx_ptr, - dg_ptr, - dw_ptr, - eps, - bs, - length, - DIM: tl.constexpr, - D: tl.constexpr, - GROUP_SIZE: tl.constexpr, - T: tl.constexpr +def rms_norm_and_smooth_quant_forward_kernel( + x_ptr, + weight_ptr, + smooth_scale_ptr, + out_ptr, + scale_ptr, + rms_ptr, + eps, + M, + T, + N: tl.constexpr, + W: tl.constexpr, + ROUND: tl.constexpr, ): - pid = tl.program_id(0) - bid = pid * T // length - sid = pid * T % length - - w = tl.load(w_ptr + tl.arange(0, DIM)) - w = tl.reshape(w, [GROUP_SIZE, D]) - - x_offs = pid * DIM * T + tl.arange(0, GROUP_SIZE)[:, None] * D + tl.arange(0, D)[ - None, :] - offs = sid * bs * DIM + bid * DIM + tl.arange(0, GROUP_SIZE)[:, None] * D + tl.arange(0, D)[ - None, :] - dw = tl.zeros((GROUP_SIZE, D), dtype=tl.float32) - for i in range(T): - x = tl.load(x_ptr + x_offs).to(tl.float32) - g = tl.load(grad_output_ptr + offs).to(tl.float32) - gate = tl.load(gate_ptr + offs).to(tl.float32) - gate = tl.sigmoid(gate) - rms = tl.sqrt(tl.sum(x * x, 1) / D + eps) - r = 1.0 / rms[:, None] - w_grad = x * g * r * gate - dw += w_grad - - dx = r * g * w * gate - r * r * r * x * tl.sum(x * g * w * gate, 1, keep_dims=True) / D - - tl.store(dx_ptr + x_offs, dx) - - dg = x * r * w * g * gate * (1 - gate) - tl.store(dg_ptr + offs, dg) - - x_offs += DIM - offs += DIM * bs - - dw = tl.reshape(dw, [DIM]) - tl.store(dw_ptr + pid * DIM + tl.arange(0, DIM), dw) - - -def triton_group_rms_norm_gate_backward(grad_output, x, gate, weight, eps=1e-6, group_size=4, transpose=True): - if transpose: - length, bs, dim = gate.shape - else: - bs, length, dim = gate.shape - assert dim <= 8192 and triton.next_power_of_2(dim) == dim and triton.next_power_of_2(group_size) == group_size - d = dim // group_size - device = x.device - dx = torch.empty_like(x) - dg = torch.empty_like(gate) - - T = 8 - g = (bs*length)//T - tmp_dw = torch.empty(g, dim, dtype=torch.float32, device=device) - grid = (g,) - group_rms_gate_backward_kernel[grid]( - grad_output, - x, - gate, - weight, - dx, - dg, - tmp_dw, - eps, - bs, - length, - dim, - d, - group_size, - T, - num_stages=3, - num_warps=8 - ) - dw = tmp_dw.sum(dim=0).to(weight.dtype) - return dx, dg, dw - - - -@triton.jit -def rms_norm_and_smooth_quant_forward_kernel(x_ptr, weight_ptr, smooth_scale_ptr, - out_ptr, scale_ptr, max_ptr, rms_ptr, - eps, - M, - T, - N: tl.constexpr, - W: tl.constexpr, - CALIBRATE: tl.constexpr, - OUTPUT: tl.constexpr, - ROUND: tl.constexpr): pid = tl.program_id(axis=0) # row-wise read, row-wise write weight = tl.load(weight_ptr + tl.arange(0, N)).to(tl.float32)[None, :] smooth_scale = tl.load(smooth_scale_ptr + tl.arange(0, N))[None, :] smooth_scale = 1.0 / tl.maximum(smooth_scale, 1e-30) - if CALIBRATE: - # triton 3.3.1 has bug with N = 2048 and calibrate=True - maxs = tl.zeros((N, ), dtype=tl.float32) - offs = pid * W * T * N + tl.arange(0, W)[:, None] * N + tl.arange(0, N)[ - None, :] + + offs = pid * W * T * N + tl.arange(0, W)[:, None] * N + tl.arange(0, N)[None, :] for i in range(T): indices = pid * W * T + i * W + tl.arange(0, W) x = tl.load(x_ptr + offs, mask=indices[:, None] < M).to(tl.float32) - rms = 1/tl.sqrt(tl.sum(x * x, axis=1) / N + eps) - if OUTPUT: - tl.store(rms_ptr + indices, rms, mask=indices < M) + rms = tl.rsqrt(tl.sum(x * x, axis=1) / N + eps) + tl.store(rms_ptr + indices, rms, mask=indices < M) x = x * rms[:, None] * weight - if CALIBRATE: - maxs = tl.maximum(maxs, tl.max(tl.abs(x),0)) - x = x * smooth_scale scale = tl.maximum(tl.max(tl.abs(x), 1) / 448.0, 1e-30) if ROUND: @@ -586,18 +540,20 @@ def rms_norm_and_smooth_quant_forward_kernel(x_ptr, weight_ptr, smooth_scale_ptr tl.store(out_ptr + offs, q, mask=indices[:, None] < M) offs += N * W - if CALIBRATE: - tl.store(max_ptr + pid * N + tl.arange(0, N), maxs) - # rms is used for moe routing, it is stored as 1/rms -def triton_rms_norm_and_smooth_quant_forward(x, weight, smooth_scale=None, - eps=1e-6, - out=None, scale=None, rms=None, - calibrate=False, - output_rms=False, - round_scale=False): +def triton_rms_norm_and_smooth_quant_forward( + x, + weight, + smooth_scale=None, + eps=1e-6, + out=None, + scale=None, + rms=None, + round_scale=False, +): """""" + assert x.is_contiguous() and weight.is_contiguous() M, N = x.shape assert N <= 8192 and 8192 % N == 0 device = x.device @@ -611,11 +567,7 @@ def triton_rms_norm_and_smooth_quant_forward(x, weight, smooth_scale=None, T = 8 if M // W >= 4096 else 4 assert M % (T * W) == 0 g = M // (T * W) - if calibrate: - maxs = torch.empty((g, N), dtype=torch.float32, device=device) - else: - maxs = None - if output_rms and rms is None: + if rms is None: rms = torch.empty((M,), dtype=torch.float32, device=device) grid = (g,) rms_norm_and_smooth_quant_forward_kernel[grid]( @@ -624,20 +576,460 @@ def triton_rms_norm_and_smooth_quant_forward(x, weight, smooth_scale=None, smooth_scale, out, scale, - maxs, rms, eps, M, T, N, W, - calibrate, - output_rms, round_scale, num_stages=3, - num_warps=2 if N == 2048 else 4 + num_warps=2 if N == 2048 else 4, ) - if calibrate: - maxs = maxs.amax(0) - return out, scale, maxs, rms + return out, scale, rms + + +@triton.jit +def rms_norm_and_mxfp8_quant_forward_n_kernel( + x_ptr, + weight_ptr, + out_ptr, + scale_ptr, + rms_ptr, + eps, + n, + M: tl.constexpr, + T: tl.constexpr, + N: tl.constexpr, + nb: tl.constexpr, + W: tl.constexpr, +): + pid = tl.program_id(axis=0) + + # row-wise read, row-wise write + weight_mask = tl.arange(0, N) < n + weight = tl.load(weight_ptr + tl.arange(0, N), mask=weight_mask).to(tl.float32)[ + None, : + ] + offs_m = tl.arange(0, W) + offs_n = tl.arange(0, N) + offs = pid * W * T * n + offs_m[:, None] * n + offs_n[None, :] + # offs = pid * W * T * N + tl.arange(0, W)[:, None] * N + tl.arange(0, N)[None, :] + + for i in range(T): + indices = pid * W * T + i * W + tl.arange(0, W) + mask = (indices[:, None] < M) & (offs_n < n) + x = tl.load(x_ptr + offs, mask=mask).to(tl.float32) + # x = tl.load(x_ptr + offs, mask=indices[:, None] < M).to(tl.float32) + rms = tl.rsqrt(tl.sum(x * x, axis=1) / n + eps) + tl.store(rms_ptr + indices, rms, mask=indices < M) + + x = x * rms[:, None] * weight + x = tl.reshape(x, [W, nb, 32]) + scale = tl.maximum(tl.max(tl.abs(x), 2) / 448.0, 1e-30) + log_scale = tl.ceil(tl.log2(scale)) + scale = tl.exp2(log_scale) + + # x = (x / scale[:,:, None]).to(out_ptr.dtype.element_ty) + # x = tl.reshape(x, [W, N]) + x = x / scale[:, :, None] + x = tl.reshape(x, [W, N]) + + # tl.store(scale_ptr + indices[:, None] * nb + tl.arange(0, nb)[None, :], scale, mask=indices[:, None] < M) + # tl.store(out_ptr + offs, x, mask=indices[:, None] < M) + scale_mask = (indices[:, None] < M) & (tl.arange(0, nb)[None, :] < (n // 32)) + tl.store( + scale_ptr + indices[:, None] * (n // 32) + tl.arange(0, nb)[None, :], + log_scale + 127, + mask=scale_mask, + ) + tl.store(out_ptr + offs, x, mask=mask) + offs += n * W + + +@triton.jit +def rms_norm_and_mxfp8_quant_forward_t_kernel( + x_ptr, + weight_ptr, + transpose_output_ptr, + transpose_scale_ptr, + rms_ptr, + n, + M, + N, + W: tl.constexpr, +): + rid = tl.program_id(axis=0) + cid = tl.program_id(axis=1) + + offs = ( + rid * 32 * n + + cid * W + + tl.arange(0, 32)[:, None] * n + + tl.arange(0, W)[None, :] + ) + # toffs = rid * 32 + cid * M * W + tl.arange(0, W)[:, None] * M + tl.arange(0, 32)[ + # None, :] + + weight = tl.load(weight_ptr + cid * W + tl.arange(0, W)).to(tl.float32) + indices = rid * 32 + tl.arange(0, 32) + rms = tl.load(rms_ptr + indices, mask=indices < M)[:, None] + x = tl.load(x_ptr + offs, mask=indices[:, None] < M).to(tl.float32) + x = x * rms * weight + scale = tl.maximum(tl.max(x.abs(), 0) / 448.0, 1e-30) + log_scale = tl.ceil(tl.log2(scale)) + scale = tl.exp2(log_scale) + + tl.store(transpose_scale_ptr + rid * n + cid * W + tl.arange(0, W), log_scale + 127) + + # x = (tl.trans(x/scale)).to(transpose_output_ptr.dtype.element_ty) + x = (x / scale).to(transpose_output_ptr.dtype.element_ty) + # tl.store(transpose_output_ptr + toffs, x, mask=indices[None, :] < M) + tl.store(transpose_output_ptr + offs, x) + + +@triton.jit +def rms_norm_and_mxfp8_quant_forward_kernel( + x_ptr, + weight_ptr, + out_ptr, + scale_ptr, + transpose_output_ptr, + transpose_scale_ptr, + rms_ptr, + eps, + n, + M, + T: tl.constexpr, + N: tl.constexpr, + nb: tl.constexpr, + W: tl.constexpr, + H: tl.constexpr, +): + pid = tl.program_id(axis=0) + + # row-wise read, row-wise write + weight_mask = tl.arange(0, N) < n + weight = tl.load(weight_ptr + tl.arange(0, N), mask=weight_mask).to(tl.float32)[ + None, : + ] + offs_m = tl.arange(0, W) + offs_n = tl.arange(0, N) + offs = pid * W * T * n + offs_m[:, None] * n + offs_n[None, :] + # offs = pid * W * T * N + tl.arange(0, W)[:, None] * N + tl.arange(0, N)[None, :] + + for i in range(T): + indices = pid * W * T + i * W + tl.arange(0, W) + mask = (indices[:, None] < M) & (offs_n < n) + x = tl.load(x_ptr + offs, mask=mask).to(tl.float32) + # x = tl.load(x_ptr + offs, mask=indices[:, None] < M).to(tl.float32) + rms = tl.rsqrt(tl.sum(x * x, axis=1) / n + eps) + tl.store(rms_ptr + indices, rms, mask=indices < M) + + x = x * rms[:, None] * weight + x = tl.reshape(x, [W, nb, 32]) + scale = tl.maximum(tl.max(tl.abs(x), 2) / 448.0, 1e-30) + log_scale = tl.ceil(tl.log2(scale)) + scale = tl.exp2(log_scale) + + # x = (x / scale[:,:, None]).to(out_ptr.dtype.element_ty) + # x = tl.reshape(x, [W, N]) + x = x / scale[:, :, None] + x = tl.reshape(x, [W, N]) + + # tl.store(scale_ptr + indices[:, None] * nb + tl.arange(0, nb)[None, :], scale, mask=indices[:, None] < M) + # tl.store(out_ptr + offs, x, mask=indices[:, None] < M) + scale_mask = (indices[:, None] < M) & (tl.arange(0, nb)[None, :] < (n // 32)) + tl.store( + scale_ptr + indices[:, None] * (n // 32) + tl.arange(0, nb)[None, :], + log_scale + 127, + mask=scale_mask, + ) + tl.store(out_ptr + offs, x, mask=mask) + offs += n * W + + offs = pid * 32 * n + tl.arange(0, 32)[:, None] * n + tl.arange(0, H)[None, :] + # toffs = pid * 32 + tl.arange(0, H)[:, None] * M + tl.arange(0, 32)[ + # None, :] + indices = pid * 32 + tl.arange(0, 32) + tl.debug_barrier() + rms = tl.load(rms_ptr + indices, mask=indices < M)[:, None] + for i in range(n // H): + x = tl.load(x_ptr + offs, mask=indices[:, None] < M).to(tl.float32) + wgt = tl.load(weight_ptr + i * H + tl.arange(0, H)).to(tl.float32) + x = x * rms * wgt + scale = tl.maximum(tl.max(x.abs(), 0) / 448.0, 1e-30) + log_scale = tl.ceil(tl.log2(scale)) + scale = tl.exp2(log_scale) + + tl.store( + transpose_scale_ptr + pid * n + i * H + tl.arange(0, H), log_scale + 127 + ) + x = (x / scale).to(transpose_output_ptr.dtype.element_ty) + # tl.store(transpose_output_ptr + toffs, tl.trans(x), mask=indices[None, :] < M) + tl.store(transpose_output_ptr + offs, x) + offs += H + # toffs += M * H + + +def triton_rms_norm_and_mxfp8_quant_forward( + x, weight, eps=1e-6, out=None, scale=None, rms=None, output_mode=2 +): + assert x.is_contiguous() and weight.is_contiguous() + M, n = x.shape + N = triton.next_power_of_2(n) + # assert N <= 8192 and 8192 % N == 0 + assert n % 32 == 0 + assert M % 128 == 0 + device = x.device + + if out is None and output_mode in (0, 2): + out = torch.empty((M, n), device=device, dtype=torch.float8_e4m3fn) + + if scale is None and output_mode in (0, 2): + # scale = torch.empty((M, n//32), device=device, dtype=torch.float32) + scale = torch.empty((M, n // 32), device=device, dtype=torch.uint8) + if rms is None: + rms = torch.empty((M,), dtype=torch.float32, device=device) + # transpose_output should be initialized, or else can not make splitted tensors + transpose_output = torch.empty((M, n), device=device, dtype=torch.float8_e4m3fn) + # transpose_scale = torch.empty(((M+31)//32, n), device=device, dtype=torch.float32) + transpose_scale = torch.empty(((M + 31) // 32, n), device=device, dtype=torch.uint8) + if output_mode == 0: # only output non-transpose tensor + W = 8192 // N + T = 16 // W + grid = (triton.cdiv(M, 16),) + rms_norm_and_mxfp8_quant_forward_n_kernel[grid]( + x, + weight, + out, + scale, + rms, + eps, + n, + M, + T, + N, + N // 32, + W, + num_stages=3, + num_warps=4, + ) + + elif output_mode == 1: # only output transposed tensor + # W = N//512 + # grid = (512,) + W = 64 + grid = (triton.cdiv(M, 32), n // W) + rms_norm_and_mxfp8_quant_forward_t_kernel[grid]( + x, + weight, + transpose_output, + transpose_scale, + rms, + n, + M, + N, + W, + num_stages=3, + num_warps=4, + ) + + elif output_mode == 2: # output non-transposed and transposed tensor together + W = 8192 // N + T = 32 // W + H = 64 + grid = (triton.cdiv(M, 32),) + rms_norm_and_mxfp8_quant_forward_kernel[grid]( + x, + weight, + out, + scale, + transpose_output, + transpose_scale, + rms, + eps, + n, + M, + T, + N, + N // 32, + W, + H, + num_stages=3, + num_warps=16, + ) + + return out, scale, rms, transpose_output, transpose_scale + + +@triton.jit +def rms_norm_fp32_gemm_block_quant_forward_n_kernel( + x_ptr, + norm_weight_ptr, + route_weight_ptr, + y_ptr, + rms_ptr, + logit_ptr, + xq_ptr, + xs_ptr, + eps, + M, + N: tl.constexpr, + K: tl.constexpr, + BLOCK_SIZE_K: tl.constexpr, + BLOCK_SIZE_M: tl.constexpr, + BLOCK_SIZE_N: tl.constexpr, + ROUND: tl.constexpr, +): + pid_m = tl.program_id(axis=0) + pid_n = tl.program_id(axis=1) + k = tl.cdiv(K, BLOCK_SIZE_K) + offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) + offs_k = tl.arange(0, BLOCK_SIZE_K) + x_offs = offs_m[:, None] * K + offs_k[None, :] + + rms = tl.zeros((BLOCK_SIZE_M,), dtype=tl.float32) + for i in range(k): + x = tl.load(x_ptr + x_offs).to(tl.float32) + rms += tl.sum(x * x, axis=1) + x_offs += BLOCK_SIZE_K + + rms = tl.rsqrt(rms / K + eps) + + tl.store(rms_ptr + pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M), rms) + + x_offs = offs_m[:, None] * K + offs_k[None, :] + w_ptrs = route_weight_ptr + offs_n[None, :] * K + offs_k[:, None] + c = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32) + for i in range(k): + norm_weight = tl.load(norm_weight_ptr + i * BLOCK_SIZE_K + offs_k).to( + tl.float32 + ) + x = tl.load(x_ptr + x_offs).to(tl.float32) + w = tl.load(w_ptrs).to(tl.float32) + + x = x * rms[:, None] * norm_weight + tl.store(y_ptr + x_offs, x) + + c = tl.dot(x, w, c) + + scale = tl.maximum(tl.max(tl.abs(x), 1) / 448.0, 1e-30) + if ROUND: + scale = tl.exp2(tl.ceil(tl.log2(scale))) + + x = x / scale[:, None] + + tl.store( + xs_ptr + M * i + pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M), scale + ) + tl.store(xq_ptr + x_offs, x) + + x_offs += BLOCK_SIZE_K + w_ptrs += BLOCK_SIZE_K + + offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) + c_ptrs = logit_ptr + offs_m[:, None] * N + offs_n[None, :] + tl.store(c_ptrs, c) + + +def triton_rms_norm_fp32_gemm_block_quant_forward( + x: torch.Tensor, + norm_weight: torch.Tensor, + route_weight: torch.Tensor, + rms: Optional[torch.Tensor] = None, + eps: float = 1e-6, + output_mode: int = 0, + round_scale=False, +): + """ + y = rms_norm(x) + logits = y@w_route + x_q, x_s, xt_q, xt_s = quantization(y) + Args: + x: input tensor + norm weight: weight tensor of rms norm + route_weight: moe router weight + eps: epsilon of rms norm + output_mode: 0 or 1 + 0: only output non-transpose quantizatino tensor + 1: only output transposed quantizatino tensor + + Returns: + - y: rms normed tensor + - rms: 1/rms + - logits: router logit + - x_q: + - x_s: + - xt_q: + - xt_s: + """ + assert ( + x.is_contiguous() + and norm_weight.is_contiguous() + and route_weight.is_contiguous() + ) + assert output_mode in (0, 1) + M, K = x.size() + N, K = route_weight.size() + assert M % 128 == 0 and K % 128 == 0 and N % 128 == 0 + device = x.device + y = torch.empty(M, K, dtype=x.dtype, device=device) + if rms is None: + assert output_mode == 0 + rms = torch.empty(M, dtype=torch.float32, device=device) + logits = torch.empty(M, N, dtype=torch.float32, device=device) + + x_q = torch.empty((M, K), device=device, dtype=torch.float8_e4m3fn) + x_s = torch.empty((K // 128, M), device=device, dtype=torch.float32) + xt_q = torch.empty((K, M), device=device, dtype=torch.float8_e4m3fn) + xt_s = torch.empty((M // 128, K), device=device, dtype=torch.float32) + + if output_mode == 0: + BLOCK_SIZE_K = 128 # MUST BE 128 (quantization block size) + BLOCK_SIZE_M = 32 + BLOCK_SIZE_N = N + num_warps = 4 + num_stages = 2 + grid = (M // BLOCK_SIZE_M, N // BLOCK_SIZE_N) + rms_norm_fp32_gemm_block_quant_forward_n_kernel[grid]( + x, + norm_weight, + route_weight, + y, + rms, + logits, + x_q, + x_s, + eps, + M, + N, + K, + BLOCK_SIZE_K, + BLOCK_SIZE_M, + BLOCK_SIZE_N, + round_scale, + num_warps=num_warps, + num_stages=num_stages, + ) + else: + W = 32 + grid = (triton.cdiv(M, 128), K // W) + rms_norm_and_block_quant_forward_t_kernel[grid]( + x, + norm_weight, + xt_q, + xt_s, + rms, + M, + K, + W, + round_scale, + num_stages=3, + num_warps=4, + ) + + return y, rms, logits, x_q, x_s, xt_q, xt_s diff --git a/linghe/utils/rearange.py b/linghe/utils/rearange.py index 868814c..ba178c9 100644 --- a/linghe/utils/rearange.py +++ b/linghe/utils/rearange.py @@ -9,9 +9,20 @@ @triton.jit -def split_and_cat_kernel(x_ptr, y_ptr, scale_ptr, scale_output_ptr, count_ptr, - accum_ptr, rev_accum_ptr, index_ptr, M, - N: tl.constexpr, SCALE: tl.constexpr, K: tl.constexpr): +def sort_chunks_by_index_kernel( + x_ptr, + y_ptr, + scale_ptr, + scale_output_ptr, + count_ptr, + accum_ptr, + rev_accum_ptr, + index_ptr, + M, + N: tl.constexpr, + SCALE: tl.constexpr, + K: tl.constexpr, +): pid = tl.program_id(axis=0) # row-wise read, row-wise write index = tl.load(index_ptr + pid) @@ -26,16 +37,21 @@ def split_and_cat_kernel(x_ptr, y_ptr, scale_ptr, scale_output_ptr, count_ptr, if SCALE: for i in range(tl.cdiv(count, K)): - scale = tl.load(scale_ptr + si + i * K + tl.arange(0, K), - mask=i * K + tl.arange(0, K) < count) - tl.store(scale_output_ptr + rev_si + i * K + tl.arange(0, K), scale, - mask=i * K + tl.arange(0, K) < count) + scale = tl.load( + scale_ptr + si + i * K + tl.arange(0, K), + mask=i * K + tl.arange(0, K) < count, + ) + tl.store( + scale_output_ptr + rev_si + i * K + tl.arange(0, K), + scale, + mask=i * K + tl.arange(0, K) < count, + ) -def triton_split_and_cat(x, counts, indices, scales=None): +def triton_sort_chunks_by_index(x, counts, indices, scales=None): """ split x to multiple tensors and cat with indices, - it is used for permutation in moe + it is used for permutation in moe with all2all communication Args: x: [bs, dim] counts: [n_split] @@ -46,6 +62,7 @@ def triton_split_and_cat(x, counts, indices, scales=None): - y: output tensor - output_scales: output scales if scales is not None """ + assert x.is_contiguous() M, N = x.shape n_split = counts.shape[0] device = x.device @@ -61,7 +78,7 @@ def triton_split_and_cat(x, counts, indices, scales=None): # TODO: adapt for n_expert <= 64 K = 256 grid = (n_split,) - split_and_cat_kernel[grid]( + sort_chunks_by_index_kernel[grid]( x, y, scales, @@ -75,6 +92,6 @@ def triton_split_and_cat(x, counts, indices, scales=None): S, K, num_stages=3, - num_warps=8 + num_warps=8, ) return y, output_scales diff --git a/linghe/utils/reduce.py b/linghe/utils/reduce.py index 72b3b3d..ccac83e 100644 --- a/linghe/utils/reduce.py +++ b/linghe/utils/reduce.py @@ -9,34 +9,37 @@ @triton.jit -def abs_max_kernel(x_ptr, - scale_ptr, - smooth_scale_ptr, - output_ptr, - min_value, - M, N, - H: tl.constexpr, - W: tl.constexpr, - EVEN: tl.constexpr, - QUANTIZED: tl.constexpr): +def abs_max_kernel( + x_ptr, + scale_ptr, + smooth_scale_ptr, + output_ptr, + min_value, + M, + N, + H: tl.constexpr, + W: tl.constexpr, + EVEN: tl.constexpr, + QUANTIZED: tl.constexpr, +): pid = tl.program_id(axis=0) # col-wise read, col-wise write x_max = tl.zeros((W,), dtype=tl.float32) m = tl.cdiv(M, H) offs = pid * W + tl.arange(0, H)[:, None] * N + tl.arange(0, W) if QUANTIZED: - smooth_scale = tl.load(smooth_scale_ptr + pid * W + tl.arange(0, W))[ - None, :] + smooth_scale = tl.load(smooth_scale_ptr + pid * W + tl.arange(0, W))[None, :] for i in range(m): if EVEN: x = tl.load(x_ptr + offs).to(tl.float32) else: - x = tl.load(x_ptr + offs, - mask=i * H + tl.arange(0, H)[:, None] < M).to( - tl.float32) + x = tl.load(x_ptr + offs, mask=i * H + tl.arange(0, H)[:, None] < M).to( + tl.float32 + ) if QUANTIZED: - scale = tl.load(scale_ptr + i * H + tl.arange(0, H), - mask=i * H + tl.arange(0, H) < M) + scale = tl.load( + scale_ptr + i * H + tl.arange(0, H), mask=i * H + tl.arange(0, H) < M + ) x = x * scale[:, None] * smooth_scale x_max = tl.maximum(x_max, tl.max(tl.abs(x), axis=0)) offs += H * N @@ -59,6 +62,7 @@ def triton_abs_max(x, scale=None, smooth_scale=None, min_value=1e-30, axis=0): Returns: max tensor """ + assert x.is_contiguous() assert axis == 0 N = x.size(-1) M = x.numel() // N @@ -76,12 +80,14 @@ def triton_abs_max(x, scale=None, smooth_scale=None, min_value=1e-30, axis=0): smooth_scale, maxs, min_value, - M, N, - H, W, + M, + N, + H, + W, EVEN, quantized, num_stages=2, - num_warps=4 + num_warps=4, ) return maxs @@ -97,7 +103,7 @@ def batch_count_zero_kernel(input_ptrs, size_ptr, count_ptr, B: tl.constexpr): input_ptr = tl.load(input_ptrs + tid).to(tl.pointer_type(tl.float32)) t = tl.cdiv(size, B * sm) offs = bid * t * B + tl.arange(0, B) - for i in range(t): + for i in tl.range(t): x = tl.load(input_ptr + offs, mask=offs < size, other=1).to(tl.float32) count += tl.sum(tl.where(x == 0, 1, 0)) offs += B @@ -114,82 +120,175 @@ def triton_batch_count_zero(xs): Returns: a single-value int64 tensor """ + assert all([x.is_contiguous() for x in xs]) device = xs[0].device - sizes = torch.tensor([x.numel() for x in xs], dtype=torch.int64, - device=device) - ptrs = torch.tensor([x.data_ptr() for x in xs], dtype=torch.int64, - device=device) + sizes = torch.tensor([x.numel() for x in xs], dtype=torch.int64).cuda( + device, non_blocking=True + ) + ptrs = torch.tensor([x.data_ptr() for x in xs], dtype=torch.int64).cuda( + device, non_blocking=True + ) - sm = torch.cuda.get_device_properties(device).multi_processor_count + block = 2048 tensor_count = len(xs) - counts = torch.empty((tensor_count, sm), device=device, dtype=torch.int64) - B = 4096 - grid = (tensor_count, sm) - batch_count_zero_kernel[grid]( - ptrs, - sizes, - counts, - B, - num_stages=2, - num_warps=4 - ) + counts = torch.empty((tensor_count, block), device=device, dtype=torch.int64) + B = 1024 + grid = (tensor_count, block) + batch_count_zero_kernel[grid](ptrs, sizes, counts, B, num_stages=2, num_warps=2) count = counts.sum() return count @triton.jit -def batch_sum_with_ord_kernel(input_ptrs, size_ptr, count_ptr, B: tl.constexpr, - ORD: tl.constexpr): +def norm_kernel(input_ptr, tmp_ptr, m, B: tl.constexpr, ORD: tl.constexpr): + pid = tl.program_id(axis=0).to(tl.int64) + + offs = pid * B + tl.arange(0, B) + x = tl.load(input_ptr + offs, mask=offs < m, other=0).to(tl.float32) + if ORD == 2: + sums = tl.sum(x * x) + elif ORD == 1: + sums = tl.sum(tl.abs(x)) + elif ORD == -1: + sums = tl.max(tl.abs(x)) + + tl.store(tmp_ptr + pid, sums) + + +def triton_norm(x, ord=2, norm=True, scalar=True, dtype=torch.float32): + """ + calculate norm. + Args: + x: input tensor. + ord: the order of tensor. -1 means 'inf' ord. + norm: + only used with ord in (1, 2) + True: (sum(sum(abs(x)**ord) x for x in xs))**(1/ord) + False: sum(sum(abs(x)**ord) x for x in xs)) + + Returns: + a scalar if scalar=True else a single-value fp32 tensor + """ + assert x.is_contiguous() + assert ord in (1, 2, -1) + # assert all([x.is_contiguous() for x in xs]) + device = x.device + m = x.numel() + B = 512 + T = triton.cdiv(m, B) + tmp = torch.empty((T,), device=device, dtype=dtype) + grid = (T,) + norm_kernel[grid](x, tmp, m, B, ord, num_stages=2, num_warps=2) + if ord == -1: + output = tmp.max() + else: + output = tmp.sum() + if ord == 2 and norm: + output = torch.sqrt(output) + if not scalar: + output = output.unsqueeze(0) + return output + + +@triton.jit +def batch_norm_kernel( + input_ptrs, + size_ptr, + tmp_ptr, + DT: tl.constexpr, + B: tl.constexpr, + ORD: tl.constexpr, + HP: tl.constexpr, +): tid = tl.program_id(axis=0) - bid = tl.program_id(axis=1) - sm = tl.num_programs(axis=1) - sums = 0.0 + bid = tl.program_id(axis=1).to(tl.int64) + T = tl.num_programs(axis=1) + if HP: + sums = tl.zeros((B,), dtype=tl.float64) + else: + sums = tl.zeros((B,), dtype=tl.float32) size = tl.load(size_ptr + tid) - input_ptr = tl.load(input_ptrs + tid).to(tl.pointer_type(tl.float32)) - t = tl.cdiv(size, B * sm) + if DT == 0: + input_ptr = tl.load(input_ptrs + tid).to(tl.pointer_type(tl.float32)) + else: + input_ptr = tl.load(input_ptrs + tid).to(tl.pointer_type(tl.bfloat16)) + t = tl.cdiv(size, B * T) offs = bid * t * B + tl.arange(0, B) for i in range(t): - x = tl.load(input_ptr + offs, mask=offs < size, other=0).to(tl.float32) + x = tl.load(input_ptr + offs, mask=offs < size, other=0) + if HP: + x = x.to(tl.float64) + else: + x = x.to(tl.float32) if ORD == 2: - sums += tl.sum(x * x) + sums += x * x elif ORD == 1: - sums += tl.sum(tl.abs(x)) + sums += tl.abs(x) + elif ORD == -1: + sums = tl.maximum(sums, tl.abs(x)) offs += B - tl.store(count_ptr + tid * sm + bid, sums) + if ORD == -1: + sums = tl.max(sums) + else: + sums = tl.sum(sums) + tl.store(tmp_ptr + tid * T + bid, sums) -def triton_batch_sum_with_ord(xs, ord=2): +def triton_batch_norm(xs, ord=2, norm=True, scalar=True, high_precision=True): """ - return sum(abs(x)**ord). + treat multiple tensors as a single tensor and calculate norm. Args: xs: Tensor lists. - ord: the order of tensor. + ord: the order of tensor. -1 means 'inf' ord. + norm: + only used with ord in (1, 2) + True: (sum(sum(abs(x)**ord) x for x in xs))**(1/ord) + False: sum(sum(abs(x)**ord) x for x in xs)) Returns: - a single-value fp32 tensor + a scalar if scalar=True else a single-value fp32 tensor """ - assert ord in (1, 2) + if len(xs) == 0: + return torch.zeros(() if scalar else (1,), device="cuda", dtype=torch.float32) + dtype = xs[0].dtype + assert dtype in (torch.float32, torch.bfloat16) + assert all([x.is_contiguous() and x.dtype == dtype for x in xs]) + assert ord in (1, 2, -1) + device = xs[0].device - sizes = torch.tensor([x.numel() for x in xs], dtype=torch.int64, - device=device) - ptrs = torch.tensor([x.data_ptr() for x in xs], dtype=torch.int64, - device=device) + sizes = torch.tensor([x.numel() for x in xs], dtype=torch.int64).cuda( + device, non_blocking=True + ) + ptrs = torch.tensor([x.data_ptr() for x in xs], dtype=torch.int64).cuda( + device, non_blocking=True + ) - sm = torch.cuda.get_device_properties(device).multi_processor_count + DT = 0 if dtype == torch.float32 else 1 + T = 256 tensor_count = len(xs) - sums = torch.empty((tensor_count, sm), device=device, dtype=torch.float32) - B = 4096 - grid = (tensor_count, sm) - batch_sum_with_ord_kernel[grid]( - ptrs, - sizes, - sums, - B, - ord, - num_stages=2, - num_warps=4 + tmp = torch.empty( + (tensor_count, T), + device=device, + dtype=torch.float64 if high_precision else torch.float32, + ) + B = 128 + grid = (tensor_count, T) + batch_norm_kernel[grid]( + ptrs, sizes, tmp, DT, B, ord, high_precision, num_stages=2, num_warps=2 ) - sums = sums.sum() - return sums + if ord == -1: + output = tmp.max() + elif ord == 1: + output = tmp.sum() + else: + if norm: + output = torch.sqrt(tmp.sum()) + else: + output = tmp.sum() + if not scalar: + output = output.unsqueeze(0) + if high_precision: + output = output.float() + return output diff --git a/linghe/utils/rope.py b/linghe/utils/rope.py index ea30d1a..8931041 100644 --- a/linghe/utils/rope.py +++ b/linghe/utils/rope.py @@ -9,135 +9,195 @@ @triton.jit -def half_rope_forward_kernel(q_ptr, k_ptr, freqs_ptr, qo_ptr, ko_ptr, B, - q_stride, - k_stride, - H: tl.constexpr, - h: tl.constexpr, - D: tl.constexpr, - d: tl.constexpr, - TRANSPOSED: tl.constexpr - ): +def half_rope_forward_kernel( + q_ptr, + k_ptr, + freqs_ptr, + qo_ptr, + ko_ptr, + B, + q_stride, + k_stride, + H: tl.constexpr, + h: tl.constexpr, + D: tl.constexpr, + d: tl.constexpr, + TRANSPOSED: tl.constexpr, +): pid = tl.program_id(0) L = tl.num_programs(0) - freqs = tl.load(freqs_ptr + pid * D + tl.arange(0, D)) + freqs = tl.load(freqs_ptr + pid * D + tl.arange(0, D)).to(tl.float32) cos = tl.cos(freqs) sin = tl.sin(freqs) signs = tl.arange(0, 2).to(tl.float32) * 2 - 1 - for i in range(B): if TRANSPOSED: # [len, bs, q_head, head_dim] q = tl.load( - q_ptr + pid * B * q_stride + i * q_stride + 2 * D * tl.arange(0, H)[ - :, - None] + tl.arange( - 0, D)[None, :]) + q_ptr + + pid * B * q_stride + + i * q_stride + + 2 * D * tl.arange(0, H)[:, None] + + tl.arange(0, D)[None, :] + ) else: # [bs, len, q_head, head_dim] q = tl.load( - q_ptr + pid * q_stride + i * L * q_stride + 2 * D * tl.arange(0, H)[ - :, - None] + tl.arange( - 0, D)[None, :]) - qr = tl.reshape(tl.permute( - tl.flip(tl.permute(tl.reshape(q, (H, 2, d)), (0, 2, 1)), - dim=2) * signs, (0, 2, 1)), (H, D)) + q_ptr + + pid * q_stride + + i * L * q_stride + + 2 * D * tl.arange(0, H)[:, None] + + tl.arange(0, D)[None, :] + ) + qr = tl.reshape( + tl.permute( + tl.flip(tl.permute(tl.reshape(q, (H, 2, d)), (0, 2, 1)), dim=2) * signs, + (0, 2, 1), + ), + (H, D), + ) q = q * cos + qr * sin if TRANSPOSED: tl.store( - qo_ptr + pid * B * H * D * 2 + i * H * D * 2 + 2 * D * tl.arange(0, - H)[ - :, - None] + tl.arange( - 0, D)[None, :], q) + qo_ptr + + pid * B * H * D * 2 + + i * H * D * 2 + + 2 * D * tl.arange(0, H)[:, None] + + tl.arange(0, D)[None, :], + q, + ) q = tl.load( - q_ptr + pid * B * q_stride + i * q_stride + D + 2 * D * tl.arange(0, - H)[ - :, - None] + tl.arange( - 0, D)[None, :]) + q_ptr + + pid * B * q_stride + + i * q_stride + + D + + 2 * D * tl.arange(0, H)[:, None] + + tl.arange(0, D)[None, :] + ) tl.store( - qo_ptr + pid * B * H * D * 2 + i * H * D * 2 + D + 2 * D * tl.arange( - 0, H)[:, None] + tl.arange(0, D)[None, :], q) + qo_ptr + + pid * B * H * D * 2 + + i * H * D * 2 + + D + + 2 * D * tl.arange(0, H)[:, None] + + tl.arange(0, D)[None, :], + q, + ) else: tl.store( - qo_ptr + pid * H * D * 2 + i * L * H * D * 2 + 2 * D * tl.arange(0, - H)[ - :, - None] + tl.arange( - 0, D)[None, :], q) + qo_ptr + + pid * H * D * 2 + + i * L * H * D * 2 + + 2 * D * tl.arange(0, H)[:, None] + + tl.arange(0, D)[None, :], + q, + ) q = tl.load( - q_ptr + pid * q_stride + i * L * q_stride + D + 2 * D * tl.arange(0, - H)[ - :, - None] + tl.arange( - 0, D)[None, :]) + q_ptr + + pid * q_stride + + i * L * q_stride + + D + + 2 * D * tl.arange(0, H)[:, None] + + tl.arange(0, D)[None, :] + ) tl.store( - qo_ptr + pid * H * D * 2 + i * L * H * D * 2 + D + 2 * D * tl.arange( - 0, H)[:, None] + tl.arange(0, D)[None, :], q) - + qo_ptr + + pid * H * D * 2 + + i * L * H * D * 2 + + D + + 2 * D * tl.arange(0, H)[:, None] + + tl.arange(0, D)[None, :], + q, + ) for i in range(B): if TRANSPOSED: k = tl.load( - k_ptr + pid * B * k_stride + i * k_stride + 2 * D * tl.arange(0, h)[ - :, - None] + tl.arange( - 0, D)[None, :]) + k_ptr + + pid * B * k_stride + + i * k_stride + + 2 * D * tl.arange(0, h)[:, None] + + tl.arange(0, D)[None, :] + ) else: k = tl.load( - k_ptr + pid * k_stride + i * L * k_stride + 2 * D * tl.arange(0, h)[ - :, - None] + tl.arange( - 0, D)[None, :]) - kr = tl.reshape(tl.permute( - tl.flip(tl.permute(tl.reshape(k, (h, 2, d)), (0, 2, 1)), - dim=2) * signs, (0, 2, 1)), (h, D)) + k_ptr + + pid * k_stride + + i * L * k_stride + + 2 * D * tl.arange(0, h)[:, None] + + tl.arange(0, D)[None, :] + ) + kr = tl.reshape( + tl.permute( + tl.flip(tl.permute(tl.reshape(k, (h, 2, d)), (0, 2, 1)), dim=2) * signs, + (0, 2, 1), + ), + (h, D), + ) k = k * cos + kr * sin if TRANSPOSED: tl.store( - ko_ptr + pid * B * h * D * 2 + i * h * D * 2 + 2 * D * tl.arange(0, - h)[ - :, - None] + tl.arange( - 0, D)[None, :], k) + ko_ptr + + pid * B * h * D * 2 + + i * h * D * 2 + + 2 * D * tl.arange(0, h)[:, None] + + tl.arange(0, D)[None, :], + k, + ) k = tl.load( - k_ptr + pid * B * k_stride + i * k_stride + D + 2 * D * tl.arange(0, - h)[ - :, - None] + tl.arange( - 0, D)[None, :]) + k_ptr + + pid * B * k_stride + + i * k_stride + + D + + 2 * D * tl.arange(0, h)[:, None] + + tl.arange(0, D)[None, :] + ) tl.store( - ko_ptr + pid * B * h * D * 2 + i * h * D * 2 + D + 2 * D * tl.arange( - 0, h)[:, None] + tl.arange(0, D)[None, :], k) + ko_ptr + + pid * B * h * D * 2 + + i * h * D * 2 + + D + + 2 * D * tl.arange(0, h)[:, None] + + tl.arange(0, D)[None, :], + k, + ) else: tl.store( - ko_ptr + pid * h * D * 2 + i * L * h * D * 2 + 2 * D * tl.arange(0, - h)[ - :, - None] + tl.arange( - 0, D)[None, :], k) + ko_ptr + + pid * h * D * 2 + + i * L * h * D * 2 + + 2 * D * tl.arange(0, h)[:, None] + + tl.arange(0, D)[None, :], + k, + ) k = tl.load( - k_ptr + pid * k_stride + i * L * k_stride + D + 2 * D * tl.arange(0, - h)[ - :, - None] + tl.arange( - 0, D)[None, :]) + k_ptr + + pid * k_stride + + i * L * k_stride + + D + + 2 * D * tl.arange(0, h)[:, None] + + tl.arange(0, D)[None, :] + ) tl.store( - ko_ptr + pid * h * D * 2 + i * L * h * D * 2 + D + 2 * D * tl.arange( - 0, h)[:, None] + tl.arange(0, D)[None, :], k) + ko_ptr + + pid * h * D * 2 + + i * L * h * D * 2 + + D + + 2 * D * tl.arange(0, h)[:, None] + + tl.arange(0, D)[None, :], + k, + ) def triton_half_rope_forward(q, k, freqs, transposed=True): """ - apply norm to qk, then apply half rope to qk + apply half rope to qk Args: q: query tensor, [len, bs, q_head, head_dim] k: key tensor, [len, bs, kv_head, head_dim] @@ -147,12 +207,15 @@ def triton_half_rope_forward(q, k, freqs, transposed=True): - qo: query output - ko: key output """ + assert q.is_contiguous() and k.is_contiguous() and freqs.is_contiguous() if transposed: L, B, H, D = q.shape else: B, L, H, D = q.shape h = k.shape[2] assert freqs.shape[1] == D // 2 + assert triton.next_power_of_2(H) == H + assert triton.next_power_of_2(D) == D num_stages = 3 num_warps = 2 @@ -166,9 +229,11 @@ def triton_half_rope_forward(q, k, freqs, transposed=True): ko = torch.empty((B, L, h, D), dtype=q.dtype, device=q.device) grid = (L,) half_rope_forward_kernel[grid]( - q, k, + q, + k, freqs, - qo, ko, + qo, + ko, B, q_stride, k_stride, @@ -178,24 +243,27 @@ def triton_half_rope_forward(q, k, freqs, transposed=True): D // 4, transposed, num_stages=num_stages, - num_warps=num_warps + num_warps=num_warps, ) return qo, ko @triton.jit -def half_rope_backward_kernel(q_ptr, k_ptr, freqs_ptr, - B, - H: tl.constexpr, - h: tl.constexpr, - D: tl.constexpr, - d: tl.constexpr, - TRANSPOSED: tl.constexpr - ): +def half_rope_backward_kernel( + q_ptr, + k_ptr, + freqs_ptr, + B, + H: tl.constexpr, + h: tl.constexpr, + D: tl.constexpr, + d: tl.constexpr, + TRANSPOSED: tl.constexpr, +): pid = tl.program_id(0) L = tl.num_programs(0) - freqs = tl.load(freqs_ptr + pid * D + tl.arange(0, D)) + freqs = tl.load(freqs_ptr + pid * D + tl.arange(0, D)).to(tl.float32) cos = tl.cos(freqs) sin = tl.sin(freqs) signs = -tl.arange(0, 2).to(tl.float32) * 2 + 1 @@ -204,73 +272,94 @@ def half_rope_backward_kernel(q_ptr, k_ptr, freqs_ptr, for i in range(B): if TRANSPOSED: q = tl.load( - q_ptr + pid * B * H * D * 2 + i * H * D * 2 + 2 * D * tl.arange(0, - H)[ - :, - None] + tl.arange( - 0, D)[None, :]) + q_ptr + + pid * B * H * D * 2 + + i * H * D * 2 + + 2 * D * tl.arange(0, H)[:, None] + + tl.arange(0, D)[None, :] + ).to(tl.float32) else: q = tl.load( - q_ptr + pid * H * D * 2 + i * L * H * D * 2 + 2 * D * tl.arange(0, - H)[ - :, - None] + tl.arange( - 0, D)[None, :]) - qr = tl.reshape(tl.permute( - tl.flip(tl.permute(tl.reshape(q, (H, 2, d)), (0, 2, 1)), - dim=2) * signs, (0, 2, 1)), (H, D)) - q = q * cos + qr * sin + q_ptr + + pid * H * D * 2 + + i * L * H * D * 2 + + 2 * D * tl.arange(0, H)[:, None] + + tl.arange(0, D)[None, :] + ).to(tl.float32) + qr = tl.reshape( + tl.permute( + tl.flip(tl.permute(tl.reshape(q, (H, 2, d)), (0, 2, 1)), dim=2) * signs, + (0, 2, 1), + ), + (H, D), + ) + qo = (q * cos + qr * sin).to(q_ptr.dtype.element_ty) if TRANSPOSED: tl.store( - q_ptr + pid * B * H * D * 2 + i * H * D * 2 + 2 * D * tl.arange(0, - H)[ - :, - None] + tl.arange( - 0, D)[None, :], q) + q_ptr + + pid * B * H * D * 2 + + i * H * D * 2 + + 2 * D * tl.arange(0, H)[:, None] + + tl.arange(0, D)[None, :], + qo, + ) else: tl.store( - q_ptr + pid * H * D * 2 + i * L * H * D * 2 + 2 * D * tl.arange(0, - H)[ - :, - None] + tl.arange( - 0, D)[None, :], q) + q_ptr + + pid * H * D * 2 + + i * L * H * D * 2 + + 2 * D * tl.arange(0, H)[:, None] + + tl.arange(0, D)[None, :], + qo, + ) for i in range(B): if TRANSPOSED: k = tl.load( - k_ptr + pid * B * h * D * 2 + i * h * D * 2 + 2 * D * tl.arange(0, - h)[ - :, - None] + tl.arange( - 0, D)[None, :]) + k_ptr + + pid * B * h * D * 2 + + i * h * D * 2 + + 2 * D * tl.arange(0, h)[:, None] + + tl.arange(0, D)[None, :] + ).to(tl.float32) else: k = tl.load( - k_ptr + pid * h * D * 2 + i * L * h * D * 2 + 2 * D * tl.arange(0, - h)[ - :, - None] + tl.arange( - 0, D)[None, :]) - kr = tl.reshape(tl.permute( - tl.flip(tl.permute(tl.reshape(k, (h, 2, d)), (0, 2, 1)), - dim=2) * signs, (0, 2, 1)), (h, D)) - k = k * cos + kr * sin + k_ptr + + pid * h * D * 2 + + i * L * h * D * 2 + + 2 * D * tl.arange(0, h)[:, None] + + tl.arange(0, D)[None, :] + ).to(tl.float32) + kr = tl.reshape( + tl.permute( + tl.flip(tl.permute(tl.reshape(k, (h, 2, d)), (0, 2, 1)), dim=2) * signs, + (0, 2, 1), + ), + (h, D), + ) + ko = (k * cos + kr * sin).to(k_ptr.dtype.element_ty) if TRANSPOSED: tl.store( - k_ptr + pid * B * h * D * 2 + i * h * D * 2 + 2 * D * tl.arange(0, - h)[ - :, - None] + tl.arange( - 0, D)[None, :], k) + k_ptr + + pid * B * h * D * 2 + + i * h * D * 2 + + 2 * D * tl.arange(0, h)[:, None] + + tl.arange(0, D)[None, :], + ko, + ) else: tl.store( - k_ptr + pid * h * D * 2 + i * L * h * D * 2 + 2 * D * tl.arange(0, - h)[ - :, - None] + tl.arange( - 0, D)[None, :], k) + k_ptr + + pid * h * D * 2 + + i * L * h * D * 2 + + 2 * D * tl.arange(0, h)[:, None] + + tl.arange(0, D)[None, :], + ko, + ) def triton_half_rope_backward(q_grad, k_grad, freqs, inplace=False, transposed=True): + assert q_grad.is_contiguous() and k_grad.is_contiguous() assert inplace if transposed: L, B, H, D = q_grad.shape @@ -283,7 +372,8 @@ def triton_half_rope_backward(q_grad, k_grad, freqs, inplace=False, transposed=T grid = (L,) half_rope_backward_kernel[grid]( - q_grad, k_grad, + q_grad, + k_grad, freqs, B, H, @@ -292,167 +382,309 @@ def triton_half_rope_backward(q_grad, k_grad, freqs, inplace=False, transposed=T D // 4, transposed, num_stages=num_stages, - num_warps=num_warps + num_warps=num_warps, ) return q_grad, k_grad @triton.jit -def qk_norm_and_half_rope_forward_kernel(qkv_ptr, - q_norm_weight_ptr, k_norm_weight_ptr, - freqs_ptr, - qo_ptr, ko_ptr, vo_ptr, - B, - stride, - eps, - H: tl.constexpr, - h: tl.constexpr, - D: tl.constexpr, - d: tl.constexpr, - INTERLEAVED: tl.constexpr, - TRANSPOSED: tl.constexpr): +def qk_norm_and_half_rope_forward_kernel( + qkv_ptr, + q_norm_weight_ptr, + k_norm_weight_ptr, + freqs_ptr, + qo_ptr, + ko_ptr, + vo_ptr, + B, + stride, + eps, + H: tl.constexpr, + h: tl.constexpr, + H_p: tl.constexpr, + h_p: tl.constexpr, + D: tl.constexpr, + d: tl.constexpr, + INTERLEAVED: tl.constexpr, + TRANSPOSED: tl.constexpr, + SILU: tl.constexpr, +): pid = tl.program_id(0) L = tl.num_programs(0) DD = D * 2 - freqs = tl.load(freqs_ptr + pid * D + tl.arange(0, D)) + freqs = tl.load(freqs_ptr + pid * D + tl.arange(0, D)).to(tl.float32) cos = tl.cos(freqs) sin = tl.sin(freqs) signs = tl.arange(0, 2).to(tl.float32) * 2 - 1 - q_weight_0 = tl.load(q_norm_weight_ptr + tl.arange(0, D)) - q_weight_1 = tl.load(q_norm_weight_ptr + D + tl.arange(0, D)) + q_weight_0 = tl.load(q_norm_weight_ptr + tl.arange(0, D)).to(tl.float32) + q_weight_1 = tl.load(q_norm_weight_ptr + D + tl.arange(0, D)).to(tl.float32) q_ptr = qkv_ptr - w = H // h + w = H // h # H =8 h = 2 w=4 # [len, bs, q_head, head_dim] -> [bs, len, q_head, head_dim] if INTERLEAVED: - row_offs = tl.arange(0, H) + tl.arange(0, H) // w * 2 + # row_offs = tl.arange(0, H) + tl.arange(0, H) // w * 2 + row_offs = tl.arange(0, H_p) + tl.arange(0, H_p) // w * 2 + row_mask = row_offs[:, None] < (H + 2 * h) else: - row_offs = tl.arange(0, H) + # row_offs = tl.arange(0, H) + row_offs = tl.arange(0, H_p) + row_mask = row_offs[:, None] < H for i in range(B): if TRANSPOSED: - q0 = tl.load(q_ptr + pid * B * stride + i * stride + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]) - q1 = tl.load(q_ptr + pid * B * stride + i * stride + D + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]) + q0 = tl.load( + q_ptr + + pid * B * stride + + i * stride + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + mask=row_mask, + ).to(tl.float32) + q1 = tl.load( + q_ptr + + pid * B * stride + + i * stride + + D + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + mask=row_mask, + ).to(tl.float32) else: - q0 = tl.load(q_ptr + i * L * stride + pid * stride + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]) - q1 = tl.load(q_ptr + i * L * stride + pid * stride + D + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]) - - rms = 1 / tl.sqrt((tl.sum(q0 * q0, 1) + tl.sum(q1 * q1, 1)) / DD + eps) + q0 = tl.load( + q_ptr + + i * L * stride + + pid * stride + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + mask=row_mask, + ).to(tl.float32) + q1 = tl.load( + q_ptr + + i * L * stride + + pid * stride + + D + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + mask=row_mask, + ).to(tl.float32) + if SILU: + q0 = q0 * tl.sigmoid(q0.to(tl.float32)) + q1 = q1 * tl.sigmoid(q1.to(tl.float32)) + rms = tl.rsqrt((tl.sum(q0 * q0, 1) + tl.sum(q1 * q1, 1)) / DD + eps) q1 *= rms[:, None] q1 *= q_weight_1 + q_mask = tl.arange(0, H_p)[:, None] < H tl.store( - qo_ptr + pid * H * DD + i * L * H * DD + D + DD * tl.arange(0, H)[:, - None] + tl.arange( - 0, D)[None, :], q1) + qo_ptr + + pid * H * DD + + i * L * H * DD + + D + + DD * tl.arange(0, H_p)[:, None] + + tl.arange(0, D)[None, :], + q1, + mask=q_mask, + ) q0 *= rms[:, None] q0 *= q_weight_0 - qr = tl.reshape(tl.permute( - tl.flip(tl.permute(tl.reshape(q0, (H, 2, d)), (0, 2, 1)), - dim=2) * signs, (0, 2, 1)), (H, D)) + qr = tl.reshape( + tl.permute( + tl.flip(tl.permute(tl.reshape(q0, (H_p, 2, d)), (0, 2, 1)), dim=2) + * signs, + (0, 2, 1), + ), + (H_p, D), + ) q0 = q0 * cos + qr * sin tl.store( - qo_ptr + pid * H * DD + i * L * H * DD + DD * tl.arange(0, H)[:, - None] + tl.arange(0, - D)[ - None, :], q0) - - k_weight_0 = tl.load(k_norm_weight_ptr + tl.arange(0, D)) - k_weight_1 = tl.load(k_norm_weight_ptr + D + tl.arange(0, D)) + qo_ptr + + pid * H * DD + + i * L * H * DD + + DD * tl.arange(0, H_p)[:, None] + + tl.arange(0, D)[None, :], + q0, + mask=q_mask, + ) + + k_weight_0 = tl.load(k_norm_weight_ptr + tl.arange(0, D)).to(tl.float32) + k_weight_1 = tl.load(k_norm_weight_ptr + D + tl.arange(0, D)).to(tl.float32) if INTERLEAVED: - row_offs = tl.arange(0, h) * (w + 2) + # row_offs = tl.arange(0, h) * (w + 2) + row_offs = tl.arange(0, h_p) * (w + 2) + row_mask = row_offs[:, None] < (h * (w + 2)) k_ptr = qkv_ptr + DD * w else: - row_offs = tl.arange(0, h) + # row_offs = tl.arange(0, h) + row_offs = tl.arange(0, h_p) + row_mask = tl.arange(0, h_p)[:, None] < h k_ptr = qkv_ptr + DD * H + for i in range(B): if TRANSPOSED: - k0 = tl.load(k_ptr + pid * B * stride + i * stride + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]) + k0 = tl.load( + k_ptr + + pid * B * stride + + i * stride + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + mask=row_mask, + ).to(tl.float32) k1 = tl.load( - k_ptr + pid * B * stride + i * stride + D + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]) + k_ptr + + pid * B * stride + + i * stride + + D + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + mask=row_mask, + ).to(tl.float32) else: - k0 = tl.load(k_ptr + i * L * stride + pid * stride + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]) + k0 = tl.load( + k_ptr + + i * L * stride + + pid * stride + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + mask=row_mask, + ).to(tl.float32) k1 = tl.load( - k_ptr + i * L * stride + pid * stride + D + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]) - rms = 1 / tl.sqrt((tl.sum(k0 * k0, 1) + tl.sum(k1 * k1, 1)) / DD + eps) + k_ptr + + i * L * stride + + pid * stride + + D + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + mask=row_mask, + ).to(tl.float32) + + if SILU: + k0 = k0 * tl.sigmoid(k0) + k1 = k1 * tl.sigmoid(k1) + rms = tl.rsqrt((tl.sum(k0 * k0, 1) + tl.sum(k1 * k1, 1)) / DD + eps) k1 *= rms[:, None] k1 *= k_weight_1 + k_mask = tl.arange(0, h_p)[:, None] < h tl.store( - ko_ptr + pid * h * DD + i * L * h * DD + D + DD * tl.arange(0, h)[:, - None] + tl.arange( - 0, D)[None, :], k1) + ko_ptr + + pid * h * DD + + i * L * h * DD + + D + + DD * tl.arange(0, h_p)[:, None] + + tl.arange(0, D)[None, :], + k1, + mask=k_mask, + ) k0 *= rms[:, None] k0 *= k_weight_0 - kr = tl.reshape(tl.permute( - tl.flip(tl.permute(tl.reshape(k0, (h, 2, d)), (0, 2, 1)), - dim=2) * signs, (0, 2, 1)), (h, D)) + kr = tl.reshape( + tl.permute( + tl.flip(tl.permute(tl.reshape(k0, (h_p, 2, d)), (0, 2, 1)), dim=2) + * signs, + (0, 2, 1), + ), + (h_p, D), + ) k0 = k0 * cos + kr * sin tl.store( - ko_ptr + pid * h * DD + i * L * h * DD + DD * tl.arange(0, h)[:, - None] + tl.arange(0, - D)[ - None, :], k0) + ko_ptr + + pid * h * DD + + i * L * h * DD + + DD * tl.arange(0, h_p)[:, None] + + tl.arange(0, D)[None, :], + k0, + mask=k_mask, + ) if INTERLEAVED: - row_offs = tl.arange(0, h) * (w + 2) + # row_offs = tl.arange(0, h) * (w + 2) + row_offs = tl.arange(0, h_p) * (w + 2) + row_mask = row_offs[:, None] < (h * (w + 2)) v_ptr = qkv_ptr + DD * w + DD else: - row_offs = tl.arange(0, h) + # row_offs = tl.arange(0, h) + row_offs = tl.arange(0, h_p) + row_mask = tl.arange(0, h_p)[:, None] < h v_ptr = qkv_ptr + DD * H + DD * h + for i in range(B): if TRANSPOSED: - v0 = tl.load(v_ptr + pid * B * stride + i * stride + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]) - else: - v0 = tl.load(v_ptr + i * L * stride + pid * stride + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]) - tl.store( - vo_ptr + pid * h * DD + i * L * h * DD + DD * tl.arange(0, h)[:, - None] + tl.arange(0, - D)[ - None, :], v0) - - if TRANSPOSED: + v0 = tl.load( + v_ptr + + pid * B * stride + + i * stride + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + mask=row_mask, + ).to(tl.float32) v1 = tl.load( - v_ptr + pid * B * stride + i * stride + D + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]) + v_ptr + + pid * B * stride + + i * stride + + D + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + mask=row_mask, + ).to(tl.float32) else: + v0 = tl.load( + v_ptr + + i * L * stride + + pid * stride + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + mask=row_mask, + ).to(tl.float32) v1 = tl.load( - v_ptr + i * L * stride + pid * stride + D + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]) + v_ptr + + i * L * stride + + pid * stride + + D + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + mask=row_mask, + ).to(tl.float32) + if SILU: + v0 = v0 * tl.sigmoid(v0) + v1 = v1 * tl.sigmoid(v1) + + v_mask = tl.arange(0, h_p)[:, None] < h tl.store( - vo_ptr + pid * h * DD + i * L * h * DD + D + DD * tl.arange(0, h)[:, - None] + tl.arange( - 0, D)[None, :], v1) - - -def triton_qk_norm_and_half_rope_forward(qkv, q_norm_weight, k_norm_weight, - freqs, H=32, h=4, eps=1e-6, - interleaved=True, transposed=True): - + vo_ptr + + pid * h * DD + + i * L * h * DD + + DD * tl.arange(0, h_p)[:, None] + + tl.arange(0, D)[None, :], + v0, + mask=v_mask, + ) + tl.store( + vo_ptr + + pid * h * DD + + i * L * h * DD + + D + + DD * tl.arange(0, h_p)[:, None] + + tl.arange(0, D)[None, :], + v1, + mask=v_mask, + ) + + +def triton_qk_norm_and_half_rope_forward( + qkv, + q_norm_weight, + k_norm_weight, + freqs, + H=32, + h=4, + eps=1e-6, + interleaved=True, + transposed=True, + silu=False, +): """ split qkv to q/k/v, apply qk norm and half rope to q/k, transpose q/k/v to flash-attention layout @@ -470,17 +702,26 @@ def triton_qk_norm_and_half_rope_forward(qkv, q_norm_weight, k_norm_weight, transposed: whether qkv is tranposed transposed: [S, B, dim] non-transposed: [B, S, dim] + silu: apply silu on qkv before qk norm and rope Returns: - qo: shape [B, S, H, head_dim] - ko: shape [B, S, h, head_dim] - vo: shape [B, S, h, head_dim] """ + assert qkv.is_contiguous() and freqs.is_contiguous() + assert k_norm_weight.is_contiguous() and q_norm_weight.is_contiguous() if transposed: L, B, Dim = qkv.shape else: B, L, Dim = qkv.shape stride = qkv.stride(1) # qkv may be a slice of a tensor - D = Dim // (H + 2 * h) + D = k_norm_weight.size(0) + tp = (H + 2 * h) * D // Dim + if tp > 1: + H = H // tp + h = h // tp + # D = Dim // (H + 2 * h) # error with tp + assert freqs.size(0) == L and freqs.size(-1) == D // 2, f"{freqs.shape=} {L=} {D=}" dtype = qkv.dtype device = qkv.device qo = torch.empty((B, L, H, D), dtype=dtype, device=device) @@ -490,50 +731,68 @@ def triton_qk_norm_and_half_rope_forward(qkv, q_norm_weight, k_norm_weight, num_stages = 5 num_warps = 2 grid = (L,) + + H_p = triton.next_power_of_2(H) + h_p = triton.next_power_of_2(h) + qk_norm_and_half_rope_forward_kernel[grid]( qkv, - q_norm_weight, k_norm_weight, + q_norm_weight, + k_norm_weight, freqs, - qo, ko, vo, + qo, + ko, + vo, B, stride, eps, H, h, + H_p, + h_p, D // 2, D // 4, interleaved, transposed, + silu, num_stages=num_stages, - num_warps=num_warps + num_warps=num_warps, ) return qo, ko, vo @triton.jit -def qk_norm_and_half_rope_backward_kernel(gq_ptr, gk_ptr, gv_ptr, - qkv_ptr, - q_norm_weight_ptr, k_norm_weight_ptr, - freqs_ptr, - dqkv_ptr, - dqw_ptr, dkw_ptr, - B, - stride, - grad_stride, - eps, - H: tl.constexpr, - h: tl.constexpr, - D: tl.constexpr, - d: tl.constexpr, - INTERLEAVED: tl.constexpr, - TRANSPOSED: tl.constexpr - ): +def qk_norm_and_half_rope_backward_kernel( + gq_ptr, + gk_ptr, + gv_ptr, + qkv_ptr, + q_norm_weight_ptr, + k_norm_weight_ptr, + freqs_ptr, + dqkv_ptr, + dqw_ptr, + dkw_ptr, + B, + stride, + grad_stride, + eps, + H: tl.constexpr, + h: tl.constexpr, + H_p: tl.constexpr, + h_p: tl.constexpr, + D: tl.constexpr, + d: tl.constexpr, + INTERLEAVED: tl.constexpr, + TRANSPOSED: tl.constexpr, + SILU: tl.constexpr, +): pid = tl.program_id(0) L = tl.num_programs(0) DD = 2 * D w = H // h - freqs = tl.load(freqs_ptr + pid * D + tl.arange(0, D)) + freqs = tl.load(freqs_ptr + pid * D + tl.arange(0, D)).to(tl.float32) cos = tl.cos(freqs) sin = tl.sin(freqs) signs = -tl.arange(0, 2).to(tl.float32) * 2 + 1 @@ -547,189 +806,457 @@ def qk_norm_and_half_rope_backward_kernel(gq_ptr, gk_ptr, gv_ptr, dq_ptr = dqkv_ptr # [bs, len, q_head, head_dim] -> [len, bs, q_head, head_dim] if INTERLEAVED: - row_offs = tl.arange(0, H) + tl.arange(0, H) // w * 2 + # row_offs = tl.arange(0, H) + tl.arange(0, H) // w * 2 + row_offs = tl.arange(0, H_p) + tl.arange(0, H_p) // w * 2 + row_mask = row_offs[:, None] < (H + 2 * h) else: - row_offs = tl.arange(0, H) + # row_offs = tl.arange(0, H) + row_offs = tl.arange(0, H_p) + row_mask = row_offs[:, None] < H for i in range(B): gq_0 = tl.load( - gq_ptr + i * L * H * DD + pid * H * DD + DD * tl.arange(0, H)[:, - None] + tl.arange(0, - D)[ - None, :]) + gq_ptr + + i * L * H * DD + + pid * H * DD + + DD * tl.arange(0, H_p)[:, None] + + tl.arange(0, D)[None, :], + mask=tl.arange(0, H_p)[:, None] < H, + ).to(tl.float32) gq_1 = tl.load( - gq_ptr + i * L * H * DD + pid * H * DD + D + DD * tl.arange(0, H)[:, - None] + tl.arange( - 0, D)[None, :]) - - gq_r = tl.reshape(tl.permute( - tl.flip(tl.permute(tl.reshape(gq_0, (H, 2, d)), (0, 2, 1)), - dim=2) * signs, (0, 2, 1)), (H, D)) + gq_ptr + + i * L * H * DD + + pid * H * DD + + D + + DD * tl.arange(0, H_p)[:, None] + + tl.arange(0, D)[None, :], + mask=tl.arange(0, H_p)[:, None] < H, + ).to(tl.float32) + + gq_r = tl.reshape( + tl.permute( + tl.flip(tl.permute(tl.reshape(gq_0, (H_p, 2, d)), (0, 2, 1)), dim=2) + * signs, + (0, 2, 1), + ), + (H_p, D), + ) gq_0 = gq_0 * cos + gq_r * sin if TRANSPOSED: - q0 = tl.load(q_ptr + pid * B * stride + i * stride + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]) + q0 = tl.load( + q_ptr + + pid * B * stride + + i * stride + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + mask=row_mask, + ).to(tl.float32) q1 = tl.load( - q_ptr + pid * B * stride + i * stride + D + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]) + q_ptr + + pid * B * stride + + i * stride + + D + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + mask=row_mask, + ).to(tl.float32) + else: - q0 = tl.load(q_ptr + pid * stride + i * L * stride + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]) + q0 = tl.load( + q_ptr + + pid * stride + + i * L * stride + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + mask=row_mask, + ).to(tl.float32) q1 = tl.load( - q_ptr + pid * stride + i * L * stride + D + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]) + q_ptr + + pid * stride + + i * L * stride + + D + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + mask=row_mask, + ).to(tl.float32) - rms = tl.sqrt((tl.sum(q0 * q0, 1) + tl.sum(q1 * q1, 1)) / DD + eps) - r = (1 / rms)[:, None] + if SILU: + s0 = tl.sigmoid(q0) + s1 = tl.sigmoid(q1) + q_0 = q0 * s0 + q_1 = q1 * s1 - dqw_0 += tl.sum(q0 * gq_0 * r, 0) - dqw_1 += tl.sum(q1 * gq_1 * r, 0) + r = tl.rsqrt((tl.sum(q_0 * q_0, 1) + tl.sum(q_1 * q_1, 1)) / DD + eps)[ + :, None + ] - s = tl.sum(q0 * gq_0 * q_w0, 1) + tl.sum(q1 * gq_1 * q_w1, 1) + dqw_0 += tl.sum(q_0 * gq_0 * r, 0) + dqw_1 += tl.sum(q_1 * gq_1 * r, 0) - dq_0 = r * gq_0 * q_w0 - r * r * r / DD * q0 * s[:, None] - dq_1 = r * gq_1 * q_w1 - r * r * r / DD * q1 * s[:, None] + s = tl.sum(q_0 * gq_0 * q_w0, 1) + tl.sum(q_1 * gq_1 * q_w1, 1) + + dq_0 = r * gq_0 * q_w0 - r * r * r / DD * q_0 * s[:, None] + dq_1 = r * gq_1 * q_w1 - r * r * r / DD * q_1 * s[:, None] + + dq_0 = dq_0 * s0 * (1 + q0 * (1 - s0)) + dq_1 = dq_1 * s1 * (1 + q1 * (1 - s1)) + + else: + r = tl.rsqrt((tl.sum(q0 * q0, 1) + tl.sum(q1 * q1, 1)) / DD + eps)[:, None] + + dqw_0 += tl.sum(q0 * gq_0 * r, 0) + dqw_1 += tl.sum(q1 * gq_1 * r, 0) + + s = tl.sum(q0 * gq_0 * q_w0, 1) + tl.sum(q1 * gq_1 * q_w1, 1) + + dq_0 = r * gq_0 * q_w0 - r * r * r / DD * q0 * s[:, None] + dq_1 = r * gq_1 * q_w1 - r * r * r / DD * q1 * s[:, None] if TRANSPOSED: - tl.store(dq_ptr + pid * B * grad_stride + i * grad_stride + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :], dq_0) - tl.store(dq_ptr + pid * B * grad_stride + i * grad_stride + D + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :], dq_1) + # tl.store(dq_ptr + pid * B * grad_stride + i * grad_stride + DD * row_offs[:,None] + tl.arange(0, D)[None, :], dq_0) + # tl.store(dq_ptr + pid * B * grad_stride + i * grad_stride + D + DD * row_offs[:,None] + tl.arange(0, D)[None, :], dq_1) + tl.store( + dq_ptr + + pid * B * grad_stride + + i * grad_stride + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + dq_0, + mask=row_mask, + ) + tl.store( + dq_ptr + + pid * B * grad_stride + + i * grad_stride + + D + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + dq_1, + mask=row_mask, + ) + else: - tl.store(dq_ptr + pid * grad_stride + i * L * grad_stride + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :], dq_0) - tl.store(dq_ptr + pid * grad_stride + i * L * grad_stride + D + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :], dq_1) + # tl.store(dq_ptr + pid * grad_stride + i * L * grad_stride + DD * row_offs[:,None] + tl.arange(0, D)[None, :], dq_0) + # tl.store(dq_ptr + pid * grad_stride + i * L * grad_stride + D + DD * row_offs[:,None] + tl.arange(0, D)[None, :], dq_1) + tl.store( + dq_ptr + + pid * grad_stride + + i * L * grad_stride + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + dq_0, + mask=row_mask, + ) + tl.store( + dq_ptr + + pid * grad_stride + + i * L * grad_stride + + D + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + dq_1, + mask=row_mask, + ) tl.store(dqw_ptr + pid * D * 2 + tl.arange(0, D), dqw_0) tl.store(dqw_ptr + pid * D * 2 + D + tl.arange(0, D), dqw_1) - k_w0 = tl.load(k_norm_weight_ptr + tl.arange(0, D)) - k_w1 = tl.load(k_norm_weight_ptr + D + tl.arange(0, D)) + k_w0 = tl.load(k_norm_weight_ptr + tl.arange(0, D)).to(tl.float32) + k_w1 = tl.load(k_norm_weight_ptr + D + tl.arange(0, D)).to(tl.float32) dkw_0 = tl.zeros((D,), dtype=tl.float32) dkw_1 = tl.zeros((D,), dtype=tl.float32) if INTERLEAVED: - row_offs = tl.arange(0, h) * (w + 2) + # row_offs = tl.arange(0, h) * (w + 2) + row_offs = tl.arange(0, h_p) * (w + 2) + row_mask = row_offs[:, None] < (h * (w + 2)) k_ptr = qkv_ptr + DD * w dk_ptr = dqkv_ptr + DD * w else: - row_offs = tl.arange(0, h) + # row_offs = tl.arange(0, h) + row_offs = tl.arange(0, h_p) + row_mask = row_offs[:, None] < h k_ptr = qkv_ptr + DD * H dk_ptr = dqkv_ptr + DD * H # [bs, len, k_head, head_dim] -> [len, bs, k_head, head_dim] for i in range(B): gk_0 = tl.load( - gk_ptr + i * L * h * DD + pid * h * DD + DD * tl.arange(0, h)[:, - None] + tl.arange(0, - D)[ - None, :]) + gk_ptr + + i * L * h * DD + + pid * h * DD + + DD * tl.arange(0, h_p)[:, None] + + tl.arange(0, D)[None, :], + mask=tl.arange(0, h_p)[:, None] < h, + ).to(tl.float32) gk_1 = tl.load( - gk_ptr + i * L * h * DD + pid * h * DD + D + DD * tl.arange(0, h)[:, - None] + tl.arange( - 0, D)[None, :]) - - gk_r = tl.reshape(tl.permute( - tl.flip(tl.permute(tl.reshape(gk_0, (h, 2, d)), (0, 2, 1)), - dim=2) * signs, (0, 2, 1)), (h, D)) + gk_ptr + + i * L * h * DD + + pid * h * DD + + D + + DD * tl.arange(0, h_p)[:, None] + + tl.arange(0, D)[None, :], + mask=tl.arange(0, h_p)[:, None] < h, + ).to(tl.float32) + + gk_r = tl.reshape( + tl.permute( + tl.flip(tl.permute(tl.reshape(gk_0, (h_p, 2, d)), (0, 2, 1)), dim=2) + * signs, + (0, 2, 1), + ), + (h_p, D), + ) gk_0 = gk_0 * cos + gk_r * sin if TRANSPOSED: - k0 = tl.load(k_ptr + pid * B * stride + i * stride + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]) + k0 = tl.load( + k_ptr + + pid * B * stride + + i * stride + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + mask=row_mask, + ).to(tl.float32) k1 = tl.load( - k_ptr + pid * B * stride + i * stride + D + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]) + k_ptr + + pid * B * stride + + i * stride + + D + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + mask=row_mask, + ).to(tl.float32) else: - k0 = tl.load(k_ptr + pid * stride + i * L * stride + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]) + k0 = tl.load( + k_ptr + + pid * stride + + i * L * stride + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + mask=row_mask, + ).to(tl.float32) k1 = tl.load( - k_ptr + pid * stride + i * L * stride + D + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]) + k_ptr + + pid * stride + + i * L * stride + + D + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + mask=row_mask, + ).to(tl.float32) - rms = tl.sqrt((tl.sum(k0 * k0, 1) + tl.sum(k1 * k1, 1)) / DD + eps) - r = (1 / rms)[:, None] + if SILU: - dkw_0 += tl.sum(k0 * gk_0 * r, 0) - dkw_1 += tl.sum(k1 * gk_1 * r, 0) + s0 = tl.sigmoid(k0) + s1 = tl.sigmoid(k1) + k_0 = k0 * s0 + k_1 = k1 * s1 - s = tl.sum(k0 * gk_0 * k_w0, 1) + tl.sum(k1 * gk_1 * k_w1, 1) + r = tl.rsqrt((tl.sum(k_0 * k_0, 1) + tl.sum(k_1 * k_1, 1)) / DD + eps)[ + :, None + ] - dk_0 = r * gk_0 * k_w0 - r * r * r / DD * k0 * s[:, None] - dk_1 = r * gk_1 * k_w1 - r * r * r / DD * k1 * s[:, None] + dkw_0 += tl.sum(k_0 * gk_0 * r, 0) + dkw_1 += tl.sum(k_1 * gk_1 * r, 0) + + s = tl.sum(k_0 * gk_0 * k_w0, 1) + tl.sum(k_1 * gk_1 * k_w1, 1) + + dk_0 = r * gk_0 * k_w0 - r * r * r / DD * k_0 * s[:, None] + dk_1 = r * gk_1 * k_w1 - r * r * r / DD * k_1 * s[:, None] + + dk_0 = dk_0 * s0 * (1 + k0 * (1 - s0)) + dk_1 = dk_1 * s1 * (1 + k1 * (1 - s1)) + + else: + r = tl.rsqrt((tl.sum(k0 * k0, 1) + tl.sum(k1 * k1, 1)) / DD + eps)[:, None] + + dkw_0 += tl.sum(k0 * gk_0 * r, 0) + dkw_1 += tl.sum(k1 * gk_1 * r, 0) + + s = tl.sum(k0 * gk_0 * k_w0, 1) + tl.sum(k1 * gk_1 * k_w1, 1) + + dk_0 = r * gk_0 * k_w0 - r * r * r / DD * k0 * s[:, None] + dk_1 = r * gk_1 * k_w1 - r * r * r / DD * k1 * s[:, None] if TRANSPOSED: - tl.store(dk_ptr + pid * B * grad_stride + i * grad_stride + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :], dk_0) - tl.store(dk_ptr + pid * B * grad_stride + i * grad_stride + D + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :], dk_1) + tl.store( + dk_ptr + + pid * B * grad_stride + + i * grad_stride + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + dk_0, + mask=row_mask, + ) + tl.store( + dk_ptr + + pid * B * grad_stride + + i * grad_stride + + D + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + dk_1, + mask=row_mask, + ) else: - tl.store(dk_ptr + pid * grad_stride + i * L * grad_stride + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :], dk_0) - tl.store(dk_ptr + pid * grad_stride + i * L * grad_stride + D + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :], dk_1) + tl.store( + dk_ptr + + pid * grad_stride + + i * L * grad_stride + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + dk_0, + mask=row_mask, + ) + tl.store( + dk_ptr + + pid * grad_stride + + i * L * grad_stride + + D + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + dk_1, + mask=row_mask, + ) + tl.store(dkw_ptr + pid * D * 2 + tl.arange(0, D), dkw_0) tl.store(dkw_ptr + pid * D * 2 + D + tl.arange(0, D), dkw_1) # [bs, len, k_head, head_dim] -> [len, bs, k_head + 2 * kv_head, head_dim] if INTERLEAVED: - row_offs = tl.arange(0, h) * (w + 2) + # row_offs = tl.arange(0, h) * (w + 2) + row_offs = tl.arange(0, h_p) * (w + 2) + row_mask = row_offs[:, None] < (h * (w + 2)) + v_ptr = qkv_ptr + DD * w + DD dv_ptr = dqkv_ptr + DD * w + DD else: - row_offs = tl.arange(0, h) + # row_offs = tl.arange(0, h) + row_offs = tl.arange(0, h_p) + row_mask = row_offs[:, None] < h + v_ptr = qkv_ptr + DD * H + DD * h dv_ptr = dqkv_ptr + DD * H + DD * h for i in range(B): - v0 = tl.load( - gv_ptr + i * L * h * DD + pid * h * DD + DD * tl.arange(0, h)[:, - None] + tl.arange(0, - D)[ - None, :]) - - if TRANSPOSED: - tl.store(dv_ptr + pid * B * grad_stride + i * grad_stride + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :], v0) - else: - tl.store(dv_ptr + pid * grad_stride + i * L * grad_stride + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :], v0) - - v1 = tl.load( - gv_ptr + i * L * h * DD + pid * h * DD + D + DD * tl.arange(0, h)[:, - None] + tl.arange( - 0, D)[None, :]) + gv_0 = tl.load( + gv_ptr + + i * L * h * DD + + pid * h * DD + + DD * tl.arange(0, h_p)[:, None] + + tl.arange(0, D)[None, :], + mask=tl.arange(0, h_p)[:, None] < h, + ).to(tl.float32) + gv_1 = tl.load( + gv_ptr + + i * L * h * DD + + pid * h * DD + + D + + DD * tl.arange(0, h_p)[:, None] + + tl.arange(0, D)[None, :], + mask=tl.arange(0, h_p)[:, None] < h, + ).to(tl.float32) + + if SILU: + if TRANSPOSED: + v0 = tl.load( + v_ptr + + pid * B * stride + + i * stride + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + mask=row_mask, + ).to(tl.float32) + v1 = tl.load( + v_ptr + + pid * B * stride + + i * stride + + D + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + mask=row_mask, + ).to(tl.float32) + else: + v0 = tl.load( + v_ptr + + i * L * stride + + pid * stride + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + mask=row_mask, + ).to(tl.float32) + v1 = tl.load( + v_ptr + + i * L * stride + + pid * stride + + D + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + mask=row_mask, + ).to(tl.float32) + + s0 = tl.sigmoid(v0) + s1 = tl.sigmoid(v1) + dv_0 = gv_0 * s0 * (1 + v0 * (1 - s0)) + dv_1 = gv_1 * s1 * (1 + v1 * (1 - s1)) + else: + dv_0 = gv_0 + dv_1 = gv_1 if TRANSPOSED: - tl.store(dv_ptr + pid * B * grad_stride + i * grad_stride + D + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :], v1) + tl.store( + dv_ptr + + pid * B * grad_stride + + i * grad_stride + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + dv_0, + mask=row_mask, + ) + tl.store( + dv_ptr + + pid * B * grad_stride + + i * grad_stride + + D + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + dv_1, + mask=row_mask, + ) else: - tl.store(dv_ptr + pid * grad_stride + i * L * grad_stride + D + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :], v1) - - -def triton_qk_norm_and_half_rope_backward(gq, gk, gv, qkv, q_norm_weight, - k_norm_weight, freqs, eps=1e-6, - interleaved=True, transposed=True): + tl.store( + dv_ptr + + pid * grad_stride + + i * L * grad_stride + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + dv_0, + mask=row_mask, + ) + tl.store( + dv_ptr + + pid * grad_stride + + i * L * grad_stride + + D + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + dv_1, + mask=row_mask, + ) + + +def triton_qk_norm_and_half_rope_backward( + gq, + gk, + gv, + qkv, + q_norm_weight, + k_norm_weight, + freqs, + eps=1e-6, + interleaved=True, + transposed=True, + silu=False, +): """ backward kernel of triton_qk_norm_and_half_rope_forward Args: @@ -747,17 +1274,17 @@ def triton_qk_norm_and_half_rope_backward(gq, gk, gv, qkv, q_norm_weight, transposed: whether qkv is tranposed transposed: [S, B, dim] non-transposed: [B, S, dim] + silu: whether silu is applied to qkv Returns: - dqkv: gradient of qkv - dqw: gradient of q_norm_weight - dkw: gradient of k_norm_weight """ + assert gq.is_contiguous() and gk.is_contiguous() and gv.is_contiguous() B, L, H, D = gq.shape - stride = qkv.stride(1) h = gk.shape[2] - num_stages = 5 - num_warps = 1 + stride = qkv.stride(1) dtype = gq.dtype device = gq.device @@ -770,27 +1297,1352 @@ def triton_qk_norm_and_half_rope_backward(gq, gk, gv, qkv, q_norm_weight, tmp_dqw = torch.empty((L, D), dtype=torch.float32, device=device) tmp_dkw = torch.empty((L, D), dtype=torch.float32, device=device) + H_p = triton.next_power_of_2(H) + h_p = triton.next_power_of_2(h) + + num_stages = 5 + num_warps = 1 grid = (L,) + qk_norm_and_half_rope_backward_kernel[grid]( - gq, gk, gv, + gq, + gk, + gv, qkv, - q_norm_weight, k_norm_weight, + q_norm_weight, + k_norm_weight, freqs, dqkv, - tmp_dqw, tmp_dkw, + tmp_dqw, + tmp_dkw, B, stride, grad_stride, eps, H, h, + H_p, + h_p, D // 2, D // 4, interleaved, transposed, + silu, + num_stages=num_stages, + num_warps=num_warps, + ) + dqw = tmp_dqw.sum(0) + dkw = tmp_dkw.sum(0) + return dqkv, dqw, dkw + + +@triton.jit +def _get_varlen_token_idx(cu_seqlens, pid_m, seq_num, block, cp_rank, cp_size): + cu = 0 + cun = 1048576 + for i in range(tl.cdiv(seq_num + 1, block)): + cus = ( + tl.load( + cu_seqlens + i * block + tl.arange(0, block), + mask=i * block + tl.arange(0, block) <= seq_num, + ) + // cp_size + ) + cu = tl.maximum(tl.max(tl.where(cus > pid_m, 0, cus), 0), cu) + for i in range(tl.cdiv(seq_num + 1, block)): + cus = ( + tl.load( + cu_seqlens + i * block + tl.arange(0, block), + mask=i * block + tl.arange(0, block) <= seq_num, + ) + // cp_size + ) + cun = tl.minimum(tl.min(tl.where(cus <= cu, 2**24, cus), 0), cun) + length = cun - cu + token_idx = pid_m - cu + + if cp_size > 1: + if token_idx < length // 2: + token_idx = token_idx + cp_rank * length // 2 + else: + token_idx = (token_idx - length // 2) + ( + 2 * cp_size - cp_rank - 1 + ) * length // 2 + return token_idx + + +# @triton.jit +# def _get_varlen_token_idx(cu_seqlens, pid_m, seq_num, padded_seq_num, cp_rank, +# cp_size): +# cus = tl.load(cu_seqlens + tl.arange(0, padded_seq_num), +# mask=tl.arange(0, padded_seq_num) <= seq_num) // cp_size +# cu = tl.max(tl.where(cus > pid_m, 0, cus), 0) +# cun = tl.min(tl.where(cus <= cu, 2 ** 24, cus), 0) +# length = cun - cu +# token_idx = pid_m - cu + +# if cp_size > 1: +# if token_idx < length // 2: +# token_idx = token_idx + cp_rank * length // 2 +# else: +# token_idx = (token_idx - length // 2) + ( +# 2 * cp_size - cp_rank - 1 +# ) * length // 2 +# return token_idx + + +# not used +@triton.jit +def _get_fixlen_token_idx(num_tokens, pid_m, seq_num, cp_rank, cp_size, transpose): + L = num_tokens // seq_num + if transpose: + token_idx = pid_m % L + else: + token_idx = pid_m // seq_num + if cp_size > 1: + if token_idx < L // 2: + token_idx = token_idx + cp_rank * L // 2 + else: + token_idx = (token_idx - L // 2) + (2 * cp_size - cp_rank - 1) * L // 2 + return token_idx + + +@triton.jit +def varlen_qk_norm_and_half_rope_forward_kernel( + qkv_ptr, + q_norm_weight_ptr, + k_norm_weight_ptr, + freqs_ptr, + cu_seqlens_q_ptr, + cu_seqlens_kv_ptr, + qo_ptr, + ko_ptr, + vo_ptr, + stride, + eps, + mscale, + cp_rank, + B, + PB: tl.constexpr, + H: tl.constexpr, + h: tl.constexpr, + PH: tl.constexpr, + ph: tl.constexpr, + D: tl.constexpr, + d: tl.constexpr, + INTERLEAVED: tl.constexpr, + SILU: tl.constexpr, + CP_SIZE: tl.constexpr, + REUSE: tl.constexpr, +): + pid = tl.program_id(0) + + pos = _get_varlen_token_idx(cu_seqlens_q_ptr, pid, B, PB, cp_rank, CP_SIZE) + + DD = D * 2 + + freqs = tl.load(freqs_ptr + pos * D + tl.arange(0, D)).to(tl.float32) + cos = tl.cos(freqs) * mscale + sin = tl.sin(freqs) * mscale + signs = tl.arange(0, 2).to(tl.float32) * 2 - 1 + + q_weight_0 = tl.load(q_norm_weight_ptr + tl.arange(0, D)).to(tl.float32) + q_weight_1 = tl.load(q_norm_weight_ptr + D + tl.arange(0, D)).to(tl.float32) + q_ptr = qkv_ptr + w = H // h + + # [len, bs, q_head, head_dim] -> [bs, len, q_head, head_dim] + if INTERLEAVED: + row_offs = tl.arange(0, PH) + tl.arange(0, PH) // w * 2 + row_mask = row_offs[:, None] < (H + 2 * h) + else: + row_offs = tl.arange(0, H) + row_mask = row_offs[:, None] < H + q_mask = tl.arange(0, PH)[:, None] < H + + q0 = tl.load( + q_ptr + pid * stride + DD * row_offs[:, None] + tl.arange(0, D)[None, :], + mask=row_mask, + ).to(tl.float32) + q1 = tl.load( + q_ptr + pid * stride + D + DD * row_offs[:, None] + tl.arange(0, D)[None, :], + mask=row_mask, + ).to(tl.float32) + + if SILU: + q0 = q0 * tl.sigmoid(q0) + q1 = q1 * tl.sigmoid(q1) + rms = tl.rsqrt((tl.sum(q0 * q0, 1) + tl.sum(q1 * q1, 1)) / DD + eps) + q1 *= rms[:, None] + q1 *= q_weight_1 + tl.store( + qo_ptr + + pid * H * DD + + D + + DD * tl.arange(0, PH)[:, None] + + tl.arange(0, D)[None, :], + q1, + mask=q_mask, + ) + + q0 *= rms[:, None] + q0 *= q_weight_0 + qr = tl.reshape( + tl.permute( + tl.flip(tl.permute(tl.reshape(q0, (PH, 2, d)), (0, 2, 1)), dim=2) * signs, + (0, 2, 1), + ), + (PH, D), + ) + q0 = q0 * cos + qr * sin + tl.store( + qo_ptr + + pid * H * DD + + DD * tl.arange(0, PH)[:, None] + + tl.arange(0, D)[None, :], + q0, + mask=q_mask, + ) + + k_weight_0 = tl.load(k_norm_weight_ptr + tl.arange(0, D)).to(tl.float32) + k_weight_1 = tl.load(k_norm_weight_ptr + D + tl.arange(0, D)).to(tl.float32) + + if not REUSE: + pos = _get_varlen_token_idx(cu_seqlens_kv_ptr, pid, B, PB, cp_rank, CP_SIZE) + freqs = tl.load(freqs_ptr + pos * D + tl.arange(0, D)).to(tl.float32) + cos = tl.cos(freqs) * mscale + sin = tl.sin(freqs) * mscale + + if INTERLEAVED: + row_offs = tl.arange(0, ph) * (w + 2) + k_ptr = qkv_ptr + DD * w + row_mask = row_offs[:, None] < (h * (w + 2)) + else: + row_offs = tl.arange(0, ph) + k_ptr = qkv_ptr + DD * H + row_mask = tl.arange(0, ph)[:, None] < h + + k0 = tl.load( + k_ptr + pid * stride + DD * row_offs[:, None] + tl.arange(0, D)[None, :], + mask=row_mask, + ).to(tl.float32) + k1 = tl.load( + k_ptr + pid * stride + D + DD * row_offs[:, None] + tl.arange(0, D)[None, :], + mask=row_mask, + ).to(tl.float32) + + if SILU: + k0 = k0 * tl.sigmoid(k0) + k1 = k1 * tl.sigmoid(k1) + rms = tl.rsqrt((tl.sum(k0 * k0, 1) + tl.sum(k1 * k1, 1)) / DD + eps) + k1 *= rms[:, None] + k1 *= k_weight_1 + k_mask = tl.arange(0, ph)[:, None] < h + tl.store( + ko_ptr + + pid * h * DD + + D + + DD * tl.arange(0, ph)[:, None] + + tl.arange(0, D)[None, :], + k1, + mask=k_mask, + ) + + k0 *= rms[:, None] + k0 *= k_weight_0 + kr = tl.reshape( + tl.permute( + tl.flip(tl.permute(tl.reshape(k0, (ph, 2, d)), (0, 2, 1)), dim=2) * signs, + (0, 2, 1), + ), + (ph, D), + ) + k0 = k0 * cos + kr * sin + tl.store( + ko_ptr + + pid * h * DD + + DD * tl.arange(0, ph)[:, None] + + tl.arange(0, D)[None, :], + k0, + mask=k_mask, + ) + + if INTERLEAVED: + row_offs = tl.arange(0, ph) * (w + 2) + row_mask = row_offs[:, None] < (h * (w + 2)) + v_ptr = qkv_ptr + DD * w + DD + else: + row_offs = tl.arange(0, ph) + row_mask = tl.arange(0, ph)[:, None] < h + v_ptr = qkv_ptr + DD * H + DD * h + + v0 = tl.load( + v_ptr + pid * stride + DD * row_offs[:, None] + tl.arange(0, D)[None, :], + mask=row_mask, + ).to(tl.float32) + v1 = tl.load( + v_ptr + pid * stride + D + DD * row_offs[:, None] + tl.arange(0, D)[None, :], + mask=row_mask, + ).to(tl.float32) + + if SILU: + v0 = v0 * tl.sigmoid(v0) + v1 = v1 * tl.sigmoid(v1) + + v_mask = tl.arange(0, ph)[:, None] < h + tl.store( + vo_ptr + + pid * h * DD + + DD * tl.arange(0, ph)[:, None] + + tl.arange(0, D)[None, :], + v0, + mask=v_mask, + ) + tl.store( + vo_ptr + + pid * h * DD + + D + + DD * tl.arange(0, ph)[:, None] + + tl.arange(0, D)[None, :], + v1, + mask=v_mask, + ) + + +def triton_varlen_qk_norm_and_half_rope_forward( + qkv, + q_norm_weight, + k_norm_weight, + freqs, + cu_seqlens_q, + cu_seqlens_kv, + H=32, + h=4, + eps=1e-6, + interleaved=True, + silu=False, + cp_rank=0, + cp_size=1, + mscale=1.0, + reuse=False, +): + """ + split qkv to q/k/v, apply qk norm and half rope to q/k, + transpose q/k/v to flash-attention layout + Args: + qkv: QKV tensor with size of [S, B, dim], heads are interleaved + q_norm_weight: rms norm weight for query + k_norm_weight: rms norm weight for key + freqs: Freqs tensor based on half dim. + H: Number of attention heads. + h: Number of key/value heads. + eps: epsilon value for L2 normalization. + interleaved: whether head of qkv is interleaved, + interleaved: [q...qkvq...qkv] + non-interleaved: [q...qk...kv...v] + silu: apply silu on qkv before qk norm and rope + Returns: + - qo: shape [B, S, H, head_dim] + - ko: shape [B, S, h, head_dim] + - vo: shape [B, S, h, head_dim] + """ + assert qkv.is_contiguous() and q_norm_weight.is_contiguous() + assert k_norm_weight.is_contiguous() and freqs.is_contiguous() + T, Dim = qkv.shape + D = k_norm_weight.size(0) + stride = qkv.stride(0) # qkv may be a slice of a tensor + + tp = (H + 2 * h) * D // Dim + if tp > 1: + H = H // tp + h = h // tp + + D = Dim // (H + 2 * h) + B = cu_seqlens_q.size(0) - 1 + PB = 128 + dtype = qkv.dtype + device = qkv.device + qo = torch.empty((T, H, D), dtype=dtype, device=device) + ko = torch.empty((T, h, D), dtype=dtype, device=device) + vo = torch.empty((T, h, D), dtype=dtype, device=device) + + num_stages = 5 + num_warps = 2 + grid = (T,) + + PH = triton.next_power_of_2(H) + ph = triton.next_power_of_2(h) + + varlen_qk_norm_and_half_rope_forward_kernel[grid]( + qkv, + q_norm_weight, + k_norm_weight, + freqs, + cu_seqlens_q, + cu_seqlens_kv, + qo, + ko, + vo, + stride, + eps, + mscale, + cp_rank, + B, + PB, + H, + h, + PH, + ph, + D // 2, + D // 4, + interleaved, + silu, + cp_size, + reuse, + num_stages=num_stages, + num_warps=num_warps, + ) + return qo, ko, vo + + +@triton.jit +def varlen_qk_norm_and_half_rope_backward_kernel( + gq_ptr, + gk_ptr, + gv_ptr, + qkv_ptr, + q_norm_weight_ptr, + k_norm_weight_ptr, + freqs_ptr, + cu_seqlens_q_ptr, + cu_seqlens_kv_ptr, + dqkv_ptr, + dqw_ptr, + dkw_ptr, + B, + stride, + grad_stride, + eps, + mscale, + cp_rank, + PB: tl.constexpr, + H: tl.constexpr, + h: tl.constexpr, + PH: tl.constexpr, + ph: tl.constexpr, + D: tl.constexpr, + d: tl.constexpr, + INTERLEAVED: tl.constexpr, + SILU: tl.constexpr, + CP_SIZE: tl.constexpr, + REUSE: tl.constexpr, +): + pid = tl.program_id(0) + DD = 2 * D + w = H // h + + pos = _get_varlen_token_idx(cu_seqlens_q_ptr, pid, B, PB, cp_rank, CP_SIZE) + + freqs = tl.load(freqs_ptr + pos * D + tl.arange(0, D)).to(tl.float32) + cos = tl.cos(freqs) * mscale + sin = tl.sin(freqs) * mscale + signs = -tl.arange(0, 2).to(tl.float32) * 2 + 1 + + q_w0 = tl.load(q_norm_weight_ptr + tl.arange(0, D)).to(tl.float32) + q_w1 = tl.load(q_norm_weight_ptr + D + tl.arange(0, D)).to(tl.float32) + + dqw_0 = tl.zeros((D,), dtype=tl.float32) + dqw_1 = tl.zeros((D,), dtype=tl.float32) + q_ptr = qkv_ptr + dq_ptr = dqkv_ptr + # [bs, len, q_head, head_dim] -> [len, bs, q_head, head_dim] + if INTERLEAVED: + row_offs = tl.arange(0, PH) + tl.arange(0, PH) // w * 2 + row_mask = row_offs[:, None] < (H + 2 * h) + else: + row_offs = tl.arange(0, PH) + row_mask = row_offs[:, None] < H + + gq_0 = tl.load( + gq_ptr + + pid * H * DD + + DD * tl.arange(0, PH)[:, None] + + tl.arange(0, D)[None, :], + mask=tl.arange(0, PH)[:, None] < H, + ).to(tl.float32) + gq_1 = tl.load( + gq_ptr + + pid * H * DD + + D + + DD * tl.arange(0, PH)[:, None] + + tl.arange(0, D)[None, :], + mask=tl.arange(0, PH)[:, None] < H, + ).to(tl.float32) + + gq_r = tl.reshape( + tl.permute( + tl.flip(tl.permute(tl.reshape(gq_0, (PH, 2, d)), (0, 2, 1)), dim=2) * signs, + (0, 2, 1), + ), + (PH, D), + ) + gq_0 = gq_0 * cos + gq_r * sin + + q0 = tl.load( + q_ptr + pid * stride + DD * row_offs[:, None] + tl.arange(0, D)[None, :], + mask=row_mask, + ).to(tl.float32) + q1 = tl.load( + q_ptr + pid * stride + D + DD * row_offs[:, None] + tl.arange(0, D)[None, :], + mask=row_mask, + ).to(tl.float32) + + if SILU: + s0 = tl.sigmoid(q0) + s1 = tl.sigmoid(q1) + q_0 = q0 * s0 + q_1 = q1 * s1 + + r = tl.rsqrt((tl.sum(q_0 * q_0, 1) + tl.sum(q_1 * q_1, 1)) / DD + eps)[:, None] + + dqw_0 += tl.sum(q_0 * gq_0 * r, 0) + dqw_1 += tl.sum(q_1 * gq_1 * r, 0) + + s = tl.sum(q_0 * gq_0 * q_w0, 1) + tl.sum(q_1 * gq_1 * q_w1, 1) + + dq_0 = r * gq_0 * q_w0 - r * r * r / DD * q_0 * s[:, None] + dq_1 = r * gq_1 * q_w1 - r * r * r / DD * q_1 * s[:, None] + + dq_0 = dq_0 * s0 * (1 + q0 * (1 - s0)) + dq_1 = dq_1 * s1 * (1 + q1 * (1 - s1)) + + else: + r = tl.rsqrt((tl.sum(q0 * q0, 1) + tl.sum(q1 * q1, 1)) / DD + eps)[:, None] + + dqw_0 += tl.sum(q0 * gq_0 * r, 0) + dqw_1 += tl.sum(q1 * gq_1 * r, 0) + + s = tl.sum(q0 * gq_0 * q_w0, 1) + tl.sum(q1 * gq_1 * q_w1, 1) + + dq_0 = r * gq_0 * q_w0 - r * r * r / DD * q0 * s[:, None] + dq_1 = r * gq_1 * q_w1 - r * r * r / DD * q1 * s[:, None] + + tl.store( + dq_ptr + pid * grad_stride + DD * row_offs[:, None] + tl.arange(0, D)[None, :], + dq_0, + mask=row_mask, + ) + tl.store( + dq_ptr + + pid * grad_stride + + D + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + dq_1, + mask=row_mask, + ) + + tl.store(dqw_ptr + pid * D * 2 + tl.arange(0, D), dqw_0) + tl.store(dqw_ptr + pid * D * 2 + D + tl.arange(0, D), dqw_1) + + if not REUSE: + pos = _get_varlen_token_idx(cu_seqlens_kv_ptr, pid, B, PB, cp_rank, CP_SIZE) + freqs = tl.load(freqs_ptr + pos * D + tl.arange(0, D)).to(tl.float32) + cos = tl.cos(freqs) * mscale + sin = tl.sin(freqs) * mscale + + k_w0 = tl.load(k_norm_weight_ptr + tl.arange(0, D)).to(tl.float32) + k_w1 = tl.load(k_norm_weight_ptr + D + tl.arange(0, D)).to(tl.float32) + + dkw_0 = tl.zeros((D,), dtype=tl.float32) + dkw_1 = tl.zeros((D,), dtype=tl.float32) + if INTERLEAVED: + row_offs = tl.arange(0, ph) * (w + 2) + row_mask = row_offs[:, None] < (h * (w + 2)) + k_ptr = qkv_ptr + DD * w + dk_ptr = dqkv_ptr + DD * w + else: + row_offs = tl.arange(0, ph) + row_mask = row_offs[:, None] < h + k_ptr = qkv_ptr + DD * H + dk_ptr = dqkv_ptr + DD * H + + gk_0 = tl.load( + gk_ptr + + pid * h * DD + + DD * tl.arange(0, ph)[:, None] + + tl.arange(0, D)[None, :], + mask=tl.arange(0, ph)[:, None] < h, + ).to(tl.float32) + gk_1 = tl.load( + gk_ptr + + pid * h * DD + + D + + DD * tl.arange(0, ph)[:, None] + + tl.arange(0, D)[None, :], + mask=tl.arange(0, ph)[:, None] < h, + ).to(tl.float32) + + gk_r = tl.reshape( + tl.permute( + tl.flip(tl.permute(tl.reshape(gk_0, (ph, 2, d)), (0, 2, 1)), dim=2) * signs, + (0, 2, 1), + ), + (ph, D), + ) + gk_0 = gk_0 * cos + gk_r * sin + + k0 = tl.load( + k_ptr + pid * stride + DD * row_offs[:, None] + tl.arange(0, D)[None, :], + mask=row_mask, + ).to(tl.float32) + k1 = tl.load( + k_ptr + pid * stride + D + DD * row_offs[:, None] + tl.arange(0, D)[None, :], + mask=row_mask, + ).to(tl.float32) + + if SILU: + + s0 = tl.sigmoid(k0) + s1 = tl.sigmoid(k1) + k_0 = k0 * s0 + k_1 = k1 * s1 + + r = tl.rsqrt((tl.sum(k_0 * k_0, 1) + tl.sum(k_1 * k_1, 1)) / DD + eps)[:, None] + + dkw_0 += tl.sum(k_0 * gk_0 * r, 0) + dkw_1 += tl.sum(k_1 * gk_1 * r, 0) + + s = tl.sum(k_0 * gk_0 * k_w0, 1) + tl.sum(k_1 * gk_1 * k_w1, 1) + + dk_0 = r * gk_0 * k_w0 - r * r * r / DD * k_0 * s[:, None] + dk_1 = r * gk_1 * k_w1 - r * r * r / DD * k_1 * s[:, None] + + dk_0 = dk_0 * s0 * (1 + k0 * (1 - s0)) + dk_1 = dk_1 * s1 * (1 + k1 * (1 - s1)) + + else: + r = tl.rsqrt((tl.sum(k0 * k0, 1) + tl.sum(k1 * k1, 1)) / DD + eps)[:, None] + + dkw_0 += tl.sum(k0 * gk_0 * r, 0) + dkw_1 += tl.sum(k1 * gk_1 * r, 0) + + s = tl.sum(k0 * gk_0 * k_w0, 1) + tl.sum(k1 * gk_1 * k_w1, 1) + + dk_0 = r * gk_0 * k_w0 - r * r * r / DD * k0 * s[:, None] + dk_1 = r * gk_1 * k_w1 - r * r * r / DD * k1 * s[:, None] + + tl.store( + dk_ptr + pid * grad_stride + DD * row_offs[:, None] + tl.arange(0, D)[None, :], + dk_0, + mask=row_mask, + ) + tl.store( + dk_ptr + + pid * grad_stride + + D + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + dk_1, + mask=row_mask, + ) + + tl.store(dkw_ptr + pid * D * 2 + tl.arange(0, D), dkw_0) + tl.store(dkw_ptr + pid * D * 2 + D + tl.arange(0, D), dkw_1) + + # [t, k_head, head_dim] -> [t, k_head + 2 * kv_head, head_dim] + if INTERLEAVED: + row_offs = tl.arange(0, ph) * (w + 2) + row_mask = row_offs[:, None] < (h * (w + 2)) + v_ptr = qkv_ptr + DD * w + DD + dv_ptr = dqkv_ptr + DD * w + DD + else: + row_offs = tl.arange(0, ph) + row_mask = row_offs[:, None] < h + v_ptr = qkv_ptr + DD * H + DD * h + dv_ptr = dqkv_ptr + DD * H + DD * h + + gv_0 = tl.load( + gv_ptr + + pid * h * DD + + DD * tl.arange(0, ph)[:, None] + + tl.arange(0, D)[None, :], + mask=tl.arange(0, ph)[:, None] < h, + ).to(tl.float32) + gv_1 = tl.load( + gv_ptr + + pid * h * DD + + D + + DD * tl.arange(0, ph)[:, None] + + tl.arange(0, D)[None, :], + mask=tl.arange(0, ph)[:, None] < h, + ).to(tl.float32) + + if SILU: + v0 = tl.load( + v_ptr + pid * stride + DD * row_offs[:, None] + tl.arange(0, D)[None, :], + mask=row_mask, + ).to(tl.float32) + v1 = tl.load( + v_ptr + + pid * stride + + D + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + mask=row_mask, + ).to(tl.float32) + + s0 = tl.sigmoid(v0) + s1 = tl.sigmoid(v1) + dv_0 = gv_0 * s0 * (1 + v0 * (1 - s0)) + dv_1 = gv_1 * s1 * (1 + v1 * (1 - s1)) + else: + dv_0 = gv_0 + dv_1 = gv_1 + + tl.store( + dv_ptr + pid * grad_stride + DD * row_offs[:, None] + tl.arange(0, D)[None, :], + dv_0, + mask=row_mask, + ) + tl.store( + dv_ptr + + pid * grad_stride + + D + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + dv_1, + mask=row_mask, + ) + + +def triton_varlen_qk_norm_and_half_rope_backward( + gq, + gk, + gv, + qkv, + q_norm_weight, + k_norm_weight, + freqs, + cu_seqlens_q, + cu_seqlens_kv, + eps=1e-6, + interleaved=True, + silu=False, + cp_rank=0, + cp_size=1, + mscale=1.0, + reuse=False, +): + """ + backward kernel of triton_qk_norm_and_half_rope_forward + Args: + gq: gradient of qo, [len, bs, q_head, head_dim] + gk: gradient of ko, [len, bs, q_head, head_dim] + gv: gradient of vo, [len, bs, q_head, head_dim] + qkv: input qkv + q_norm_weight: rms norm weight for query + k_norm_weight: rms norm weight for key + freqs: Freqs tensor based on half dim. + eps: epsilon value for L2 normalization. + interleaved: whether head of qkv is interleaved, + interleaved: [q...qkvq...qkv] + non-interleaved: [q...qk...kv...v] + silu: whether silu is applied to qkv + + Returns: + - dqkv: gradient of qkv + - dqw: gradient of q_norm_weight + - dkw: gradient of k_norm_weight + """ + assert gq.is_contiguous() and gk.is_contiguous() and gv.is_contiguous() + T, H, D = gq.shape + stride = qkv.stride(0) + h = gk.shape[1] + B = cu_seqlens_q.size(0) - 1 + PB = 128 + num_stages = 5 + num_warps = 1 + + dtype = gq.dtype + device = gq.device + dqkv = torch.empty((T, (H + 2 * h) * D), dtype=dtype, device=device) + grad_stride = dqkv.stride(0) # for potential fused kernel + + tmp_dqw = torch.empty((T, D), dtype=torch.float32, device=device) + tmp_dkw = torch.empty((T, D), dtype=torch.float32, device=device) + + grid = (T,) + + PH = triton.next_power_of_2(H) + ph = triton.next_power_of_2(h) + + varlen_qk_norm_and_half_rope_backward_kernel[grid]( + gq, + gk, + gv, + qkv, + q_norm_weight, + k_norm_weight, + freqs, + cu_seqlens_q, + cu_seqlens_kv, + dqkv, + tmp_dqw, + tmp_dkw, + B, + stride, + grad_stride, + eps, + mscale, + cp_rank, + PB, + H, + h, + PH, + ph, + D // 2, + D // 4, + interleaved, + silu, + cp_size, + reuse, num_stages=num_stages, - num_warps=num_warps + num_warps=num_warps, ) - dqw = tmp_dqw.sum(0).to(dtype) - dkw = tmp_dkw.sum(0).to(dtype) + dqw = tmp_dqw.sum(0) + dkw = tmp_dkw.sum(0) return dqkv, dqw, dkw + + +@triton.jit +def mla_rope_forward_kernel( + q_ptr, + kv_ptr, + k_pos_emb_ptr, + freqs_ptr, + qo_ptr, + ko_ptr, + vo_ptr, + cu_seqlens_q_ptr, + cu_seqlens_kv_ptr, + mscale, + kpe_stride, + B, + cp_rank, + PB: tl.constexpr, + cp_size: tl.constexpr, + H: tl.constexpr, + VARLEN: tl.constexpr, + TRANSPOSE: tl.constexpr, + REUSE: tl.constexpr, +): + pid = tl.program_id(0) + num_tokens = tl.num_programs(0) + + if VARLEN: + pos = _get_varlen_token_idx(cu_seqlens_q_ptr, pid, B, PB, cp_rank, cp_size) + L = num_tokens + bid = 0 + else: + pos = pid // B + L = num_tokens // B + bid = pid % B + + freqs = tl.load(freqs_ptr + pos * 64 + tl.arange(0, 64)).to(tl.float32) + + cos = tl.cos(freqs) * mscale + sin = tl.sin(freqs) * mscale + signs = tl.arange(0, 2).to(tl.float32) * 2 - 1 + + q = tl.load( + q_ptr + + pid * H * 192 + + 128 + + 192 * tl.arange(0, H)[:, None] + + tl.arange(0, 64)[None, :] + ).to(tl.float32) + qt = tl.permute(tl.reshape(q, (H, 32, 2)), (0, 2, 1)) + q = tl.reshape(qt, (H, 64)) + qr = tl.reshape( + tl.permute(tl.flip(tl.permute(qt, (0, 2, 1)), dim=2) * signs, (0, 2, 1)), + (H, 64), + ) + q = q * cos + qr * sin + + # q0, q1 = tl.split(tl.reshape(q, (H, 32, 2))) + # qo0 = q0 * cos0 - q1 * sin0 + # qo1 = q1 * cos1 + q0 * sin1 + if TRANSPOSE: + # [L, B, H, D] -> [B, L, H, D] + qn = tl.load( + q_ptr + + pid * H * 192 + + 192 * tl.arange(0, H)[:, None] + + tl.arange(0, 128)[None, :] + ).to(tl.float32) + tl.store( + qo_ptr + + (bid * L + pos) * H * 192 + + 128 + + 192 * tl.arange(0, H)[:, None] + + tl.arange(0, 64)[None, :], + q, + ) + tl.store( + qo_ptr + + (bid * L + pos) * H * 192 + + 192 * tl.arange(0, H)[:, None] + + tl.arange(0, 128)[None, :], + qn, + ) + else: + tl.store( + q_ptr + + pid * H * 192 + + 128 + + 192 * tl.arange(0, H)[:, None] + + tl.arange(0, 64)[None, :], + q, + ) + + k = tl.load(k_pos_emb_ptr + pid * kpe_stride + tl.arange(0, 64)).to(tl.float32) + + if VARLEN and not REUSE: + pos = _get_varlen_token_idx(cu_seqlens_kv_ptr, pid, B, PB, cp_rank, cp_size) + freqs = tl.load(freqs_ptr + pos * 64 + tl.arange(0, 64)) + cos = tl.cos(freqs) * mscale + sin = tl.sin(freqs) * mscale + + kt = tl.permute(tl.reshape(k, (32, 2)), (1, 0)) + k = tl.reshape(kt, (64,)) + kr = tl.reshape( + tl.permute(tl.flip(tl.permute(kt, (1, 0)), dim=1) * signs, (1, 0)), (64,) + ) + k = k * cos + kr * sin + if TRANSPOSE: + tl.store( + ko_ptr + + (bid * L + pos) * H * 192 + + 128 + + 192 * tl.arange(0, H)[:, None] + + tl.arange(0, 64)[None, :], + k[None, :], + ) + else: + tl.store( + ko_ptr + + pid * H * 192 + + 128 + + 192 * tl.arange(0, H)[:, None] + + tl.arange(0, 64)[None, :], + k[None, :], + ) + + k = tl.load( + kv_ptr + + pid * H * 256 + + 256 * tl.arange(0, H)[:, None] + + tl.arange(0, 128)[None, :] + ).to(tl.float32) + if TRANSPOSE: + tl.store( + ko_ptr + + (bid * L + pos) * H * 192 + + 192 * tl.arange(0, H)[:, None] + + tl.arange(0, 128)[None, :], + k, + ) + else: + tl.store( + ko_ptr + + pid * H * 192 + + 192 * tl.arange(0, H)[:, None] + + tl.arange(0, 128)[None, :], + k, + ) + + v = tl.load( + kv_ptr + + pid * H * 256 + + 128 + + 256 * tl.arange(0, H)[:, None] + + tl.arange(0, 128)[None, :] + ).to(tl.float32) + if TRANSPOSE: + tl.store( + vo_ptr + + (bid * L + pos) * H * 128 + + 128 * tl.arange(0, H)[:, None] + + tl.arange(0, 128)[None, :], + v, + ) + else: + tl.store( + vo_ptr + + pid * H * 128 + + 128 * tl.arange(0, H)[:, None] + + tl.arange(0, 128)[None, :], + v, + ) + + +def triton_mla_rope_forward( + q, + kv, + k_pos_emb, + freqs, + mscale=1.0, + transpose=False, + cu_seqlens_q=None, + cu_seqlens_kv=None, + cp_rank=0, + cp_size=1, + reuse=False, +): + """ + apply MLA-type rope to qkv + Args: + q: query tensor, [len, bs, n_heads, 192] + kv: key-value tensor, [len, bs, n_heads, 256] + k_pos_emb: k pos emb, [len, bs, 1, 64] + freqs: rope freqs, [len, 64] + mscale: mscale for rope + transpose: whether transpose the output to [bs, len, n_heads, dim] layout + cu_seqlens_q: accummulated query length + cu_seqlens_kv: accummulated kv length + cp_rank: rank of context parallel + cp_size: size of context parallel + + Returns: + - qo: inplace updated query, [len, bs, n_heads, 192] if not transpose + else [bs, len, n_heads, 192] + - ko: key output, [len, bs, n_heads, 192] if not transpose + else [bs, len, n_heads, 192] + - vo: value output, [len, bs, n_heads, 128] if not transpose + else [bs, len, n_heads, 128] + """ + + assert q.is_contiguous() and freqs.is_contiguous() + VARLEN = cu_seqlens_q is not None + + dtype = q.dtype + device = q.device + if VARLEN: + assert cu_seqlens_kv is not None + N, H, D = q.shape + B = cu_seqlens_q.shape[0] - 1 + PB = 128 + qo = None + ko = torch.empty((N, H, 192), dtype=dtype, device=device) + vo = torch.empty((N, H, 128), dtype=dtype, device=device) + kpe_stride = k_pos_emb.stride(0) + else: + L, B, H, D = q.shape + PB = 1 + if transpose: + qo = torch.empty((B, L, H, 192), dtype=dtype, device=device) + ko = torch.empty((B, L, H, 192), dtype=dtype, device=device) + vo = torch.empty((B, L, H, 128), dtype=dtype, device=device) + else: + qo = None + ko = torch.empty((L, B, H, 192), dtype=dtype, device=device) + vo = torch.empty((L, B, H, 128), dtype=dtype, device=device) + N = L * B + kpe_stride = k_pos_emb.stride(0) if B == 1 else k_pos_emb.stride(1) + assert D == 192 and kv.shape[-1] == 256 and k_pos_emb.shape[-1] == 64 + assert ( + kv.stride(-2) == 256 and k_pos_emb.stride(-2) == 64 + ), f"{kv.stride()=} {k_pos_emb.stride()=}" + num_stages = 2 + num_warps = 2 + + grid = (N,) + mla_rope_forward_kernel[grid]( + q, + kv, + k_pos_emb, + freqs, + qo, + ko, + vo, + cu_seqlens_q, + cu_seqlens_kv, + mscale, + kpe_stride, + B, + cp_rank, + PB, + cp_size, + H, + VARLEN, + False if VARLEN else transpose, + reuse, + num_stages=num_stages, + num_warps=num_warps, + ) + if VARLEN or not transpose: + qo = q + return qo, ko, vo + + +@triton.jit +def mla_rope_backward_kernel( + q_ptr, + k_ptr, + v_ptr, + freqs_ptr, + dq_ptr, + dkv_ptr, + dp_ptr, + cu_seqlens_q_ptr, + cu_seqlens_kv_ptr, + mscale, + B, + cp_rank, + PB: tl.constexpr, + cp_size: tl.constexpr, + H: tl.constexpr, + VARLEN: tl.constexpr, + TRANSPOSED: tl.constexpr, + REUSE: tl.constexpr, +): + pid = tl.program_id(0) + num_tokens = tl.num_programs(0) + + if VARLEN: + pos = _get_varlen_token_idx(cu_seqlens_q_ptr, pid, B, PB, cp_rank, cp_size) + L = num_tokens + bid = 0 + else: + L = num_tokens // B + if TRANSPOSED: + # [B,L,H,D] + pos = pid % L + bid = pid // L + else: + # [L,B,H,D] + pos = pid // B + bid = pid % B + + freqs0 = tl.load(freqs_ptr + pos * 64 + tl.arange(0, 32)).to(tl.float32) + freqs1 = tl.load(freqs_ptr + pos * 64 + 32 + tl.arange(0, 32)).to(tl.float32) + + q0 = tl.load( + q_ptr + + pid * H * 192 + + 128 + + 192 * tl.arange(0, H)[:, None] + + tl.arange(0, 32)[None, :] + ).to(tl.float32) + q1 = tl.load( + q_ptr + + pid * H * 192 + + 160 + + 192 * tl.arange(0, H)[:, None] + + tl.arange(0, 32)[None, :] + ).to(tl.float32) + + cos0 = tl.cos(freqs0) * mscale + sin0 = tl.sin(freqs0) * mscale + + cos1 = tl.cos(freqs1) * mscale + sin1 = tl.sin(freqs1) * mscale + + dq0 = q0 * cos0 + q1 * sin1 + dq1 = q1 * cos1 - q0 * sin0 + dq = tl.reshape(tl.join(dq0, dq1), (H, 64)) + + if TRANSPOSED: + # [B,L,H,D] -> [L,B,H,D] + dqn = tl.load( + q_ptr + + pid * H * 192 + + 192 * tl.arange(0, H)[:, None] + + tl.arange(0, 128)[None, :] + ) + tl.store( + dq_ptr + + (pos * B + bid) * H * 192 + + 128 + + 192 * tl.arange(0, H)[:, None] + + tl.arange(0, 64)[None, :], + dq, + ) + tl.store( + dq_ptr + + (pos * B + bid) * H * 192 + + 192 * tl.arange(0, H)[:, None] + + tl.arange(0, 128)[None, :], + dqn, + ) + else: + tl.store( + q_ptr + + pid * H * 192 + + 128 + + 192 * tl.arange(0, H)[:, None] + + tl.arange(0, 64)[None, :], + dq, + ) + + # qr = tl.reshape(tl.permute( + # tl.flip(tl.permute(tl.reshape(q, (H, 2, 32)), (0, 2, 1)), + # dim=2) * signs, (0, 2, 1)), (H, 64)) + # q = q * cos + qr * sin + # q = tl.reshape(tl.permute(tl.reshape(q, (H, 2, 32)), (0, 2, 1)), (H, 64)) + + if VARLEN and not REUSE: + pos = _get_varlen_token_idx(cu_seqlens_kv_ptr, pid, B, PB, cp_rank, cp_size) + + freqs0 = tl.load(freqs_ptr + pos * 64 + tl.arange(0, 32)).to(tl.float32) + freqs1 = tl.load(freqs_ptr + pos * 64 + 32 + tl.arange(0, 32)).to(tl.float32) + + cos0 = tl.cos(freqs0) * mscale + sin0 = tl.sin(freqs0) * mscale + + cos1 = tl.cos(freqs1) * mscale + sin1 = tl.sin(freqs1) * mscale + + kp0 = tl.load( + k_ptr + + pid * H * 192 + + 128 + + 192 * tl.arange(0, H)[:, None] + + tl.arange(0, 32)[None, :] + ).to(tl.float32) + kp1 = tl.load( + k_ptr + + pid * H * 192 + + 160 + + 192 * tl.arange(0, H)[:, None] + + tl.arange(0, 32)[None, :] + ).to(tl.float32) + dkp0 = tl.sum(kp0 * cos0 + kp1 * sin1, 0) + dkp1 = tl.sum(kp1 * cos1 - kp0 * sin0, 0) + dkp = tl.reshape(tl.join(dkp0, dkp1), (64,)) + + if TRANSPOSED: + tl.store(dp_ptr + (pos * B + bid) * 64 + tl.arange(0, 64), dkp) + else: + tl.store(dp_ptr + pid * 64 + tl.arange(0, 64), dkp) + + k = tl.load( + k_ptr + + pid * H * 192 + + 192 * tl.arange(0, H)[:, None] + + tl.arange(0, 128)[None, :] + ) + if TRANSPOSED: + tl.store( + dkv_ptr + + (pos * B + bid) * H * 256 + + 256 * tl.arange(0, H)[:, None] + + tl.arange(0, 128)[None, :], + k, + ) + else: + tl.store( + dkv_ptr + + pid * H * 256 + + 256 * tl.arange(0, H)[:, None] + + tl.arange(0, 128)[None, :], + k, + ) + + v = tl.load( + v_ptr + + pid * H * 128 + + 128 * tl.arange(0, H)[:, None] + + tl.arange(0, 128)[None, :] + ) + + if TRANSPOSED: + tl.store( + dkv_ptr + + (pos * B + bid) * H * 256 + + 128 + + 256 * tl.arange(0, H)[:, None] + + tl.arange(0, 128)[None, :], + v, + ) + else: + tl.store( + dkv_ptr + + pid * H * 256 + + 128 + + 256 * tl.arange(0, H)[:, None] + + tl.arange(0, 128)[None, :], + v, + ) + + +def triton_mla_rope_backward( + q_grad, + k_grad, + v_grad, + freqs, + mscale=1.0, + transposed=False, + cu_seqlens_q=None, + cu_seqlens_kv=None, + cp_rank=0, + cp_size=1, + reuse=False, +): + assert q_grad.is_contiguous() and k_grad.is_contiguous() and v_grad.is_contiguous() + VARLEN = cu_seqlens_q is not None + + dtype = q_grad.dtype + device = q_grad.device + if VARLEN: + assert cu_seqlens_kv is not None + N, H, D = q_grad.shape + B = cu_seqlens_q.shape[0] - 1 + PB = 128 + dq = None + dkv = torch.empty((N, H, 256), dtype=dtype, device=device) + dp = torch.empty((N, 1, 64), dtype=dtype, device=device) + else: + if transposed: + B, L, H, D = q_grad.shape + N = L * B + dq = torch.empty((L, B, H, 192), dtype=dtype, device=device) + dkv = torch.empty((L, B, H, 256), dtype=dtype, device=device) + dp = torch.empty((L, B, 1, 64), dtype=dtype, device=device) + else: + L, B, H, D = q_grad.shape + N = L * B + dq = None + dkv = torch.empty((L, B, H, 256), dtype=dtype, device=device) + dp = torch.empty((L, B, 1, 64), dtype=dtype, device=device) + PB = 1 + + num_stages = 2 + num_warps = 4 + grid = (N,) + mla_rope_backward_kernel[grid]( + q_grad, + k_grad, + v_grad, + freqs, + dq, + dkv, + dp, + cu_seqlens_q, + cu_seqlens_kv, + mscale, + B, + cp_rank, + PB, + cp_size, + H, + VARLEN, + False if VARLEN else transposed, + reuse, + num_stages=num_stages, + num_warps=num_warps, + ) + if VARLEN or not transposed: + dq = q_grad + return dq, dkv, dp diff --git a/linghe/utils/scatter.py b/linghe/utils/scatter.py index 2a4e465..27a57ea 100644 --- a/linghe/utils/scatter.py +++ b/linghe/utils/scatter.py @@ -4,16 +4,23 @@ """ from typing import Optional + import torch import triton import triton.language as tl - @triton.jit -def aligned_scatter_add_kernel(x_ptr, o_ptr, indices_ptr, weights_ptr, M, - N: tl.constexpr, K: tl.constexpr, - SCALE: tl.constexpr): +def aligned_scatter_add_kernel( + x_ptr, + o_ptr, + indices_ptr, + weights_ptr, + M, + N: tl.constexpr, + K: tl.constexpr, + SCALE: tl.constexpr, +): pid = tl.program_id(axis=0) offs = tl.arange(0, N) @@ -30,10 +37,12 @@ def aligned_scatter_add_kernel(x_ptr, o_ptr, indices_ptr, weights_ptr, M, tl.store(o_ptr + pid * N + offs, sums) -def triton_aligned_scatter_add(x: torch.Tensor, - outputs: torch.Tensor, - indices: torch.Tensor, - weights: Optional[torch.Tensor] = None): +def triton_aligned_scatter_add( + x: torch.Tensor, + outputs: torch.Tensor, + indices: torch.Tensor, + weights: Optional[torch.Tensor] = None, +): """ scatter_add for megatron 0.11 Args: @@ -45,12 +54,13 @@ def triton_aligned_scatter_add(x: torch.Tensor, Returns: output tensor """ + assert x.is_contiguous() and outputs.is_contiguous() and indices.is_contiguous() M, N = x.shape m = outputs.size(0) indices = torch.argsort(indices) K = M // m - assert K * m == M + assert K * m == M and triton.next_power_of_2(N) == N SCALE = 1 if weights is not None else 0 num_stages = 5 @@ -58,19 +68,23 @@ def triton_aligned_scatter_add(x: torch.Tensor, grid = (m,) aligned_scatter_add_kernel[grid]( - x, outputs, + x, + outputs, indices, weights, - M, N, K, + M, + N, + K, SCALE, num_stages=num_stages, - num_warps=num_warps + num_warps=num_warps, ) return outputs # for deepep scatter_add -# atomic_add supports fp16 and fp32, but not bf16 +# atomic_add supports fp16 and fp32, but not bf16 + @triton.jit def scatter_add_kernel(x_ptr, o_ptr, indices_ptr, M, T, N: tl.constexpr): @@ -81,7 +95,7 @@ def scatter_add_kernel(x_ptr, o_ptr, indices_ptr, M, T, N: tl.constexpr): src_idx = pid * T + i dst_idx = tl.load(indices_ptr + src_idx, mask=src_idx < M) x = tl.load(x_ptr + src_idx * N + offs, mask=src_idx < M).to(tl.float32) - tl.atomic_add(o_ptr + dst_idx * N + offs, x) + tl.atomic_add(o_ptr + dst_idx * N + offs, x, sem="relaxed") @triton.jit @@ -106,12 +120,15 @@ def triton_scatter_add(x, outputs, indices): Returns: output tensor """ + assert x.is_contiguous() and outputs.is_contiguous() and indices.is_contiguous() M, N = x.shape + assert triton.next_power_of_2(N) == N - float_outputs = torch.zeros(outputs.shape, dtype=torch.float32, - device=outputs.device) + float_outputs = torch.zeros( + outputs.shape, dtype=torch.float32, device=outputs.device + ) - sm = torch.cuda.get_device_properties(x.device).multi_processor_count + sm = 512 T = triton.cdiv(M, sm) num_stages = 5 @@ -119,64 +136,64 @@ def triton_scatter_add(x, outputs, indices): grid = (sm,) scatter_add_kernel[grid]( - x, float_outputs, - indices, - M, T, N, - num_stages=num_stages, - num_warps=num_warps + x, float_outputs, indices, M, T, N, num_stages=num_stages, num_warps=num_warps ) m = outputs.shape[0] T = triton.cdiv(m, sm) grid = (sm,) fp32_to_bf16_kernel[grid]( - float_outputs, outputs, - m, T, N, - num_stages=num_stages, - num_warps=num_warps + float_outputs, outputs, m, T, N, num_stages=num_stages, num_warps=num_warps ) return outputs @triton.jit -def unpermute_with_mask_map_kernel(grads_ptr, probs_ptr, mask_map_ptr, - output_ptr, output_probs_ptr, - num_experts: tl.constexpr, N: tl.constexpr, - PROB: tl.constexpr): +def unpermute_with_mask_map_kernel( + grads_ptr, + probs_ptr, + mask_map_ptr, + output_ptr, + output_probs_ptr, + n, + N: tl.constexpr, + num_experts: tl.constexpr, + PROB: tl.constexpr, +): pid = tl.program_id(axis=0) - + n = n.to(tl.int64) + # sums = tl.zeros((N,), dtype=tl.float32) sums = tl.zeros((N,), dtype=tl.float32) - indices = tl.load( - mask_map_ptr + pid * num_experts + tl.arange(0, num_experts)) + indices = tl.load(mask_map_ptr + pid * num_experts + tl.arange(0, num_experts)) count = tl.sum(tl.where(indices >= 0, 1, 0)) - mask_indices = tl.where(indices < 0, 2 ** 20, indices) + mask_indices = tl.where(indices < 0, 2**24, indices) idx = tl.argmin(mask_indices, 0) index = tl.min(mask_indices) for i in range(count): - - mask = index >= 0 - sums += tl.load(grads_ptr + index * N + tl.arange(0, N), mask=mask).to( - tl.float32) + load_mask = (index >= 0) & (tl.arange(0, N) < n) + sums += tl.load(grads_ptr + index * n + tl.arange(0, N), mask=load_mask).to( + tl.float32 + ) if PROB: + mask = index >= 0 prob = tl.load(probs_ptr + index, mask=mask) - tl.store(output_probs_ptr + pid * num_experts + idx, prob, - mask=mask) + tl.store(output_probs_ptr + pid * num_experts + idx, prob, mask=mask) - mask_indices = tl.where(indices <= index, 2 ** 20, indices) + mask_indices = tl.where(indices <= index, 2**24, indices) idx = tl.argmin(mask_indices, 0) index = tl.min(mask_indices) - tl.store(output_ptr + pid * N + tl.arange(0, N), sums) + tl.store(output_ptr + pid * n + tl.arange(0, N), sums, mask=tl.arange(0, N) < n) def triton_unpermute_with_mask_map( - grad: torch.Tensor, - row_id_map: torch.Tensor, - probs: torch.Tensor, + grad: torch.Tensor, + row_id_map: torch.Tensor, + probs: torch.Tensor, ): """ scatter add with row id map @@ -189,16 +206,19 @@ def triton_unpermute_with_mask_map( - output: [num_tokens, hidden_size] - restore_probs: [num_tokens, num_experts] """ - hidden_size = grad.shape[1] + assert grad.is_contiguous() and row_id_map.is_contiguous() + n = grad.shape[1] + N = triton.next_power_of_2(n) num_tokens, num_experts = row_id_map.shape # not transposed - output = torch.empty((num_tokens, hidden_size), dtype=grad.dtype, - device="cuda") + output = torch.empty((num_tokens, n), dtype=grad.dtype, device="cuda") PROB = probs is not None if PROB: - restore_probs = torch.zeros((num_tokens, num_experts), - dtype=probs.dtype, device="cuda") + assert probs.is_contiguous() + restore_probs = torch.zeros( + (num_tokens, num_experts), dtype=probs.dtype, device="cuda" + ) else: restore_probs = None @@ -212,10 +232,99 @@ def triton_unpermute_with_mask_map( row_id_map, output, restore_probs, + n, + N, num_experts, - hidden_size, PROB, num_stages=4, - num_warps=4 + num_warps=4, ) return output, restore_probs + + +@triton.jit +def triton_unpermute_with_reverse_map_kernel( + input_ptr, + output_ptr, + map_ptr, + prob_ptr, + prob_output_ptr, + M, + m, + N, + stride_im, + stride_in, + stride_om, + stride_on, + BLOCK_SIZE_M: tl.constexpr, + BLOCK_SIZE_N: tl.constexpr, + PROB: tl.constexpr, +): + pid_m = tl.program_id(0) + pid_n = tl.program_id(1) + + rm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + rn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) + + mask_m = rm < m + src_row_indices = tl.load(map_ptr + rm, mask=mask_m, other=0) + + input_offsets = src_row_indices[:, None] * stride_im + rn[None, :] * stride_in + output_offsets = rm[:, None] * stride_om + rn[None, :] * stride_on + + mask = (rm[:, None] < m) & (rn[None, :] < N) + + data = tl.load(input_ptr + input_offsets, mask=mask) + tl.store(output_ptr + output_offsets, data, mask=mask) + if PROB: + if pid_n == 0: + prob = tl.load(prob_ptr + src_row_indices[:, None]) + tl.store(prob_output_ptr + rm[:, None], prob) + + +def triton_unpermute_with_reverse_map(input_tensor, row_id_map, probs=None): + """ + input_tensor: [M, N] + row_id_map: [m], + """ + M, N = input_tensor.shape + m = row_id_map.shape[0] + + output_tensor = torch.empty( + (m, N), device=input_tensor.device, dtype=input_tensor.dtype + ) + + BLOCK_SIZE_M = 128 + BLOCK_SIZE_N = 1024 + + PROB = probs is not None + if PROB: + probs_output = torch.empty((m,), device=input_tensor.device, dtype=probs.dtype) + else: + probs_output = None + + if M == 0: + return output_tensor, probs_output + + grid = (triton.cdiv(m, BLOCK_SIZE_M), triton.cdiv(N, BLOCK_SIZE_N)) + + triton_unpermute_with_reverse_map_kernel[grid]( + input_tensor, + output_tensor, + row_id_map, + probs, + probs_output, + M, + m, + N, + input_tensor.stride(0), + input_tensor.stride(1), + output_tensor.stride(0), + output_tensor.stride(1), + BLOCK_SIZE_M=BLOCK_SIZE_M, + BLOCK_SIZE_N=BLOCK_SIZE_N, + PROB=PROB, + num_warps=8, + ) + + return output_tensor, probs_output diff --git a/linghe/utils/silu.py b/linghe/utils/silu.py index 90545b3..af58d5e 100644 --- a/linghe/utils/silu.py +++ b/linghe/utils/silu.py @@ -4,41 +4,98 @@ """ from typing import Optional + import torch import triton import triton.language as tl +@triton.jit +def exp2(x): + return tl.inline_asm_elementwise( + "ex2.approx.ftz.f32 $0, $1;", + "=r, r", + [x], + dtype=tl.float32, + is_pure=True, + pack=1, + ) + @triton.jit -def weighted_silu_forward_kernel(x_ptr, weight_ptr, out_ptr, M, T, - N: tl.constexpr, - n: tl.constexpr, - W: tl.constexpr, - WEIGHT: tl.constexpr): - pid = tl.program_id(axis=0) +def weighted_silu_forward_asm_kernel( + x_ptr, + weight_ptr, + out_ptr, + M, + N, + H: tl.constexpr, + W: tl.constexpr, + WEIGHT: tl.constexpr, +): + rid = tl.program_id(axis=0) + cid = tl.program_id(axis=1) + n = N // 2 - row_offs = pid * W * T * n + tl.arange(0, W)[:, None] * n - col_offs = tl.arange(0, n)[None, :] + offs = ( + rid * H * N + cid * W + tl.arange(0, H)[:, None] * N + tl.arange(0, W)[None, :] + ) + indices = rid * H + tl.arange(0, H) + mask = indices[:, None] < M + x1 = tl.load(x_ptr + offs, mask=mask).to(tl.float32) + x2 = tl.load(x_ptr + n + offs, mask=mask).to(tl.float32) + if WEIGHT: + w = tl.load(weight_ptr + indices, mask=indices < M).to(tl.float32)[:, None] + # x = x1 * tl.sigmoid(x1) * x2 * w + log2_e: tl.constexpr = 1.4426950408889634 + sigx = x1 / (1 + exp2((-log2_e) * x1)) + x = sigx * x2 * w + else: + # x = x1 * tl.sigmoid(x1) * x2 + log2_e: tl.constexpr = 1.4426950408889634 + sigx = x1 / (1 + exp2((-log2_e) * x1)) + x = sigx * x2 + offs = ( + rid * H * n + cid * W + tl.arange(0, H)[:, None] * n + tl.arange(0, W)[None, :] + ) + tl.store(out_ptr + offs, x, mask=mask) - for i in range(T): - indices = pid * W * T + i * W + tl.arange(0, W) - mask = indices[:, None] < M - x1 = tl.load(x_ptr + row_offs * 2 + col_offs, mask=mask).to(tl.float32) - x2 = tl.load(x_ptr + n + row_offs * 2 + col_offs, mask=mask).to( - tl.float32) - if WEIGHT: - w = tl.load(weight_ptr + indices, mask=indices < M).to(tl.float32)[:, - None] - x = x1 / (1 + tl.exp(-x1)) * x2 * w - else: - x = x1 / (1 + tl.exp(-x1)) * x2 - tl.store(out_ptr + row_offs + col_offs, x, mask=mask) - row_offs += n * W + +@triton.jit +def weighted_silu_forward_kernel( + x_ptr, + weight_ptr, + out_ptr, + M, + N, + H: tl.constexpr, + W: tl.constexpr, + WEIGHT: tl.constexpr, +): + rid = tl.program_id(axis=0) + cid = tl.program_id(axis=1) + n = (N // 2).to(tl.int64) + + offs = ( + rid * H * N + cid * W + tl.arange(0, H)[:, None] * N + tl.arange(0, W)[None, :] + ) + indices = rid * H + tl.arange(0, H) + mask = indices[:, None] < M + x1 = tl.load(x_ptr + offs, mask=mask).to(tl.float32) + x2 = tl.load(x_ptr + n + offs, mask=mask).to(tl.float32) + if WEIGHT: + w = tl.load(weight_ptr + indices, mask=indices < M).to(tl.float32)[:, None] + x = x1 * tl.sigmoid(x1) * x2 * w + else: + x = x1 * tl.sigmoid(x1) * x2 + offs = ( + rid * H * n + cid * W + tl.arange(0, H)[:, None] * n + tl.arange(0, W)[None, :] + ) + tl.store(out_ptr + offs, x, mask=mask) # used in bf16 moe -def triton_weighted_silu_forward(x, weight=None, out=None): +def triton_weighted_silu_forward(x, weight=None, out=None, asm=False): """ compute silu(x)*weight, used in bf16/fp16 training with MoE Args: @@ -47,73 +104,79 @@ def triton_weighted_silu_forward(x, weight=None, out=None): Returns: out: output tensor """ - # row-wise read, row-wise write + assert x.is_contiguous() M, N = x.shape - assert N <= 8192 device = x.device if out is None: out = torch.empty((M, N // 2), device=device, dtype=x.dtype) WEIGHT = weight is not None - W = 8192 // N - T = 8 - grid = (triton.cdiv(M, T * W),) - weighted_silu_forward_kernel[grid]( - x, - weight, - out, - M, T, - N, - N // 2, - W, - WEIGHT, - num_stages=3, - num_warps=8 - ) + if WEIGHT: + assert weight.is_contiguous() + H = 32 + W = 128 + assert N % (W * 2) == 0 + grid = (triton.cdiv(M, H), N // W // 2) + if asm: + weighted_silu_forward_asm_kernel[grid]( + x, weight, out, M, N, H, W, WEIGHT, num_stages=3, num_warps=8 + ) + else: + weighted_silu_forward_kernel[grid]( + x, weight, out, M, N, H, W, WEIGHT, num_stages=3, num_warps=8 + ) return out @triton.jit -def weighted_silu_backward_kernel(g_ptr, x_ptr, weight_ptr, dx_ptr, dw_ptr, M, - T, - N: tl.constexpr, - n: tl.constexpr, - W: tl.constexpr, - WEIGHT: tl.constexpr): +def weighted_silu_backward_kernel( + g_ptr, + x_ptr, + weight_ptr, + dx_ptr, + dw_ptr, + M, + N, + H: tl.constexpr, + W: tl.constexpr, + WEIGHT: tl.constexpr, +): pid = tl.program_id(axis=0) + n = (N // 2).to(tl.int64) - offs = pid * W * T * N + tl.arange(0, W)[:, None] * N + tl.arange(0, n)[ - None, :] - hoffs = pid * W * T * n + tl.arange(0, W)[:, None] * n + tl.arange(0, n)[ - None, :] - for i in range(T): - mask = pid * W * T + i * W + tl.arange(0, W) + offs = pid * H * N + tl.arange(0, H)[:, None] * N + tl.arange(0, W)[None, :] + hoffs = pid * H * n + tl.arange(0, H)[:, None] * n + tl.arange(0, W)[None, :] + mask = pid * H + tl.arange(0, H) + if WEIGHT: + w = tl.load(weight_ptr + mask, mask=mask < M).to(tl.float32)[:, None] + + dw = tl.zeros((H,), dtype=tl.float32) + for i in range(n // W): x1 = tl.load(x_ptr + offs, mask=mask[:, None] < M).to(tl.float32) x2 = tl.load(x_ptr + offs + n, mask=mask[:, None] < M).to(tl.float32) g = tl.load(g_ptr + hoffs, mask=mask[:, None] < M).to(tl.float32) if WEIGHT: - w = tl.load(weight_ptr + mask, mask=mask < M).to(tl.float32)[:, None] - sigmoid = 1 / (1 + tl.exp(-x1)) - dw = tl.sum(x1 * sigmoid * x2 * g, 1) + sigmoid = tl.sigmoid(x1) + dw += tl.sum(x1 * sigmoid * x2 * g, 1) tl.store(dw_ptr + mask, dw, mask=mask < M) - dx1 = g * x2 * w * sigmoid * (1 + x1 * tl.exp(-x1) * sigmoid) + dx1 = g * x2 * w * sigmoid * (1 + x1 * (1 - sigmoid)) tl.store(dx_ptr + offs, dx1, mask=mask[:, None] < M) dx2 = g * x1 * sigmoid * w tl.store(dx_ptr + offs + n, dx2, mask=mask[:, None] < M) else: - sigmoid = 1 / (1 + tl.exp(-x1)) - dx1 = g * x2 * sigmoid * (1 + x1 * tl.exp(-x1) * sigmoid) + sigmoid = tl.sigmoid(x1) + dx1 = g * x2 * sigmoid * (1 + x1 * (1 - sigmoid)) tl.store(dx_ptr + offs, dx1, mask=mask[:, None] < M) dx2 = g * x1 * sigmoid tl.store(dx_ptr + offs + n, dx2, mask=mask[:, None] < M) - offs += N * W - hoffs += n * W + offs += W + hoffs += W -def triton_weighted_silu_backward(g: torch.Tensor, - x: torch.Tensor, - weight: Optional[torch.Tensor] = None): +def triton_weighted_silu_backward( + g: torch.Tensor, x: torch.Tensor, weight: Optional[torch.Tensor] = None +): """ backward of triton_weighted_silu_forward Args: @@ -125,100 +188,1418 @@ def triton_weighted_silu_backward(g: torch.Tensor, - dx: gradient of x - dw: gradient of weight """ - # row-wise read, row-wise write + assert g.is_contiguous() and x.is_contiguous() M, N = x.shape - assert N <= 8192 + H = 8 if M <= 4096 else 16 + W = 128 + assert N % (W * 2) == 0 device = x.device if weight is not None: - dw = torch.empty(weight.shape, device=device, dtype=x.dtype) + assert weight.is_contiguous() + dw = torch.empty(weight.shape, device=device, dtype=weight.dtype) WEIGHT = True else: dw = None WEIGHT = False dx = torch.empty((M, N), device=device, dtype=x.dtype) - W = 8192 // N - T = 8 - grid = (triton.cdiv(M, W*T),) + + grid = (triton.cdiv(M, H),) weighted_silu_backward_kernel[grid]( + g, x, weight, dx, dw, M, N, H, W, WEIGHT, num_stages=3, num_warps=8 + ) + return dx, dw + + +@triton.jit +def silu_and_block_quant_forward_kernel( + x_ptr, + out_ptr, + scale_ptr, + transpose_output_ptr, + transpose_scale_ptr, + limit, + M, + n: tl.constexpr, + H: tl.constexpr, + W: tl.constexpr, + CUTOFF: tl.constexpr, + ROUND: tl.constexpr, + OUTPUT_MODE: tl.constexpr, +): + rid = tl.program_id(axis=0) + cid = tl.program_id(axis=1) + + offs = ( + rid * H * n * 2 + + cid * W + + tl.arange(0, H)[:, None] * n * 2 + + tl.arange(0, W)[None, :] + ) + indices = rid * H + tl.arange(0, H) + mask = indices[:, None] < M + + if OUTPUT_MODE % 2 == 0: + x1 = tl.load(x_ptr + offs, mask=mask).to(tl.float32) + x2 = tl.load(x_ptr + n + offs, mask=mask).to(tl.float32) + if CUTOFF: + x = tl.minimum(x1 * tl.sigmoid(x1), limit) * tl.clamp(x2, -limit, limit) + else: + x = x1 * tl.sigmoid(x1) * x2 + + scale = tl.maximum(tl.max(x.abs(), 1) / 448, 1e-30) + if ROUND: + scale = tl.exp2(tl.ceil(tl.log2(scale))) + + tl.store( + scale_ptr + rid * H + cid * M + tl.arange(0, H), scale, mask=indices < M + ) + xq = (x / scale[:, None]).to(out_ptr.dtype.element_ty) + tl.store( + out_ptr + + rid * H * n + + cid * W + + tl.arange(0, H)[:, None] * n + + tl.arange(0, W)[None, :], + xq, + mask=mask, + ) + + if OUTPUT_MODE > 0: + # reload is faster than transpose + x1 = tl.load(x_ptr + offs, mask=mask).to(tl.float32) + x2 = tl.load(x_ptr + n + offs, mask=mask).to(tl.float32) + if CUTOFF: + x = tl.minimum(x1 * tl.sigmoid(x1), limit) * tl.clamp(x2, -limit, limit) + else: + x = x1 * tl.sigmoid(x1) * x2 + + scale = tl.maximum(tl.max(x.abs(), 0) / 448, 1e-30) + if ROUND: + scale = tl.exp2(tl.ceil(tl.log2(scale))) + tl.store(transpose_scale_ptr + rid * n + cid * W + tl.arange(0, W), scale) + xq = (x / scale).to(out_ptr.dtype.element_ty) + tl.store( + transpose_output_ptr + + rid * H + + cid * W * M + + tl.arange(0, W)[:, None] * M + + tl.arange(0, H)[None, :], + tl.trans(xq), + mask=indices[None, :] < M, + ) + + +def triton_silu_and_block_quant_forward( + x, out=None, scale=None, limit=None, round_scale=False, output_mode=2 +): + """ + fused silu and blockwise quantization, used in shared expert + Args: + x: input tensor + round_scale: whether round scale to power of 2 + output_mode: one of {0, 1, 2} + 0: only output non-transposed quantized tensor + 1: only output transposed quantized tensor + 2: output both + + Returns: + - out: quantized tensor + - scale: quantization scale + - transpose_output: quantized tensor of transposed output + - transpose_scale: quantization scale of transposed output + """ + assert x.is_contiguous() + M, N = x.shape + n = N // 2 + device = x.device + if out is None: + out = torch.empty((M, n), device=device, dtype=torch.float8_e4m3fn) + if scale is None: + scale = torch.empty((n // 128, M), device=device, dtype=torch.float32) + transpose_output = torch.empty((n, M), device=device, dtype=torch.float8_e4m3fn) + transpose_scale = torch.empty( + (triton.cdiv(M, 128), n), device=device, dtype=torch.float32 + ) + if output_mode == 0: + H, W, num_warps = 64, 128, 8 + elif output_mode == 1: + H, W, num_warps = 128, 64, 8 + else: + H, W, num_warps = 128, 128, 8 + if limit is None: + CUTOFF = False + limit = 0.0 + else: + CUTOFF = True + assert n % W == 0 + grid = (triton.cdiv(M, H), n // W) + silu_and_block_quant_forward_kernel[grid]( + x, + out, + scale, + transpose_output, + transpose_scale, + limit, + M, + n, + H, + W, + CUTOFF, + round_scale, + output_mode, + num_stages=3, + num_warps=num_warps, + ) + + return out, scale, transpose_output, transpose_scale + + +@triton.jit +def silu_and_block_quant_backward_kernel( + g_ptr, + x_ptr, + dx_ptr, + dx_scale_ptr, + transpose_dx_ptr, + transpose_dx_scale_ptr, + limit, + M, + n: tl.constexpr, + CUTOFF: tl.constexpr, + ROUND: tl.constexpr, +): + rid = tl.program_id(axis=0) + cid = tl.program_id(axis=1) + nb = n // 128 + offs = ( + rid * 128 * n * 2 + + cid * 128 + + tl.arange(0, 128)[:, None] * n * 2 + + tl.arange(0, 128)[None, :] + ) + toffs = ( + rid * 128 + + cid * M * 128 + + tl.arange(0, 128)[:, None] * M + + tl.arange(0, 128)[None, :] + ) + idx = rid * 128 + tl.arange(0, 128) + mask = idx[:, None] < M + x1 = tl.load(x_ptr + offs, mask=mask).to(tl.float32) + x2 = tl.load(x_ptr + n + offs, mask=mask).to(tl.float32) + if CUTOFF: + dx2_mask = (x2 <= limit) & (x2 >= -limit) + x2 = tl.clamp(x2, -limit, limit) + g = tl.load( + g_ptr + + rid * 128 * n + + cid * 128 + + tl.arange(0, 128)[:, None] * n + + tl.arange(0, 128)[None, :], + mask=mask, + ).to(tl.float32) + sigmoid = tl.sigmoid(x1) + dx1 = sigmoid * g * x2 * (1 + x1 * (1 - sigmoid)) + if CUTOFF: + dx1_mask = (sigmoid * x1 <= limit).to(tl.float32) + dx1 *= dx1_mask + scale1 = tl.maximum(tl.max(dx1.abs(), 1) / 448, 1e-30) + if ROUND: + scale1 = tl.exp2(tl.ceil(tl.log2(scale1))) + tl.store( + dx_scale_ptr + cid * M + rid * 128 + tl.arange(0, 128), scale1, mask=idx < M + ) + + qdx1 = (dx1 / scale1[:, None]).to(dx_ptr.dtype.element_ty) + tl.store(dx_ptr + offs, qdx1, mask=mask) + + scale1 = tl.maximum(tl.max(dx1.abs(), 0) / 448, 1e-30) + if ROUND: + scale1 = tl.exp2(tl.ceil(tl.log2(scale1))) + tl.store( + transpose_dx_scale_ptr + rid * n * 2 + cid * 128 + tl.arange(0, 128), scale1 + ) + + qdx1 = (dx1 / scale1[None, :]).to(transpose_dx_ptr.dtype.element_ty) + tl.store(transpose_dx_ptr + toffs, tl.trans(qdx1), mask=idx[None, :] < M) + + if CUTOFF: + dx2 = g * tl.minimum(sigmoid * x1, limit) * dx2_mask.to(tl.float32) + else: + dx2 = g * sigmoid * x1 + scale2 = tl.maximum(tl.max(dx2.abs(), 1) / 448, 1e-30) + if ROUND: + scale2 = tl.exp2(tl.ceil(tl.log2(scale2))) + tl.store( + dx_scale_ptr + cid * M + rid * 128 + M * nb + tl.arange(0, 128), + scale2, + mask=idx < M, + ) + + qdx2 = (dx2 / scale2[:, None]).to(dx_ptr.dtype.element_ty) + tl.store(dx_ptr + offs + n, qdx2, mask=idx[:, None] < M) + + scale2 = tl.maximum(tl.max(dx2.abs(), 0) / 448, 1e-30) + if ROUND: + scale2 = tl.exp2(tl.ceil(tl.log2(scale2))) + tl.store( + transpose_dx_scale_ptr + rid * n * 2 + n + cid * 128 + tl.arange(0, 128), scale2 + ) + + qdx2 = (dx2 / scale2[None, :]).to(transpose_dx_ptr.dtype.element_ty) + tl.store(transpose_dx_ptr + M * n + toffs, tl.trans(qdx2), mask=idx[None, :] < M) + + +# used in shared expert +def triton_silu_and_block_quant_backward(g, x, limit=None, round_scale=False): + """ + backward of triton_silu_and_block_quant_forward + Args: + g: gradient + x: input tensor + round_scale: whether round to power of 2 + + Returns: + - dx: quantized non-transposed gradient + - dx_scale: scales of quantization non-transposed gradient + - transpose_dx: quantized transposed gradient + - transpose_dx_scale: scales of quantization transposed gradient + """ + assert g.is_contiguous() + M, N = x.shape + n = N // 2 + device = x.device + dx = torch.empty((M, N), device=device, dtype=torch.float8_e4m3fn) + + dx_scale = torch.empty((N // 128, M), device=device, dtype=torch.float32) + scale_shape = (triton.cdiv(M, 128), N) + transpose_dx = torch.empty((N, M), device=device, dtype=torch.float8_e4m3fn) + transpose_dx_scale = torch.empty(scale_shape, device=device, dtype=torch.float32) + + if limit is None: + CUTOFF = False + limit = 0.0 + else: + CUTOFF = True + + assert N % 256 == 0 + grid = (triton.cdiv(M, 128), N // 256) + silu_and_block_quant_backward_kernel[grid]( + g, + x, + dx, + dx_scale, + transpose_dx, + transpose_dx_scale, + limit, + M, + n, + CUTOFF, + round_scale, + num_stages=2, + num_warps=8, + ) + return dx, dx_scale, transpose_dx, transpose_dx_scale + + +@triton.jit +def batch_weighted_silu_and_block_quant_forward_kernel( + x_ptr, + weight_ptr, + out_ptr, + scale_ptr, + transpose_output_ptr, + transpose_scale_ptr, + count_ptr, + accum_ptr, + limit, + n, + E: tl.constexpr, + CUTOFF: tl.constexpr, + ROUND: tl.constexpr, +): + eid = tl.program_id(axis=0) + rid = tl.program_id(axis=1) + cid = tl.program_id(axis=2) + + count = tl.load(count_ptr + eid) + ei = tl.load(accum_ptr + eid) + si = ei - count + c = tl.cdiv(count, 128) + + if rid >= c: + return + + n = n.to(tl.int64) + nb = n // 128 + + counts = tl.load(count_ptr + tl.arange(0, E)) + n_blocks = tl.cdiv(counts, 128) + transpose_scale_off = tl.sum(tl.where(tl.arange(0, E) < eid, n_blocks, 0)) + + offs = ( + si * n * 2 + + rid * 128 * n * 2 + + cid * 128 + + tl.arange(0, 128)[:, None] * n * 2 + + tl.arange(0, 128)[None, :] + ) + hoffs = ( + si * n + + rid * 128 * n + + cid * 128 + + tl.arange(0, 128)[:, None] * n + + tl.arange(0, 128)[None, :] + ) + + indices = rid * 128 + tl.arange(0, 128) + mask = indices[:, None] < count + w = tl.load(weight_ptr + si + indices, mask=indices < count).to(tl.float32) + x1 = tl.load(x_ptr + offs, mask=mask).to(tl.float32) + x2 = tl.load(x_ptr + n + offs, mask=mask).to(tl.float32) + + if CUTOFF: + x = ( + tl.minimum(x1 * tl.sigmoid(x1), limit) + * tl.clamp(x2, -limit, limit) + * w[:, None] + ) + else: + x = x1 * tl.sigmoid(x1) * x2 * w[:, None] + + scale1 = tl.maximum(tl.max(tl.abs(x), 1) / 448, 1e-30) + scale2 = tl.maximum(tl.max(tl.abs(x), 0) / 448, 1e-30) + if ROUND: + scale1 = tl.exp2(tl.ceil(tl.log2(scale1))) + scale2 = tl.exp2(tl.ceil(tl.log2(scale2))) + + tl.store( + scale_ptr + si * nb + cid * count + rid * 128 + tl.arange(0, 128), + scale1, + mask=indices < count, + ) + xq1 = (x / scale1[:, None]).to(out_ptr.dtype.element_ty) + tl.store(out_ptr + hoffs, xq1, mask=mask) + + tl.store( + transpose_scale_ptr + + transpose_scale_off * n + + rid * n + + cid * 128 + + tl.arange(0, 128), + scale2, + ) + toffs = ( + si * n + + rid * 128 + + cid * count * 128 + + tl.arange(0, 128)[:, None] * count + + tl.arange(0, 128)[None, :] + ) + xq2 = tl.trans((x / scale2).to(out_ptr.dtype.element_ty)) + tl.store(transpose_output_ptr + toffs, xq2, mask=indices[None, :] < count) + + +@triton.jit +def batch_weighted_silu_and_block_quant_forward_nt_kernel( + x_ptr, + weight_ptr, + out_ptr, + scale_ptr, + transpose_output_ptr, + transpose_scale_ptr, + count_ptr, + accum_ptr, + limit, + n, + B: tl.constexpr, + E: tl.constexpr, + CUTOFF: tl.constexpr, + ROUND: tl.constexpr, +): + eid = tl.program_id(axis=0) + rid = tl.program_id(axis=1) + cid = tl.program_id(axis=2) + + count = tl.load(count_ptr + eid) + ei = tl.load(accum_ptr + eid) + si = ei - count + c = tl.cdiv(count, 128) + + if rid >= c: + return + + n = n.to(tl.int64) + nb = n // 128 + I: tl.constexpr = 128 // B + + offs = ( + si * n * 2 + + rid * 128 * n * 2 + + cid * 128 + + tl.arange(0, B)[:, None] * n * 2 + + tl.arange(0, 128)[None, :] + ) + hoffs = ( + si * n + + rid * 128 * n + + cid * 128 + + tl.arange(0, B)[:, None] * n + + tl.arange(0, 128)[None, :] + ) + soffs = si * nb + cid * count + rid * 128 + tl.arange(0, B) + indices = rid * 128 + tl.arange(0, B) + maxs = tl.zeros((B, 128), dtype=tl.float32) + for i in range(I): + mask = indices[:, None] < count + w = tl.load(weight_ptr + si + indices, mask=indices < count).to(tl.float32) + x1 = tl.load(x_ptr + offs, mask=mask).to(tl.float32) + x2 = tl.load(x_ptr + n + offs, mask=mask).to(tl.float32) + if CUTOFF: + x = ( + tl.minimum(x1 * tl.sigmoid(x1), limit) + * tl.clamp(x2, -limit, limit) + * w[:, None] + ) + else: + x = x1 * tl.sigmoid(x1) * x2 * w[:, None] + + maxs = tl.maximum(maxs, tl.abs(x)) + + scale = tl.maximum(tl.max(tl.abs(x), 1) / 448, 1e-30) + if ROUND: + scale = tl.exp2(tl.ceil(tl.log2(scale))) + tl.store(scale_ptr + soffs, scale, mask=indices < count) + + xq = (x / scale[:, None]).to(out_ptr.dtype.element_ty) + + tl.store(out_ptr + hoffs, xq, mask=mask) + offs += B * n * 2 + hoffs += B * n + soffs += B + indices += B + + offs = ( + si * n * 2 + + rid * 128 * n * 2 + + cid * 128 + + tl.arange(0, B)[:, None] * n * 2 + + tl.arange(0, 128)[None, :] + ) + hoffs = ( + si * n + + rid * 128 * n + + cid * 128 + + tl.arange(0, B)[:, None] * n + + tl.arange(0, 128)[None, :] + ) + indices = rid * 128 + tl.arange(0, B) + scale = tl.maximum(tl.max(maxs, 0) / 448, 1e-30) + if ROUND: + scale = tl.exp2(tl.ceil(tl.log2(scale))) + + counts = tl.load(count_ptr + tl.arange(0, E)) + n_blocks = tl.cdiv(counts, 128) + transpose_scale_off = tl.sum(tl.where(tl.arange(0, E) < eid, n_blocks, 0)) + tl.store( + transpose_scale_ptr + + transpose_scale_off * n + + rid * n + + cid * 128 + + tl.arange(0, 128), + scale, + ) + toffs = ( + si * n + + rid * 128 + + cid * count * 128 + + tl.arange(0, 128)[:, None] * count + + tl.arange(0, B)[None, :] + ) + for i in range(I): + mask = indices[:, None] < count + w = tl.load(weight_ptr + si + indices, mask=indices < count).to(tl.float32) + x1 = tl.load(x_ptr + offs, mask=mask).to(tl.float32) + x2 = tl.load(x_ptr + n + offs, mask=mask).to(tl.float32) + if CUTOFF: + x = ( + tl.minimum(x1 * tl.sigmoid(x1), limit) + * tl.clamp(x2, -limit, limit) + * w[:, None] + ) + else: + x = x1 * tl.sigmoid(x1) * x2 * w[:, None] + + xq = tl.trans((x / scale[None, :])) + + tl.store(transpose_output_ptr + toffs, xq, mask=indices[None, :] < count) + offs += B * n * 2 + hoffs += B * n + indices += B + toffs += B + + +@triton.jit +def batch_weighted_silu_and_block_quant_forward_n_kernel( + x_ptr, + weight_ptr, + out_ptr, + scale_ptr, + count_ptr, + accum_ptr, + limit, + n, + B: tl.constexpr, + E: tl.constexpr, + CUTOFF: tl.constexpr, + ROUND: tl.constexpr, +): + eid = tl.program_id(axis=0) + rid = tl.program_id(axis=1) + cid = tl.program_id(axis=2) + + count = tl.load(count_ptr + eid) + ei = tl.load(accum_ptr + eid) + si = ei - count + c = tl.cdiv(count, B) + + if rid >= c: + return + n = n.to(tl.int64) + nb = n // 128 + + offs = ( + si * n * 2 + + rid * B * n * 2 + + cid * 128 + + tl.arange(0, B)[:, None] * n * 2 + + tl.arange(0, 128)[None, :] + ) + indices = rid * B + tl.arange(0, B) + mask = indices[:, None] < count + w = tl.load(weight_ptr + si + indices, mask=indices < count).to(tl.float32) + x1 = tl.load(x_ptr + offs, mask=mask).to(tl.float32) + x2 = tl.load(x_ptr + n + offs, mask=mask).to(tl.float32) + + if CUTOFF: + x = ( + tl.minimum(x1 * tl.sigmoid(x1), limit) + * tl.clamp(x2, -limit, limit) + * w[:, None] + ) + else: + x = x1 * tl.sigmoid(x1) * x2 * w[:, None] + + scale = tl.maximum(tl.max(tl.abs(x), 1) / 448, 1e-30) + if ROUND: + scale = tl.exp2(tl.ceil(tl.log2(scale))) + tl.store( + scale_ptr + si * nb + cid * count + rid * B + tl.arange(0, B), + scale, + mask=indices < count, + ) + + xq = (x / scale[:, None]).to(out_ptr.dtype.element_ty) + hoffs = ( + si * n + + rid * B * n + + cid * 128 + + tl.arange(0, B)[:, None] * n + + tl.arange(0, 128)[None, :] + ) + tl.store(out_ptr + hoffs, xq, mask=mask) + + +@triton.jit +def batch_weighted_silu_and_block_quant_forward_t_kernel( + x_ptr, + weight_ptr, + transpose_output_ptr, + transpose_scale_ptr, + count_ptr, + accum_ptr, + limit, + n, + B: tl.constexpr, + E: tl.constexpr, + CUTOFF: tl.constexpr, + ROUND: tl.constexpr, +): + eid = tl.program_id(axis=0) + rid = tl.program_id(axis=1) + cid = tl.program_id(axis=2) + + count = tl.load(count_ptr + eid) + ei = tl.load(accum_ptr + eid) + si = ei - count + c = tl.cdiv(count, 128) + + if rid >= c: + return + + n = n.to(tl.int64) + counts = tl.load(count_ptr + tl.arange(0, E)) + n_blocks = tl.cdiv(counts, 128) + transpose_scale_off = tl.sum(tl.where(tl.arange(0, E) < eid, n_blocks, 0)) + + offs = ( + si * n * 2 + + rid * 128 * n * 2 + + cid * B + + tl.arange(0, 128)[:, None] * n * 2 + + tl.arange(0, B)[None, :] + ) + indices = rid * 128 + tl.arange(0, 128) + mask = indices[:, None] < count + w = tl.load(weight_ptr + si + indices, mask=indices < count).to(tl.float32) + x1 = tl.load(x_ptr + offs, mask=mask).to(tl.float32) + x2 = tl.load(x_ptr + n + offs, mask=mask).to(tl.float32) + + if CUTOFF: + x = ( + tl.minimum(x1 * tl.sigmoid(x1), limit) + * tl.clamp(x2, -limit, limit) + * w[:, None] + ) + else: + x = x1 * tl.sigmoid(x1) * x2 * w[:, None] + + scale = tl.maximum(tl.max(tl.abs(x), 0) / 448, 1e-30) + if ROUND: + scale = tl.exp2(tl.ceil(tl.log2(scale))) + tl.store( + transpose_scale_ptr + + transpose_scale_off * n + + rid * n + + cid * B + + tl.arange(0, B), + scale, + ) + + xq = tl.trans((x / scale).to(transpose_output_ptr.dtype.element_ty)) + + toffs = ( + si * n + + rid * 128 + + cid * count * B + + tl.arange(0, B)[:, None] * count + + tl.arange(0, 128)[None, :] + ) + tl.store(transpose_output_ptr + toffs, xq, mask=indices[None, :] < count) + + +def triton_batch_weighted_silu_and_block_quant_forward( + x, + weight, + counts, + splits=None, + out=None, + scale=None, + limit=None, + round_scale=False, + output_mode=2, +): + """ + silu and blockwise quantize activation in routed experts + Args: + x: activation tensor in routed experts + weight: router prob tensor + counts: cuda tensor of token count per expert + splits: python int list of token count per expert + round_scale: whether round scale to power of 2 + output_mode: one of {0, 1, 2} + 0: only output non-transposed quantized tensor + 1: only output transposed quantized tensor + 2: output both + + Returns: + - out: quantized tensor + - scale: quantization scale + - transpose_output: quantized tensor of transposed output + - transpose_scale: quantization scale of transposed output + """ + assert splits is not None, "batch mode need splits to launch kernels" + assert x.is_contiguous() and weight.is_contiguous() + M, N = x.shape + n = N // 2 + n_experts = counts.shape[0] + assert N <= 8192 + device = x.device + if out is None: + out = torch.empty((M, n), device=device, dtype=torch.float8_e4m3fn) + if scale is None: + scale = torch.empty((M, n // 128), device=device, dtype=torch.float32) + + blocks = sum([(x + 127) // 128 for x in splits]) + transpose_output = torch.empty((M, n), device=device, dtype=torch.float8_e4m3fn) + transpose_scale = torch.empty((blocks, n), device=device, dtype=torch.float32) + + if M == 0: + return out, scale, transpose_output, transpose_scale + + if limit is None: + CUTOFF = False + limit = 0.0 + else: + CUTOFF = True + + accums = torch.cumsum(counts, 0) + + if output_mode == 0: + B = 32 + grid = (n_experts, triton.cdiv(max(splits), B), n // 128) + batch_weighted_silu_and_block_quant_forward_n_kernel[grid]( + x, + weight, + out, + scale, + counts, + accums, + limit, + n, + B, + len(splits), + CUTOFF, + round_scale, + num_stages=2, + num_warps=2, + ) + elif output_mode == 1: + B = 32 + grid = (n_experts, triton.cdiv(max(splits), 128), n // B) + batch_weighted_silu_and_block_quant_forward_t_kernel[grid]( + x, + weight, + transpose_output, + transpose_scale, + counts, + accums, + limit, + n, + B, + len(splits), + CUTOFF, + round_scale, + num_stages=2, + num_warps=2, + ) + else: + B = 16 + grid = (n_experts, triton.cdiv(max(splits), 128), n // 128) + batch_weighted_silu_and_block_quant_forward_nt_kernel[grid]( + x, + weight, + out, + scale, + transpose_output, + transpose_scale, + counts, + accums, + limit, + n, + B, + len(splits), + CUTOFF, + round_scale, + num_stages=3, + num_warps=4, + ) + return out, scale, transpose_output, transpose_scale + + +@triton.jit +def batch_weighted_silu_and_block_quant_backward_kernel( + g_ptr, + x_ptr, + weight_ptr, + count_ptr, + accum_ptr, + dx_ptr, + dx_scale_ptr, + transpose_dx_ptr, + transpose_dx_scale_ptr, + dw_ptr, + limit, + n, + E: tl.constexpr, + CUTOFF: tl.constexpr, + ROUND: tl.constexpr, +): + eid = tl.program_id(axis=0) + rid = tl.program_id(axis=1) + cid = tl.program_id(axis=2) + + count = tl.load(count_ptr + eid) + si = tl.load(accum_ptr + eid) - count + # very slow with triton 3.3.1, fix in 3.5.1 + # counts = tl.load(count_ptr + tl.arange(0, E)) + # si = tl.sum(tl.where(tl.arange(0, E) < eid, counts, 0)) + + if rid >= tl.cdiv(count, 128): + return + + n = n.to(tl.int64) + nb = n // 128 + transpose_off = tl.sum( + tl.where( + tl.arange(0, E) < eid, tl.cdiv(tl.load(count_ptr + tl.arange(0, E)), 128), 0 + ) + ) + + offs = ( + si * n * 2 + + rid * 128 * n * 2 + + cid * 128 + + tl.arange(0, 128)[:, None] * n * 2 + + tl.arange(0, 128)[None, :] + ) + # hoffs = si * n + tid * 128 * n + tl.arange(0, 128)[:, None] * n + tl.arange(0, 128)[None, :] + # toffs = si * n * 2 + tid * 128 + tl.arange(0, 128)[:, None] * count + tl.arange(0, 128)[None, :] + idx = rid * 128 + tl.arange(0, 128) + w = tl.load(weight_ptr + si + idx, mask=idx < count).to(tl.float32)[:, None] + + x1 = tl.load(x_ptr + offs, mask=idx[:, None] < count).to(tl.float32) + x2 = tl.load(x_ptr + n + offs, mask=idx[:, None] < count).to(tl.float32) + sigmoid = tl.sigmoid(x1) + gate = sigmoid * x1 + + if CUTOFF: + dx2_mask = (x2 <= limit) & (x2 >= -limit) + x2 = tl.clamp(x2, -limit, limit) + dx1_mask = gate <= limit + gate = tl.minimum(gate, limit) + + g = tl.load( + g_ptr + + si * n + + rid * 128 * n + + 128 * cid + + tl.arange(0, 128)[:, None] * n + + tl.arange(0, 128)[None, :], + mask=idx[:, None] < count, + ).to(tl.float32) + + dw = tl.sum(gate * x2 * g, 1) + tl.store(dw_ptr + si * nb + cid + idx * nb, dw, mask=idx < count) + + dx = sigmoid * g * x2 * w * (1 + x1 * (1 - sigmoid)) + + if CUTOFF: + dx *= dx1_mask + + scale1 = tl.maximum(tl.max(dx.abs(), 1) / 448, 1e-30) + scale2 = tl.maximum(tl.max(dx.abs(), 0) / 448, 1e-30) + if ROUND: + scale1 = tl.exp2(tl.ceil(tl.log2(scale1))) + scale2 = tl.exp2(tl.ceil(tl.log2(scale2))) + tl.store( + dx_scale_ptr + si * nb * 2 + cid * count + rid * 128 + tl.arange(0, 128), + scale1, + mask=idx < count, + ) + + tl.store(dx_ptr + offs, dx / scale1[:, None], mask=idx[:, None] < count) + + tl.store( + transpose_dx_scale_ptr + + transpose_off * n * 2 + + rid * n * 2 + + cid * 128 + + tl.arange(0, 128), + scale2, + ) + + qdx = tl.trans((dx / scale2[None, :]).to(dx_ptr.dtype.element_ty)) + # tl.store(transpose_dx_ptr + toffs, qdx, mask=idx[None, :] < count) + tl.store( + transpose_dx_ptr + + si * n * 2 + + rid * 128 + + cid * 128 * count + + tl.arange(0, 128)[:, None] * count + + tl.arange(0, 128)[None, :], + qdx, + mask=idx[None, :] < count, + ) + + dx = g * gate * w + + if CUTOFF: + dx *= dx2_mask + + scale3 = tl.maximum(tl.max(dx.abs(), 1) / 448, 1e-30) + + if ROUND: + scale3 = tl.exp2(tl.ceil(tl.log2(scale3))) + tl.store( + dx_scale_ptr + + si * nb * 2 + + cid * count + + rid * 128 + + count * nb + + tl.arange(0, 128), + scale3, + mask=idx < count, + ) + tl.store(dx_ptr + n + offs, dx / scale3[:, None], mask=idx[:, None] < count) + + scale4 = tl.maximum(tl.max(dx.abs(), 0) / 448, 1e-30) + if ROUND: + scale4 = tl.exp2(tl.ceil(tl.log2(scale4))) + qdx = tl.trans((dx / scale4[None, :]).to(dx_ptr.dtype.element_ty)) + tl.store( + transpose_dx_scale_ptr + + transpose_off * n * 2 + + rid * n * 2 + + n + + cid * 128 + + tl.arange(0, 128), + scale4, + ) + tl.store( + transpose_dx_ptr + + count * n + + si * n * 2 + + rid * 128 + + cid * 128 * count + + tl.arange(0, 128)[:, None] * count + + tl.arange(0, 128)[None, :], + qdx, + mask=idx[None, :] < count, + ) + + +@triton.jit +def batch_weighted_silu_and_block_quant_backward_n_kernel( + g_ptr, + x_ptr, + weight_ptr, + count_ptr, + accum_ptr, + dx_ptr, + dx_scale_ptr, + dw_ptr, + limit, + n, + B: tl.constexpr, + E: tl.constexpr, + CUTOFF: tl.constexpr, + ROUND: tl.constexpr, +): + eid = tl.program_id(axis=0) + rid = tl.program_id(axis=1) + cid = tl.program_id(axis=2) + + count = tl.load(count_ptr + eid) + si = tl.load(accum_ptr + eid) - count + # counts = tl.load(count_ptr + tl.arange(0, E)) + # si = tl.sum(tl.where(tl.arange(0, E) < eid, counts, 0)) + + if rid >= tl.cdiv(count, B): + return + + n = n.to(tl.int64) + nb = n // 128 + + offs = ( + si * n * 2 + + rid * B * n * 2 + + cid * 128 + + tl.arange(0, B)[:, None] * n * 2 + + tl.arange(0, 128)[None, :] + ) + idx = rid * B + tl.arange(0, B) + w = tl.load(weight_ptr + si + idx, mask=idx < count).to(tl.float32)[:, None] + + x1 = tl.load(x_ptr + offs, mask=idx[:, None] < count).to(tl.float32) + x2 = tl.load(x_ptr + n + offs, mask=idx[:, None] < count).to(tl.float32) + sigmoid = tl.sigmoid(x1) + gate = sigmoid * x1 + + if CUTOFF: + dx2_mask = (x2 <= limit) & (x2 >= -limit) + x2 = tl.clamp(x2, -limit, limit) + dx1_mask = gate <= limit + gate = tl.minimum(gate, limit) + + g = tl.load( + g_ptr + + si * n + + rid * B * n + + 128 * cid + + tl.arange(0, B)[:, None] * n + + tl.arange(0, 128)[None, :], + mask=idx[:, None] < count, + ).to(tl.float32) + + dw = tl.sum(gate * x2 * g, 1) + tl.store(dw_ptr + si * nb + cid + idx * nb, dw, mask=idx < count) + + dx = sigmoid * g * x2 * w * (1 + x1 * (1 - sigmoid)) + + if CUTOFF: + dx *= dx1_mask + + scale = tl.maximum(tl.max(dx.abs(), 1) / 448, 1e-30) + if ROUND: + scale = tl.exp2(tl.ceil(tl.log2(scale))) + tl.store( + dx_scale_ptr + si * nb * 2 + cid * count + rid * B + tl.arange(0, B), + scale, + mask=idx < count, + ) + + tl.store(dx_ptr + offs, dx / scale[:, None], mask=idx[:, None] < count) + + dx = g * gate * w + + if CUTOFF: + dx *= dx2_mask + + scale = tl.maximum(tl.max(dx.abs(), 1) / 448, 1e-30) + if ROUND: + scale = tl.exp2(tl.ceil(tl.log2(scale))) + tl.store( + dx_scale_ptr + + si * nb * 2 + + cid * count + + rid * B + + count * nb + + tl.arange(0, B), + scale, + mask=idx < count, + ) + tl.store(dx_ptr + n + offs, dx / scale[:, None], mask=idx[:, None] < count) + + +@triton.jit +def batch_weighted_silu_and_block_quant_backward_t_kernel( + g_ptr, + x_ptr, + weight_ptr, + count_ptr, + accum_ptr, + transpose_dx_ptr, + transpose_dx_scale_ptr, + limit, + n, + B: tl.constexpr, + E: tl.constexpr, + CUTOFF: tl.constexpr, + ROUND: tl.constexpr, +): + eid = tl.program_id(axis=0) + rid = tl.program_id(axis=1) + cid = tl.program_id(axis=2) + + count = tl.load(count_ptr + eid) + si = tl.load(accum_ptr + eid) - count + # counts = tl.load(count_ptr + tl.arange(0, E)) + # si = tl.sum(tl.where(tl.arange(0, E) < eid, counts, 0)) + + if rid >= tl.cdiv(count, 128): + return + + n = n.to(tl.int64) + transpose_off = tl.sum( + tl.where( + tl.arange(0, E) < eid, tl.cdiv(tl.load(count_ptr + tl.arange(0, E)), 128), 0 + ) + ) + + offs = ( + si * n * 2 + + rid * 128 * n * 2 + + cid * B + + tl.arange(0, 128)[:, None] * n * 2 + + tl.arange(0, B)[None, :] + ) + idx = rid * 128 + tl.arange(0, 128) + w = tl.load(weight_ptr + si + idx, mask=idx < count).to(tl.float32)[:, None] + + x1 = tl.load(x_ptr + offs, mask=idx[:, None] < count).to(tl.float32) + x2 = tl.load(x_ptr + n + offs, mask=idx[:, None] < count).to(tl.float32) + + sigmoid = tl.sigmoid(x1) + gate = sigmoid * x1 + + if CUTOFF: + dx2_mask = (x2 <= limit) & (x2 >= -limit) + x2 = tl.clamp(x2, -limit, limit) + dx1_mask = gate <= limit + gate = tl.minimum(gate, limit) + + g = tl.load( + g_ptr + + si * n + + rid * 128 * n + + B * cid + + tl.arange(0, 128)[:, None] * n + + tl.arange(0, B)[None, :], + mask=idx[:, None] < count, + ).to(tl.float32) + + dx = sigmoid * g * x2 * w * (1 + x1 * (1 - sigmoid)) + + if CUTOFF: + dx *= dx1_mask + + scale = tl.maximum(tl.max(dx.abs(), 0) / 448, 1e-30) + if ROUND: + scale = tl.exp2(tl.ceil(tl.log2(scale))) + tl.store( + transpose_dx_scale_ptr + + transpose_off * n * 2 + + rid * n * 2 + + cid * B + + tl.arange(0, B), + scale, + ) + + qdx = tl.trans((dx / scale[None, :]).to(transpose_dx_ptr.dtype.element_ty)) + tl.store( + transpose_dx_ptr + + si * n * 2 + + rid * 128 + + cid * B * count + + tl.arange(0, B)[:, None] * count + + tl.arange(0, 128)[None, :], + qdx, + mask=idx[None, :] < count, + ) + + dx = g * gate * w + + if CUTOFF: + dx *= dx2_mask + + scale = tl.maximum(tl.max(dx.abs(), 0) / 448, 1e-30) + if ROUND: + scale = tl.exp2(tl.ceil(tl.log2(scale))) + qdx = tl.trans((dx / scale[None, :]).to(transpose_dx_ptr.dtype.element_ty)) + tl.store( + transpose_dx_scale_ptr + + transpose_off * n * 2 + + rid * n * 2 + + n + + cid * B + + tl.arange(0, B), + scale, + ) + tl.store( + transpose_dx_ptr + + count * n + + si * n * 2 + + rid * 128 + + cid * B * count + + tl.arange(0, B)[:, None] * count + + tl.arange(0, 128)[None, :], + qdx, + mask=idx[None, :] < count, + ) + + +# used in routed experts +def triton_batch_weighted_silu_and_block_quant_backward( + g, x, weight, counts, splits=None, limit=None, round_scale=False +): + """ + backward of triton_batch_weighted_silu_and_block_quant_forward + Args: + g: gradient + x: input tensor + weight: router prob tensor + counts: cuda tensor of token count per expert + splits: python int list of token count per expert + round_scale: whether round scale to power of 2 + Returns: + - dx: quantized non-transposed gradient + - dx_scale: scales of quantization non-transposed gradient + - dw: gradient of weight + - transpose_dx: quantized transposed gradient + - transpose_dx_scale: scales of quantization transposed gradient + """ + M, N = x.shape + n = N // 2 + n_expert = counts.shape[0] + assert n % 128 == 0 + assert g.is_contiguous() + assert splits is not None, "batch mode need splits to launch kernels" + + device = x.device + + accums = torch.cumsum(counts, 0) + + dx = torch.empty((M, N), device=device, dtype=torch.float8_e4m3fn) + dx_scale = torch.empty((M, N // 128), device=device, dtype=torch.float32) + + s = sum([(x + 127) // 128 for x in splits]) + transpose_dx = torch.empty((M, N), device=device, dtype=torch.float8_e4m3fn) + transpose_dx_scale = torch.empty((s, N), device=device, dtype=torch.float32) + if s == 0: + dw = torch.empty_like(weight) + return dx, dx_scale, dw, transpose_dx, transpose_dx_scale + + if limit is None: + CUTOFF = False + limit = 0.0 + else: + CUTOFF = True + + # grid = (n_expert, triton.cdiv(max(splits), 128), n // 128) + # dws = torch.empty((M, N // 256), device=device, dtype=torch.float32) + # batch_weighted_silu_and_block_quant_backward_kernel[grid]( + # g, + # x, + # weight, + # counts, + # accums, + # dx, + # dx_scale, + # transpose_dx, + # transpose_dx_scale, + # dws, + # limit, + # n, + # n_expert, + # CUTOFF, + # round_scale, + # num_stages=2, + # num_warps=8 + # ) + # dw = dws.sum(1, keepdim=True).to(weight.dtype) + + B = 32 + grid = (n_expert, triton.cdiv(max(splits), B), n // 128) + dws = torch.empty((M, N // 256), device=device, dtype=torch.float32) + batch_weighted_silu_and_block_quant_backward_n_kernel[grid]( g, x, weight, + counts, + accums, dx, - dw, - M, T, - N, - N // 2, - W, - WEIGHT, - num_stages=3, - num_warps=8 + dx_scale, + dws, + limit, + n, + B, + n_expert, + CUTOFF, + round_scale, + num_stages=2, + num_warps=4, ) - return dx, dw + dw = dws.sum(1, keepdim=True).to(weight.dtype) + + B = 32 + grid = (n_expert, triton.cdiv(max(splits), 128), N // 2 // B) + batch_weighted_silu_and_block_quant_backward_t_kernel[grid]( + g, + x, + weight, + counts, + accums, + transpose_dx, + transpose_dx_scale, + limit, + n, + B, + n_expert, + CUTOFF, + round_scale, + num_stages=2, + num_warps=4, + ) + + return dx, dx_scale, dw, transpose_dx, transpose_dx_scale @triton.jit -def silu_and_block_quant_forward_kernel(x_ptr, - out_ptr, scale_ptr, - transpose_output_ptr, - transpose_scale_ptr, - M, - n: tl.constexpr, - ROUND: tl.constexpr, - OUTPUT_MODE: tl.constexpr): +def silu_and_mxfp8_quant_forward_kernel( + x_ptr, + out_ptr, + scale_ptr, + transpose_output_ptr, + transpose_scale_ptr, + limit, + M, + m, + n: tl.constexpr, + B: tl.constexpr, + CUTOFF: tl.constexpr, + OUTPUT_MODE: tl.constexpr, +): rid = tl.program_id(axis=0) cid = tl.program_id(axis=1) - - offs = rid * 128 * n * 2 + cid * 128 + tl.arange(0, 128)[:, - None] * n * 2 + tl.arange(0, 128)[ - None, :] - indices = rid * 128 + tl.arange(0, 128) - mask = indices[:, None] < M + N: tl.constexpr = n * 2 + sb: tl.constexpr = B // 32 + nb: tl.constexpr = n // 32 + offs = ( + rid * 32 * N + + cid * B + + tl.arange(0, 32)[:, None] * N + + tl.arange(0, B)[None, :] + ) + indices = rid * 32 + tl.arange(0, 32) + mask = indices[:, None] < m x1 = tl.load(x_ptr + offs, mask=mask).to(tl.float32) x2 = tl.load(x_ptr + n + offs, mask=mask).to(tl.float32) - x = x1 * tl.sigmoid(x1) * x2 - # x1 = tl.load(x_ptr + offs, mask=mask) - # x2 = tl.load(x_ptr + n + offs, mask=mask) - # x = tl.sigmoid(x1.to(tl.float32)) * x1 * x2 + if CUTOFF: + x = tl.minimum(x1 * tl.sigmoid(x1), limit) * tl.clamp(x2, -limit, limit) + else: + x = x1 * tl.sigmoid(x1) * x2 if OUTPUT_MODE % 2 == 0: - scale = tl.maximum(tl.max(x.abs(), 1) / 448, 1e-30) - if ROUND: - scale = tl.exp2(tl.ceil(tl.log2(scale))) + xr = tl.reshape(x, [32, sb, 32]) + scale = tl.maximum(tl.max(xr.abs(), 2) / 448, 1e-30) + log_scale = tl.ceil(tl.log2(scale)) + scale = tl.exp2(log_scale) - tl.store(scale_ptr + rid * 128 + cid * M + tl.arange(0, 128), scale, - mask=indices < M) - xq = (x / scale[:, None]).to(out_ptr.dtype.element_ty) - tl.store(out_ptr + rid * 128 * n + cid * 128 + tl.arange(0, 128)[:, - None] * n + tl.arange(0, - 128)[ - None, :], xq, - mask=mask) + tl.store( + scale_ptr + + rid * n + + cid * sb + + tl.arange(0, 32)[:, None] * nb + + tl.arange(0, sb), + log_scale + 127, + ) + xq = tl.reshape(xr / scale[:, :, None], [32, B]).to(out_ptr.dtype.element_ty) + tl.store( + out_ptr + + rid * 32 * n + + cid * B + + tl.arange(0, 32)[:, None] * n + + tl.arange(0, B)[None, :], + xq, + mask=mask, + ) if OUTPUT_MODE > 0: scale = tl.maximum(tl.max(x.abs(), 0) / 448, 1e-30) - if ROUND: - scale = tl.exp2(tl.ceil(tl.log2(scale))) - tl.store(transpose_scale_ptr + rid * n + cid * 128 + tl.arange(0, 128), - scale) - xq = (x / scale).to(out_ptr.dtype.element_ty) - tl.store(transpose_output_ptr + rid * 128 + cid * 128 * M + tl.arange(0, - 128)[ - :, - None] * M + tl.arange( - 0, 128)[ - None, - :], - tl.trans(xq), mask=indices[None, :] < M) - - -def triton_silu_and_block_quant_forward(x, - out=None, - scale=None, - round_scale=False, - output_mode=2): + log_scale = tl.ceil(tl.log2(scale)) + scale = tl.exp2(log_scale) + tl.store( + transpose_scale_ptr + rid * n + cid * B + tl.arange(0, B), log_scale + 127 + ) + xq = (x / scale).to(transpose_output_ptr.dtype.element_ty) + tl.store( + transpose_output_ptr + + rid * 32 * n + + cid * B + + tl.arange(0, 32)[:, None] * n + + tl.arange(0, B)[None, :], + xq, + mask=mask, + ) + + +def triton_silu_and_mxfp8_quant_forward( + x, out=None, scale=None, limit=None, output_mode=2 +): """ - fused silu and blockwise quantization, used in shared expert + fused silu and mxfp8 quantization, used in shared expert Args: x: input tensor round_scale: whether round scale to power of 2 @@ -233,168 +1614,214 @@ def triton_silu_and_block_quant_forward(x, - transpose_output: quantized tensor of transposed output - transpose_scale: quantization scale of transposed output """ - M, N = x.shape + assert x.is_contiguous() + m, N = x.shape + M = (m + 127) // 128 * 128 n = N // 2 + assert n % 128 == 0 # transposed scaled should be multiplier of 128 device = x.device if out is None: - out = torch.empty((M, N // 2), device=device, dtype=torch.float8_e4m3fn) + out = torch.empty((m, n), device=device, dtype=torch.float8_e4m3fn) if scale is None: - scale = torch.empty((N // 2 // 128, M), device=device, - dtype=torch.float32) + scale = torch.empty((M, n // 32), device=device, dtype=torch.uint8) - transpose_output = torch.empty((N // 2, M), device=device, - dtype=torch.float8_e4m3fn) - transpose_scale = torch.empty((triton.cdiv(M, 128), N // 2), device=device, - dtype=torch.float32) + if limit is None: + CUTOFF = False + limit = 0.0 + else: + CUTOFF = True - grid = (triton.cdiv(M, 128), n // 128) - silu_and_block_quant_forward_kernel[grid]( + transpose_output = torch.empty((m, n), device=device, dtype=torch.float8_e4m3fn) + transpose_scale = torch.empty((M // 32, n), device=device, dtype=torch.uint8) + B = 32 # larger B does not perform better, so we use fixed 32 in other kernels + grid = (M // 32, n // B) + silu_and_mxfp8_quant_forward_kernel[grid]( x, out, scale, transpose_output, transpose_scale, + limit, M, + m, n, - round_scale, + B, + CUTOFF, output_mode, num_stages=2, - num_warps=16 + num_warps=1, ) return out, scale, transpose_output, transpose_scale @triton.jit -def silu_and_block_quant_backward_kernel(g_ptr, x_ptr, - dx_ptr, - dx_scale_ptr, - transpose_dx_ptr, - transpose_dx_scale_ptr, - M, - n: tl.constexpr, - ROUND: tl.constexpr): +def silu_and_mxfp8_quant_backward_kernel( + g_ptr, + x_ptr, + dx_ptr, + dx_scale_ptr, + transpose_dx_ptr, + transpose_dx_scale_ptr, + limit, + M, + m, + n: tl.constexpr, + CUTOFF: tl.constexpr, +): rid = tl.program_id(axis=0) cid = tl.program_id(axis=1) - nb = n // 128 - offs = rid * 128 * n * 2 + cid * 128 + tl.arange(0, 128)[:, - None] * n * 2 + tl.arange(0, 128)[ - None, :] - toffs = rid * 128 + cid * M * 128 + tl.arange(0, 128)[:, - None] * M + tl.arange(0, 128)[ - None, :] - idx = rid * 128 + tl.arange(0, 128) - mask = idx[:, None] < M - x1 = tl.load(x_ptr + offs, mask=mask) # .to(tl.float32) - x2 = tl.load(x_ptr + n + offs, mask=mask) # .to(tl.float32) - g = tl.load(g_ptr + rid * 128 * n + cid * 128 + - tl.arange(0, 128)[:, None] * n + - tl.arange(0, 128)[None, :], mask=mask) # .to(tl.float32) + nb = n // 32 + offs = ( + rid * 32 * n * 2 + + cid * 32 + + tl.arange(0, 32)[:, None] * n * 2 + + tl.arange(0, 32)[None, :] + ) + idx = rid * 32 + tl.arange(0, 32) + mask = idx[:, None] < m + x1 = tl.load(x_ptr + offs, mask=mask).to(tl.float32) + x2 = tl.load(x_ptr + n + offs, mask=mask).to(tl.float32) + if CUTOFF: + dx2_mask = (x2 <= limit) & (x2 >= -limit) + x2 = tl.clamp(x2, -limit, limit) + g = tl.load( + g_ptr + + rid * 32 * n + + cid * 32 + + tl.arange(0, 32)[:, None] * n + + tl.arange(0, 32)[None, :], + mask=mask, + ) # .to(tl.float32) sigmoid = tl.sigmoid(x1.to(tl.float32)) - dx1 = sigmoid * g * x2 * ( - 1 + x1 * (1 - sigmoid)) # change order to trigger autocast - scale1 = tl.maximum( - tl.max(dx1.abs(), 1) / 448, 1e-30) - if ROUND: - scale1 = tl.exp2(tl.ceil(tl.log2(scale1))) - tl.store(dx_scale_ptr + cid * M + rid * 128 + tl.arange(0, 128), scale1, - mask=idx < M) + dx1 = ( + sigmoid * g * x2 * (1 + x1 * (1 - sigmoid)) + ) # change order to trigger autocast + if CUTOFF: + dx1_mask = (sigmoid * x1 <= limit).to(tl.float32) + dx1 *= dx1_mask + scale1 = tl.maximum(tl.max(dx1.abs(), 1) / 448, 1e-30) + + log_scale1 = tl.ceil(tl.log2(scale1)) + scale1 = tl.exp2(log_scale1) + + # rid * nb * 32 = rid * n + # write padding zeros to scale tensor + tl.store( + dx_scale_ptr + rid * n * 2 + cid + tl.arange(0, 32) * nb * 2, log_scale1 + 127 + ) qdx1 = (dx1 / scale1[:, None]).to(dx_ptr.dtype.element_ty) tl.store(dx_ptr + offs, qdx1, mask=mask) - scale1 = tl.maximum( - tl.max(dx1.abs(), 0) / 448, 1e-30) - if ROUND: - scale1 = tl.exp2(tl.ceil(tl.log2(scale1))) + scale1 = tl.maximum(tl.max(dx1.abs(), 0) / 448, 1e-30) + + log_scale1 = tl.ceil(tl.log2(scale1)) + scale1 = tl.exp2(log_scale1) + tl.store( - transpose_dx_scale_ptr + rid * n * 2 + cid * 128 + tl.arange(0, 128), - scale1) + transpose_dx_scale_ptr + rid * n * 2 + cid * 32 + tl.arange(0, 32), + log_scale1 + 127, + ) - qdx1 = (dx1 / scale1[None, :]).to(dx_ptr.dtype.element_ty) - tl.store(transpose_dx_ptr + toffs, tl.trans(qdx1), mask=idx[None, :] < M) + qdx1 = (dx1 / scale1[None, :]).to(transpose_dx_ptr.dtype.element_ty) + tl.store(transpose_dx_ptr + offs, qdx1, mask=mask) - dx2 = sigmoid * g * x1 - scale2 = tl.maximum( - tl.max(dx2.abs(), 1) / 448, 1e-30) - if ROUND: - scale2 = tl.exp2(tl.ceil(tl.log2(scale2))) - tl.store(dx_scale_ptr + cid * M + rid * 128 + M * nb + tl.arange(0, 128), - scale2, mask=idx < M) + # dx2 = sigmoid * g * x1 + if CUTOFF: + dx2 = g * tl.minimum(sigmoid * x1, limit) * dx2_mask.to(tl.float32) + else: + dx2 = g * sigmoid * x1 + scale2 = tl.maximum(tl.max(dx2.abs(), 1) / 448, 1e-30) + log_scale2 = tl.ceil(tl.log2(scale2)) + scale2 = tl.exp2(log_scale2) + + tl.store( + dx_scale_ptr + rid * n * 2 + cid + nb + tl.arange(0, 32) * nb * 2, + log_scale2 + 127, + ) qdx2 = (dx2 / scale2[:, None]).to(dx_ptr.dtype.element_ty) - tl.store(dx_ptr + offs + n, qdx2, mask=idx[:, None] < M) + tl.store(dx_ptr + offs + n, qdx2, mask=mask) - scale2 = tl.maximum( - tl.max(dx2.abs(), 0) / 448, 1e-30) - if ROUND: - scale2 = tl.exp2(tl.ceil(tl.log2(scale2))) - tl.store(transpose_dx_scale_ptr + rid * n * 2 + n + cid * 128 + tl.arange(0, - 128), - scale2) + scale2 = tl.maximum(tl.max(dx2.abs(), 0) / 448, 1e-30) + log_scale2 = tl.ceil(tl.log2(scale2)) + scale2 = tl.exp2(log_scale2) + tl.store( + transpose_dx_scale_ptr + rid * n * 2 + n + cid * 32 + tl.arange(0, 32), + log_scale2 + 127, + ) - qdx2 = (dx2 / scale2[None, :]).to(dx_ptr.dtype.element_ty) - tl.store(transpose_dx_ptr + M * n + toffs, tl.trans(qdx2), - mask=idx[None, :] < M) + qdx2 = (dx2 / scale2[None, :]).to(transpose_dx_ptr.dtype.element_ty) + tl.store(transpose_dx_ptr + n + offs, qdx2, mask=mask) # used in shared expert -def triton_silu_and_block_quant_backward(g, x, - round_scale=False): +def triton_silu_and_mxfp8_quant_backward(g, x, limit=None): """ - backward of triton_silu_and_block_quant_forward + backward of triton_silu_and_mxfp8_quant_forward Args: g: gradient x: input tensor - round_scale: whether round to power of 2 Returns: - - dx: quantized non-transposed gradient - - dx_scale: scales of quantization non-transposed gradient - - transpose_dx: quantized transposed gradient - - transpose_dx_scale: scales of quantization transposed gradient + - dx: rowwise quantized gradient + - dx_scale: scales of rowwise quantized gradient + - transpose_dx: columnwise quantized gradient + - transpose_dx_scale: scales of columnwise quantized gradient """ - M, N = x.shape + assert g.is_contiguous() + m, N = x.shape + M = (m + 127) // 128 * 128 n = N // 2 + assert N % 128 == 0 device = x.device - dx = torch.empty((M, N), device=device, dtype=torch.float8_e4m3fn) - - dx_scale = torch.empty((N // 128, M), device=device, dtype=torch.float32) - scale_shape = (triton.cdiv(M, 128), N) - transpose_dx = torch.empty((N, M), device=device, dtype=torch.float8_e4m3fn) - transpose_dx_scale = torch.empty(scale_shape, device=device, - dtype=torch.float32) + dx = torch.empty((m, N), device=device, dtype=torch.float8_e4m3fn) + dx_scale = torch.empty((M, N // 32), device=device, dtype=torch.uint8) + transpose_dx = torch.empty((m, N), device=device, dtype=torch.float8_e4m3fn) + transpose_dx_scale = torch.empty((M // 32, N), device=device, dtype=torch.uint8) + + if limit is None: + CUTOFF = False + limit = 0.0 + else: + CUTOFF = True - assert M % 128 == 0 - grid = (M // 128, N // 256) - silu_and_block_quant_backward_kernel[grid]( + grid = (M // 32, n // 32) + silu_and_mxfp8_quant_backward_kernel[grid]( g, x, dx, dx_scale, transpose_dx, transpose_dx_scale, + limit, M, + m, n, - round_scale, - num_stages=2, - num_warps=8 + CUTOFF, + num_stages=3, + num_warps=4, ) return dx, dx_scale, transpose_dx, transpose_dx_scale @triton.jit -def batch_weighted_silu_and_block_quant_forward_kernel(x_ptr, weight_ptr, - out_ptr, - scale_ptr, - transpose_output_ptr, - transpose_scale_ptr, - count_ptr, - accum_ptr, - n: tl.constexpr, - E: tl.constexpr, - ROUND: tl.constexpr, - OUTPUT_MODE: tl.constexpr): +def batch_weighted_silu_and_mxfp8_quant_forward_kernel( + x_ptr, + weight_ptr, + out_ptr, + scale_ptr, + transpose_output_ptr, + transpose_scale_ptr, + count_ptr, + accum_ptr, + limit, + n, + E: tl.constexpr, + CUTOFF: tl.constexpr, + OUTPUT_MODE: tl.constexpr, +): eid = tl.program_id(axis=0) rid = tl.program_id(axis=1) cid = tl.program_id(axis=2) @@ -404,68 +1831,78 @@ def batch_weighted_silu_and_block_quant_forward_kernel(x_ptr, weight_ptr, si = ei - count c = tl.cdiv(count, 128) - if rid >= c: + if rid >= c * 4: return - nb = n // 128 + n = n.to(tl.int64) + nb = n // 32 counts = tl.load(count_ptr + tl.arange(0, E)) n_blocks = tl.cdiv(counts, 128) - transpose_scale_off = tl.sum(tl.where(tl.arange(0, E) < eid, n_blocks, 0)) - - offs = si * n * 2 + rid * 128 * n * 2 + cid * 128 + tl.arange(0, 128)[:, - None] * n * 2 + tl.arange( - 0, 128)[None, :] - hoffs = si * n + rid * 128 * n + cid * 128 + tl.arange(0, 128)[:, - None] * n + tl.arange(0, 128)[ - None, :] - toffs = si * n + rid * 128 + cid * count * 128 + tl.arange(0, 128)[:, - None] * count + tl.arange( - 0, 128)[ - None, :] - indices = rid * 128 + tl.arange(0, 128) + scale_off = tl.sum(tl.where(tl.arange(0, E) < eid, n_blocks, 0)) + + offs = ( + si * n * 2 + + rid * 32 * n * 2 + + cid * 32 + + tl.arange(0, 32)[:, None] * n * 2 + + tl.arange(0, 32)[None, :] + ) + hoffs = ( + si * n + + rid * 32 * n + + cid * 32 + + tl.arange(0, 32)[:, None] * n + + tl.arange(0, 32)[None, :] + ) + indices = rid * 32 + tl.arange(0, 32) mask = indices[:, None] < count - w = tl.load(weight_ptr + si + indices, mask=indices < count).to( - tl.float32) + w = tl.load(weight_ptr + si + indices, mask=indices < count).to(tl.float32) x1 = tl.load(x_ptr + offs, mask=mask).to(tl.float32) - x2 = tl.load(x_ptr + n + offs, mask=mask).to( - tl.float32) + x2 = tl.load(x_ptr + n + offs, mask=mask).to(tl.float32) - x = x1 * tl.sigmoid(x1) * x2 * w[:, None] + if CUTOFF: + x = ( + tl.minimum(x1 * tl.sigmoid(x1), limit) + * tl.clamp(x2, -limit, limit) + * w[:, None] + ) + else: + x = x1 * tl.sigmoid(x1) * x2 * w[:, None] if OUTPUT_MODE % 2 == 0: scale = tl.maximum(tl.max(tl.abs(x), 1) / 448, 1e-30) - if ROUND: - scale = tl.exp2(tl.ceil(tl.log2(scale))) + log_scale = tl.ceil(tl.log2(scale)) + scale = tl.exp2(log_scale) + # 4 = 128 // 32 tl.store( - scale_ptr + si * nb + cid * count + rid * 128 + tl.arange(0, 128), - scale, mask=indices < count) + scale_ptr + scale_off * 4 * n + rid * n + cid + tl.arange(0, 32) * nb, + log_scale + 127, + ) xq = (x / scale[:, None]).to(out_ptr.dtype.element_ty) tl.store(out_ptr + hoffs, xq, mask=mask) if OUTPUT_MODE > 0: scale = tl.maximum(tl.max(tl.abs(x), 0) / 448, 1e-30) - if ROUND: - scale = tl.exp2(tl.ceil(tl.log2(scale))) + log_scale = tl.ceil(tl.log2(scale)) + scale = tl.exp2(log_scale) tl.store( - transpose_scale_ptr + transpose_scale_off * n + rid * n + cid * 128 + tl.arange( - 0, 128), scale) - - xq = tl.trans((x / scale).to(out_ptr.dtype.element_ty)) - tl.store(transpose_output_ptr + toffs, xq, - mask=indices[None, :] < count) + transpose_scale_ptr + + scale_off * 4 * n + + rid * n + + cid * 32 + + tl.arange(0, 32), + log_scale + 127, + ) + xq = (x / scale).to(transpose_output_ptr.dtype.element_ty) + tl.store(transpose_output_ptr + hoffs, xq, mask=mask) -def triton_batch_weighted_silu_and_block_quant_forward(x, - weight, - counts, - splits=None, - out=None, - scale=None, - round_scale=False, - output_mode=2): +def triton_batch_weighted_silu_and_mxfp8_quant_forward( + x, weight, counts, splits=None, out=None, scale=None, limit=None, output_mode=2 +): """ silu and blockwise quantize activation in routed experts Args: @@ -473,7 +1910,6 @@ def triton_batch_weighted_silu_and_block_quant_forward(x, weight: router prob tensor counts: cuda tensor of token count per expert splits: python int list of token count per expert - round_scale: whether round scale to power of 2 output_mode: one of {0, 1, 2} 0: only output non-transposed quantized tensor 1: only output transposed quantized tensor @@ -485,32 +1921,37 @@ def triton_batch_weighted_silu_and_block_quant_forward(x, - transpose_output: quantized tensor of transposed output - transpose_scale: quantization scale of transposed output """ - M, N = x.shape + assert x.is_contiguous() and weight.is_contiguous() + m, N = x.shape n = N // 2 n_experts = counts.shape[0] - assert N <= 8192 + assert N <= 8192 and n % 128 == 0 + assert splits is not None, "batch mode need splits to launch kernels" + M = sum([(x + 127) // 128 for x in splits]) * 128 + device = x.device if out is None: - out = torch.empty((M, n), device=device, dtype=torch.float8_e4m3fn) + out = torch.empty((m, n), device=device, dtype=torch.float8_e4m3fn) - assert splits is not None, 'batch mode need splits to launch kernels' - blocks = sum([(x + 127) // 128 for x in splits]) - transpose_output = torch.empty((M * n), device=device, - dtype=torch.float8_e4m3fn) - transpose_scale = torch.empty((blocks * n), device=device, - dtype=torch.float32) - # intra layout and inner layput are not consist, - # tensors will be viewed after splitting if scale is None: - scale = torch.empty((M * n // 128,), device=device, dtype=torch.float32) + scale = torch.empty((M, n // 32), device=device, dtype=torch.uint8) + + transpose_output = torch.empty((m, n), device=device, dtype=torch.float8_e4m3fn) + transpose_scale = torch.empty((M // 32, n), device=device, dtype=torch.uint8) if M == 0: return out, scale, transpose_output, transpose_scale + if limit is None: + CUTOFF = False + limit = 0.0 + else: + CUTOFF = True + accums = torch.cumsum(counts, 0) - grid = (n_experts, triton.cdiv(max(splits), 128), n // 128) - batch_weighted_silu_and_block_quant_forward_kernel[grid]( + grid = (n_experts, triton.cdiv(max(splits), 128) * 4, n // 32) + batch_weighted_silu_and_mxfp8_quant_forward_kernel[grid]( x, weight, out, @@ -519,30 +1960,35 @@ def triton_batch_weighted_silu_and_block_quant_forward(x, transpose_scale, counts, accums, + limit, n, len(splits), - round_scale, + CUTOFF, output_mode, - num_stages=2, - num_warps=8 + num_stages=3, + num_warps=1, ) return out, scale, transpose_output, transpose_scale @triton.jit -def batch_weighted_silu_and_block_quant_backward_kernel(g_ptr, x_ptr, - weight_ptr, - count_ptr, - accum_ptr, - dx_ptr, - dx_scale_ptr, - transpose_dx_ptr, - transpose_dx_scale_ptr, - dw_ptr, - n: tl.constexpr, - E: tl.constexpr, - ROUND: tl.constexpr): +def batch_weighted_silu_and_mxfp8_quant_backward_kernel( + g_ptr, + x_ptr, + weight_ptr, + count_ptr, + accum_ptr, + dx_ptr, + dx_scale_ptr, + transpose_dx_ptr, + transpose_dx_scale_ptr, + dw_ptr, + limit, + n, + E: tl.constexpr, + CUTOFF: tl.constexpr, +): eid = tl.program_id(axis=0) rid = tl.program_id(axis=1) cid = tl.program_id(axis=2) @@ -550,97 +1996,157 @@ def batch_weighted_silu_and_block_quant_backward_kernel(g_ptr, x_ptr, count = tl.load(count_ptr + eid) si = tl.load(accum_ptr + eid) - count - if rid >= tl.cdiv(count, 128): + if rid >= tl.cdiv(count, 128) * 4: return - nb = n // 128 - transpose_off = tl.sum(tl.where(tl.arange(0, E) < eid, tl.cdiv( - tl.load(count_ptr + tl.arange(0, E)), 128), 0)) + n = n.to(tl.int64) + + nb = n // 32 + scale_off = tl.sum( + tl.where( + tl.arange(0, E) < eid, tl.cdiv(tl.load(count_ptr + tl.arange(0, E)), 128), 0 + ) + ) - offs = si * n * 2 + rid * 128 * n * 2 + cid * 128 + tl.arange(0, 128)[:, - None] * n * 2 + tl.arange( - 0, 128)[None, :] + offs = ( + si * n * 2 + + rid * 32 * n * 2 + + cid * 32 + + tl.arange(0, 32)[:, None] * n * 2 + + tl.arange(0, 32)[None, :] + ) # hoffs = si * n + tid * 128 * n + tl.arange(0, 128)[:, None] * n + tl.arange(0, 128)[None, :] # toffs = si * n * 2 + tid * 128 + tl.arange(0, 128)[:, None] * count + tl.arange(0, 128)[None, :] - idx = rid * 128 + tl.arange(0, 128) + idx = rid * 32 + tl.arange(0, 32) w = tl.load(weight_ptr + si + idx, mask=idx < count).to(tl.float32)[:, None] - x1 = tl.load(x_ptr + offs, mask=idx[:, None] < count) # .to(tl.float32) - x2 = tl.load(x_ptr + n + offs, mask=idx[:, None] < count) # .to(tl.float32) - g = tl.load(g_ptr + si * n + rid * 128 * n + 128 * cid + - tl.arange(0, 128)[:, None] * n + - tl.arange(0, 128)[None, :], - mask=idx[:, None] < count) # .to(tl.float32) - sigmoid = tl.sigmoid(x1.to(tl.float32)) - - dw = tl.sum(sigmoid * x1 * x2 * g, 1) + x1 = tl.load(x_ptr + offs, mask=idx[:, None] < count).to(tl.float32) + x2 = tl.load(x_ptr + n + offs, mask=idx[:, None] < count).to(tl.float32) + sigmoid = tl.sigmoid(x1) + gate = sigmoid * x1 + + if CUTOFF: + dx2_mask = (x2 <= limit) & (x2 >= -limit) + x2 = tl.clamp(x2, -limit, limit) + dx1_mask = gate <= limit + gate = tl.minimum(gate, limit) + + g = tl.load( + g_ptr + + si * n + + rid * 32 * n + + 32 * cid + + tl.arange(0, 32)[:, None] * n + + tl.arange(0, 32)[None, :], + mask=idx[:, None] < count, + ).to(tl.float32) + # sigmoid = tl.sigmoid(x1.to(tl.float32)) + + dw = tl.sum(gate * x2 * g, 1) tl.store(dw_ptr + si * nb + cid + idx * nb, dw, mask=idx < count) dx = sigmoid * g * x2 * w * (1 + x1 * (1 - sigmoid)) - scale = tl.maximum( - tl.max(dx.abs(), 1) / 448, 1e-30) - if ROUND: - scale = tl.exp2(tl.ceil(tl.log2(scale))) - tl.store(dx_scale_ptr + si * nb * 2 + cid * count + rid * 128 + tl.arange(0, - 128), - scale, mask=idx < count) + + if CUTOFF: + dx *= dx1_mask + + scale = tl.maximum(tl.max(dx.abs(), 1) / 448, 1e-30) + + log_scale = tl.ceil(tl.log2(scale)) + scale = tl.exp2(log_scale) + tl.store( + dx_scale_ptr + + scale_off * n * 8 + + rid * n * 2 + + cid + + tl.arange(0, 32) * nb * 2, + log_scale + 127, + ) tl.store(dx_ptr + offs, dx / scale[:, None], mask=idx[:, None] < count) - scale = tl.maximum( - tl.max(dx.abs(), 0) / 448, 1e-30) - if ROUND: - scale = tl.exp2(tl.ceil(tl.log2(scale))) + scale = tl.maximum(tl.max(dx.abs(), 0) / 448, 1e-30) + log_scale = tl.ceil(tl.log2(scale)) + scale = tl.exp2(log_scale) tl.store( - transpose_dx_scale_ptr + transpose_off * n * 2 + rid * n * 2 + cid * 128 + tl.arange( - 0, 128), scale) + transpose_dx_scale_ptr + + scale_off * n * 8 + + rid * n * 2 + + cid * 32 + + tl.arange(0, 32), + log_scale + 127, + ) - qdx = tl.trans((dx / scale[None, :]).to(dx_ptr.dtype.element_ty)) + qdx = (dx / scale[None, :]).to(transpose_dx_ptr.dtype.element_ty) # tl.store(transpose_dx_ptr + toffs, qdx, mask=idx[None, :] < count) - tl.store(transpose_dx_ptr + si * n * 2 + rid * 128 + cid * 128 * count + - tl.arange(0, 128)[:, None] * count + - tl.arange(0, 128)[None, :], - qdx, - mask=idx[None, :] < count) + tl.store( + transpose_dx_ptr + + si * n * 2 + + rid * 32 * n * 2 + + cid * 32 + + tl.arange(0, 32)[:, None] * n * 2 + + tl.arange(0, 32)[None, :], + qdx, + mask=idx[:, None] < count, + ) dx = sigmoid * g * x1 * w - scale = tl.maximum( - tl.max(dx.abs(), 1) / 448, 1e-30) - if ROUND: - scale = tl.exp2(tl.ceil(tl.log2(scale))) + + if CUTOFF: + dx *= dx2_mask + + scale = tl.maximum(tl.max(dx.abs(), 1) / 448, 1e-30) + log_scale = tl.ceil(tl.log2(scale)) + scale = tl.exp2(log_scale) tl.store( - dx_scale_ptr + si * nb * 2 + cid * count + rid * 128 + count * nb + tl.arange( - 0, 128), scale, mask=idx < count) + dx_scale_ptr + + scale_off * n * 8 + + rid * n * 2 + + cid + + nb + + tl.arange(0, 32) * nb * 2, + log_scale + 127, + ) tl.store(dx_ptr + n + offs, dx / scale[:, None], mask=idx[:, None] < count) - scale = tl.maximum( - tl.max(dx.abs(), 0) / 448, 1e-30) - if ROUND: - scale = tl.exp2(tl.ceil(tl.log2(scale))) - qdx = tl.trans((dx / scale[None, :]).to(dx_ptr.dtype.element_ty)) + scale = tl.maximum(tl.max(dx.abs(), 0) / 448, 1e-30) + log_scale = tl.ceil(tl.log2(scale)) + scale = tl.exp2(log_scale) + qdx = (dx / scale[None, :]).to(transpose_dx_ptr.dtype.element_ty) tl.store( - transpose_dx_scale_ptr + transpose_off * n * 2 + rid * n * 2 + n + cid * 128 + tl.arange( - 0, 128), scale) + transpose_dx_scale_ptr + + scale_off * n * 8 + + rid * n * 2 + + n + + cid * 32 + + tl.arange(0, 32), + log_scale + 127, + ) tl.store( - transpose_dx_ptr + count * n + si * n * 2 + rid * 128 + cid * 128 * count + tl.arange( - 0, 128)[:, None] * count + tl.arange(0, 128)[None, :], qdx, - mask=idx[None, :] < count) + transpose_dx_ptr + + si * n * 2 + + rid * 32 * n * 2 + + cid * 32 + + n + + tl.arange(0, 32)[:, None] * n * 2 + + tl.arange(0, 32)[None, :], + qdx, + mask=idx[:, None] < count, + ) # used in routed experts -def triton_batch_weighted_silu_and_block_quant_backward(g, x, weight, - counts, - splits=None, - round_scale=False): +def triton_batch_weighted_silu_and_mxfp8_quant_backward( + g, x, weight, counts, splits=None, limit=None +): """ - backward of triton_batch_weighted_silu_and_block_quant_forward + backward of triton_batch_weighted_silu_and_mxfp8_quant_forward Args: g: gradient x: input tensor weight: router prob tensor counts: cuda tensor of token count per expert splits: python int list of token count per expert - round_scale: whether round scale to power of 2 Returns: - dx: quantized non-transposed gradient - dx_scale: scales of quantization non-transposed gradient @@ -648,36 +2154,36 @@ def triton_batch_weighted_silu_and_block_quant_backward(g, x, weight, - transpose_dx: quantized transposed gradient - transpose_dx_scale: scales of quantization transposed gradient """ - # row-wise read, row-wise write - M, N = x.shape + assert g.is_contiguous() + m, N = x.shape n = N // 2 - n_expert = counts.shape[0] - assert N <= 8192 and 8192 % N == 0 - assert splits is not None, 'batch mode need splits to launch kernels' + n_experts = counts.shape[0] + assert N <= 8192 and N % 128 == 0 + assert splits is not None, "batch mode need splits to launch kernels" + M = sum([(x + 127) // 128 for x in splits]) * 128 device = x.device accums = torch.cumsum(counts, 0) - dx = torch.empty((M, N), device=device, dtype=torch.float8_e4m3fn) - - # intra layout and inner layput are not consist, - # tensors will be viewed after splitting - dx_scale = torch.empty((N // 128 * M), device=device, dtype=torch.float32) + dx = torch.empty((m, N), device=device, dtype=torch.float8_e4m3fn) + dx_scale = torch.empty((M, N // 32), device=device, dtype=torch.uint8) - s = sum([(x + 127) // 128 for x in splits]) - transpose_dx = torch.empty((N * M), device=device, - dtype=torch.float8_e4m3fn) - transpose_dx_scale = torch.empty((s * N), device=device, - dtype=torch.float32) - if s == 0: + transpose_dx = torch.empty((m, N), device=device, dtype=torch.float8_e4m3fn) + transpose_dx_scale = torch.empty((M // 32, N), device=device, dtype=torch.uint8) + if M == 0: dw = torch.empty_like(weight) return dx, dx_scale, dw, transpose_dx, transpose_dx_scale - # grid = (n_expert, triton.cdiv(max(splits), 128)) - grid = (n_expert, triton.cdiv(max(splits), 128), N // 256) - dws = torch.empty((M, N // 256), device=device, dtype=torch.float32) - batch_weighted_silu_and_block_quant_backward_kernel[grid]( + if limit is None: + CUTOFF = False + limit = 0.0 + else: + CUTOFF = True + + grid = (n_experts, triton.cdiv(max(splits), 128) * 4, N // 64) + dws = torch.empty((m, N // 64), device=device, dtype=torch.float32) + batch_weighted_silu_and_mxfp8_quant_backward_kernel[grid]( g, x, weight, @@ -688,43 +2194,43 @@ def triton_batch_weighted_silu_and_block_quant_backward(g, x, weight, transpose_dx, transpose_dx_scale, dws, + limit, n, - n_expert, - round_scale, + n_experts, + CUTOFF, num_stages=3, - num_warps=16 + num_warps=4, ) dw = dws.sum(1, keepdim=True).to(weight.dtype) return dx, dx_scale, dw, transpose_dx, transpose_dx_scale - - - # n is power of 2 @triton.jit -def silu_and_smooth_quant_forward_kernel(x_ptr, smooth_scale_ptr, out_ptr, scale_ptr, - max_ptr, M, T, n: tl.constexpr, - W: tl.constexpr, ROUND: tl.constexpr, - CALIBRATE: tl.constexpr): +def silu_and_smooth_quant_forward_kernel( + x_ptr, + smooth_scale_ptr, + out_ptr, + scale_ptr, + M, + T, + n: tl.constexpr, + W: tl.constexpr, + ROUND: tl.constexpr, +): pid = tl.program_id(axis=0) row_offs = pid * T * W * n + tl.arange(0, W)[:, None] * n col_offs = tl.arange(0, n)[None, :] smooth_scale = tl.load(smooth_scale_ptr + tl.arange(0, n)) smooth_scale = 1.0 / smooth_scale - if CALIBRATE: - maxs = tl.zeros((W, n), dtype=tl.float32) for i in range(T): indices = pid * T * W + i * W + tl.arange(0, W) mask = indices[:, None] < M x1 = tl.load(x_ptr + row_offs * 2 + col_offs, mask=mask).to(tl.float32) - x2 = tl.load(x_ptr + n + row_offs * 2 + col_offs, mask=mask).to( - tl.float32) - x = x1 / (1 + tl.exp(-x1)) * x2 - if CALIBRATE: - maxs = tl.maximum(x.abs(), maxs) + x2 = tl.load(x_ptr + n + row_offs * 2 + col_offs, mask=mask).to(tl.float32) + x = x1 * tl.sigmoid(x1) * x2 x = x * smooth_scale scale = tl.maximum(tl.max(x.abs(), 1) / 448, 1e-30) if ROUND: @@ -734,19 +2240,20 @@ def silu_and_smooth_quant_forward_kernel(x_ptr, smooth_scale_ptr, out_ptr, scale tl.store(out_ptr + row_offs + col_offs, x, mask=mask) row_offs += n * W - if CALIBRATE: - maxs = tl.max(maxs, 0) - tl.store(max_ptr + pid * n + tl.arange(0, n), maxs) - # n is NOT power of 2 @triton.jit -def compatible_silu_and_smooth_quant_forward_kernel(x_ptr, smooth_scale_ptr, out_ptr, - scale_ptr, max_ptr, M, - T: tl.constexpr, n: tl.constexpr, - B: tl.constexpr, - ROUND: tl.constexpr, - CALIBRATE: tl.constexpr): +def compatible_silu_and_smooth_quant_forward_kernel( + x_ptr, + smooth_scale_ptr, + out_ptr, + scale_ptr, + M, + T: tl.constexpr, + n: tl.constexpr, + B: tl.constexpr, + ROUND: tl.constexpr, +): pid = tl.program_id(axis=0) # rowwise read with block size [T, B] @@ -756,14 +2263,10 @@ def compatible_silu_and_smooth_quant_forward_kernel(x_ptr, smooth_scale_ptr, out nb = n // B maxs = tl.zeros((T,), dtype=tl.float32) for i in range(nb): - smooth_scale = tl.load(smooth_scale_ptr + i * B + tl.arange(0, B)) x1 = tl.load(x_ptr + row_offs * 2 + col_offs).to(tl.float32) x2 = tl.load(x_ptr + n + row_offs * 2 + col_offs).to(tl.float32) - x = x1 / (1 + tl.exp(-x1)) * x2 - if CALIBRATE: - x_maxs = tl.max(x.abs(), 0) - tl.store(max_ptr + pid * n + i * B + tl.arange(0, B), x_maxs) + x = x1 * tl.sigmoid(x1) * x2 x = x / smooth_scale maxs = tl.maximum(tl.max(x.abs(), 1), maxs) col_offs += B @@ -779,7 +2282,7 @@ def compatible_silu_and_smooth_quant_forward_kernel(x_ptr, smooth_scale_ptr, out x1 = tl.load(x_ptr + row_offs * 2 + col_offs).to(tl.float32) x2 = tl.load(x_ptr + n + row_offs * 2 + col_offs).to(tl.float32) - x = x1 / (1 + tl.exp(-x1)) * x2 + x = x1 * tl.sigmoid(x1) * x2 x = x / smooth_scale x = (x / scale[:, None]).to(out_ptr.dtype.element_ty) @@ -787,13 +2290,12 @@ def compatible_silu_and_smooth_quant_forward_kernel(x_ptr, smooth_scale_ptr, out col_offs += B - - # used in shared expert -def triton_silu_and_smooth_quant_forward(x, smooth_scale=None, out=None, scale=None, - maxs=None, round_scale=False, - calibrate=False): +def triton_silu_and_smooth_quant_forward( + x, smooth_scale, out=None, scale=None, round_scale=False +): """""" + assert x.is_contiguous() M, N = x.shape n = N // 2 device = x.device @@ -803,86 +2305,74 @@ def triton_silu_and_smooth_quant_forward(x, smooth_scale=None, out=None, scale=N scale = torch.empty((M,), device=device, dtype=torch.float32) if triton.next_power_of_2(N) == N and N <= 8192: - # sm = torch.cuda.get_device_properties(device).multi_processor_count W = 8192 // N - T = 8 if M//W >= 1024 else 4 - assert M % (T*W) == 0 - g = M//(T*W) + T = 8 if M // W >= 1024 else 4 + assert M % (T * W) == 0 + g = M // (T * W) # T = triton.cdiv(M, sm * W) - if maxs is None and calibrate: - maxs = torch.empty((g, n), device=device, dtype=torch.float32) grid = (g,) silu_and_smooth_quant_forward_kernel[grid]( x, smooth_scale, out, scale, - maxs, M, T, n, W, round_scale, - calibrate, num_stages=2, - num_warps=16 + num_warps=16, ) else: B = 512 T = 16 assert n % B == 0 and M % T == 0 grid = (M // T,) - if maxs is None and calibrate: - maxs = torch.empty((M // T, n), device=device, dtype=torch.float32) compatible_silu_and_smooth_quant_forward_kernel[grid]( x, smooth_scale, out, scale, - maxs, M, T, N // 2, B, round_scale, - calibrate, num_stages=2, - num_warps=16 + num_warps=16, ) - if calibrate: - maxs = maxs.amax(0) - - - return out, scale, maxs - - - + return out, scale @triton.jit -def silu_and_smooth_quant_backward_kernel(g_ptr, x_ptr, - smooth_scale_ptr, - transpose_smooth_scale_ptr, - dx_ptr, dx_scale_ptr, - transpose_dx_ptr, - transpose_dx_scale_ptr, - M, - n: tl.constexpr, - T: tl.constexpr, - B: tl.constexpr, - REVERSE: tl.constexpr, - ROUND: tl.constexpr): +def silu_and_smooth_quant_backward_kernel( + g_ptr, + x_ptr, + smooth_scale_ptr, + transpose_smooth_scale_ptr, + dx_ptr, + dx_scale_ptr, + transpose_dx_ptr, + transpose_dx_scale_ptr, + M, + n: tl.constexpr, + T: tl.constexpr, + B: tl.constexpr, + REVERSE: tl.constexpr, + ROUND: tl.constexpr, +): pid = tl.program_id(axis=0) - offs = pid * T * n * 2 + tl.arange(0, T)[:, None] * n * 2 + tl.arange(0, B)[ - None, :] - hoffs = pid * T * n + tl.arange(0, T)[:, None] * n + tl.arange(0, B)[None, - :] + offs = pid * T * n * 2 + tl.arange(0, T)[:, None] * n * 2 + tl.arange(0, B)[None, :] + hoffs = pid * T * n + tl.arange(0, T)[:, None] * n + tl.arange(0, B)[None, :] toffs = pid * T + tl.arange(0, B)[:, None] * M + tl.arange(0, T)[None, :] nb = n // B - maxs = tl.zeros((T, ), dtype=tl.float32) - transpose_smooth_scale = tl.load(transpose_smooth_scale_ptr + pid * T + tl.arange(0, T))[:, None] + maxs = tl.zeros((T,), dtype=tl.float32) + transpose_smooth_scale = tl.load( + transpose_smooth_scale_ptr + pid * T + tl.arange(0, T) + )[:, None] for i in range(nb): smooth_scale_1 = tl.load(smooth_scale_ptr + i * B + tl.arange(0, B)) smooth_scale_2 = tl.load(smooth_scale_ptr + n + i * B + tl.arange(0, B)) @@ -893,40 +2383,46 @@ def silu_and_smooth_quant_backward_kernel(g_ptr, x_ptr, x1 = tl.load(x_ptr + offs).to(tl.float32) x2 = tl.load(x_ptr + offs + n).to(tl.float32) g = tl.load(g_ptr + hoffs).to(tl.float32) - sigmoid = 1 / (1 + tl.exp(-x1)) + sigmoid = tl.sigmoid(x1) # x1 = tl.load(x_ptr + offs) # x2 = tl.load(x_ptr + offs + n) # g = tl.load(g_ptr + hoffs) - # sigmoid = 1 / (1 + tl.exp(-x1.to(tl.float32))) + # sigmoid = tl.sigmoid(x1.to(tl.float32)) - dx1 = g * x2 * sigmoid * ( - 1 + x1 * (1 - sigmoid)) + dx1 = g * x2 * sigmoid * (1 + x1 * (1 - sigmoid)) dx2 = g * x1 * sigmoid t_dx = dx1 * transpose_smooth_scale t_s = tl.maximum(tl.max(tl.abs(t_dx), 0) / 448, 1e-30) if ROUND: t_s = tl.exp2(tl.ceil(tl.log2(t_s))) - t_dx = t_dx/t_s - tl.store(transpose_dx_ptr + toffs, tl.trans(t_dx.to(transpose_dx_ptr.dtype.element_ty))) + t_dx = t_dx / t_s + tl.store( + transpose_dx_ptr + toffs, + tl.trans(t_dx.to(transpose_dx_ptr.dtype.element_ty)), + ) tl.store(transpose_dx_scale_ptr + pid * n * 2 + i * B + tl.arange(0, B), t_s) t_dx = dx2 * transpose_smooth_scale t_s = tl.maximum(tl.max(tl.abs(t_dx), 0) / 448, 1e-30) if ROUND: t_s = tl.exp2(tl.ceil(tl.log2(t_s))) - t_dx = t_dx/t_s - tl.store(transpose_dx_ptr + M * n + toffs, tl.trans(t_dx.to(transpose_dx_ptr.dtype.element_ty))) - tl.store(transpose_dx_scale_ptr + pid * n * 2 + n + i * B + tl.arange(0, B), t_s) + t_dx = t_dx / t_s + tl.store( + transpose_dx_ptr + M * n + toffs, + tl.trans(t_dx.to(transpose_dx_ptr.dtype.element_ty)), + ) + tl.store( + transpose_dx_scale_ptr + pid * n * 2 + n + i * B + tl.arange(0, B), t_s + ) dx1 = dx1 * smooth_scale_1 dx2 = dx2 * smooth_scale_2 # maxs = tl.maximum( # tl.maximum(dx1.abs(), dx2.abs()), maxs) - maxs = tl.maximum( - tl.maximum(tl.max(dx1.abs(), 1), tl.max(dx2.abs(), 1)), maxs) + maxs = tl.maximum(tl.maximum(tl.max(dx1.abs(), 1), tl.max(dx2.abs(), 1)), maxs) offs += B hoffs += B @@ -940,10 +2436,8 @@ def silu_and_smooth_quant_backward_kernel(g_ptr, x_ptr, tl.store(dx_scale_ptr + pid * T + tl.arange(0, T), scale) s = 1 / scale[:, None] - offs = pid * T * n * 2 + tl.arange(0, T)[:, None] * n * 2 + tl.arange(0, B)[ - None, :] - hoffs = pid * T * n + tl.arange(0, T)[:, None] * n + tl.arange(0, B)[None, - :] + offs = pid * T * n * 2 + tl.arange(0, T)[:, None] * n * 2 + tl.arange(0, B)[None, :] + hoffs = pid * T * n + tl.arange(0, T)[:, None] * n + tl.arange(0, B)[None, :] for i in range(nb): smooth_scale_1 = tl.load(smooth_scale_ptr + i * B + tl.arange(0, B)) smooth_scale_2 = tl.load(smooth_scale_ptr + n + i * B + tl.arange(0, B)) @@ -954,9 +2448,8 @@ def silu_and_smooth_quant_backward_kernel(g_ptr, x_ptr, x1 = tl.load(x_ptr + offs).to(tl.float32) x2 = tl.load(x_ptr + offs + n).to(tl.float32) g = tl.load(g_ptr + hoffs).to(tl.float32) - sigmoid = 1 / (1 + tl.exp(-x1)) - dx1 = g * x2 * sigmoid * ( - 1 + x1 * (1 - sigmoid)) * smooth_scale_1 + sigmoid = tl.sigmoid(x1) + dx1 = g * x2 * sigmoid * (1 + x1 * (1 - sigmoid)) * smooth_scale_1 dx2 = g * x1 * sigmoid * smooth_scale_2 dx1 = (dx1 * s).to(dx_ptr.dtype.element_ty) @@ -967,40 +2460,44 @@ def silu_and_smooth_quant_backward_kernel(g_ptr, x_ptr, offs += B hoffs += B + # requant multi-column quantized tensor @triton.jit -def _requant_kernel(x_ptr, scale_ptr, scales_ptr, - M, - N, - H: tl.constexpr, - W: tl.constexpr - ): +def _requant_kernel( + x_ptr, scale_ptr, scales_ptr, M, N, H: tl.constexpr, W: tl.constexpr +): rid = tl.program_id(axis=0) cid = tl.program_id(axis=1) - offs = rid * H * N + cid * W + tl.arange(0, H)[:, None] * N + tl.arange(0, W)[None, :] + offs = ( + rid * H * N + cid * W + tl.arange(0, H)[:, None] * N + tl.arange(0, W)[None, :] + ) global_scale = tl.load(scale_ptr + rid * H + tl.arange(0, H)) # scales is stored with column-major format local_scale = tl.load(scales_ptr + cid * M + rid * H + tl.arange(0, H)) - x = tl.load(x_ptr+offs).to(tl.float32) - rescale = local_scale/global_scale - x = x * rescale[:,None] - tl.store(x_ptr+offs, x) + x = tl.load(x_ptr + offs).to(tl.float32) + rescale = local_scale / global_scale + x = x * rescale[:, None] + tl.store(x_ptr + offs, x) # used in shared expert -def triton_silu_and_smooth_quant_backward(g, x, - smooth_scale=None, - transpose_smooth_scale=None, - reverse=True, - round_scale=False): +def triton_silu_and_smooth_quant_backward( + g, + x, + smooth_scale=None, + transpose_smooth_scale=None, + reverse=True, + round_scale=False, +): """""" + assert g.is_contiguous() assert round_scale M, N = x.shape n = N // 2 device = x.device dx = torch.empty((M, N), device=device, dtype=torch.float8_e4m3fn) dx_scale = torch.empty((M,), device=device, dtype=torch.float32) - scale_shape = (N, ) + scale_shape = (N,) transpose_dx = torch.empty((N, M), device=device, dtype=torch.float8_e4m3fn) transpose_dx_scale = torch.empty(scale_shape, device=device, dtype=torch.float32) @@ -1025,38 +2522,39 @@ def triton_silu_and_smooth_quant_backward(g, x, reverse, round_scale, num_stages=3, - num_warps=2 + num_warps=2, ) transpose_dx_scale = transpose_dx_scales.amax(0) grid = (N // B, M // T) - _requant_kernel[grid](transpose_dx, transpose_dx_scale, transpose_dx_scales, - N, - M, - B, - T) + _requant_kernel[grid]( + transpose_dx, transpose_dx_scale, transpose_dx_scales, N, M, B, T + ) return dx, dx_scale, transpose_dx, transpose_dx_scale @triton.jit -def batch_weighted_silu_and_smooth_quant_forward_kernel(x_ptr, weight_ptr, - smooth_scale_ptr, - out_ptr, - scale_ptr, max_ptr, - count_ptr, - accum_ptr, M, - n: tl.constexpr, - W: tl.constexpr, - ROUND: tl.constexpr, - REVERSE: tl.constexpr, - CALIBRATE: tl.constexpr): +def batch_weighted_silu_and_smooth_quant_forward_kernel( + x_ptr, + weight_ptr, + smooth_scale_ptr, + out_ptr, + scale_ptr, + count_ptr, + accum_ptr, + M, + n: tl.constexpr, + W: tl.constexpr, + ROUND: tl.constexpr, + REVERSE: tl.constexpr, +): eid = tl.program_id(axis=0) tid = tl.program_id(axis=1) sm = tl.num_programs(axis=1) count = tl.load(count_ptr + eid) ei = tl.load(accum_ptr + eid) - si = ei - count + si = (ei - count).to(tl.int64) c = tl.cdiv(count, sm * W) row_offs = si * n + tid * c * W * n + tl.arange(0, W)[:, None] * n @@ -1065,23 +2563,16 @@ def batch_weighted_silu_and_smooth_quant_forward_kernel(x_ptr, weight_ptr, if not REVERSE: smooth_scale = 1.0 / smooth_scale - if CALIBRATE: - maxs = tl.zeros((W, n), dtype=tl.float32) - for i in range(c): indices = tid * c * W + i * W + tl.arange(0, W) mask = indices[:, None] < count x1 = tl.load(x_ptr + row_offs * 2 + col_offs, mask=mask).to(tl.float32) - x2 = tl.load(x_ptr + n + row_offs * 2 + col_offs, mask=mask).to( - tl.float32) + x2 = tl.load(x_ptr + n + row_offs * 2 + col_offs, mask=mask).to(tl.float32) - w = tl.load(weight_ptr + si + indices, mask=indices < count).to( - tl.float32)[:, - None] - x = x1 / (1 + tl.exp(-x1)) * x2 - - if CALIBRATE: - maxs = tl.maximum(x.abs(), maxs) + w = tl.load(weight_ptr + si + indices, mask=indices < count).to(tl.float32)[ + :, None + ] + x = x1 * tl.sigmoid(x1) * x2 x *= w * smooth_scale scale = tl.maximum(tl.max(x.abs(), 1) / 448, 1e-30) @@ -1092,49 +2583,35 @@ def batch_weighted_silu_and_smooth_quant_forward_kernel(x_ptr, weight_ptr, tl.store(out_ptr + row_offs + col_offs, x, mask=mask) row_offs += n * W - if CALIBRATE: - maxs = tl.max(maxs, 0) - tl.store(max_ptr + eid * sm * n + tid * n + tl.arange(0, n), maxs) - # used in routed experts -def triton_batch_weighted_silu_and_smooth_quant_forward(x, - weight, - counts, - smooth_scale=None, - splits=None, - out=None, - scale=None, - round_scale=False, - reverse=False, - calibrate=False): +def triton_batch_weighted_silu_and_smooth_quant_forward( + x, + weight, + counts, + smooth_scale=None, + splits=None, + out=None, + scale=None, + round_scale=False, + reverse=False, +): """""" + assert x.is_contiguous() and weight.is_contiguous() M, N = x.shape n = N // 2 n_experts = counts.shape[0] - assert N <= 8192 + assert N <= 8192 and triton.next_power_of_2(N) == N device = x.device if out is None: out = torch.empty((M, n), device=device, dtype=torch.float8_e4m3fn) - sm = torch.cuda.get_device_properties(device).multi_processor_count - tmp_maxs = None + sm = 128 if scale is None: scale = torch.empty((M,), device=device, dtype=torch.float32) - if M == 0: - maxs = torch.zeros((n_experts, n), device=device, - dtype=torch.float32) - - elif calibrate: - tmp_maxs = torch.empty((n_experts, sm, n), device=device, - dtype=torch.float32) - maxs = torch.empty((n_experts, n), device=device, - dtype=torch.float32) - else: - maxs = None if M == 0: - return out, scale, maxs + return out, scale accums = torch.cumsum(counts, 0) W = 8192 // N @@ -1145,7 +2622,6 @@ def triton_batch_weighted_silu_and_smooth_quant_forward(x, smooth_scale, out, scale, - tmp_maxs, counts, accums, M, @@ -1153,91 +2629,102 @@ def triton_batch_weighted_silu_and_smooth_quant_forward(x, W, round_scale, reverse, - calibrate, num_stages=3, - num_warps=16 + num_warps=16, ) - if calibrate: - maxs = tmp_maxs.amax(1) - return out, scale, maxs + return out, scale @triton.jit -def batch_weighted_silu_and_smooth_quant_backward_kernel(g_ptr, x_ptr, - weight_ptr, - smooth_scale_ptr, - transpose_smooth_scale_ptr, - count_ptr, - accum_ptr, - dx_ptr, - dx_scale_ptr, - transpose_dx_ptr, - transpose_dx_scale_ptr, - dw_ptr, - n: tl.constexpr, - T: tl.constexpr, - B: tl.constexpr, - E: tl.constexpr, - REVERSE: tl.constexpr, - ROUND: tl.constexpr): +def batch_weighted_silu_and_smooth_quant_backward_kernel( + g_ptr, + x_ptr, + weight_ptr, + smooth_scale_ptr, + transpose_smooth_scale_ptr, + count_ptr, + accum_ptr, + dx_ptr, + dx_scale_ptr, + transpose_dx_ptr, + transpose_dx_scale_ptr, + dw_ptr, + n: tl.constexpr, + T: tl.constexpr, + B: tl.constexpr, + E: tl.constexpr, + REVERSE: tl.constexpr, + ROUND: tl.constexpr, +): eid = tl.program_id(axis=0) pid = tl.program_id(axis=1) max_block = tl.num_programs(axis=1) count = tl.load(count_ptr + eid) round_count = tl.cdiv(count, 32) * 32 - si = tl.load(accum_ptr + eid) - count + si = (tl.load(accum_ptr + eid) - count).to(tl.int64) if pid >= tl.cdiv(count, T): return - round_off = tl.sum(tl.where(tl.arange(0, E) < eid, - tl.cdiv(tl.load(count_ptr + tl.arange(0, E)), - 32), 0)) * 32 - - offs = si * n * 2 + pid * T * n * 2 + tl.arange(0, T)[:, - None] * n * 2 + tl.arange(0, B)[ - None, :] - hoffs = si * n + pid * T * n + tl.arange(0, T)[:, None] * n + tl.arange(0, - B)[ - None, - :] - toffs = round_off * n * 2 + pid * T + tl.arange(0, B)[:, - None] * round_count + tl.arange(0, T)[ - None, :] + counts = tl.load(count_ptr + tl.arange(0, E)) + round_off = tl.sum(tl.where(tl.arange(0, E) < eid, tl.cdiv(counts, 32), 0)) * 32 + + offs = ( + si * n * 2 + + pid * T * n * 2 + + tl.arange(0, T)[:, None] * n * 2 + + tl.arange(0, B)[None, :] + ) + hoffs = ( + si * n + pid * T * n + tl.arange(0, T)[:, None] * n + tl.arange(0, B)[None, :] + ) + toffs = ( + round_off * n * 2 + + pid * T + + tl.arange(0, B)[:, None] * round_count + + tl.arange(0, T)[None, :] + ) nb = n // B maxs = tl.zeros((T,), dtype=tl.float32) indices = pid * T + tl.arange(0, T) if REVERSE: transpose_smooth_scale = tl.load( transpose_smooth_scale_ptr + si + pid * T + tl.arange(0, T), - mask=indices < count)[:, None] + mask=indices < count, + )[:, None] else: - transpose_smooth_scale = 1 / tl.load( - transpose_smooth_scale_ptr + si + pid * T + tl.arange(0, T), - mask=indices < count, other=1e-30)[:, None] + transpose_smooth_scale = ( + 1 + / tl.load( + transpose_smooth_scale_ptr + si + pid * T + tl.arange(0, T), + mask=indices < count, + other=1e-30, + )[:, None] + ) - w = tl.load(weight_ptr + si + pid * T + tl.arange(0, T), - mask=indices < count)[:, None] + w = tl.load(weight_ptr + si + pid * T + tl.arange(0, T), mask=indices < count)[ + :, None + ] dw = tl.zeros((T,), dtype=tl.float32) qdtype = transpose_dx_ptr.dtype.element_ty for i in range(nb): smooth_scale_1 = tl.load( - smooth_scale_ptr + eid * n * 2 + i * B + tl.arange(0, B)) + smooth_scale_ptr + eid * n * 2 + i * B + tl.arange(0, B) + ) smooth_scale_2 = tl.load( - smooth_scale_ptr + eid * n * 2 + n + i * B + tl.arange(0, B)) + smooth_scale_ptr + eid * n * 2 + n + i * B + tl.arange(0, B) + ) if not REVERSE: smooth_scale_1 = 1 / smooth_scale_1 smooth_scale_2 = 1 / smooth_scale_2 x1 = tl.load(x_ptr + offs, mask=indices[:, None] < count).to(tl.float32) - x2 = tl.load(x_ptr + offs + n, mask=indices[:, None] < count).to( - tl.float32) + x2 = tl.load(x_ptr + offs + n, mask=indices[:, None] < count).to(tl.float32) g = tl.load(g_ptr + hoffs, mask=indices[:, None] < count).to(tl.float32) - sigmoid = 1 / (1 + tl.exp(-x1)) - dx1 = g * x2 * sigmoid * ( - 1 + x1 * (1 - sigmoid)) * w + sigmoid = tl.sigmoid(x1) + dx1 = g * x2 * sigmoid * (1 + x1 * (1 - sigmoid)) * w dx2 = g * x1 * sigmoid * w dw += tl.sum(x1 * sigmoid * x2 * g, 1) @@ -1247,27 +2734,43 @@ def batch_weighted_silu_and_smooth_quant_backward_kernel(g_ptr, x_ptr, if ROUND: t_s = tl.exp2(tl.ceil(tl.log2(t_s))) t_dx = t_dx / t_s - tl.store(transpose_dx_ptr + toffs, tl.trans(t_dx.to(qdtype)), - mask=indices[None, :] < round_count) tl.store( - transpose_dx_scale_ptr + eid * max_block * n * 2 + pid * n * 2 + i * B + tl.arange( - 0, B), t_s) + transpose_dx_ptr + toffs, + tl.trans(t_dx.to(qdtype)), + mask=indices[None, :] < round_count, + ) + tl.store( + transpose_dx_scale_ptr + + eid * max_block * n * 2 + + pid * n * 2 + + i * B + + tl.arange(0, B), + t_s, + ) t_dx = dx2 * transpose_smooth_scale t_s = tl.maximum(tl.max(tl.abs(t_dx), 0) / 448, 1e-30) if ROUND: t_s = tl.exp2(tl.ceil(tl.log2(t_s))) t_dx = t_dx / t_s - tl.store(transpose_dx_ptr + round_count * n + toffs, - tl.trans(t_dx.to(qdtype)), mask=indices[None, :] < round_count) tl.store( - transpose_dx_scale_ptr + eid * max_block * n * 2 + pid * n * 2 + n + i * B + tl.arange( - 0, B), t_s) + transpose_dx_ptr + round_count * n + toffs, + tl.trans(t_dx.to(qdtype)), + mask=indices[None, :] < round_count, + ) + tl.store( + transpose_dx_scale_ptr + + eid * max_block * n * 2 + + pid * n * 2 + + n + + i * B + + tl.arange(0, B), + t_s, + ) dx1 = dx1 * smooth_scale_1 dx2 = dx2 * smooth_scale_2 - maxs = tl.maximum( - tl.maximum(tl.max(dx1.abs(), 1), tl.max(dx2.abs(), 1)), maxs) + maxs = tl.maximum(tl.maximum(tl.max(dx1.abs(), 1), tl.max(dx2.abs(), 1)), maxs) offs += B hoffs += B @@ -1277,33 +2780,34 @@ def batch_weighted_silu_and_smooth_quant_backward_kernel(g_ptr, x_ptr, scale = tl.maximum(maxs / 448, 1e-30) if ROUND: scale = tl.exp2(tl.ceil(tl.log2(scale))) - tl.store(dx_scale_ptr + si + pid * T + tl.arange(0, T), scale, - mask=indices < count) + tl.store(dx_scale_ptr + si + pid * T + tl.arange(0, T), scale, mask=indices < count) s = 1 / scale[:, None] - offs = si * n * 2 + pid * T * n * 2 + tl.arange(0, T)[:, - None] * n * 2 + tl.arange(0, B)[ - None, :] - hoffs = si * n + pid * T * n + tl.arange(0, T)[:, None] * n + tl.arange(0, - B)[ - None, - :] + offs = ( + si * n * 2 + + pid * T * n * 2 + + tl.arange(0, T)[:, None] * n * 2 + + tl.arange(0, B)[None, :] + ) + hoffs = ( + si * n + pid * T * n + tl.arange(0, T)[:, None] * n + tl.arange(0, B)[None, :] + ) for i in range(nb): smooth_scale_1 = tl.load( - smooth_scale_ptr + eid * n * 2 + i * B + tl.arange(0, B)) + smooth_scale_ptr + eid * n * 2 + i * B + tl.arange(0, B) + ) smooth_scale_2 = tl.load( - smooth_scale_ptr + eid * n * 2 + n + i * B + tl.arange(0, B)) + smooth_scale_ptr + eid * n * 2 + n + i * B + tl.arange(0, B) + ) if not REVERSE: smooth_scale_1 = 1 / smooth_scale_1 smooth_scale_2 = 1 / smooth_scale_2 x1 = tl.load(x_ptr + offs, mask=indices[:, None] < count).to(tl.float32) - x2 = tl.load(x_ptr + offs + n, mask=indices[:, None] < count).to( - tl.float32) + x2 = tl.load(x_ptr + offs + n, mask=indices[:, None] < count).to(tl.float32) g = tl.load(g_ptr + hoffs, mask=indices[:, None] < count).to(tl.float32) - sigmoid = 1 / (1 + tl.exp(-x1)) - dx1 = g * x2 * sigmoid * ( - 1 + x1 * (1 - sigmoid)) * smooth_scale_1 * w + sigmoid = tl.sigmoid(x1) + dx1 = g * x2 * sigmoid * (1 + x1 * (1 - sigmoid)) * smooth_scale_1 * w dx2 = g * x1 * sigmoid * smooth_scale_2 * w dx1 = (dx1 * s).to(dx_ptr.dtype.element_ty) @@ -1317,13 +2821,16 @@ def batch_weighted_silu_and_smooth_quant_backward_kernel(g_ptr, x_ptr, # requant multi-column quantized tensor @triton.jit -def _batch_requant_kernel(x_ptr, scale_ptr, scales_ptr, - count_ptr, - N, - H: tl.constexpr, - W: tl.constexpr, - E: tl.constexpr - ): +def _batch_requant_kernel( + x_ptr, + scale_ptr, + scales_ptr, + count_ptr, + N, + H: tl.constexpr, + W: tl.constexpr, + E: tl.constexpr, +): eid = tl.program_id(axis=0) rid = tl.program_id(axis=1) cid = tl.program_id(axis=2) @@ -1334,17 +2841,21 @@ def _batch_requant_kernel(x_ptr, scale_ptr, scales_ptr, if cid >= tl.cdiv(round_count, W): return - round_off = tl.sum(tl.where(tl.arange(0, E) < eid, - tl.cdiv(tl.load(count_ptr + tl.arange(0, E)), - 32) * 32, 0)) - - offs = round_off * N + rid * H * round_count + cid * W + tl.arange(0, H)[:, - None] * round_count + tl.arange( - 0, W)[None, :] + counts = tl.load(count_ptr + tl.arange(0, E)) + round_off = tl.sum(tl.where(tl.arange(0, E) < eid, tl.cdiv(counts, 32) * 32, 0)) + + offs = ( + round_off * N + + rid * H * round_count + + cid * W + + tl.arange(0, H)[:, None] * round_count + + tl.arange(0, W)[None, :] + ) global_scale = tl.load(scale_ptr + eid * N + rid * H + tl.arange(0, H)) # scales is stored with column-major format local_scale = tl.load( - scales_ptr + max_block * N * eid + cid * N + rid * H + tl.arange(0, H)) + scales_ptr + max_block * N * eid + cid * N + rid * H + tl.arange(0, H) + ) x = tl.load(x_ptr + offs).to(tl.float32) rescale = local_scale / tl.maximum(global_scale, 1e-30) x = x * rescale[:, None] @@ -1352,20 +2863,25 @@ def _batch_requant_kernel(x_ptr, scale_ptr, scales_ptr, # used in routed experts -def triton_batch_weighted_silu_and_smooth_quant_backward(g, x, weight, - counts, - smooth_scale=None, - transpose_smooth_scale=None, - splits=None, - reverse=True, - round_scale=False): +def triton_batch_weighted_silu_and_smooth_quant_backward( + g, + x, + weight, + counts, + smooth_scale=None, + transpose_smooth_scale=None, + splits=None, + reverse=True, + round_scale=False, +): """""" + assert g.is_contiguous() assert round_scale M, N = x.shape n = N // 2 n_expert = counts.shape[0] assert N <= 8192 and 8192 % N == 0 - assert splits is not None, 'batch mode need splits to launch kernels' + assert splits is not None, "batch mode need splits to launch kernels" device = x.device @@ -1381,16 +2897,17 @@ def triton_batch_weighted_silu_and_smooth_quant_backward(g, x, weight, assert n % B == 0 and T == 32 max_block = triton.cdiv(max(splits), T) s = sum([(x + 31) // 32 for x in splits]) * 32 - transpose_dx = torch.empty((N * s,), device=device, - dtype=torch.float8_e4m3fn) + transpose_dx = torch.empty((N * s,), device=device, dtype=torch.float8_e4m3fn) if s == 0: - transpose_dx_scale = torch.zeros((n_expert, N), device=device, - dtype=torch.float32) + transpose_dx_scale = torch.zeros( + (n_expert, N), device=device, dtype=torch.float32 + ) return dx, dx_scale, dw, transpose_dx, transpose_dx_scale else: - transpose_dx_scales = torch.zeros((n_expert, max_block, N), - device=device, dtype=torch.bfloat16) + transpose_dx_scales = torch.zeros( + (n_expert, max_block, N), device=device, dtype=torch.bfloat16 + ) grid = (n_expert, max_block) batch_weighted_silu_and_smooth_quant_backward_kernel[grid]( @@ -1413,19 +2930,21 @@ def triton_batch_weighted_silu_and_smooth_quant_backward(g, x, weight, reverse, round_scale, num_stages=5, - num_warps=4 + num_warps=4, ) transpose_dx_scale = transpose_dx_scales.amax(1).float() grid = (n_expert, N // B, max_block) - _batch_requant_kernel[grid](transpose_dx, transpose_dx_scale, - transpose_dx_scales, - counts, - N, - B, - T, - n_expert, - num_stages=3, - num_warps=2) + _batch_requant_kernel[grid]( + transpose_dx, + transpose_dx_scale, + transpose_dx_scales, + counts, + N, + B, + T, + n_expert, + num_stages=3, + num_warps=2, + ) return dx, dx_scale, dw, transpose_dx, transpose_dx_scale - diff --git a/linghe/utils/topk.py b/linghe/utils/topk.py new file mode 100644 index 0000000..1c32262 --- /dev/null +++ b/linghe/utils/topk.py @@ -0,0 +1,410 @@ +# -*- coding: utf-8 -*- +""" +Copyright (c) Ant Financial Service Group and its affiliates. +""" + +import torch +import triton +import triton.language as tl + + +@triton.jit +def topk_forward_kernel( + input_ptr, value_ptr, index_ptr, N: tl.constexpr, K: tl.constexpr +): + pid = tl.program_id(axis=0) + + xo = tl.load(input_ptr + pid * N + tl.arange(0, N)).to(tl.float32) + + x = xo + for i in range(K): + val = tl.max(x, 0) + idx = tl.argmax(x, 0) + tl.store(value_ptr + pid * K + i, val) + tl.store(index_ptr + pid * K + i, idx) + x = tl.where(x == val, -2e38, x) + + if tl.sum(tl.where(x < -1e38, 1, 0)) > K: + y = xo.to(tl.float64) - tl.arange(0, N).to(tl.float64) * 1e-12 + for i in range(K): + val = tl.max(y, 0) + idx = tl.argmax(y, 0) + tl.store(value_ptr + pid * K + i, val) + tl.store(index_ptr + pid * K + i, idx) + y = tl.where(y == val, -2e38, y) + + +@triton.jit +def unsorted_topk_forward_kernel( + input_ptr, value_ptr, index_ptr, N: tl.constexpr, K: tl.constexpr +): + pid = tl.program_id(axis=0) + + x = tl.load(input_ptr + pid * N + tl.arange(0, N)).to(tl.float32) + + xt = tl.topk(x, K, dim=0) + min_xt = tl.min(xt) + mask = x >= min_xt + + if tl.sum(mask) > K: + y = x.to(tl.float64) * (1 - tl.arange(0, N).to(tl.float64) * 1e-12) + yt = tl.topk(y, K, dim=0) + min_yt = tl.min(yt) + masks = y >= min_yt + else: + masks = mask + + acc = tl.cumsum(masks, 0) - 1 + + tl.store(value_ptr + pid * K + acc, x, mask=masks) + tl.store(index_ptr + pid * K + acc, tl.arange(0, N), mask=masks) + + +@triton.jit +def sorted_topk_forward_kernel( + input_ptr, value_ptr, index_ptr, N: tl.constexpr, K: tl.constexpr +): + pid = tl.program_id(axis=0) + + x = ( + tl.load(input_ptr + pid * N + tl.arange(0, N)) + .to(tl.float32) + .to(tl.uint32, bitcast=True) + ) + x = tl.where(x >= 2**31, x - 2**31, x + 2**31) + + x = x.to(tl.uint64) + + x = (x << 32) + tl.arange(0, N).to(tl.uint64) + + xt = tl.topk(x, K, dim=0) + + value = (xt >> 32).to(tl.uint32) + value = tl.where(value >= 2**31, value - 2**31, value + 2**31) + value = value.to(tl.float32, bitcast=True) + index = (xt % (2**32)).to(tl.int32) + + tl.store(value_ptr + pid * K + tl.arange(0, K), value) + tl.store(index_ptr + pid * K + tl.arange(0, K), index) + + +def triton_topk_forward(x, k, dim=-1, sorted=True, impl="iter"): + """ + calculate topk. + Args: + x: input tensor. + k: topk + Returns: + values: topk values + indices: topk indices + """ + device = x.device + shape = x.shape + assert dim == -1 + assert x.is_contiguous() and x.dtype in ( + torch.float32, + torch.bfloat16, + torch.float16, + ) + N = shape[-1] + g = x.numel() // N + values = torch.empty(shape[:-1] + (k,), device=device, dtype=x.dtype) + indices = torch.empty(shape[:-1] + (k,), device=device, dtype=torch.int64) + grid = (g,) + + if impl == "iter": + topk_forward_kernel[grid](x, values, indices, N, k, num_stages=2, num_warps=1) + else: + if sorted: + sorted_topk_forward_kernel[grid]( + x, values, indices, N, k, num_stages=2, num_warps=1 + ) + else: + unsorted_topk_forward_kernel[grid]( + x, values, indices, N, k, num_stages=2, num_warps=1 + ) + + return values, indices + + +@triton.jit +def topk_backward_kernel(grad_ptr, index_ptr, dx_ptr, N: tl.constexpr, K: tl.constexpr): + pid = tl.program_id(axis=0) + + grad = tl.load(grad_ptr + pid * K + tl.arange(0, K)) + index = tl.load(index_ptr + pid * K + tl.arange(0, K)) + tl.store(dx_ptr + pid * N + index, grad) + + +def triton_topk_backward(grad_output, indices, N, dim=-1): + """ + topk backward. + Args: + grad_output: grad tensor of values. + indices: topk indices + N: dim + Returns: + dx + """ + device = grad_output.device + shape = grad_output.shape + assert dim == -1 + assert grad_output.is_contiguous() + k = shape[-1] + g = grad_output.numel() // k + dx = torch.zeros(shape[:-1] + (N,), device=device, dtype=grad_output.dtype) + grid = (g,) + topk_backward_kernel[grid]( + grad_output, indices, dx, N, k, num_stages=2, num_warps=2 + ) + return dx + + +@triton.jit +def deprecated_group_topk_score_forward_kernel( + input_ptr, + bias_ptr, + prob_ptr, + map_ptr, + scale, + eps, + N: tl.constexpr, + K: tl.constexpr, + G: tl.constexpr, + GK: tl.constexpr, + BIAS: tl.constexpr, +): + pid = tl.program_id(axis=0) + GS: tl.constexpr = N // G + k: tl.constexpr = K // GK + + logit = tl.load(input_ptr + pid * N + tl.arange(0, N)) + x = tl.sigmoid(logit) + if BIAS: + b = tl.load(bias_ptr + tl.arange(0, N)) + else: + b = 0.0 + xb = tl.reshape(x + b, (G, GS)) + xbsort = tl.sort(xb, dim=1, descending=True) + array = tl.arange(0, GS) + xbsum = tl.sum(tl.where(array < k, xbsort, 0.0), 1) + xbsumsort = tl.sort(xbsum, dim=0, descending=True) + + arr = tl.arange(0, G) + group_min_value = tl.min(tl.where(arr < GK, xbsumsort, 2e38)) + + xb_group_mask = tl.where(xbsum[:, None] >= group_min_value, xb, -1e38) + xb_group_mask = tl.reshape(xb_group_mask, (N,)) + x_group_mask_sort = tl.sort(xb_group_mask, dim=0, descending=True) + expert_array = tl.arange(0, N) + min_value = tl.min(tl.where(expert_array < K, x_group_mask_sort, 1e38)) + score = tl.where(xb_group_mask >= min_value, x, 0) + + score = score / (tl.sum(score) + eps) * scale + map_idx = tl.where(xb_group_mask >= min_value, 1, 0) + + if tl.sum(map_idx) > K: + y = x.to(tl.float64) + b.to(tl.float64) - tl.arange(0, N).to(tl.float64) * 1e-12 + yb = tl.reshape(y, (G, GS)) + ybsort = tl.sort(yb, dim=1, descending=True) + ysortmask = tl.where(array < k, ybsort, 0) + + ybsum = tl.sum(ysortmask, 1) + ybsumsort = tl.sort(ybsum, dim=0, descending=True) + + yb_group_min_value = tl.min(tl.where(arr < GK, ybsumsort, 2e38)) + + y_group_mask = tl.where(ybsum[:, None] >= yb_group_min_value, yb, -1e38) + y_group_mask = tl.reshape(y_group_mask, (N,)) + y_group_mask_sort = tl.sort(y_group_mask, dim=0, descending=True) + y_min_value = tl.min(tl.where(expert_array < K, y_group_mask_sort, 1e38)) + double_score = tl.where(y_group_mask >= y_min_value, y, 0) + + double_score = double_score / (tl.sum(double_score) + eps) * scale + + tl.store(prob_ptr + pid * N + tl.arange(0, N), double_score) + tl.store( + map_ptr + pid * N + tl.arange(0, N), + tl.where(y_group_mask >= y_min_value, 1, 0), + ) + else: + tl.store(prob_ptr + pid * N + tl.arange(0, N), score) + tl.store(map_ptr + pid * N + tl.arange(0, N), map_idx) + + +@triton.jit +def group_topk_score_forward_kernel( + input_ptr, + bias_ptr, + prob_ptr, + map_ptr, + scale, + eps, + N: tl.constexpr, + K: tl.constexpr, + G: tl.constexpr, + GK: tl.constexpr, + BIAS: tl.constexpr, +): + pid = tl.program_id(axis=0) + GS: tl.constexpr = N // G + k: tl.constexpr = K // GK + + logit = tl.load(input_ptr + pid * N + tl.arange(0, N)) + x = tl.sigmoid(logit) + if BIAS: + b = tl.load(bias_ptr + tl.arange(0, N)) + else: + b = 0.0 + m = tl.reshape(x + b, (G, GS)) + + gt = tl.topk(m, k, dim=1) + + gts = tl.sum(gt, 1) + gtst = tl.topk(gts, GK, dim=0) + sum_min_value = tl.min(gtst) + + group_filling = tl.where((gts[:, None] >= sum_min_value), m, -1.0) + group_filling = tl.reshape(group_filling, [N]) + t = tl.min(tl.topk(group_filling, K, dim=0)) + mask = group_filling >= t + + if tl.sum(mask) > K: + group_fillings = ( + group_filling.to(tl.float64) - tl.arange(0, N).to(tl.float64) * 1e-12 + ) + ts = tl.min(tl.topk(group_fillings, K, dim=0)) + masks = group_fillings >= ts + else: + masks = mask + + filling = tl.where(mask, x, 0.0) + score = filling / (tl.sum(filling) + eps) + score = score * scale + tl.store(prob_ptr + pid * N + tl.arange(0, N), score) + tl.store(map_ptr + pid * N + tl.arange(0, N), masks) + + +def triton_group_topk_score_forward( + x, + k, + expert_bias=None, + num_groups=32, + group_topk=4, + scaling_factor=1.0, + score_function="sigmoid", + eps=1e-20, +): + """ + calculate topk. + Args: + x: input tensor. + expert_bias: expert bias + k: topk + Returns: + probs: + routing_map: + tokens_per_expert: + """ + device = x.device + shape = x.shape + assert len(shape) <= 3 and x.is_contiguous() and score_function == "sigmoid" + if len(shape) == 3: + M, B, N = shape + g = M * B + probs = torch.empty((M, B, N), device=device, dtype=x.dtype) + routing_map = torch.empty((M, B, N), device=device, dtype=torch.bool) + else: + M, N = shape + g = M + probs = torch.empty((M, N), device=device, dtype=x.dtype) + routing_map = torch.empty((M, N), device=device, dtype=torch.bool) + BIAS = expert_bias is not None + grid = (g,) + if hasattr(tl, "topk"): + group_topk_score_forward_kernel[grid]( + x, + expert_bias, + probs, + routing_map, + scaling_factor, + eps, + N, + k, + num_groups, + group_topk, + BIAS, + num_stages=2, + num_warps=1, + ) + else: + deprecated_group_topk_score_forward_kernel[grid]( + x, + expert_bias, + probs, + routing_map, + scaling_factor, + eps, + N, + k, + num_groups, + group_topk, + BIAS, + num_stages=2, + num_warps=1, + ) + return probs, routing_map, routing_map.sum(0) + + +@triton.jit +def group_topk_score_backward_kernel( + grad_ptr, input_ptr, map_ptr, dx_ptr, scale, eps, N: tl.constexpr +): + pid = tl.program_id(axis=0) + grad = tl.load(grad_ptr + pid * N + tl.arange(0, N)) + logit = tl.load(input_ptr + pid * N + tl.arange(0, N)) + mask = tl.load(map_ptr + pid * N + tl.arange(0, N)).to(tl.float32) + + s = tl.sigmoid(logit) + z = tl.sum(s * mask) + eps + dx = scale * mask * s * (1 - s) / z * (grad - tl.sum(s * grad * mask) / z) + tl.store(dx_ptr + pid * N + tl.arange(0, N), dx) + + +def triton_group_topk_score_backward( + grad_output, input, routing_map, scaling_factor=1.0, eps=1e-20 +): + """ + topk backward. + Args: + grad_output: grad tensor of prob. + routing_map: topk indices + Returns: + dx: grad of logits + """ + device = grad_output.device + shape = grad_output.shape + assert ( + len(shape) <= 3 and grad_output.is_contiguous() and routing_map.is_contiguous() + ) + if len(shape) == 3: + M, B, N = shape + g = M * B + dx = torch.empty((M, B, N), device=device, dtype=grad_output.dtype) + else: + M, N = shape + g = M + dx = torch.empty((M, N), device=device, dtype=grad_output.dtype) + grid = (g,) + group_topk_score_backward_kernel[grid]( + grad_output, + input, + routing_map, + dx, + scaling_factor, + eps, + N, + num_stages=2, + num_warps=1, + ) + return dx diff --git a/linghe/utils/transpose.py b/linghe/utils/transpose.py index bfc1fcf..81ae4e9 100644 --- a/linghe/utils/transpose.py +++ b/linghe/utils/transpose.py @@ -4,7 +4,7 @@ """ import itertools -from typing import Optional + import torch import triton import triton.language as tl @@ -17,32 +17,35 @@ @triton.jit -def transpose_kernel(x_ptr, t_ptr, M, N, H: tl.constexpr, W: tl.constexpr, - EVEN: tl.constexpr): +def transpose_kernel( + x_ptr, t_ptr, M, N, H: tl.constexpr, W: tl.constexpr, EVEN: tl.constexpr +): rid = tl.program_id(axis=0) cid = tl.program_id(axis=1) - offs = rid * H * N + cid * W + tl.arange(0, H)[:, None] * N + tl.arange(0, - W)[ - None, :] - toffs = rid * H + cid * M * W + tl.arange(0, W)[:, None] * M + tl.arange(0, - H)[ - None, :] + offs = ( + rid * H * N + cid * W + tl.arange(0, H)[:, None] * N + tl.arange(0, W)[None, :] + ) + toffs = ( + rid * H + cid * M * W + tl.arange(0, W)[:, None] * M + tl.arange(0, H)[None, :] + ) if EVEN: y = tl.trans(tl.load(x_ptr + offs)) tl.store(t_ptr + toffs, y) else: - y = tl.trans(tl.load(x_ptr + offs, - mask=(cid * W + tl.arange(0, W)[None, :] < N) & ( - rid * H + tl.arange(0, H)[:, - None] < M))) - tl.store(t_ptr + toffs, y, - mask=(cid * W + tl.arange(0, W)[:, None] < N) & ( - rid * H + tl.arange(0, H)[None, :] < M)) + mask = (cid * W + tl.arange(0, W)[None, :] < N) & ( + rid * H + tl.arange(0, H)[:, None] < M + ) + y = tl.trans(tl.load(x_ptr + offs, mask=mask)) + mask = (cid * W + tl.arange(0, W)[:, None] < N) & ( + rid * H + tl.arange(0, H)[None, :] < M + ) + tl.store(t_ptr + toffs, y, mask=mask) @triton.jit -def transpose_dim_0_1_kernel(x_ptr, t_ptr, B, M, b_stride, m_stride, - N: tl.constexpr): +def transpose_inner_dims_kernel( + x_ptr, t_ptr, B, M, b_stride, m_stride, N: tl.constexpr +): rid = tl.program_id(axis=0) cid = tl.program_id(axis=1) offs = rid * b_stride + cid * m_stride + tl.arange(0, N) @@ -51,15 +54,47 @@ def transpose_dim_0_1_kernel(x_ptr, t_ptr, B, M, b_stride, m_stride, tl.store(t_ptr + toffs, y) -def triton_transpose(x: torch.Tensor, - dim0: Optional[int] = None, - dim1: Optional[int] = None): +@triton.jit +def transpose_outer_dims_kernel( + x_ptr, t_ptr, M, N, H: tl.constexpr, W: tl.constexpr, EVEN: tl.constexpr +): + bid = tl.program_id(axis=0) + rid = tl.program_id(axis=1) + cid = tl.program_id(axis=2) + offs = ( + bid * M * N + + rid * H * N + + cid * W + + tl.arange(0, H)[:, None] * N + + tl.arange(0, W)[None, :] + ) + toffs = ( + bid * M * N + + rid * H + + cid * M * W + + tl.arange(0, W)[:, None] * M + + tl.arange(0, H)[None, :] + ) + if EVEN: + y = tl.trans(tl.load(x_ptr + offs)) + tl.store(t_ptr + toffs, y) + else: + mask = (cid * W + tl.arange(0, W)[None, :] < N) & ( + rid * H + tl.arange(0, H)[:, None] < M + ) + y = tl.trans(tl.load(x_ptr + offs, mask=mask)) + mask = (cid * W + tl.arange(0, W)[:, None] < N) & ( + rid * H + tl.arange(0, H)[None, :] < M + ) + tl.store(t_ptr + toffs, y, mask=mask) + + +def triton_transpose(x: torch.Tensor, inner=True): """ transpose x with dim0 and dim1 Args: x: input tensor - dim0: dim 0 - dim1: dim 1 + inner: inner dim if True, outer dim if False Returns: transposed tensor @@ -79,20 +114,14 @@ def triton_transpose(x: torch.Tensor, grid = (triton.cdiv(M, H), triton.cdiv(N, W)) transpose_kernel[grid]( - x, t, - M, N, - H, W, - EVEN, - num_stages=num_stages, - num_warps=num_warps + x, t, M, N, H, W, EVEN, num_stages=num_stages, num_warps=num_warps ) - elif dim0 == 0 and dim1 == 1: + elif inner: stride = x.stride() if rank == 4: B, M, N = shape[0], shape[1], shape[2] * shape[3] - assert stride[2] == shape[3], 'must be contiguous in last two dims' - t = torch.empty((M, B, shape[2], shape[3]), device=x.device, - dtype=x.dtype) + assert stride[2] == shape[3], "must be contiguous in last two dims" + t = torch.empty((M, B, shape[2], shape[3]), device=x.device, dtype=x.dtype) else: B, M, N = shape t = torch.empty((M, B, N), device=x.device, dtype=x.dtype) @@ -101,51 +130,63 @@ def triton_transpose(x: torch.Tensor, num_stages = 5 num_warps = 2 grid = (B, M) - transpose_dim_0_1_kernel[grid](x, - t, - B, - M, - b_stride, - m_stride, - N, - num_stages=num_stages, - num_warps=num_warps - ) + transpose_inner_dims_kernel[grid]( + x, + t, + B, + M, + b_stride, + m_stride, + N, + num_stages=num_stages, + num_warps=num_warps, + ) else: - raise NotImplementedError() + if rank == 4: + B, M, N = shape[0] * shape[1], shape[2], shape[3] + t = torch.empty((shape[0], shape[1], N, M), device=x.device, dtype=x.dtype) + else: + B, M, N = shape + t = torch.empty((B, N, M), device=x.device, dtype=x.dtype) + + H = 64 + W = 32 if x.dtype.itemsize == 1 else 16 + EVEN = M % H == 0 and N % W == 0 + num_stages = 5 + num_warps = 2 + + grid = (B, triton.cdiv(M, H), triton.cdiv(N, W)) + transpose_outer_dims_kernel[grid]( + x, t, M, N, H, W, EVEN, num_stages=num_stages, num_warps=num_warps + ) return t @triton.jit -def transpose_and_pad_kernel(x_ptr, t_ptr, - M, N, P, - H: tl.constexpr, - W: tl.constexpr, - EVEN: tl.constexpr): +def pad_transpose_kernel( + x_ptr, t_ptr, M, N, P, H: tl.constexpr, W: tl.constexpr, EVEN: tl.constexpr +): rid = tl.program_id(axis=0) cid = tl.program_id(axis=1) - offs = rid * H * N + cid * W + tl.arange(0, H)[:, None] * N + tl.arange(0, - W)[ - None, :] - toffs = rid * H + cid * P * W + tl.arange(0, W)[:, None] * P + tl.arange(0, - H)[ - None, :] + offs = ( + rid * H * N + cid * W + tl.arange(0, H)[:, None] * N + tl.arange(0, W)[None, :] + ) + toffs = ( + rid * H + cid * P * W + tl.arange(0, W)[:, None] * P + tl.arange(0, H)[None, :] + ) if EVEN: y = tl.load(x_ptr + offs) else: - y = tl.load(x_ptr + offs, - mask=(rid * H + tl.arange(0, H)[:, None] < M)) + y = tl.load(x_ptr + offs, mask=(rid * H + tl.arange(0, H)[:, None] < M)) y = tl.trans(y) if EVEN: tl.store(t_ptr + toffs, y) else: # paddings are filled with 0 - tl.store(t_ptr + toffs, y, - mask=(rid * H + tl.arange(0, H)[None, :] < P)) - + tl.store(t_ptr + toffs, y, mask=(rid * H + tl.arange(0, H)[None, :] < P)) -def triton_transpose_and_pad(x, out=None, pad=True): +def triton_pad_transpose(x, out=None, multiple=32): """ transpose x and padding the column size to be mutiplier of 32, it is used for calculated gradient of weight with torch._scaled__mm @@ -158,8 +199,9 @@ def triton_transpose_and_pad(x, out=None, pad=True): out: output tensor """ # fat block, shape:[H,W] + assert x.is_contiguous() M, N = x.shape - P = round_up(M, b=32) if pad else M + P = round_up(M, b=multiple) device = x.device if out is None: out = torch.empty((N, P), device=device, dtype=x.dtype) @@ -171,26 +213,24 @@ def triton_transpose_and_pad(x, out=None, pad=True): assert N % W == 0 EVEN = M % H == 0 and M == P grid = (triton.cdiv(P, H), triton.cdiv(N, W)) - transpose_and_pad_kernel[grid]( - x, out, - M, N, P, - H, W, - EVEN, - num_stages=num_stages, - num_warps=num_warps + pad_transpose_kernel[grid]( + x, out, M, N, P, H, W, EVEN, num_stages=num_stages, num_warps=num_warps ) return out @triton.jit -def batch_transpose_kernel(xs_ptr, xts_ptr, M, N, H: tl.constexpr, - W: tl.constexpr): +def batch_transpose_kernel(xs_ptr, xts_ptr, M, N, H: tl.constexpr, W: tl.constexpr): eid = tl.program_id(axis=0) cid = tl.program_id(axis=1) x_ptr = tl.load(xs_ptr + eid).to(tl.pointer_type(xts_ptr.dtype.element_ty)) offs = cid * W + tl.arange(0, H)[:, None] * N + tl.arange(0, W)[None, :] - toffs = eid * M * N + cid * W * M + tl.arange(0, W)[:, - None] * M + tl.arange(0, H)[None, :] + toffs = ( + eid * M * N + + cid * W * M + + tl.arange(0, W)[:, None] * M + + tl.arange(0, H)[None, :] + ) for i in range(0, M, H): y = tl.trans(tl.load(x_ptr + offs)) tl.store(xts_ptr + toffs, y) @@ -206,36 +246,39 @@ def triton_batch_transpose(xs, xts=None): Returns: xts: output tensor list, [N,M]*expert """ + assert all([x.is_contiguous() for x in xs]) M, N = xs[0].shape n_experts = len(xs) + device = xs[0].device if xts is None: - xts = torch.empty((M * n_experts, N), device=xs[0].device, - dtype=xs[0].dtype) - pointers = torch.tensor([x.data_ptr() for x in xs], device=xs[0].device) - H = 64 + xts = torch.empty((M * n_experts, N), device=device, dtype=xs[0].dtype) + pointers = torch.tensor([x.data_ptr() for x in xs], dtype=torch.int64).cuda( + device, non_blocking=True + ) + + H = 32 W = 64 - num_stages = 3 + num_stages = 2 num_warps = 8 grid = (n_experts, N // W) batch_transpose_kernel[grid]( - pointers, xts, - M, N, - H, W, - num_stages=num_stages, - num_warps=num_warps + pointers, xts, M, N, H, W, num_stages=num_stages, num_warps=num_warps ) - # outputs = torch.split(xts, n_experts) # very slow outputs = torch.split(xts, [M] * n_experts) - # outputs = [] - # for i in range(n_experts): - # outputs.append(xts[i*M:(i+1)*M]) return outputs @triton.jit -def batch_transpose_and_pad_kernel(x_ptr, t_ptr, count_ptr, accum_ptr, - pad_accum_ptr, N, H: tl.constexpr, - W: tl.constexpr): +def batch_pad_transpose_kernel( + x_ptr, + t_ptr, + count_ptr, + accum_ptr, + pad_accum_ptr, + N, + H: tl.constexpr, + W: tl.constexpr, +): eid = tl.program_id(axis=0) cid = tl.program_id(axis=1) count = tl.load(count_ptr + eid) @@ -243,20 +286,22 @@ def batch_transpose_and_pad_kernel(x_ptr, t_ptr, count_ptr, accum_ptr, si = ei - count pad_si = tl.load(pad_accum_ptr + eid) P = tl.cdiv(count, 32) * 32 - offs = si * N + cid * W + tl.arange(0, H)[:, None] * N + tl.arange(0, W)[ - None, :] - toffs = pad_si * N + cid * W * P + tl.arange(0, W)[:, None] * P + tl.arange( - 0, H)[None, :] + offs = si * N + cid * W + tl.arange(0, H)[:, None] * N + tl.arange(0, W)[None, :] + toffs = ( + pad_si * N + + cid * W * P + + tl.arange(0, W)[:, None] * P + + tl.arange(0, H)[None, :] + ) for i in range(0, P, H): - y = tl.trans( - tl.load(x_ptr + offs, mask=i + tl.arange(0, H)[:, None] < count)) + y = tl.trans(tl.load(x_ptr + offs, mask=i + tl.arange(0, H)[:, None] < count)) # paddings are filled with 0 tl.store(t_ptr + toffs, y, mask=i + tl.arange(0, H)[None, :] < P) offs += N * H toffs += H -def triton_batch_transpose_and_pad(x, count_list, x_t=None, pad=True): +def triton_batch_pad_transpose(x, count_list, x_t=None, multiple=32): """ transpose and pad each tensor stored in x Args: @@ -268,16 +313,18 @@ def triton_batch_transpose_and_pad(x, count_list, x_t=None, pad=True): Returns: x_t: output tensor """ - assert pad + assert x.is_contiguous() # block shape:[H,W] M, N = x.shape n_experts = len(count_list) # NOTE: b must be 32, or kernel will miscalculated padding size with original size - pad_sizes = [round_up(x, b=32) for x in count_list] + pad_sizes = [round_up(x, b=multiple) for x in count_list] counts = torch.tensor(count_list, dtype=torch.int32, device=x.device) pad_accum_sizes = torch.tensor( - list(itertools.accumulate(pad_sizes, initial=0)), dtype=torch.int32, - device=x.device) + list(itertools.accumulate(pad_sizes, initial=0)), + dtype=torch.int32, + device=x.device, + ) accums = torch.cumsum(counts, 0) device = x.device if x_t is None: @@ -287,14 +334,17 @@ def triton_batch_transpose_and_pad(x, count_list, x_t=None, pad=True): num_stages = 2 num_warps = 8 grid = (n_experts, N // W) - batch_transpose_and_pad_kernel[grid]( - x, x_t, - counts, accums, + batch_pad_transpose_kernel[grid]( + x, + x_t, + counts, + accums, pad_accum_sizes, N, - H, W, + H, + W, num_stages=num_stages, - num_warps=num_warps + num_warps=num_warps, ) split_size = [x * N for x in pad_sizes] chunks = torch.split(x_t.view(torch.uint8), split_size) @@ -315,13 +365,11 @@ def triton_batch_transpose_and_pad(x, count_list, x_t=None, pad=True): @triton.autotune(configs=configs, key=["M", "N", "D"]) @triton.jit -def opt_transpose_kernel(x_ptr, t_ptr, M, N, D, H: tl.constexpr, - W: tl.constexpr): +def opt_transpose_kernel(x_ptr, t_ptr, M, N, D, H: tl.constexpr, W: tl.constexpr): pid = tl.program_id(axis=0) # row-wise read, col-wise write offs = pid * W + tl.arange(0, H)[:, None] * N + tl.arange(0, W)[None, :] - toffs = pid * W * M + tl.arange(0, W)[:, None] * M + tl.arange(0, H)[None, - :] + toffs = pid * W * M + tl.arange(0, W)[:, None] * M + tl.arange(0, H)[None, :] m = tl.cdiv(M, H) for i in range(m): y = tl.trans(tl.load(x_ptr + offs)) @@ -331,13 +379,11 @@ def opt_transpose_kernel(x_ptr, t_ptr, M, N, D, H: tl.constexpr, def triton_opt_transpose(x): + assert x.is_contiguous() M, N = x.shape device = x.device D = 0 if x.dtype.itemsize == 1 else 1 t = torch.empty((N, M), device=device, dtype=x.dtype) grid = lambda META: (N // META["W"],) # noqa - opt_transpose_kernel[grid]( - x, t, - M, N, D - ) + opt_transpose_kernel[grid](x, t, M, N, D) return t diff --git a/linghe/utils/unary.py b/linghe/utils/unary.py index d3dd682..fbd194b 100644 --- a/linghe/utils/unary.py +++ b/linghe/utils/unary.py @@ -9,31 +9,39 @@ @triton.jit -def calculate_smooth_scale_kernel(x_ptr, y_ptr, min_value, smooth_coef, - N, - B: tl.constexpr, - EVEN: tl.constexpr, - ROUND: tl.constexpr): +def calculate_smooth_scale_kernel( + x_ptr, + y_ptr, + min_value, + smooth_coef, + N, + B: tl.constexpr, + EVEN: tl.constexpr, + ROUND: tl.constexpr, +): pid = tl.program_id(axis=0) offs = pid * B + tl.arange(0, B) if EVEN: - x = tl.load(x_ptr+offs).to(tl.float32) + x = tl.load(x_ptr + offs).to(tl.float32) else: - x = tl.load(x_ptr+offs, mask=offs clip_value)) + else: + x = tl.load(x_ptr + offs, mask=offs < N) + xc = tl.minimum(tl.maximum(x, -clip_value), clip_value) + tl.store(x_ptr + offs, xc, mask=(offs < N) & (tl.abs(x) > clip_value)) + + +def triton_clip(x, clip_value=100.0): + """ + clip(x, -clip_value, clip_value) + used to clip gradient. + Args: + x: Tensor. + clip_value: a python float scale + Returns: + updated x + """ + N = x.numel() + B = 512 + EVEN = N % N == 0 + grid = (triton.cdiv(N, B),) + clip_kernel[grid](x, clip_value, N, B, EVEN, num_stages=2, num_warps=4) + return x + + +@triton.jit +def batch_clip_kernel( + input_ptrs, size_ptr, clip_value, DT: tl.constexpr, B: tl.constexpr +): + tid = tl.program_id(axis=0).to(tl.int64) + bid = tl.program_id(axis=1).to(tl.int64) + T = tl.num_programs(axis=1) + + size = tl.load(size_ptr + tid) + if DT == 0: + input_ptr = tl.load(input_ptrs + tid).to(tl.pointer_type(tl.float32)) + elif DT == 1: + input_ptr = tl.load(input_ptrs + tid).to(tl.pointer_type(tl.bfloat16)) + else: + input_ptr = tl.load(input_ptrs + tid).to(tl.pointer_type(tl.float16)) + t = tl.cdiv(size, B * T) + offs = bid * t * B + tl.arange(0, B) + for i in range(t): + x = tl.load(input_ptr + offs, mask=offs < size) + xc = tl.minimum(tl.maximum(x, -clip_value), clip_value) + tl.store(input_ptr + offs, xc, mask=(offs < size) & (tl.abs(x) > clip_value)) + offs += B + + +def triton_batch_clip(xs, clip_value=100.0): + """ + return [clip(x, -clip_value, clip_value) for x in xs], + used to clip gradient. + Args: + xs: Tensor lists. + clip_value: a python float scale + Returns: + updated xs + """ + if len(xs) == 0: + return + dtype = xs[0].dtype + assert dtype in (torch.float32, torch.bfloat16, torch.float16) + assert all([x.is_contiguous() and x.dtype == dtype for x in xs]) + + device = xs[0].device + sizes = torch.tensor([x.numel() for x in xs], dtype=torch.int64).cuda( + device, non_blocking=True + ) + ptrs = torch.tensor([x.data_ptr() for x in xs], dtype=torch.int64).cuda( + device, non_blocking=True + ) + if dtype == torch.float32: + DT = 0 + elif dtype == torch.bfloat16: + DT = 1 + else: + DT = 2 + T = 256 + tensor_count = len(xs) + B = 512 + grid = (tensor_count, T) + batch_clip_kernel[grid](ptrs, sizes, clip_value, DT, B, num_stages=2, num_warps=2) + return xs diff --git a/scripts/plot_input_output.py b/scripts/plot_input_output.py index 4d0395c..f8a7b2e 100644 --- a/scripts/plot_input_output.py +++ b/scripts/plot_input_output.py @@ -1,59 +1,58 @@ -import numpy as np -import torch import matplotlib.pyplot as plt +import torch +def read_bf16_inputs(prefix="fc2"): + idx = {"qkv": 0, "out": 1, "fc1s": 2, "fc2s": 3, "fc1": 4, "fc2": 5}[prefix] -def read_bf16_inputs(prefix='fc2'): - idx = {"qkv":0, "out":1, "fc1s":2, "fc2s":3, "fc1":4, "fc2":5}[prefix] - - d = torch.load(f'/tmp/deepseek/bf16_forward_{idx}.bin', weights_only=True) - N, K = d['w'].shape - M = d['x'].numel()//K + d = torch.load(f"/tmp/deepseek/bf16_forward_{idx}.bin", weights_only=True) + N, K = d["w"].shape + M = d["x"].numel() // K - x = d['x'].detach().float().view(M,K) - w = d['w'].data.float().view(N,K) + x = d["x"].detach().float().view(M, K) + w = d["w"].data.float().view(N, K) - d = torch.load(f'/tmp/deepseek/bf16_backward_{idx}.bin', weights_only=True) - dy = d['dy'].detach().float().view(M,N) - dx = d['dx'].detach().float().view(M,K) + d = torch.load(f"/tmp/deepseek/bf16_backward_{idx}.bin", weights_only=True) + dy = d["dy"].detach().float().view(M, N) + dx = d["dx"].detach().float().view(M, K) - d = torch.load(f'/tmp/deepseek/bf16_update_{idx}.bin', weights_only=True) - dw = d['dw'].detach().float().view(N,K) + d = torch.load(f"/tmp/deepseek/bf16_update_{idx}.bin", weights_only=True) + dw = d["dw"].detach().float().view(N, K) x = x.cuda() w = w.cuda().transpose().contiguous() dy = dy.cuda() dx = dx.cuda() dw = dw.cuda().transpose().contiguous() - return x,w,dy,dx,dw - -def read_fp8_inputs(prefix='fc2'): - idx = {"qkv":0, "out":1, "fc1s":2, "fc2s":3, "fc1":4, "fc2":5}[prefix] - d = torch.load(f'/tmp/deepseek/fp8_forward_{idx}.bin', weights_only=True) - N,K = d['w'].shape - M = d['x'].numel()//K - - xq = d['x'].float().view(M,K) - xs = d['xs'].float() - xm = d['x_smooth_scale'] - x = xq*xm*xs[:,None] - - wq = d['w'].float() - ws = d['ws'].float() - wm = d['w_smooth_scale'] - w = wq*wm*ws[:,None] - - d = torch.load(f'/tmp/deepseek/fp8_backward_{idx}.bin', weights_only=True) - dyq = d['dy'].float() - dys = d['dys'] - dym = d['dy_smooth_scale'] - dy = dyq/dym*dys[:,None] - - dytq = d['dyt'].float().t() - dyts = d['dyts'] - dytm=d['dyt_smooth_scale'] - dyt = dytq/dytm*dyts[:,None] + return x, w, dy, dx, dw + + +def read_fp8_inputs(prefix="fc2"): + idx = {"qkv": 0, "out": 1, "fc1s": 2, "fc2s": 3, "fc1": 4, "fc2": 5}[prefix] + d = torch.load(f"/tmp/deepseek/fp8_forward_{idx}.bin", weights_only=True) + N, K = d["w"].shape + M = d["x"].numel() // K + + xq = d["x"].float().view(M, K) + xs = d["xs"].float() + xm = d["x_smooth_scale"] + x = xq * xm * xs[:, None] + + wq = d["w"].float() + ws = d["ws"].float() + wm = d["w_smooth_scale"] + w = wq * wm * ws[:, None] + + d = torch.load(f"/tmp/deepseek/fp8_backward_{idx}.bin", weights_only=True) + dyq = d["dy"].float() + dys = d["dys"] + dym = d["dy_smooth_scale"] + dy = dyq / dym * dys[:, None] + + dytq = d["dyt"].float().t() + dyts = d["dyts"] + dytm = d["dyt_smooth_scale"] + dyt = dytq / dytm * dyts[:, None] x = x.cuda() w = w.cuda().transpose().contiguous() @@ -61,88 +60,85 @@ def read_fp8_inputs(prefix='fc2'): dyt = dyt.cuda() xm = xm.cuda() wm = wm.cuda() - return x,w,dy,dyt,xm,wm + return x, w, dy, dyt, xm, wm # bf16 if True: - prefix = 'out' - x,w,dy,dx,dw = read_bf16_inputs(prefix=prefix) + prefix = "out" + x, w, dy, dx, dw = read_bf16_inputs(prefix=prefix) r = 5 - xb = torch.nn.functional.max_pool2d(x.abs()[None],r).cpu().numpy()[0] - wb = torch.nn.functional.max_pool2d(w.abs()[None],r).cpu().numpy()[0] - dyb = torch.nn.functional.max_pool2d(dy.abs()[None],r).cpu().numpy()[0] - dxb = torch.nn.functional.max_pool2d(dx.abs()[None],r).cpu().numpy()[0] - dwb = torch.nn.functional.max_pool2d(dw.abs()[None],r).cpu().numpy()[0] + xb = torch.nn.functional.max_pool2d(x.abs()[None], r).cpu().numpy()[0] + wb = torch.nn.functional.max_pool2d(w.abs()[None], r).cpu().numpy()[0] + dyb = torch.nn.functional.max_pool2d(dy.abs()[None], r).cpu().numpy()[0] + dxb = torch.nn.functional.max_pool2d(dx.abs()[None], r).cpu().numpy()[0] + dwb = torch.nn.functional.max_pool2d(dw.abs()[None], r).cpu().numpy()[0] - fmt = 'png' + fmt = "png" fig, ax = plt.subplots(figsize=(8, 12)) - ax.imshow(xb, cmap='gray') + ax.imshow(xb, cmap="gray") # plt.show() - plt.axis('off') - plt.savefig(f"figures/{prefix}_x.{fmt}", bbox_inches='tight',dpi=600) - plt.close('all') + plt.axis("off") + plt.savefig(f"figures/{prefix}_x.{fmt}", bbox_inches="tight", dpi=600) + plt.close("all") fig, ax = plt.subplots(figsize=(8, 12)) - ax.imshow(wb, cmap='gray') + ax.imshow(wb, cmap="gray") # plt.show() - plt.axis('off') - plt.savefig(f"figures/{prefix}_w.{fmt}", bbox_inches='tight',dpi=600) - plt.close('all') + plt.axis("off") + plt.savefig(f"figures/{prefix}_w.{fmt}", bbox_inches="tight", dpi=600) + plt.close("all") fig, ax = plt.subplots(figsize=(8, 12)) - ax.imshow(dyb, cmap='gray') + ax.imshow(dyb, cmap="gray") # plt.show() - plt.axis('off') - plt.savefig(f"figures/{prefix}_dy.{fmt}", bbox_inches='tight', dpi=600) - plt.close('all') + plt.axis("off") + plt.savefig(f"figures/{prefix}_dy.{fmt}", bbox_inches="tight", dpi=600) + plt.close("all") fig, ax = plt.subplots(figsize=(8, 12)) - ax.imshow(dxb, cmap='gray') + ax.imshow(dxb, cmap="gray") # plt.show() - plt.axis('off') - plt.savefig(f"figures/{prefix}_dx.{fmt}", bbox_inches='tight', dpi=600) - plt.close('all') + plt.axis("off") + plt.savefig(f"figures/{prefix}_dx.{fmt}", bbox_inches="tight", dpi=600) + plt.close("all") fig, ax = plt.subplots(figsize=(8, 12)) - ax.imshow(dwb, cmap='gray') + ax.imshow(dwb, cmap="gray") # plt.show() - plt.axis('off') - plt.savefig(f"figures/{prefix}_dw.{fmt}", bbox_inches='tight', dpi=600) - plt.close('all') + plt.axis("off") + plt.savefig(f"figures/{prefix}_dw.{fmt}", bbox_inches="tight", dpi=600) + plt.close("all") # fp8 if False: - prefix = 'fc2' - x,w,dy,dyt,xm,wm = read_fp8_inputs(prefix=prefix) + prefix = "fc2" + x, w, dy, dyt, xm, wm = read_fp8_inputs(prefix=prefix) r = 5 - xb = torch.nn.functional.max_pool2d(x.abs()[None],r).cpu().numpy()[0] - wb = torch.nn.functional.max_pool2d(w.abs()[None],r).cpu().numpy()[0] - dyb = torch.nn.functional.max_pool2d(dy.abs()[None],r).cpu().numpy()[0] + xb = torch.nn.functional.max_pool2d(x.abs()[None], r).cpu().numpy()[0] + wb = torch.nn.functional.max_pool2d(w.abs()[None], r).cpu().numpy()[0] + dyb = torch.nn.functional.max_pool2d(dy.abs()[None], r).cpu().numpy()[0] - fmt = 'png' + fmt = "png" fig, ax = plt.subplots(figsize=(8, 12)) - ax.imshow(xb, cmap='gray') + ax.imshow(xb, cmap="gray") # plt.show() - plt.axis('off') - plt.savefig(f"figures/{prefix}_x.{fmt}", bbox_inches='tight',dpi=600) - plt.close('all') + plt.axis("off") + plt.savefig(f"figures/{prefix}_x.{fmt}", bbox_inches="tight", dpi=600) + plt.close("all") fig, ax = plt.subplots(figsize=(8, 12)) - ax.imshow(wb, cmap='gray') + ax.imshow(wb, cmap="gray") # plt.show() - plt.axis('off') - plt.savefig(f"figures/{prefix}_w.{fmt}", bbox_inches='tight',dpi=600) - plt.close('all') + plt.axis("off") + plt.savefig(f"figures/{prefix}_w.{fmt}", bbox_inches="tight", dpi=600) + plt.close("all") fig, ax = plt.subplots(figsize=(8, 12)) - ax.imshow(dyb, cmap='gray') + ax.imshow(dyb, cmap="gray") # plt.show() - plt.axis('off') - plt.savefig(f"figures/{prefix}_dy.{fmt}", bbox_inches='tight', dpi=600) - plt.close('all') - - - + plt.axis("off") + plt.savefig(f"figures/{prefix}_dy.{fmt}", bbox_inches="tight", dpi=600) + plt.close("all") diff --git a/scripts/reproduce_triton_bug.py b/scripts/reproduce_triton_bug.py index 11bc3c2..ff567b0 100644 --- a/scripts/reproduce_triton_bug.py +++ b/scripts/reproduce_triton_bug.py @@ -3,43 +3,46 @@ Copyright (c) Ant Financial Service Group and its affiliates. """ - -import triton -import triton.language as tl import torch - - +import triton +import triton.language as tl @triton.jit -def rms_norm_and_smooth_quant_forward_kernel(x_ptr, weight_ptr, smooth_scale_ptr, - out_ptr, scale_ptr, max_ptr, rms_ptr, - eps, - M, - T, - N: tl.constexpr, - W: tl.constexpr, - CALIBRATE: tl.constexpr, - ROUND: tl.constexpr): +def rms_norm_and_smooth_quant_forward_kernel( + x_ptr, + weight_ptr, + smooth_scale_ptr, + out_ptr, + scale_ptr, + max_ptr, + rms_ptr, + eps, + M, + T, + N: tl.constexpr, + W: tl.constexpr, + CALIBRATE: tl.constexpr, + ROUND: tl.constexpr, +): pid = tl.program_id(axis=0) # row-wise read, row-wise write weight = tl.load(weight_ptr + tl.arange(0, N)).to(tl.float32)[None, :] smooth_scale = tl.load(smooth_scale_ptr + tl.arange(0, N))[None, :] smooth_scale = 1.0 / tl.maximum(smooth_scale, 1e-30) - # triton 3.3.1 has bug with N = 2048 and calibrate=True + # triton 3.3.1 has bug with N = 2048 and calibrate=True if CALIBRATE: - maxs = tl.zeros((N, ), dtype=tl.float32) - offs = pid * W * T * N + tl.arange(0, W)[:, None] * N + tl.arange(0, N)[ - None, :] + maxs = tl.zeros((N,), dtype=tl.float32) + offs = pid * W * T * N + tl.arange(0, W)[:, None] * N + tl.arange(0, N)[None, :] for i in range(T): indices = pid * W * T + i * W + tl.arange(0, W) x = tl.load(x_ptr + offs, mask=indices[:, None] < M).to(tl.float32) - rms = 1/tl.sqrt(tl.sum(x * x, axis=1) / N + eps) + rms = 1 / tl.sqrt(tl.sum(x * x, axis=1) / N + eps) tl.store(rms_ptr + indices, rms, mask=indices < M) x = x * rms[:, None] * weight if CALIBRATE: - maxs = tl.maximum(maxs, tl.max(tl.abs(x),0)) + maxs = tl.maximum(maxs, tl.max(tl.abs(x), 0)) x = x * smooth_scale scale = tl.maximum(tl.max(tl.abs(x), 1) / 448.0, 1e-30) @@ -54,11 +57,10 @@ def rms_norm_and_smooth_quant_forward_kernel(x_ptr, weight_ptr, smooth_scale_ptr tl.store(max_ptr + pid * N + tl.arange(0, N), maxs) - # rms is used for moe routing, it is stored as 1/rms -def triton_rms_norm_and_quant_forward(x, weight, smooth_scale, eps=1e-6, - calibrate=False, - round_scale=False): +def triton_rms_norm_and_quant_forward( + x, weight, smooth_scale, eps=1e-6, calibrate=False, round_scale=False, num_warps=4 +): # row-wise read, row-wise write M, N = x.shape assert N <= 8192 and 8192 % N == 0 @@ -70,7 +72,7 @@ def triton_rms_norm_and_quant_forward(x, weight, smooth_scale, eps=1e-6, W = 8192 // N T = 8 if M // W >= 4096 else 4 assert M % (T * W) == 0 - g = M // (T*W) + g = M // (T * W) if calibrate: maxs = torch.empty((g, N), dtype=torch.float32, device=device) else: @@ -93,29 +95,44 @@ def triton_rms_norm_and_quant_forward(x, weight, smooth_scale, eps=1e-6, calibrate, round_scale, num_stages=3, - num_warps=4 + num_warps=num_warps, ) if calibrate: maxs = maxs.amax(0) - return out, scale, maxs, rms - + return out, scale, maxs, rms -if __name__ == '__main__': +if __name__ == "__main__": M = 4096 N = 2048 dtype = torch.bfloat16 - device = 'cuda:0' - calibrate = True - # bug condition: triton=3.3.1 N=2048 calibrate=True num_warps=4 + device = "cuda:0" + calibrate = True x = torch.randn(M, N, dtype=dtype, requires_grad=True, device=device) weight = torch.randn(N, dtype=dtype, requires_grad=True, device=device) - smooth_scale = torch.rand(N, dtype=torch.float32, requires_grad=False, - device=device) + 0.1 + smooth_scale = ( + torch.rand(N, dtype=torch.float32, requires_grad=False, device=device) + 0.1 + ) - q, scale, maxs, rms = triton_rms_norm_and_quant_forward(x, weight, - smooth_scale=smooth_scale, - calibrate=calibrate, - round_scale=True) - print(f'{q=}\n{scale=}') \ No newline at end of file + # bug condition: triton=3.3.1 N=2048 calibrate=True num_warps=4 + q, scale, maxs, rms = triton_rms_norm_and_quant_forward( + x, + weight, + smooth_scale=smooth_scale, + calibrate=calibrate, + round_scale=True, + num_warps=4, + ) + print(f"bug_result: {q=}\n{scale=}") + + # no bug condition: triton=3.3.1 N=2048 calibrate=True num_warps=2 + q, scale, maxs, rms = triton_rms_norm_and_quant_forward( + x, + weight, + smooth_scale=smooth_scale, + calibrate=calibrate, + round_scale=True, + num_warps=2, + ) + print(f"correct_result: {q=}\n{scale=}") diff --git a/scripts/test.sh b/scripts/test.sh index be937e4..bdaae0a 100644 --- a/scripts/test.sh +++ b/scripts/test.sh @@ -1,18 +1,29 @@ -cd ../tests && -echo "test_add.py" && python test_add.py -echo "test_channel_quant.py" && python test_channel_quant.py -echo "test_dot.py" && python test_dot.py -echo "test_gather.py" && python test_gather.py -echo "test_gemm.py" && python test_gemm.py -echo "test_group_quant.py" && python test_group_quant.py -echo "test_loss.py" && python test_loss.py -echo "test_norm.py" && python test_norm.py -echo "test_rearange.py" && python test_rearange.py -echo "test_reduce.py" && python test_reduce.py -echo "test_scatter.py" && python test_scatter.py -echo "test_silu.py" && python test_silu.py -echo "test_smooth_quant.py" && python test_smooth_quant.py -echo "test_transpose.py" && python test_transpose.py -echo "test_unary.py" && python test_unary.py +cd tests && +# echo "test_add.py" && python test_add.py && +# echo "test_blockwise_fp8_gemm.py" && python test_blockwise_fp8_gemm.py && +# echo "test_blockwise_quant.py" && python test_blockwise_quant.py && +# echo "test_channel_quant.py" && python test_channel_quant.py && +# echo "test_channelwise_fp8_gemm.py" && python test_channelwise_fp8_gemm.py && +# echo "test_embedding.py" && python test_embedding.py && +# echo "test_fp32_gemm.py" && python test_fp32_gemm.py && +echo "test_gate.py" && python test_gate.py && +echo "test_gather.py" && python test_gather.py && +echo "test_group_quant.py" && python test_group_quant.py && +echo "test_hadamard_quant.py" && python test_hadamard_quant.py && +# echo "test_la.py" && python test_la.py && +echo "test_loss.py" && python test_loss.py && +# echo "test_mla.py" && python test_mla.py && +echo "test_mul.py" && python test_mul.py && +echo "test_norm.py" && python test_norm.py && +echo "test_rearange.py" && python test_rearange.py && +echo "test_reduce.py" && python test_reduce.py && +echo "test_rope.py" && python test_rope.py && +echo "test_scatter.py" && python test_scatter.py && +echo "test_silu.py" && python test_silu.py && +echo "test_smooth_quant.py" && python test_smooth_quant.py && +echo "test_topk.py" && python test_topk.py && +echo "test_transpose.py" && python test_transpose.py && +echo "test_unary.py" && python test_unary.py && +echo "success!" diff --git a/setup.py b/setup.py index bc21ae2..5e1df95 100644 --- a/setup.py +++ b/setup.py @@ -9,12 +9,13 @@ from setuptools import find_packages, setup with pathlib.Path("requirements.txt").open() as f: - install_requires = [str(requirement) for requirement in - pkg_resources.parse_requirements(f)] + install_requires = [ + str(requirement) for requirement in pkg_resources.parse_requirements(f) + ] setup( name="linghe", - version="0.0.2", + version="0.3.0", license="MIT", license_files=("LICENSE",), description="LLM traning kernels", diff --git a/tests/test_add.py b/tests/test_add.py index c6f0b0c..40264fd 100644 --- a/tests/test_add.py +++ b/tests/test_add.py @@ -3,48 +3,94 @@ Copyright (c) Ant Financial Service Group and its affiliates. """ +import pytest import torch -from linghe.tools.benchmark import benchmark_func -from linghe.tools.util import output_check -from linghe.utils.add import triton_inplace_add +from linghe.tools.check import output_check +from linghe.utils.add import triton_inplace_add, triton_batch_inplace_add -def torch_add(x, outputs, accum=True): +def torch_add(x, y, accum=True): if accum: - x += outputs + x += y return x else: - return x.float() + return x.copy_(y) -def test_triton_inplace_add(M=4096, N=4096, bench=False): - dtype = torch.bfloat16 - device = 'cuda:0' +def torch_batch_add(xs, ys, accum=True): + if accum: + for i, x in enumerate(xs): + x += ys[i] + return xs + else: + for i, x in enumerate(xs): + x.copy_(ys[i]) + return xs + + +@pytest.mark.parametrize( + "M,N", + [ + (4096, 4096), + (4096, 3467), + (3467, 3467), + ], +) +@pytest.mark.parametrize("accum", [False, True]) +def test_triton_inplace_add(M, N, accum, benchmark): + x_dtype = torch.bfloat16 + y_dtype = torch.float32 + device = "cuda:0" + + x = torch.randn(M, N, dtype=x_dtype, device=device) + y = torch.randn(M, N, dtype=y_dtype, device=device) + + out = x.clone() + triton_inplace_add(out, y, accum=accum) + out_ref = x.clone() + out_ref = torch_add(out_ref, y, accum=accum) + output_check(out_ref, out, "out") - outputs = torch.randn(M, N, dtype=dtype, device=device) - x = torch.randn(M, N, dtype=dtype, device=device) + ref_bytes = M * N * (x_dtype.itemsize * (2 if accum else 1) + x_dtype.itemsize) + ref_time = benchmark(torch_add, x, y, accum=accum, ref_bytes=ref_bytes) + benchmark( + triton_inplace_add, x, y, accum=accum, ref_time=ref_time, ref_bytes=ref_bytes + ) - out = outputs.clone() - triton_inplace_add(out, x) - out_ref = outputs + x - output_check(out_ref, out, 'sum') - n_repeat = 100 +@pytest.mark.parametrize( + "B,N", + [ + (8, 4096), + (8, 4097), + ], +) +@pytest.mark.parametrize("accum", [False, True]) +def test_triton_batch_inplace_add(B, N, accum, benchmark, M=4096): + x_dtype = torch.bfloat16 + y_dtype = torch.float32 + device = "cuda:0" - if bench: - ref_time = benchmark_func(torch_add, x, out, accum=False, - n_repeat=n_repeat) - benchmark_func(triton_inplace_add, out, x, accum=False, - n_repeat=n_repeat, - ref_time=ref_time, ref_bytes=M * N * 4) + xs = [torch.randn(M, N, dtype=x_dtype, device=device) for _ in range(B)] + ys = [torch.randn(M, N, dtype=y_dtype, device=device) for _ in range(B)] - ref_time = benchmark_func(torch_add, x, out, accum=True, - n_repeat=n_repeat) - benchmark_func(triton_inplace_add, out, x, accum=True, - n_repeat=n_repeat, - ref_time=ref_time, ref_bytes=M * N * 6) + out_ref = [x.clone() for x in xs] + out_ref = torch_batch_add(out_ref, ys, accum=accum) + out_ref = torch.cat([x.view(-1) for x in out_ref], 0) + out = [x.clone() for x in xs] + triton_batch_inplace_add(out, ys, accum=accum) + out = torch.cat([x.view(-1) for x in out], 0) + output_check(out_ref, out, "batch_add") -if __name__ == '__main__': - test_triton_inplace_add(M=4096, N=4096) + ref_bytes = M * N * (x_dtype.itemsize * (2 if accum else 1) + x_dtype.itemsize) + ref_time = benchmark(torch_batch_add, xs, ys, accum=accum, ref_bytes=ref_bytes) + benchmark( + triton_batch_inplace_add, + xs, + ys, + accum=accum, + ref_bytes=ref_bytes, + ref_time=ref_time, + ) diff --git a/tests/test_blockwise_fp8_gemm.py b/tests/test_blockwise_fp8_gemm.py index 497b791..243d531 100644 --- a/tests/test_blockwise_fp8_gemm.py +++ b/tests/test_blockwise_fp8_gemm.py @@ -3,18 +3,22 @@ Copyright (c) Ant Financial Service Group and its affiliates. """ +import pytest import torch -from linghe.gemm.blockwise_fp8_gemm import triton_bb_fp8_gemm, triton_tt_fp8_gemm +from linghe.gemm.blockwise_fp8_gemm import triton_blockwise_fp8_gemm +from linghe.tools.check import output_check -from linghe.tools.benchmark import benchmark_func -from linghe.tools.util import output_check - - -def test_triton_bb_gemm(M=4096, N=4096, K=4096, bench=False): +@pytest.mark.parametrize( + "M,N,K", + [ + (4096, 8192, 2048), + ], +) +def test_blockwise_gemm(M, N, K, benchmark): dtype = torch.bfloat16 - device = 'cuda:0' + device = "cuda:0" B = 64 x = torch.randn(M, K, dtype=dtype, device=device) @@ -26,60 +30,30 @@ def test_triton_bb_gemm(M=4096, N=4096, K=4096, bench=False): x_q = x.to(torch.float8_e4m3fn) w_q = w.to(torch.float8_e4m3fn) - x_dq = (x_q.float().view(M//B, B, K//B, B)*x_scales[:,None,:,None]).view(M,K) - w_dq = (w_q.float().view(N//B, B, K//B, B)*w_scales[:,None,:,None]).view(N,K) - - y_ref = x_dq@w_dq.t() - y = triton_bb_fp8_gemm(x_q, w_q, x_scales, w_scales, - out_dtype=dtype, block_size=B) - output_check(y_ref, y, 'y') - - - if bench: - n_repeat = 100 - ref_flops = M * N * K * 2 - - benchmark_func(triton_bb_fp8_gemm, x_q, w_q, x_scales, w_scales, - out_dtype=dtype, block_size=B, - n_repeat=n_repeat, ref_flops=ref_flops) - - - -def test_triton_tt_gemm(M=4096, N=4096, K=4096, bench=False): - dtype = torch.bfloat16 - device = 'cuda:0' - B = 64 - - x = torch.randn(M, K, dtype=dtype, device=device) - w = torch.randn(N, K, dtype=dtype, device=device) - - x_scales = torch.rand((M, K // B), dtype=torch.float32, device=device) - w_scales = torch.rand((N, K // B), dtype=torch.float32, device=device) - - x_q = x.to(torch.float8_e4m3fn) - w_q = w.to(torch.float8_e4m3fn) - - x_dq = (x_q.float().view(M, K//B, B)*x_scales[:,:,None]).view(M,K) - w_dq = (w_q.float().view(N, K//B, B)*w_scales[:,:,None]).view(N,K) - - y_ref = x_dq@w_dq.t() - y = triton_tt_fp8_gemm(x_q, w_q, x_scales, w_scales, - out_dtype=dtype, block_size=B) - output_check(y_ref, y, 'y') - - - if bench: - n_repeat = 100 - ref_flops = M * N * K * 2 - - benchmark_func(triton_tt_fp8_gemm, x_q, w_q, x_scales, w_scales, - out_dtype=dtype, block_size=B, - n_repeat=n_repeat, ref_flops=ref_flops) - - - -if __name__ == '__main__': - test_triton_bb_gemm(M=4096, N=8192, K=2048, bench=True) - test_triton_tt_gemm(M=4096, N=8192, K=2048, bench=True) - - + x_dq = (x_q.float().view(M // B, B, K // B, B) * x_scales[:, None, :, None]).view( + M, K + ) + w_dq = (w_q.float().view(N // B, B, K // B, B) * w_scales[:, None, :, None]).view( + N, K + ) + + y_ref = x_dq @ w_dq.t() + y = triton_blockwise_fp8_gemm( + x_q, w_q, x_scales, w_scales, out_dtype=dtype, block_size=B + ) + output_check(y_ref.to(dtype), y, name="y", rtol=0.05, atol=1.0) + + n_repeat = 100 + ref_flops = M * N * K * 2 + + benchmark( + triton_blockwise_fp8_gemm, + x_q, + w_q, + x_scales, + w_scales, + out_dtype=dtype, + block_size=B, + n_repeat=n_repeat, + ref_flops=ref_flops, + ) diff --git a/tests/test_blockwise_quant.py b/tests/test_blockwise_quant.py new file mode 100644 index 0000000..f97a749 --- /dev/null +++ b/tests/test_blockwise_quant.py @@ -0,0 +1,124 @@ +import pytest +import torch + +from linghe.quant.block import ( + triton_block_quant, + triton_blockwise_quant, + triton_batch_blockwise_quant, +) +from linghe.tools.check import output_check +from linghe.tools.util import ( + torch_block_quant, + torch_blockwise_quant, + torch_make_indices, +) + + +def torch_batch_blockwise_quant(x, token_count_per_expert_list, round_scale=True): + M, DIM = x.shape + q_refs = [] + s_refs = [] + qt_refs = [] + st_refs = [] + s = 0 + for i, c in enumerate(token_count_per_expert_list): + c = token_count_per_expert_list[i] + if c == 0: + continue + y = x[s : s + c] + y = y.float() + + y_q, y_scale, yt_q, yt_scale = torch_blockwise_quant( + y, round_scale=round_scale, padding=False + ) + q_refs.append(y_q.view(-1)) + s_refs.append(y_scale.view(-1)) + qt_refs.append(yt_q.view(-1)) + st_refs.append(yt_scale.view(-1)) + s += c + q_ref = torch.cat(q_refs, 0) + s_ref = torch.cat(s_refs, 0) + qt_ref = torch.cat(qt_refs, 0) + st_ref = torch.cat(st_refs, 0) + return q_ref, s_ref, qt_ref, st_ref + + +@pytest.mark.parametrize( + "M,N", + [ + (8192, 4096), + ], +) +def test_block_quant(M, N, benchmark): + device = "cuda:0" + x = torch.randn((M, N), dtype=torch.bfloat16, device=device) ** 3 + + x_q_ref, x_s_ref = torch_block_quant(x, round_scale=True) + x_q, x_s = triton_block_quant(x, round_scale=True) + output_check(x_q_ref.float(), x_q.float(), "data") + output_check(x_s_ref.float(), x_s.float(), "scale") + + benchmark(triton_block_quant, x, round_scale=True, ref_bytes=M * N * 4) + + +@pytest.mark.parametrize( + "M,N", + [ + (8192, 4096), + ], +) +def test_blockwise_quant(M, N, benchmark): + device = "cuda:0" + x = torch.randn((M, N), dtype=torch.bfloat16, device=device) ** 3 + + x_q_ref, x_s_ref, xt_q_ref, xt_s_ref = torch_blockwise_quant( + x, round_scale=True, padding=False + ) + x_q, x_s, xt_q, xt_s = triton_blockwise_quant(x, round_scale=True) + output_check(x_q_ref.float(), x_q.float(), "data") + output_check(x_s_ref.float(), x_s.float(), "scale") + output_check(xt_q_ref.float(), xt_q.float(), "t.data") + output_check(xt_s_ref.float(), xt_s.float(), "t.scale") + + benchmark(triton_blockwise_quant, x, round_scale=True, ref_bytes=M * N * 4) + + +@pytest.mark.parametrize( + "M,N,n_experts,topk", + [ + (16384, 2048, 32, 2), + ], +) +def test_batch_block_quant(M, N, n_experts, topk, benchmark): + device = "cuda:0" + logits = torch.randn((M, n_experts), dtype=torch.float32, device=device) ** 3 + logits[:, 0] -= 1000 + logits[:, 2] -= 100 + probs, mask_map, token_count_per_expert, indices, row_id_map = torch_make_indices( + logits, topk=topk, bias=-0.01 + ) + token_count_per_expert_list = token_count_per_expert.tolist() + + x = torch.randn((M, N), dtype=torch.bfloat16, device=device) + x = x[indices] + + x_q_ref, x_s_ref, xt_q_ref, xt_s_ref = torch_batch_blockwise_quant( + x, token_count_per_expert_list, round_scale=True + ) + + x_q, x_s, xt_q, xt_s = triton_batch_blockwise_quant( + x, token_count_per_expert, token_count_per_expert_list, round_scale=True + ) + output_check(x_q_ref.float(), x_q.view(-1).float(), "data") + output_check(x_s_ref.float(), x_s.view(-1).float(), "scale") + output_check(xt_q_ref.float(), xt_q.view(-1).float(), "t.data") + output_check(xt_s_ref.float(), xt_s.view(-1).float(), "t.scale") + + benchmark( + triton_batch_blockwise_quant, + x, + token_count_per_expert, + token_count_per_expert_list, + round_scale=True, + ref_bytes=M * N * 4, + ) diff --git a/tests/test_channel_quant.py b/tests/test_channel_quant.py index 5b4fd49..47d15d2 100644 --- a/tests/test_channel_quant.py +++ b/tests/test_channel_quant.py @@ -3,45 +3,57 @@ Copyright (c) Ant Financial Service Group and its affiliates. """ +import pytest import torch -from linghe.quant.channel import (triton_deprecated_tokenwise_row_quant, - triton_row_quant, - triton_tokenwise_row_quant) -from linghe.tools.benchmark import benchmark_func -from linghe.tools.util import (output_check, - torch_row_quant) - - -def test_row_quant(M=4096, N=4096, round_scale=True, bench=False): - device = 'cuda:0' +from linghe.quant.channel import ( + triton_deprecated_tokenwise_row_quant, + triton_row_quant, + triton_tokenwise_row_quant, +) +from linghe.tools.check import output_check +from linghe.tools.util import torch_row_quant + + +@pytest.mark.parametrize( + "M,N", + [ + (4096, 4096), + (4090, 4096), + (4096, 8192), + (3456, 2048), + (1, 2048), + ], +) +@pytest.mark.parametrize("round_scale", [False, True]) +def test_row_quant(M, N, round_scale, benchmark): + device = "cuda:0" dtype = torch.bfloat16 x = torch.randn((M, N), dtype=dtype, device=device) ** 3 x_q_ref, x_scale_ref = torch_row_quant(x, round_scale=round_scale) x_q, x_scale = triton_row_quant(x, round_scale=round_scale) - output_check(x_q_ref.float(), x_q.float(), mode='data') - output_check(x_scale_ref, x_scale, mode='scale') + output_check(x_q_ref, x_q, name="data") + output_check(x_scale_ref, x_scale, name="scale") x_q, x_scale = triton_tokenwise_row_quant(x, round_scale=round_scale) - output_check(x_q_ref.float(), x_q.float(), mode='data') - output_check(x_scale_ref, x_scale, mode='scale') - - if bench: - ref_time = benchmark_func(torch_row_quant, x, n_repeat=100, - ref_bytes=M * N * 3) - benchmark_func(triton_row_quant, x, n_repeat=100, ref_bytes=M * N * 3, - ref_time=ref_time) - benchmark_func(triton_deprecated_tokenwise_row_quant, x, n_repeat=100, - ref_bytes=M * N * 3, ref_time=ref_time) - benchmark_func(triton_tokenwise_row_quant, x, n_repeat=100, - ref_bytes=M * N * 3, ref_time=ref_time) - - -if __name__ == '__main__': - test_row_quant(M=4096, N=4096, round_scale=False) - test_row_quant(M=4090, N=4096, round_scale=True) - test_row_quant(M=4096, N=8192, round_scale=True) - test_row_quant(M=3456, N=2048, round_scale=True) - test_row_quant(M=1, N=2048, round_scale=True) + output_check(x_q_ref, x_q, name="data") + output_check(x_scale_ref, x_scale, name="scale") + + ref_time = benchmark(torch_row_quant, x, n_repeat=100, ref_bytes=M * N * 3) + benchmark(triton_row_quant, x, n_repeat=100, ref_bytes=M * N * 3, ref_time=ref_time) + benchmark( + triton_deprecated_tokenwise_row_quant, + x, + n_repeat=100, + ref_bytes=M * N * 3, + ref_time=ref_time, + ) + benchmark( + triton_tokenwise_row_quant, + x, + n_repeat=100, + ref_bytes=M * N * 3, + ref_time=ref_time, + ) diff --git a/tests/test_channelwise_fp8_gemm.py b/tests/test_channelwise_fp8_gemm.py index ce5a195..8d5417d 100644 --- a/tests/test_channelwise_fp8_gemm.py +++ b/tests/test_channelwise_fp8_gemm.py @@ -3,22 +3,23 @@ Copyright (c) Ant Financial Service Group and its affiliates. """ +import pytest import torch from linghe.gemm.channelwise_fp8_gemm import triton_scaled_mm - +from linghe.tools.check import output_check from linghe.utils.add import triton_inplace_add -from linghe.tools.benchmark import benchmark_func -from linghe.tools.util import output_check def scaled_gemm_and_update(x_q, w_q, x_scales, w_scales, c=None, accum=False): - o = torch._scaled_mm(x_q, - w_q.t(), - scale_a=x_scales.view(-1, 1), - scale_b=w_scales.view(1, -1), - out_dtype=torch.bfloat16, - use_fast_accum=True) + o = torch._scaled_mm( + x_q, + w_q.t(), + scale_a=x_scales.view(-1, 1), + scale_b=w_scales.view(1, -1), + out_dtype=torch.bfloat16, + use_fast_accum=True, + ) if accum: assert c is not None triton_inplace_add(c, o, accum=accum) @@ -27,12 +28,15 @@ def scaled_gemm_and_update(x_q, w_q, x_scales, w_scales, c=None, accum=False): return c - - -def test_triton_channelwise_gemm(M=4096, N=4096, K=4096, bench=False): - +@pytest.mark.parametrize( + "M,N,K", + [ + (4096, 4096, 4096), + ], +) +def test_triton_channelwise_gemm(M, N, K, benchmark): dtype = torch.bfloat16 - device = 'cuda:0' + device = "cuda:0" x = torch.randn(M, K, dtype=dtype, device=device) x_scales = torch.rand((M,), dtype=torch.float32, device=device) @@ -42,39 +46,94 @@ def test_triton_channelwise_gemm(M=4096, N=4096, K=4096, bench=False): w_scales = torch.rand((N,), dtype=torch.float32, device=device) w_q = w.to(torch.float8_e4m3fn) - y_ref = (x_q.float()*x_scales[:,None])@(w_q.float()*w_scales[:,None]).t() - y = triton_scaled_mm(x_q, w_q, x_scales, w_scales, c=None, - accum=False) - - output_check(y_ref, y, 'y') - - - if bench: - y_bf16 = torch.randn(M, N, dtype=dtype, device=device) - y_fp16 = y.to(torch.float16) - y_fp32 = y.to(torch.float32) - - n_repeat = 100 - ref_flops = M * N * K * 2 - - benchmark_func(scaled_gemm_and_update, x_q, w_q, x_scales, w_scales, c=y_bf16, - accum=False, n_repeat=n_repeat, ref_flops=ref_flops) - - benchmark_func(scaled_gemm_and_update, x_q, w_q, x_scales, w_scales, c=y_bf16, - accum=True, n_repeat=n_repeat, ref_flops=ref_flops) - benchmark_func(scaled_gemm_and_update, x_q, w_q, x_scales, w_scales, c=y_fp16, - accum=True, n_repeat=n_repeat, ref_flops=ref_flops) - benchmark_func(scaled_gemm_and_update, x_q, w_q, x_scales, w_scales, c=y_fp32, - accum=True, n_repeat=n_repeat, ref_flops=ref_flops) - - benchmark_func(triton_scaled_mm, x_q, w_q, x_scales, w_scales, c=y_bf16, - accum=True, n_repeat=n_repeat, ref_flops=ref_flops) - benchmark_func(triton_scaled_mm, x_q, w_q, x_scales, w_scales, c=y_fp16, - accum=True, n_repeat=n_repeat, ref_flops=ref_flops) - benchmark_func(triton_scaled_mm, x_q, w_q, x_scales, w_scales, c=y_fp32, - accum=True, n_repeat=n_repeat, ref_flops=ref_flops) - - -if __name__ == '__main__': - test_triton_channelwise_gemm(M=4096, N=4096, K=4096, bench=True) - + y_ref = (x_q.float() * x_scales[:, None]) @ (w_q.float() * w_scales[:, None]).t() + y = triton_scaled_mm(x_q, w_q, x_scales, w_scales, c=None, accum=False) + + output_check(y_ref, y, name="y", atol=-1) + + y_bf16 = torch.randn(M, N, dtype=dtype, device=device) + y_fp16 = y.to(torch.float16) + y_fp32 = y.to(torch.float32) + + n_repeat = 100 + ref_flops = M * N * K * 2 + + benchmark( + scaled_gemm_and_update, + x_q, + w_q, + x_scales, + w_scales, + c=y_bf16, + accum=False, + n_repeat=n_repeat, + ref_flops=ref_flops, + ) + + benchmark( + scaled_gemm_and_update, + x_q, + w_q, + x_scales, + w_scales, + c=y_bf16, + accum=True, + n_repeat=n_repeat, + ref_flops=ref_flops, + ) + benchmark( + scaled_gemm_and_update, + x_q, + w_q, + x_scales, + w_scales, + c=y_fp16, + accum=True, + n_repeat=n_repeat, + ref_flops=ref_flops, + ) + benchmark( + scaled_gemm_and_update, + x_q, + w_q, + x_scales, + w_scales, + c=y_fp32, + accum=True, + n_repeat=n_repeat, + ref_flops=ref_flops, + ) + + benchmark( + triton_scaled_mm, + x_q, + w_q, + x_scales, + w_scales, + c=y_bf16, + accum=True, + n_repeat=n_repeat, + ref_flops=ref_flops, + ) + benchmark( + triton_scaled_mm, + x_q, + w_q, + x_scales, + w_scales, + c=y_fp16, + accum=True, + n_repeat=n_repeat, + ref_flops=ref_flops, + ) + benchmark( + triton_scaled_mm, + x_q, + w_q, + x_scales, + w_scales, + c=y_fp32, + accum=True, + n_repeat=n_repeat, + ref_flops=ref_flops, + ) diff --git a/tests/test_dist_loss.py b/tests/test_dist_loss.py new file mode 100644 index 0000000..516b142 --- /dev/null +++ b/tests/test_dist_loss.py @@ -0,0 +1,218 @@ +# -*- coding: utf-8 -*- +""" +Copyright (c) Ant Financial Service Group and its affiliates. +""" + +import os +import random +from datetime import timedelta + +import torch +import torch.distributed as dist + +from linghe.tools.benchmark import benchmark_func +from linghe.tools.check import output_check +from linghe.utils.loss import ( + triton_parallel_softmax_cross_entropy_forward, + triton_parallel_softmax_cross_entropy_backward, + triton_softmax_cross_entropy_forward, +) +# from megatron.core.fusions.fused_cross_entropy import fused_vocab_parallel_cross_entropy + +def torch_cross_entropy(logits, targets, ignore_index=-100, reduction="none"): + float_logits = logits.to(torch.float32) + losses = torch.nn.functional.cross_entropy( + float_logits.view(-1, logits.size()[-1]), + targets.view(-1), + reduction=reduction, + ignore_index=ignore_index, + ) + return losses + + +def test_triton_softmax_cross_entropy( + M=4096, + N=157184, + coef=1.0, + grad_coef=1.0, + ignore_index=None, + fill=False, + inplace=False, + group=None, + bench=False, +): + group_size = group.size() + group_rank = group.rank() + + device_module = torch.get_device_module("cuda") + device_module.set_device(torch.device(f"cuda:{group_rank}")) + + device = "cuda" + dtype = torch.bfloat16 + local_logits = torch.randn((M, N), dtype=dtype, device=device, requires_grad=False) + + select = True + if select: + top_indices = torch.topk(local_logits, 1)[1].tolist() + targets = [] + for i, idx in enumerate(top_indices): + targets.append(random.choice(idx) * group_size) + + targets = torch.tensor(targets, dtype=torch.long, device=device) + else: + targets = torch.randint( + 0, N * group_size, (M,), dtype=torch.long, device=device + ) + + if ignore_index is not None: + targets[:10] = ignore_index + + if fill: + local_logits[:, :8192] = -10000 + + ignore_index = -100 if ignore_index is None else ignore_index + local_logits = (local_logits * coef).detach().clone().requires_grad_() + + global_logits = torch.empty((group_size, M, N), dtype=dtype, device=device) + dist.all_gather_into_tensor(global_logits, local_logits.detach(), group=group) + global_logits = ( + torch.reshape(torch.permute(global_logits, (1, 0, 2)), (M, group_size * N)) + .contiguous() + .requires_grad_() + ) + + global_targets = torch.empty((group_size, M), dtype=torch.long, device=device) + dist.all_gather_into_tensor(global_targets, targets, group=group) + global_targets = global_targets[0] + + local_output_grad = ( + torch.randn((M,), dtype=torch.float32, device=device) * grad_coef + ) + global_output_grad = torch.empty( + (group_size, M), dtype=torch.float32, device=device + ) + dist.all_gather_into_tensor(global_output_grad, local_output_grad, group=group) + global_output_grad = global_output_grad[0] + + loss_ref = torch_cross_entropy( + global_logits, global_targets, ignore_index=ignore_index, reduction="none" + ) + loss_ref.backward(global_output_grad, retain_graph=True) + grad_ref = global_logits.grad + global_logits.grad = None + + loss_sa, sum_exp_sa, max_logit_sa = triton_softmax_cross_entropy_forward( + global_logits.detach().clone(), global_targets, ignore_index=ignore_index + ) + + loss, sum_exp, max_logit = triton_parallel_softmax_cross_entropy_forward( + local_logits.detach().clone(), global_targets, group, ignore_index=ignore_index + ) + output_check(loss_ref, loss, name=f"ref_loss:{group_rank}", atol=1e-4, rtol=1e-5) + + grad = triton_parallel_softmax_cross_entropy_backward( + local_logits.detach().clone(), + global_targets, + sum_exp, + max_logit, + global_output_grad, + group, + ignore_index=ignore_index, + inplace=inplace, + ) + + output_check( + grad_ref[:, group_rank * N : (group_rank + 1) * N], + grad, + name=f"grad:{group_rank}", + digest=10, + ) + + # loss_native = fused_vocab_parallel_cross_entropy(local_logits, global_targets, group) + # grad_native = loss_native.backward(global_output_grad) + # grad_native = local_logits.grad + # local_logits.grad = None + # output_check(loss_ref, loss_native, name=f'native_loss:{group_rank}', atol=1e-4, rtol=1e-5) + # output_check(grad_ref[:, group_rank*N:(group_rank+1)*N], grad_native, name=f'native_grad:{group_rank}', atol=1e-4, rtol=1e-5) + + if bench: + benchmark_func( + torch_cross_entropy, + global_logits.requires_grad_(), + global_targets, + ref_bytes=M * N * 2, + ) + benchmark_func( + triton_parallel_softmax_cross_entropy_forward, + local_logits, + global_targets, + group, + ignore_index=ignore_index, + ref_bytes=M * N * 2, + ) + benchmark_func( + loss_ref.backward, + global_output_grad, + retain_graph=True, + ref_bytes=M * N * 4, + ) + benchmark_func( + triton_parallel_softmax_cross_entropy_backward, + local_logits.detach().clone(), + global_targets, + sum_exp, + max_logit, + global_output_grad, + group, + ignore_index=ignore_index, + inplace=True, + ref_bytes=M * N * 4, + ) + + +if __name__ == "__main__": + # torchrun --nproc_per_node=2 test_dist_loss.py + world_size = int(os.environ["WORLD_SIZE"]) + local_rank = int(os.environ["LOCAL_RANK"]) + print(f"{world_size=} {local_rank=}") + dist.init_process_group( + backend="nccl", + init_method="env://", + world_size=world_size, + rank=local_rank, + timeout=timedelta(seconds=30), + ) + pg = dist.distributed_c10d._get_default_group() + test_triton_softmax_cross_entropy( + M=8192, N=157184, coef=1.0, grad_coef=1.0, inplace=False, group=pg, bench=True + ) + test_triton_softmax_cross_entropy( + M=8192, + N=157184, + coef=1.0, + grad_coef=1.0, + ignore_index=-100, + inplace=False, + group=pg, + bench=False, + ) + test_triton_softmax_cross_entropy( + M=8192, + N=157184, + coef=1.0, + grad_coef=1.0, + fill=True, + inplace=False, + group=pg, + bench=False, + ) + test_triton_softmax_cross_entropy( + M=8192, + N=157184, + coef=1.0, + grad_coef=1.0, + fill=True, + inplace=True, + group=pg, + bench=False, + ) diff --git a/tests/test_dot.py b/tests/test_dot.py deleted file mode 100644 index fcedbd6..0000000 --- a/tests/test_dot.py +++ /dev/null @@ -1,42 +0,0 @@ -# -*- coding: utf-8 -*- -""" -Copyright (c) Ant Financial Service Group and its affiliates. -""" - -import torch - -from linghe.tools.benchmark import benchmark_func -from linghe.tools.util import output_check -from linghe.utils.dot import triton_dot - - -def torch_fp16_dot(x, y): - return (x * y).sum(1) - - -def test_dot(M=4096, N=4096, bench=False): - dtype = torch.bfloat16 - device = 'cuda:0' - - n_repeat = 100 - - x = torch.randn(M, N, dtype=dtype, device=device) - y = torch.randn(M, N, dtype=dtype, device=device) - q = torch.randn(M, N, dtype=dtype, device=device).to(torch.float8_e4m3fn) - quant_scale = torch.randn(M, dtype=torch.float32, device=device).abs() - smooth_scale = torch.randn(N, dtype=torch.float32, device=device).abs() - - sums = triton_dot(x, q) - sums_ref = torch_fp16_dot(x, q.float().to(dtype)) - output_check(sums_ref, sums, 'sum') - - sums_ref = (x.float() * ( - q.to(torch.float32) * quant_scale[:, None] * smooth_scale[None, - :])).sum(dim=1) - - if bench: - ref_time = benchmark_func(torch_fp16_dot, x, y, n_repeat=n_repeat) - - -if __name__ == '__main__': - test_dot(M=4096, N=4096) diff --git a/tests/test_embedding.py b/tests/test_embedding.py new file mode 100644 index 0000000..89ec033 --- /dev/null +++ b/tests/test_embedding.py @@ -0,0 +1,200 @@ +# -*- coding: utf-8 -*- +""" +Copyright (c) Ant Financial Service Group and its affiliates. +""" + +import pytest +import torch + +from linghe.facade.emb import embedding_lookup, fused_accumulation_embedding_lookup +from linghe.tools.check import output_check +from linghe.utils.emb import ( + triton_embedding_forward, + triton_embedding_backward, + triton_scan_and_count, + triton_sync_embedding_backward, + triton_atomic_embedding_backward, +) + + +@pytest.mark.parametrize( + "B,M", + [ + (None, 8192), + (None, 4097), + (4, 8192), + (2, 4097), + ], +) +def test_scan(B, M, benchmark): + device = "cuda:0" + if B is None or B == 0: + input_ids = torch.randint(0, 10000, (M,), dtype=torch.int32, device=device) + + sorted_ids, sorted_indices = torch.sort(input_ids, stable=False) + unique_ids_ref, unique_counts_ref = torch.unique_consecutive( + sorted_ids, return_counts=True + ) + accum_counts_ref = torch.cumsum( + torch.tensor( + [0] + unique_counts_ref.tolist(), device=unique_counts_ref.device + ), + 0, + ) + size = accum_counts_ref.size(0) + + accum_counts = triton_scan_and_count(sorted_ids) + output_check(accum_counts_ref, accum_counts[:size], name="accum_counts") + + benchmark(triton_scan_and_count, sorted_ids) + else: + input_ids = torch.randint(0, 10000, (B, M), dtype=torch.int32, device=device) + + sorted_ids, sorted_indices = torch.sort(input_ids, dim=-1, stable=False) + + accum_counts = triton_scan_and_count(sorted_ids) + for b in range(B): + unique_ids_ref, unique_counts_ref = torch.unique_consecutive( + sorted_ids[b], return_counts=True + ) + accum_counts_ref = torch.cumsum( + torch.tensor( + [0] + unique_counts_ref.tolist(), device=unique_counts_ref.device + ), + 0, + ) + size = accum_counts_ref.size(0) + output_check( + accum_counts_ref, + accum_counts[b, :size], + name=f"accum_counts (2D, B={B}, M={M}, batch={b})", + ) + + benchmark(triton_scan_and_count, sorted_ids) + + +@pytest.mark.parametrize( + "B,M,V,D,transpose", + [ + (1, 8192, 150000, 8192, False), + (2, 4096, 150000, 4096, True), + (1, 4097, 150000, 4096, False), + (3, 4097, 150000, 4096, False), + ], +) +def test_embedding(B, M, V, D, transpose, benchmark): + dtype = torch.bfloat16 + device = "cuda:0" + + embedding = torch.nn.Embedding(V, D, dtype=dtype, device=device) + input_ids = torch.randint(0, V // 15, (B, M), dtype=torch.int32, device=device) + weights = embedding.weight + grad_ref = torch.randn((V, D), dtype=dtype, device=device) + weights.grad = grad_ref + grad = grad_ref.clone().detach() + + y_ref = embedding(input_ids) + if transpose: + dy = torch.randn((M, B, D), device=device, dtype=dtype).permute(1, 0, 2) + else: + dy = torch.randn((B, M, D), device=device, dtype=dtype) + y_ref.backward(dy, retain_graph=True) + + weights.grad = grad + y = embedding_lookup(input_ids, weights) + y.backward(dy, retain_graph=True) + output_check(y_ref, y, name="y") + output_check(grad_ref, grad.to(dtype), name="grad") + + ref_bytes = B * M * D * 4 + ref_time = benchmark(embedding.forward, input_ids, ref_bytes=ref_bytes) + benchmark( + embedding_lookup, input_ids, weights, ref_time=ref_time, ref_bytes=ref_bytes + ) + + ref_time = benchmark(y_ref.backward, dy, retain_graph=True, ref_bytes=ref_bytes) + benchmark(y.backward, dy, retain_graph=True, ref_time=ref_time, ref_bytes=ref_bytes) + + +@pytest.mark.parametrize( + "B,M,V,D,transpose", + [ + (1, 8192, 150000, 8192, False), + (1, 4096, 150000, 8192, False), + (2, 4096, 150000, 8192, True), + (2, 4097, 150000, 4096, True), + (0, 4096, 150000, 8192, True), + (1, 8100, 150000, 8192, False), + ], +) +def test_fused_embedding(B, M, V, D, transpose, benchmark, use_main_grad=True): + dtype = torch.bfloat16 + device = "cuda:0" + + embedding = torch.nn.Embedding(V, D, dtype=dtype, device=device) + input_ids = torch.randint(0, V // 15, (B, M), dtype=torch.int32, device=device) + weights = embedding.weight + main_grad = torch.randn((V, D), dtype=torch.float32, device=device) + grad_ref = main_grad.to(dtype) + grad = main_grad.to(dtype) + + weights.grad = grad_ref + y_ref = embedding(input_ids) + if transpose: + dy = torch.randn((M, B, D), device=device, dtype=dtype).permute(1, 0, 2) + else: + dy = torch.randn((B, M, D), device=device, dtype=dtype) + y_ref.backward(dy, retain_graph=True) + + weights.grad = grad + weights.main_grad = main_grad + grad_name = "main_grad" if use_main_grad else "grad" + y = fused_accumulation_embedding_lookup(input_ids, weights, grad_name=grad_name) + y.backward(dy, retain_graph=True) + output_check(y_ref, y, name="y") + output_check( + grad_ref, main_grad.to(dtype) if use_main_grad else grad, name="grad", atol=0.05 + ) + + ref_bytes = B * M * D * 4 + ref_time = benchmark(embedding.forward, input_ids) + benchmark( + fused_accumulation_embedding_lookup, + input_ids, + weights, + grad_name=grad_name, + ref_time=ref_time, + ref_bytes=ref_bytes, + ) + + ref_time = benchmark(y_ref.backward, dy, retain_graph=True) + benchmark(y.backward, dy, retain_graph=True, ref_time=ref_time, ref_bytes=ref_bytes) + grad_ptr = main_grad.data_ptr() if use_main_grad else grad.data_ptr() + grad_dtype = main_grad.dtype if use_main_grad else grad.dtype + benchmark( + triton_atomic_embedding_backward, + dy, + input_ids, + grad_ptr, + grad_dtype, + ref_time=ref_time, + ref_bytes=ref_bytes, + ) + benchmark( + triton_sync_embedding_backward, + dy, + input_ids, + grad_ptr, + grad_dtype, + ref_time=ref_time, + ref_bytes=ref_bytes, + ) + benchmark( + triton_embedding_backward, + dy, + input_ids, + grad_ptr, + grad_dtype, + ref_time=ref_time, + ref_bytes=ref_bytes, + ) diff --git a/tests/test_fp32_gemm.py b/tests/test_fp32_gemm.py index 3f60d18..1333971 100644 --- a/tests/test_fp32_gemm.py +++ b/tests/test_fp32_gemm.py @@ -3,93 +3,299 @@ Copyright (c) Ant Financial Service Group and its affiliates. """ +import pytest + import torch -from linghe.gemm.fp32_gemm import (triton_fp32_gemm, - triton_fp32_gemm_for_backward, - triton_fp32_gemm_for_update, - triton_scaled_fp32_gemm, - triton_scaled_fp32_gemm_for_update) -from linghe.tools.benchmark import benchmark_func -from linghe.tools.util import output_check +import triton +from packaging.version import Version as _Version +from linghe.facade.fp32_gemm import fp32_gemm +from linghe.gemm.fp32_gemm import ( + triton_fp32_gemm, + triton_fp32_gemm_for_backward, + triton_fp32_gemm_for_update, + triton_split_fp32_gemm, + triton_split_fp32_gemm_for_backward, + triton_split_fp32_gemm_for_update, +) +from linghe.experimental.gemm import triton_tma_persistent_matmul +from linghe.tools.check import output_check + +def torch_fp64_matmul(x, w): + return torch.nn.functional.linear(x.to(torch.float64), w.to(torch.float64)).to( + torch.float32 + ) def torch_fp32_matmul(x, w): - return torch.nn.functional.linear(x.float(), w.float()) + return torch.nn.functional.linear(x.to(torch.float32), w.to(torch.float32)) + def torch_fp32_matmul_backward(dy, w): return (dy @ w).to(torch.bfloat16) -def torch_fp32_matmul_update(y, x): - return (y.t() @ x).to(torch.bfloat16) +def torch_fp32_matmul_update(dy, x): + return (dy.transpose(-2, -1) @ x).to(torch.bfloat16) -def test_fp32_matmul(M=2048, N=256, K=8192, bench=False): - # M, N, K = 4096, 256, 8192 + +@pytest.mark.parametrize( + "M,N,K", + [ + (4096, 256, 8192), + (16384, 128, 1536), + (16384, 256, 2048), + (16384 - 32, 256, 2048), + (128, 16, 128), + (32, 16, 128), + ], +) +def test_fp32_matmul(M, N, K, benchmark): dtype = torch.bfloat16 - device = 'cuda:0' - n_repeat = 100 + device = "cuda:0" - x = torch.randn(M, K, dtype=dtype, device=device) - w = torch.randn(N, K, dtype=dtype, device=device) - scale = torch.randn(M, dtype=torch.float32, device=device) + x = torch.randn(M, K, dtype=dtype, device=device, requires_grad=True) + w = (torch.randn(N, K, dtype=dtype, device=device) * 0.1).requires_grad_() dy = torch.randn(M, N, dtype=torch.float32, device=device) y_ref = torch_fp32_matmul(x, w) + y_ref.backward(gradient=dy) + dx_ref = x.grad + dw_ref = w.grad + y = triton_fp32_gemm(x, w) - output_check(y_ref, y.float(), mode='fp32_gemm') + dx = triton_fp32_gemm_for_backward(dy, w) + dw = triton_fp32_gemm_for_update(dy, x) - y_ref = torch_fp32_matmul(x * scale[:, None], w) - y = triton_scaled_fp32_gemm(x, w, scale) - output_check(y_ref, y.float(), mode='scaled_fp32_gemm') + output_check(y_ref, y, name="y", atol=5e-3, rtol=2e-3) + output_check(dx_ref, dx, name="dx", atol=2e-2, rtol=2e-2) + output_check(dw_ref, dw.to(dtype), name="dw", atol=2e-1, rtol=2e-2) - dx = torch.zeros(M, K, dtype=dtype, device=device) - dx = triton_fp32_gemm_for_backward(dy, w) - dx_ref = dy @ w.float() - output_check(dx_ref, dx.float(), mode='backward') - - main_grad = triton_fp32_gemm_for_update(y, x) - main_grad_ref = y.t() @ (x.float()) - output_check(main_grad_ref, main_grad.float(), mode='update') - - main_grad = triton_scaled_fp32_gemm_for_update(y, x, scale) - main_grad_ref = y.t() @ (x.float() * scale[:, None]) - output_check(main_grad_ref, main_grad.float(), mode='scaled_update') - - if bench: - print('\nbenchmark\n') - ref_time = benchmark_func(torch_fp32_matmul, x, w, n_repeat=n_repeat, - ref_bytes=M * K * 6 + N * K * 6 + M * N * 4, - ref_linghe=2 * M * N * K) - benchmark_func(triton_fp32_gemm, x, w, n_repeat=n_repeat, - ref_bytes=M * K * 6 + N * K * 6 + M * N * 4, - ref_linghe=2 * M * N * K, ref_time=ref_time) - benchmark_func(triton_scaled_fp32_gemm, x, w, scale, n_repeat=n_repeat, - ref_bytes=M * K * 6 + N * K * 6 + M * N * 4, - ref_linghe=2 * M * N * K, ref_time=ref_time) - - ref_time = benchmark_func(torch_fp32_matmul_backward, dy, w.float(), - n_repeat=n_repeat, - ref_bytes=M * K * 10 + N * K * 4 + M * N * 4, - ref_linghe=2 * M * N * K) - benchmark_func(triton_fp32_gemm_for_backward, dy, w, - n_repeat=n_repeat, - ref_bytes=M * K * 2 + N * K * 2 + M * N * 4, - ref_linghe=2 * M * N * K, ref_time=ref_time) - - ref_time = benchmark_func(torch_fp32_matmul_update, dy, x.float(), - n_repeat=n_repeat, - ref_bytes=M * K * 4 + N * K * 12 + M * N * 4, - ref_linghe=2 * M * N * K) - benchmark_func(triton_fp32_gemm_for_update, dy, x, n_repeat=n_repeat, - ref_bytes=M * K * 2 + N * K * 8 + M * N * 4, - ref_linghe=2 * M * N * K, ref_time=ref_time) - benchmark_func(triton_scaled_fp32_gemm_for_update, dy, x, scale, - n_repeat=n_repeat, - ref_bytes=M * K * 2 + N * K * 8 + M * N * 4, - ref_linghe=2 * M * N * K, ref_time=ref_time) - - -if __name__ == '__main__': - test_fp32_matmul(M=2048, N=256, K=8192) + print("") + ref_bytes = M * K * 6 + N * K * 6 + M * N * 4 + ref_flops = 2 * M * N * K + ref_time = benchmark( + torch_fp32_matmul, x, w, ref_bytes=ref_bytes, ref_flops=ref_flops + ) + benchmark( + triton_fp32_gemm, + x, + w, + ref_bytes=ref_bytes, + ref_flops=ref_flops, + ref_time=ref_time, + ) + + print("") + ref_bytes = M * K * 10 + N * K * 4 + M * N * 4 + ref_time = benchmark( + torch_fp32_matmul_backward, + dy, + w.float(), + ref_bytes=ref_bytes, + ref_flops=ref_flops, + ) + benchmark( + triton_fp32_gemm_for_backward, + dy, + w, + ref_bytes=ref_bytes, + ref_flops=ref_flops, + ref_time=ref_time, + ) + + print("") + ref_bytes = M * K * 4 + N * K * 12 + M * N * 4 + ref_time = benchmark( + torch_fp32_matmul_update, + dy, + x.float(), + ref_bytes=ref_bytes, + ref_flops=ref_flops, + ) + benchmark( + triton_fp32_gemm_for_update, + dy, + x, + ref_bytes=ref_bytes, + ref_flops=ref_flops, + ref_time=ref_time, + ) + + +@pytest.mark.parametrize( + "M,N,K", + [ + (16384, 128, 1536), + (16384, 256, 2048), + (16384, 256, 4096), + ], +) +def test_split_fp32_matmul(M, N, K, benchmark): + dtype = torch.bfloat16 + device = "cuda:0" + + x = torch.randn(M, K, dtype=dtype, device=device, requires_grad=True) + w = (torch.randn(N, K, dtype=dtype, device=device) * 0.1).requires_grad_() + dy = torch.randn(M, N, dtype=torch.float32, device=device) + + y_ref = torch_fp32_matmul(x, w) + y_ref.backward(gradient=dy) + dx_ref = x.grad + dw_ref = w.grad + + y = triton_split_fp32_gemm(x, w) # may inplace update + dx = triton_split_fp32_gemm_for_backward(dy, w) + dw = triton_split_fp32_gemm_for_update(dy, x) + output_check(y_ref, y, name="split.y", atol=5e-3, rtol=2e-3) + output_check(dx_ref, dx, name="split.dx", atol=2e-2, rtol=2e-2) + output_check(dw_ref, dw.to(dtype), name="split.dw", atol=2e-1, rtol=2e-1) + + print("") + ref_bytes = M * K * 6 + N * K * 6 + M * N * 4 + ref_flops = 2 * M * N * K + ref_time = benchmark( + torch_fp32_matmul, x, w, ref_bytes=ref_bytes, ref_flops=ref_flops + ) + benchmark( + triton_split_fp32_gemm, + x, + w, + ref_bytes=ref_bytes, + ref_flops=ref_flops, + ref_time=ref_time, + ) + + print("") + ref_bytes = M * K * 10 + N * K * 4 + M * N * 4 + ref_time = benchmark( + torch_fp32_matmul_backward, + dy, + w.float(), + ref_bytes=ref_bytes, + ref_flops=ref_flops, + ) + benchmark( + triton_split_fp32_gemm_for_backward, + dy, + w, + ref_bytes=ref_bytes, + ref_flops=ref_flops, + ref_time=ref_time, + ) + + print("") + ref_bytes = M * K * 4 + N * K * 12 + M * N * 4 + ref_time = benchmark( + torch_fp32_matmul_update, + dy, + x.float(), + ref_bytes=ref_bytes, + ref_flops=ref_flops, + ) + benchmark( + triton_split_fp32_gemm_for_update, + dy, + x, + ref_bytes=ref_bytes, + ref_flops=ref_flops, + ref_time=ref_time, + n_profile=0, + ) + + +@pytest.mark.skipif( + _Version(triton.__version__) <= _Version("3.4.0"), + reason="requires Triton > 3.4.0", +) +@pytest.mark.parametrize( + "M,N,K", + [ + (16384, 256, 2048), + ], +) +def test_tma_fp32_matmul(M, N, K, benchmark): + dtype = torch.bfloat16 + device = "cuda:0" + + x = torch.randn(M, K, dtype=dtype, device=device, requires_grad=True) + w = (torch.randn(N, K, dtype=dtype, device=device) * 0.1).requires_grad_() + dy = torch.randn(M, N, dtype=torch.float32, device=device) + + y_ref = torch_fp32_matmul(x, w) + y = triton_tma_persistent_matmul(x, w) + output_check(y_ref, y, name="persist.y", atol=5e-3, rtol=2e-3) + + print("") + ref_bytes = M * K * 6 + N * K * 6 + M * N * 4 + ref_flops = 2 * M * N * K + ref_time = benchmark( + torch_fp32_matmul, x, w, ref_bytes=ref_bytes, ref_flops=ref_flops + ) + benchmark( + triton_tma_persistent_matmul, + x, + w, + ref_bytes=ref_bytes, + ref_flops=ref_flops, + ref_time=ref_time, + ) + + +@pytest.mark.skipif( + _Version(triton.__version__) <= _Version("3.4.0"), + reason="requires Triton > 3.4.0", +) +@pytest.mark.parametrize( + "B,M,N,K,impl", + [ + (2, 8192, 128, 1536, "native"), + (2, 8192 - 32, 128, 1536, "native"), + (2, 8192, 128, 1536, "tma"), + (2, 8192, 128, 1536, "split"), + ], +) +def test_fp32_facade(B, M, N, K, impl, benchmark): + # M, N, K = 4096, 256, 8192 + dtype = torch.bfloat16 + device = "cuda:0" + n_repeat = 100 + + x = torch.randn(B, M, K, dtype=dtype, device=device, requires_grad=True) + w = torch.randn(N, K, dtype=dtype, device=device, requires_grad=True) + dy = torch.randn(B, M, N, dtype=torch.float32, device=device) + + y_ref = torch_fp32_matmul(x, w) + y_ref.backward(gradient=dy) + dx_ref = x.grad + dw_ref = w.grad + + x.grad = None + w.grad = None + y = fp32_gemm(x, w, impl=impl) + y.backward(gradient=dy) + dx = x.grad + dw = w.grad + output_check(y_ref, y, name="forward", atol=5e-3, rtol=2e-3) + output_check(dx_ref, dx, name="backward", atol=2e-1, rtol=2e-2) + output_check(dw_ref, dw, name="update", atol=2e-1, rtol=2e-2) + + ref_time = benchmark( + torch_fp32_matmul, + x, + w, + n_repeat=n_repeat, + ref_bytes=M * K * 6 + N * K * 6 + M * N * 4, + ref_flops=2 * M * N * K, + ) + benchmark( + fp32_gemm, + x, + w, + impl=impl, + n_repeat=n_repeat, + ref_bytes=M * K * 6 + N * K * 6 + M * N * 4, + ref_flops=2 * M * N * K, + ref_time=ref_time, + ) diff --git a/tests/test_gate.py b/tests/test_gate.py new file mode 100644 index 0000000..9761a3a --- /dev/null +++ b/tests/test_gate.py @@ -0,0 +1,283 @@ +# -*- coding: utf-8 -*- +""" +Copyright (c) Ant Financial Service Group and its affiliates. +""" + +import pytest +import torch +import torch.nn.functional as F + +from linghe.tools.check import output_check +from linghe.utils.gate import ( + triton_group_rms_norm_gate_forward, + triton_group_rms_norm_gate_backward, + triton_group_rms_norm_gate_and_mxfp8_quant_forward, + triton_group_rms_norm_gate_and_mxfp8_quant_backward, +) +from linghe.tools.util import torch_mxfp8_quant + + +# @torch.compile +def torch_group_rms_norm_gate_forward( + x, gate, weight, eps=1e-6, group_size=4, native=True, high_precison=False +): + dtype = x.dtype + if not native: + x = torch.permute(x, [1, 0, 2]) + x = x.float() + gate = gate.float() + weight = weight.float() + length, bs, dim = gate.shape + d = dim // group_size + attn_output = x.view(bs, length, group_size, d) + outputs = [] + for i in range(group_size): + if weight.size(0) == dim: + o = F.rms_norm( + attn_output[:, :, i], [d], weight=weight[i * d : (i + 1) * d], eps=eps + ) + else: + o = F.rms_norm(attn_output[:, :, i], [d], weight=weight, eps=eps) + outputs.append(o) + outputs = torch.stack(outputs, 2).view(bs, length, dim) + outputs = outputs.transpose(0, 1) + gate = F.sigmoid(gate) + if high_precison: + outputs = outputs * gate + else: + outputs = (outputs * gate).to(dtype) + return outputs + + +def torch_group_rms_norm_gate_mxfp8_quant_forward( + x, gate, weight, eps=1e-6, group_size=4, native=True +): + out = torch_group_rms_norm_gate_forward( + x, gate, weight, eps, group_size, native, high_precison=True + ) + out = out.reshape(-1, x.size(-1)) + x_q_ref, x_scale_ref, xt_q_ref, xt_scale_ref = torch_mxfp8_quant(out) + return x_q_ref, x_scale_ref, xt_q_ref, xt_scale_ref + + +def torch_group_rms_norm_gate_mxfp8_quant_backward( + grad_output, x, gate, weight, eps=1e-6, group_size=4, native=True +): + dtype = grad_output.dtype + grad_output = grad_output.float() + x = x.float().clone().detach().requires_grad_() + gate = gate.float().clone().detach().requires_grad_() + weight = weight.float().clone().detach().requires_grad_() + y = torch_group_rms_norm_gate_forward( + x, + gate, + weight, + eps=eps, + group_size=group_size, + native=native, + high_precison=True, + ) + y.backward(gradient=grad_output) + + dx = x.grad.to(dtype) + dg = gate.grad + dw = weight.grad.to(dtype) + dg = dg.reshape(-1, gate.size(-1)) + g_q, g_scale, gt_q, gt_scale = torch_mxfp8_quant(dg) + + return dx, g_q, g_scale, gt_q, gt_scale, dw + + +def torch_group_rms_norm_gate_backward( + grad_output, x, gate, weight, eps=1e-6, group_size=4, native=True +): + dtype = grad_output.dtype + grad_output = grad_output.float() + x = x.float().clone().detach().requires_grad_() + gate = gate.float().clone().detach().requires_grad_() + weight = weight.float().clone().detach().requires_grad_() + y = torch_group_rms_norm_gate_forward( + x, gate, weight, eps=eps, group_size=group_size + ) + y.backward(gradient=grad_output) + return x.grad.to(dtype), gate.grad.to(dtype), weight.grad.to(dtype) + + +@pytest.mark.parametrize( + "bs,length,dim,group_size,contiguous,share,coef,grad_coef,native", + [ + (2, 4096, 2048, 4, True, False, 1.0, 1.0, True), + (2, 4096, 2048, 4, True, True, 1.0, 1.0, True), + (1, 4096, 4096, 4, True, False, 1.0, 1.0, True), + (2, 4096, 1536, 4, True, False, 1.0, 1.0, True), + (2, 4096, 1536, 4, True, False, 10000.0, 10000.0, True), + (2, 4096, 1536, 4, True, False, 0.0, 0.0, True), + (2, 4096, 1536, 4, False, False, 1.0, 1.0, True), + (2, 4096, 1536, 4, False, False, 1.0, 1.0, False), + ], +) +def test_group_rms_norm_gate( + bs, length, dim, group_size, contiguous, share, coef, grad_coef, native, benchmark +): + dtype = torch.bfloat16 + device = "cuda:0" + if native: + x = torch.randn(bs, length, dim, dtype=dtype, requires_grad=True, device=device) + else: + x = torch.randn(length, bs, dim, dtype=dtype, requires_grad=True, device=device) + weight = torch.randn( + dim // group_size if share else dim, + dtype=dtype, + requires_grad=True, + device=device, + ) + if contiguous: + gate = ( + torch.randn(length, bs, dim, dtype=dtype, device=device) * coef + ).requires_grad_() + else: + tmp = torch.randn(length, bs, 3 * dim, dtype=dtype, device=device) * coef + split_sizes = [dim, dim, dim] + _, _, gate = torch.split(tmp, split_sizes, dim=-1) + gate = gate.requires_grad_() + + grad_output = torch.randn(length, bs, dim, dtype=dtype, device=device) * grad_coef + + output_ref = torch_group_rms_norm_gate_forward( + x, gate, weight, group_size=group_size, native=native + ) + output_ref.backward(gradient=grad_output) + dx_ref = x.grad.to(dtype) + dg_ref = gate.grad.to(dtype) + dw_ref = weight.grad.to(dtype) + + output = triton_group_rms_norm_gate_forward(x, gate, weight, group_size=group_size) + output_check(output_ref, output, name="group_norm_gate.y") + + dx, dg, dw = triton_group_rms_norm_gate_backward( + grad_output, x, gate, weight, group_size=group_size + ) + output_check(dx_ref, dx, name="group_norm_gate.dx") + output_check(dg_ref, dg, name="group_norm_gate.dg") + output_check(dw_ref, dw.to(dtype), name="group_norm_gate.dw") + + benchmark( + torch_group_rms_norm_gate_forward, + x, + gate, + weight, + group_size=group_size, + ref_bytes=bs * length * dim * 6, + ) + + benchmark( + triton_group_rms_norm_gate_forward, + x, + gate, + weight, + group_size=group_size, + ref_bytes=bs * length * dim * 6, + ) + + benchmark( + triton_group_rms_norm_gate_backward, + grad_output, + x, + gate, + weight, + group_size=group_size, + ref_bytes=bs * length * dim * 10, + ) + + +@pytest.mark.parametrize( + "bs,length,dim,group_size,contiguous,share,coef,grad_coef,native", + [ + (2, 4096, 2048, 4, False, False, 1.0, 1.0, False), + (2, 4096, 2048, 4, False, False, 1.0, 1.0, True), + (2, 4096, 2048, 4, True, False, 1.0, 1.0, False), + (2, 4096, 2048, 4, False, False, 1.0, 1.0, False), + (2, 4096, 2048, 4, True, False, 1.0, 1.0, True), + ], +) +def test_group_rms_norm_gate_quant( + bs, length, dim, group_size, contiguous, share, coef, grad_coef, native, benchmark +): + + dtype = torch.bfloat16 + device = "cuda:0" + if native: + x = torch.randn(bs, length, dim, dtype=dtype, requires_grad=True, device=device) + else: + x = torch.randn(length, bs, dim, dtype=dtype, requires_grad=True, device=device) + weight = torch.randn( + dim // group_size if share else dim, + dtype=dtype, + requires_grad=True, + device=device, + ) + if contiguous: + gate = ( + torch.randn(length, bs, dim, dtype=dtype, device=device) * coef + ).requires_grad_() + else: + tmp = torch.randn(length, bs, 3 * dim, dtype=dtype, device=device) * coef + split_sizes = [dim, dim, dim] + _, _, gate = torch.split(tmp, split_sizes, dim=-1) + gate = gate.requires_grad_() + + grad_output = torch.randn(length, bs, dim, dtype=dtype, device=device) * grad_coef + + x_q_ref, x_scale_ref, xt_q_ref, xt_scale_ref = ( + torch_group_rms_norm_gate_mxfp8_quant_forward( + x, gate, weight, group_size=group_size, native=native + ) + ) + x_q, x_scale, xt_q, xt_scale = triton_group_rms_norm_gate_and_mxfp8_quant_forward( + x, gate, weight, group_size=group_size + ) + + 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") + + dx_ref, g_q_ref, g_scale_ref, gt_q_ref, gt_scale_ref, dw_ref = ( + torch_group_rms_norm_gate_mxfp8_quant_backward( + grad_output, x, gate, weight, group_size=group_size, native=native + ) + ) + + dx, g_q, g_scale, gt_q, gt_scale, dw = ( + triton_group_rms_norm_gate_and_mxfp8_quant_backward( + grad_output, x, gate, weight, group_size=group_size + ) + ) + + output_check(dx_ref, dx, "dx") + output_check(dw_ref, dw, "dw") + output_check(g_q_ref, g_q, "gq") + output_check(g_scale_ref, g_scale, "x_scale") + output_check(gt_q_ref, gt_q, "xt_q") + output_check(gt_scale_ref, gt_scale, "xt_scale") + + benchmark( + triton_group_rms_norm_gate_and_mxfp8_quant_forward, + x, + gate, + weight, + 1e-6, + group_size, + n_repeat=100, + ) + + benchmark( + triton_group_rms_norm_gate_and_mxfp8_quant_backward, + grad_output, + x, + gate, + weight, + 1e-6, + group_size, + n_repeat=100, + ) diff --git a/tests/test_gather.py b/tests/test_gather.py index 4606284..032aa62 100644 --- a/tests/test_gather.py +++ b/tests/test_gather.py @@ -3,50 +3,65 @@ Copyright (c) Ant Financial Service Group and its affiliates. """ +import pytest import torch -from linghe.utils.gather import (triton_make_row_id_map, - triton_make_row_id_map_and_indices, - triton_index_select, - triton_permute_with_mask_map, - triton_smooth_permute_with_indices, - triton_smooth_permute_with_mask_map, - triton_smooth_weighted_permute_with_indices, - triton_batch_transpose_smooth_permute_with_indices) -from linghe.tools.util import (output_check, - torch_batch_smooth_quant, - torch_make_indices, - torch_smooth_quant) -from linghe.tools.benchmark import benchmark_func +from linghe.quant.block import triton_batch_blockwise_quant +from linghe.quant.mxfp8 import triton_batch_mxfp8_quant +from linghe.tools.check import output_check +from linghe.tools.util import ( + torch_batch_smooth_quant, + torch_blockwise_quant, + torch_make_indices, + torch_smooth_quant, + torch_mxfp8_quant, + torch_make_chunk_sort_map, +) +from linghe.utils.gather import ( + triton_make_row_id_map, + triton_make_row_id_map_and_index, + triton_permute_with_indices, + triton_permute_with_mask_map, + triton_batch_smooth_permute_with_indices, + triton_batch_transpose_smooth_permute_with_indices, + triton_batch_smooth_fused_permute_with_indices, + triton_batch_transpose_smooth_fused_permute_with_indices, + triton_batch_block_pad_permute_with_indices, + triton_batch_mxfp8_permute_with_indices, + triton_make_chunk_sort_map, +) def torch_index_select(y, indices): output = y.index_select(0, indices) return output + def torch_select_with_padded_map_mask(y, mask_map, out_tokens): E = mask_map.shape[1] if y.ndim > 1: output = torch.zeros((out_tokens, y.shape[1]), dtype=y.dtype, device=y.device) else: - output = torch.zeros((out_tokens, ), dtype=y.dtype, device=y.device) + output = torch.zeros((out_tokens,), dtype=y.dtype, device=y.device) for i in range(E): - indices = mask_map[:,i] - src_idx = torch.nonzero(indices>-1) + indices = mask_map[:, i] + src_idx = torch.nonzero(indices > -1) dst_idx = indices[src_idx] output[dst_idx] = y[src_idx] return output + def torch_ravel_with_padded_map_mask(y, mask_map, out_tokens): E = mask_map.shape[1] - output = torch.zeros((out_tokens, ), dtype=y.dtype, device=y.device) + output = torch.zeros((out_tokens,), dtype=y.dtype, device=y.device) for i in range(E): - indices = mask_map[:,i] - src_idx = torch.nonzero(indices>-1) + indices = mask_map[:, i] + src_idx = torch.nonzero(indices > -1) dst_idx = indices[src_idx] - output[dst_idx] = y[src_idx,i] + output[dst_idx] = y[src_idx, i] return output + def torch_fp16_index_select(x, scales, indices): return x.index_select(0, indices), scales.index_select(0, indices) @@ -55,48 +70,112 @@ def torch_scatter(logits, routing_map, weights): logits[routing_map] = weights -# optional dequant and smooth and quant -def torch_smooth_permute_with_indices(grad_data, grad_scale, indices, - smooth_scales, - token_count_per_expert_list, - round_scale=True): +# smooth and quant +def torch_smooth_permute_with_indices( + x, indices, smooth_scales, token_count_per_expert_list, probs=None, round_scale=True +): + M, N = x.shape + q_refs = [] + scale_refs = [] + prob_refs = [] + s = 0 + for i, c in enumerate(token_count_per_expert_list): + c = token_count_per_expert_list[i] + data_slice = x[indices[s : s + c]] + prob_slice = probs[:, i][indices[s : s + c]] if probs is not None else None + y_smooth = data_slice.float() / smooth_scales[i] + scale = y_smooth.abs().amax(1) / 448 + if round_scale: + scale = torch.exp2(torch.ceil(torch.log2(scale))) + scale_refs.append(scale) + q = (y_smooth / scale[:, None]).to(torch.float8_e4m3fn) + q_refs.append(q) + prob_refs.append(prob_slice) + s += c + q_ref = torch.cat(q_refs, 0) + scale_ref = torch.cat(scale_refs, 0) + if probs is not None: + prob_ref = torch.cat(prob_refs) + return q_ref, scale_ref, prob_ref + + +# dequant, desmooth, smooth, quant +def torch_smooth_fused_permute_with_indices( + grad_data, + grad_scale, + indices, + org_smooth_scales, + smooth_scales, + token_count_per_expert_list, + round_scale=True, +): M, N = grad_data.shape - if grad_scale is not None: - B = grad_data.shape[1] // ( - 1 if grad_scale.ndim == 1 else grad_scale.shape[1]) q_refs = [] scale_refs = [] s = 0 + smooth_scales = smooth_scales * org_smooth_scales for i, c in enumerate(token_count_per_expert_list): c = token_count_per_expert_list[i] - data_slice = grad_data.view(torch.uint8)[indices[s:s + c]].view( - torch.float8_e4m3fn) - if grad_scale is not None: - scale_slice = grad_scale[indices[s:s + c]] - y_smooth = (data_slice.float().view(c, N // B, B) * scale_slice[:, :, - None]).view(c, N) / \ - smooth_scales[i] - else: - y_smooth = data_slice.float() / smooth_scales[i] + data_slice = grad_data[indices[s : s + c]] + scale_slice = grad_scale[indices[s : s + c]] + dsr = data_slice.float() + y_smooth = (dsr * scale_slice[:, None]).view(c, N) / smooth_scales[i] + scale = y_smooth.abs().amax(1) / 448 if round_scale: scale = torch.exp2(torch.ceil(torch.log2(scale))) scale_refs.append(scale) q = (y_smooth / scale[:, None]).to(torch.float8_e4m3fn) - q_refs.append(q.view(torch.uint8)) + q_refs.append(q) s += c - q_ref = torch.cat(q_refs, 0).view(torch.float8_e4m3fn) + q_ref = torch.cat(q_refs, 0) scale_ref = torch.cat(scale_refs, 0) return q_ref, scale_ref +# gather, pad, transpose, smooth, quant +def torch_batch_transpose_smooth_permute_with_indices( + x, smooth_scales, indices, token_count_per_expert_list, round_scale=True +): + M, DIM = x.shape + q_refs = [] + scale_refs = [] + s = 0 + for i, c in enumerate(token_count_per_expert_list): + c = token_count_per_expert_list[i] + if c == 0: + y_scale = torch.zeros((DIM,), dtype=torch.float32, device=x.device) + scale_refs.append(y_scale.view(-1)) + continue + N = (c + 31) // 32 * 32 + data_slice = x[indices[s : s + c]] + y = data_slice.float() + smooth_scale = smooth_scales[s : s + c] + if N > c: + y = torch.nn.functional.pad(y, (0, 0, 0, N - c)) + smooth_scale = torch.nn.functional.pad(smooth_scale, (0, N - c)) + y_q, y_scale = torch_smooth_quant( + y.t().contiguous(), smooth_scale, reverse=True, round_scale=round_scale + ) + scale_refs.append(y_scale) + q_refs.append(y_q.view(N, DIM)) + s += c + q_ref = torch.cat(q_refs, 0) + scale_ref = torch.stack(scale_refs, 0) + return q_ref, scale_ref + -# desmooth,dequant, gather, pad, transpose, smooth, quant -def torch_batch_transpose_smooth_permute_with_indices(x_q, x_scale, org_smooth_scale, smooth_scales, - indices, - token_count_per_expert_list, - round_scale=True): +# desmooth, dequant, gather, pad, transpose, smooth, quant +def torch_batch_transpose_smooth_fused_permute_with_indices( + x_q, + x_scale, + org_smooth_scale, + smooth_scales, + indices, + token_count_per_expert_list, + round_scale=True, +): M, DIM = x_q.shape q_refs = [] scale_refs = [] @@ -107,18 +186,20 @@ def torch_batch_transpose_smooth_permute_with_indices(x_q, x_scale, org_smooth_s y_scale = torch.zeros((DIM,), dtype=torch.float32, device=x_q.device) scale_refs.append(y_scale.view(-1)) continue - N = (c + 31)//32 * 32 - data_slice = x_q[indices[s:s + c]] + N = (c + 31) // 32 * 32 + data_slice = x_q[indices[s : s + c]] if x_scale is not None: - scale_slice = x_scale[indices[s:s + c]] + scale_slice = x_scale[indices[s : s + c]] y = data_slice.float() * scale_slice[:, None] * org_smooth_scale else: y = data_slice.float() - smooth_scale = smooth_scales[s:s+c] + smooth_scale = smooth_scales[s : s + c] if N > c: - y = torch.nn.functional.pad(y, (0,0,0, N-c)) - smooth_scale = torch.nn.functional.pad(smooth_scale, (0, N-c)) - y_q, y_scale, y_max= torch_smooth_quant(y.t().contiguous(), smooth_scale, reverse=True, round_scale=round_scale) + y = torch.nn.functional.pad(y, (0, 0, 0, N - c)) + smooth_scale = torch.nn.functional.pad(smooth_scale, (0, N - c)) + y_q, y_scale = torch_smooth_quant( + y.t().contiguous(), smooth_scale, reverse=True, round_scale=round_scale + ) scale_refs.append(y_scale.view(-1)) q_refs.append(y_q.view(-1)) s += c @@ -127,67 +208,158 @@ def torch_batch_transpose_smooth_permute_with_indices(x_q, x_scale, org_smooth_s return q_ref, scale_ref -def test_make_id_map(M=4098, n_experts=32, topk=2, bias=0.0, bench=False): +def torch_batch_block_pad_permute_with_indices( + x, indices, probs, token_count_per_expert_list, round_scale=True +): + M, DIM = x.shape + if M == 0: + device = x.device + q_ref = torch.empty((0,), device=device, dtype=torch.float8_e4m3fn) + s_ref = torch.empty((0,), device=device, dtype=torch.float32) + qt_ref = torch.empty((0,), device=device, dtype=torch.float8_e4m3fn) + st_ref = torch.empty((0,), device=device, dtype=torch.float32) + probs_refs = torch.empty((0,), device=device, dtype=torch.float32) + return q_ref, s_ref, qt_ref, st_ref, probs_refs + + q_refs = [] + s_refs = [] + qt_refs = [] + st_refs = [] + probs_refs = [] + s = 0 + for i, c in enumerate(token_count_per_expert_list): + c = token_count_per_expert_list[i] + if c == 0: + continue + index = indices[s : s + c] + assert len(index) == c + y = x[index] + y = y.float() + p_slice = probs[:, i][index] + + padding_size = (c + 15) // 16 * 16 - c + if padding_size > 0: + p_slice = torch.nn.functional.pad(p_slice, (0, padding_size)) + + y_q, y_scale, yt_q, yt_scale = torch_blockwise_quant( + y, round_scale=round_scale, padding=True + ) + q_refs.append(y_q.view(-1)) + s_refs.append(y_scale.view(-1)) + qt_refs.append(yt_q.view(-1)) + st_refs.append(yt_scale.view(-1)) + probs_refs.append(p_slice) + s += c + q_ref = torch.cat(q_refs, 0) + s_ref = torch.cat(s_refs, 0) + qt_ref = torch.cat(qt_refs, 0) + st_ref = torch.cat(st_refs, 0) + probs_refs = torch.cat(probs_refs, 0) + return q_ref, s_ref, qt_ref, st_ref, probs_refs + + +def torch_batch_mxfp8_permute_with_indices( + x, indices, probs, token_count_per_expert_list +): + M, DIM = x.shape + if M == 0: + device = x.device + q_ref = torch.empty((0, DIM), device=device, dtype=torch.float8_e4m3fn) + s_ref = torch.empty((0, DIM // 32), device=device, dtype=torch.float32) + qt_ref = torch.empty((0, DIM), device=device, dtype=torch.float8_e4m3fn) + st_ref = torch.empty((0, DIM), device=device, dtype=torch.float32) + probs_refs = torch.empty((0,), device=device, dtype=torch.float32) + return q_ref, s_ref, qt_ref, st_ref, probs_refs + + q_refs = [] + s_refs = [] + qt_refs = [] + st_refs = [] + probs_refs = [] + s = 0 + for i, c in enumerate(token_count_per_expert_list): + c = token_count_per_expert_list[i] + if c == 0: + continue + index = indices[s : s + c] + assert len(index) == c + y = x[index] + y = y.float() + p_slice = probs[:, i][index] + + padding_size = (c + 31) // 32 * 32 - c + if padding_size > 0: + p_slice = torch.nn.functional.pad(p_slice, (0, padding_size)) + + y_q, y_scale, yt_q, yt_scale = torch_mxfp8_quant(y, padding=True, zero=True) + q_refs.append(y_q) + s_refs.append(y_scale) + qt_refs.append(yt_q) + st_refs.append(yt_scale) + probs_refs.append(p_slice) + s += c + q_ref = torch.cat(q_refs, 0) + s_ref = torch.cat(s_refs, 0) + qt_ref = torch.cat(qt_refs, 0) + st_ref = torch.cat(st_refs, 0) + probs_refs = torch.cat(probs_refs, 0) + return q_ref, s_ref, qt_ref, st_ref, probs_refs + + +@pytest.mark.parametrize( + "M,n_experts,topk,bias", + [ + (4098, 32, 2, 0.0), + ], +) +def test_make_id_map(M, n_experts, topk, bias, benchmark, ep=4): dtype = torch.bfloat16 - device = 'cuda:0' + device = "cuda:0" 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=bias) + logits, topk=topk, bias=bias + ) token_count_per_expert_list = token_count_per_expert.tolist() out_tokens = sum(token_count_per_expert_list) - row_id_map_output = triton_make_row_id_map(mask_map) assert (row_id_map - row_id_map_output).abs().sum().item() == 0 - _, row_id_indices = triton_make_row_id_map_and_indices(mask_map, out_tokens) + _, row_id_indices = triton_make_row_id_map_and_index(mask_map, out_tokens) assert (row_id_indices - indices).abs().sum().item() == 0 - - -def test_triton_smooth_weighted_permute_with_indices(M=4096, N=4096, - n_experts=256, - topk=8, - round_scale=True, - bench=False): - device = 'cuda:0' - reverse = True - y = torch.randn((M, N), dtype=torch.bfloat16, device=device) - logits = torch.randn((M, n_experts), dtype=torch.float32, device=device) - smooth_scales = 1 + 10 * torch.rand((n_experts, N), device=device, - dtype=torch.float32) - probs, mask_map, token_count_per_expert, indices, row_id_map = torch_make_indices( - logits, topk=topk, bias=0.0) - - tokens = torch.randn((indices.shape[0], N), dtype=torch.bfloat16, - device=device) - y_q, y_scale, y_sum = triton_smooth_weighted_permute_with_indices( - y, tokens, smooth_scales, token_count_per_expert, indices, x_q=None, - x_scale=None, reverse=reverse, round_scale=round_scale) - - y_q_ref, y_scale_ref = torch_batch_smooth_quant(y, smooth_scales, indices, - token_count_per_expert, - reverse=reverse, - round_scale=round_scale) - sum_ref = (tokens * y[indices]).sum(1) - - output_check(y_q_ref.float(), y_q.float(), 'data') - output_check(y_scale_ref.float(), y_scale.float(), 'scale') - output_check(sum_ref.float(), y_sum.float(), 'sum') - - if bench: - n_repeat = 100 - benchmark_func(triton_smooth_weighted_permute_with_indices, - y, tokens, smooth_scales, token_count_per_expert, - indices, reverse=reverse, round_scale=round_scale, - n_repeat=n_repeat) - - -def test_triton_permute_with_mask_map(M=4096, N=4096, n_experts=256, topk=8, - bench=False): - device = 'cuda:0' + max_tokens_per_chunk = 100 + num_global_tokens_per_local_expert = torch.randint( + 0, + max_tokens_per_chunk, + (ep, n_experts // ep), + dtype=torch.int32, + device="cuda:0", + ) + token_per_expert = num_global_tokens_per_local_expert.sum(0).tolist() + row_id_map, resort_row_id_map = triton_make_chunk_sort_map( + num_global_tokens_per_local_expert, token_per_expert + ) + row_id_map_torch, resort_row_id_map_torch = torch_make_chunk_sort_map( + num_global_tokens_per_local_expert + ) + + assert torch.equal(row_id_map_torch, row_id_map) + assert torch.equal(resort_row_id_map_torch, resort_row_id_map) + + +@pytest.mark.parametrize( + "M,N,n_experts,topk", + [ + (16384, 2048, 256, 8), + (8192, 4096, 256, 8), + (7628, 2048, 256, 8), + ], +) +def test_triton_permute_with_mask_map(M, N, n_experts, topk, benchmark): + device = "cuda:0" dtype = torch.bfloat16 x = torch.randn(M, N, dtype=dtype, device=device) ** 3 x_q = x.to(torch.float8_e4m3fn) @@ -196,209 +368,521 @@ def test_triton_permute_with_mask_map(M=4096, N=4096, n_experts=256, topk=8, 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.01) + logits, topk=topk, bias=-0.01 + ) out_tokens = sum(token_count_per_expert.tolist()) - x_out, scale_out = triton_index_select(x, indices, scale=scales) + x_out, scale_out = triton_permute_with_indices(x, indices, scale=scales) x_out_ref, scale_out_ref = torch_fp16_index_select(x, scales, indices) - output_check(x_out_ref, x_out, 'x_out') - output_check(scale_out_ref, scale_out, 'scale_out') - - probs_out_ref = probs.T.contiguous().masked_select( - (probs > 0).T.contiguous()) - x_out, scale_out, probs_out = triton_permute_with_mask_map(x, scales, probs, - row_id_map, - out_tokens, - contiguous=True) - output_check(x_out_ref, x_out, 'x_out') - output_check(scale_out_ref, scale_out, 'scale_out') - output_check(probs_out_ref, probs_out, 'prob_out') - - nzs = torch.sum(row_id_map>=0, 0) - bias = torch.cumsum((nzs + 15)//16*16 - nzs, 0) + output_check(x_out_ref, x_out, "x_out") + output_check(scale_out_ref, scale_out, "scale_out") + + probs_out_ref = probs.T.contiguous().masked_select((probs > 0).T.contiguous()) + x_out, scale_out, probs_out = triton_permute_with_mask_map( + x, scales, probs, row_id_map, out_tokens, contiguous=True + ) + output_check(x_out_ref, x_out, "x_out") + output_check(scale_out_ref, scale_out, "scale_out") + output_check(probs_out_ref, probs_out, "prob_out") + + nzs = torch.sum(row_id_map >= 0, 0) + bias = torch.cumsum((nzs + 15) // 16 * 16 - nzs, 0) row_id_map_clone = row_id_map.clone().detach() row_id_map_clone[:, 1:] += bias[:-1] - round_row_id_map = torch.where(row_id_map>=0, row_id_map_clone, -1) - padded_out_tokens = sum([(x+15)//16*16 for x in token_count_per_expert.tolist()]) - x_out_ref = torch_select_with_padded_map_mask(x, round_row_id_map, padded_out_tokens) - scale_out_ref = torch_select_with_padded_map_mask(scales, round_row_id_map, padded_out_tokens) - prob_out_ref = torch_ravel_with_padded_map_mask(probs, round_row_id_map, padded_out_tokens) - x_out, scale_out, probs_out = triton_permute_with_mask_map(x, scales, probs, - round_row_id_map, - padded_out_tokens, - contiguous=False, - tokens_per_expert=token_count_per_expert) - output_check(x_out_ref, x_out, 'noncontiguous.x_out') - output_check(scale_out_ref, scale_out, 'noncontiguous.scale_out') - output_check(prob_out_ref, probs_out, 'noncontiguous.prob') - - if bench: - n_repeat = 100 - ref_bytes = out_tokens * N * 2 - ref_time = benchmark_func(torch_fp16_index_select, x, scales, indices, - n_repeat=n_repeat, ref_bytes=ref_bytes) - benchmark_func(triton_index_select, x, indices, scale=scales, - n_repeat=n_repeat, ref_time=ref_time, ref_bytes=ref_bytes) - benchmark_func(triton_permute_with_mask_map, x, scales, probs, - row_id_map, out_tokens, contiguous=True, n_repeat=n_repeat, - ref_time=ref_time, ref_bytes=ref_bytes) - benchmark_func(triton_permute_with_mask_map, x, scales, probs, - row_id_map, out_tokens, contiguous=False, - token_per_expert=token_count_per_expert, - n_repeat=n_repeat, - ref_time=ref_time, ref_bytes=ref_bytes) - - -def test_triton_smooth_permute_with_mask_map(M=4096, N=4096, n_experts=32, - topk=8, round_scale=True, - bench=False): - device = 'cuda:0' + round_row_id_map = torch.where(row_id_map >= 0, row_id_map_clone, -1) + padded_out_tokens = sum( + [(x + 15) // 16 * 16 for x in token_count_per_expert.tolist()] + ) + x_out_ref = torch_select_with_padded_map_mask( + x, round_row_id_map, padded_out_tokens + ) + scale_out_ref = torch_select_with_padded_map_mask( + scales, round_row_id_map, padded_out_tokens + ) + prob_out_ref = torch_ravel_with_padded_map_mask( + probs, round_row_id_map, padded_out_tokens + ) + x_out, scale_out, probs_out = triton_permute_with_mask_map( + x, + scales, + probs, + round_row_id_map, + padded_out_tokens, + contiguous=False, + tokens_per_expert=token_count_per_expert, + ) + output_check(x_out_ref, x_out, "noncontiguous.x_out") + output_check(scale_out_ref, scale_out, "noncontiguous.scale_out") + output_check(prob_out_ref, probs_out, "noncontiguous.prob") + + n_repeat = 100 + ref_bytes = out_tokens * N * 2 + ref_time = benchmark( + torch_fp16_index_select, + x, + scales, + indices, + n_repeat=n_repeat, + ref_bytes=ref_bytes, + ) + benchmark( + triton_permute_with_indices, + x, + indices, + scale=scales, + n_repeat=n_repeat, + ref_time=ref_time, + ref_bytes=ref_bytes, + ) + benchmark( + triton_permute_with_mask_map, + x, + scales, + probs, + row_id_map, + out_tokens, + contiguous=True, + n_repeat=n_repeat, + ref_time=ref_time, + ref_bytes=ref_bytes, + ) + benchmark( + triton_permute_with_mask_map, + x, + scales, + probs, + row_id_map, + out_tokens, + contiguous=False, + tokens_per_expert=token_count_per_expert, + n_repeat=n_repeat, + ref_time=ref_time, + ref_bytes=ref_bytes, + ) + + +@pytest.mark.parametrize( + "M,N,n_experts,topk", + [ + (16384, 2048, 32, 2), + (8192, 4096, 32, 2), + ], +) +def test_batch_smooth_permute_with_indices( + M, N, n_experts, topk, benchmark, round_scale=True +): + device = "cuda:0" + dtype = torch.bfloat16 + smooth_scales = 1 + 10 * torch.rand( + (n_experts, N), device=device, dtype=torch.float32 + ) + logits = torch.randn((M, n_experts), dtype=torch.float32, device=device) ** 3 + probs, mask_map, token_count_per_expert, indices, row_id_map = torch_make_indices( + logits, topk=topk, bias=-0.01 + ) + + token_count_per_expert_list = token_count_per_expert.tolist() + out_tokens = sum(token_count_per_expert_list) + + grad_output = torch.randn((M, N), dtype=dtype, device=device) + q_ref, scale_ref, prob_ref = torch_smooth_permute_with_indices( + grad_output, + indices, + smooth_scales, + token_count_per_expert_list, + probs=probs, + round_scale=round_scale, + ) + y_q, y_scale, prob = triton_batch_smooth_permute_with_indices( + grad_output, + smooth_scales, + token_count_per_expert, + indices, + x_q=None, + x_scale=None, + reverse=False, + probs=probs, + round_scale=round_scale, + ) + output_check(q_ref, y_q, name="data", rtol=0.125) + output_check(scale_ref, y_scale, name="scale") + output_check(prob_ref, prob, name="prob") + + benchmark( + triton_batch_smooth_permute_with_indices, + grad_output, + smooth_scales, + token_count_per_expert, + indices, + round_scale=round_scale, + probs=probs, + ref_bytes=out_tokens * N * 3, + ) + + +@pytest.mark.parametrize( + "M,N,n_experts,topk", + [ + (16384, 2048, 32, 2), + (8192, 4096, 32, 2), + ], +) +def test_batch_transpose_smooth_permute_with_indices(M, N, n_experts, topk, benchmark): + device = "cuda:0" + logits = torch.randn((M, n_experts), dtype=torch.float32, device=device) ** 3 + logits[:, 0] -= 1000 + logits[:, 2] -= 100 + probs, mask_map, token_count_per_expert, indices, row_id_map = torch_make_indices( + logits, topk=topk, bias=-0.01 + ) + + token_count_per_expert_list = token_count_per_expert.tolist() + out_tokens = sum(token_count_per_expert_list) + + x = torch.randn((M, N), dtype=torch.bfloat16, device=device) + smooth_scales = torch.rand((out_tokens,), dtype=torch.float32, device=device) + 0.1 + + x_q_ref, x_scale_ref = torch_batch_transpose_smooth_permute_with_indices( + x, smooth_scales, indices, token_count_per_expert_list, round_scale=True + ) + + x_q, x_scale = triton_batch_transpose_smooth_permute_with_indices( + x, + smooth_scales, + indices, + token_count_per_expert, + token_count_per_expert_list, + round_scale=True, + ) + output_check(x_q_ref, x_q, name="smoothed.data", rtol=0.125) + output_check(x_scale_ref.float(), x_scale.float(), "smoothed.scale") + + benchmark( + torch_batch_transpose_smooth_permute_with_indices, + x, + smooth_scales, + indices, + token_count_per_expert_list, + round_scale=True, + ref_bytes=out_tokens * N * 2, + ) + benchmark( + triton_batch_transpose_smooth_permute_with_indices, + x, + smooth_scales, + indices, + token_count_per_expert, + token_count_per_expert_list, + round_scale=True, + ref_bytes=out_tokens * N * 2, + ) + + +@pytest.mark.parametrize( + "M,N,n_experts,topk", + [ + (16384, 2048, 32, 2), + (8192, 4096, 32, 2), + ], +) +def test_batch_smooth_fused_permute_with_indices( + M, N, n_experts, topk, benchmark, round_scale=True +): + device = "cuda:0" dtype = torch.bfloat16 - smooth_scales = 1 + 10 * torch.rand((n_experts, N), device=device, - dtype=torch.float32) - logits = torch.randn((M, n_experts), dtype=torch.float32, - device=device) ** 3 + smooth_scales = 1 + 1 * torch.rand( + (n_experts, N), device=device, dtype=torch.float32 + ) + org_smooth_scales = 1 + 1 * torch.rand( + (n_experts, N), device=device, dtype=torch.float32 + ) + 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.01) + logits, topk=topk, bias=-0.01 + ) token_count_per_expert_list = token_count_per_expert.tolist() out_tokens = sum(token_count_per_expert_list) - B = 128 - grad_data = torch.randn((M, N), dtype=torch.bfloat16, device=device).to( - torch.float8_e4m3fn) - grad_scale = 1 + torch.rand((M, N // B), dtype=torch.float32, device=device) - q_ref, scale_ref = torch_smooth_permute_with_indices(grad_data, grad_scale, - indices, smooth_scales, - token_count_per_expert_list, - round_scale=round_scale) - y_q, y_scale = triton_smooth_permute_with_indices(grad_data, - grad_scale, - smooth_scales, - token_count_per_expert, - indices, - x_q=None, - x_scale=None, - reverse=False, - round_scale=round_scale) - output_check(q_ref.float(), y_q.float(), 'data') - output_check(scale_ref.float(), y_scale.float(), 'scale') - - - - # smooth_scale_ptrs = torch.tensor([x.data_ptr() for x in torch.split(smooth_scales,1)], device=device) - permuted_data, permuted_scale = triton_smooth_permute_with_mask_map( - grad_data, row_id_map, grad_scale, M, n_experts, out_tokens, N, - smooth_scales, reverse=False, round_scale=round_scale) - output_check(q_ref.float(), permuted_data.float(), 'smoothed.data') - output_check(scale_ref.float(), permuted_scale.float(), 'smoothed.scale') - - q_ref, scale_ref = torch_smooth_permute_with_indices(grad_data, None, - indices, smooth_scales, - token_count_per_expert_list, - round_scale=round_scale) - permuted_data, permuted_scale = triton_smooth_permute_with_mask_map( - grad_data, row_id_map, None, M, n_experts, out_tokens, N, - smooth_scales, reverse=False, round_scale=round_scale) - output_check(q_ref.float(), permuted_data.float(), 'smoothed.data') - output_check(scale_ref.float(), permuted_scale.float(), 'smoothed.scale') - - - - if bench: - benchmark_func(triton_smooth_permute_with_indices, grad_data, - grad_scale, smooth_scales, token_count_per_expert, - indices, round_scale=round_scale, n_repeat=100, - ref_bytes=out_tokens * N * 2) - benchmark_func(triton_smooth_permute_with_mask_map, grad_data, - row_id_map, grad_scale, M, n_experts, out_tokens, N, - smooth_scales, reverse=False, round_scale=round_scale, - n_repeat=100, ref_bytes=out_tokens * N * 2) - - - - -def test_triton_batch_transpose_smooth_permute_with_indices(M=1024, N=2048, n_experts=32, topk=8, bench=False): - - device = 'cuda:0' + grad_data = torch.randn((M, N), dtype=dtype, device=device).to(torch.float8_e4m3fn) + grad_scale = 1 + torch.rand((M,), dtype=torch.float32, device=device) + q_ref, scale_ref = torch_smooth_fused_permute_with_indices( + grad_data, + grad_scale, + indices, + org_smooth_scales, + smooth_scales, + token_count_per_expert_list, + round_scale=round_scale, + ) + y_q, y_scale = triton_batch_smooth_fused_permute_with_indices( + grad_data, + grad_scale, + org_smooth_scales, + smooth_scales, + token_count_per_expert, + indices, + x_q=None, + x_scale=None, + reverse=False, + round_scale=round_scale, + ) + output_check(q_ref, y_q, name="data", rtol=0.125) + output_check(scale_ref, y_scale, name="scale") + + benchmark( + triton_batch_smooth_fused_permute_with_indices, + grad_data, + grad_scale, + smooth_scales, + token_count_per_expert, + indices, + round_scale=round_scale, + n_repeat=100, + ref_bytes=out_tokens * N * 2, + ) + + +@pytest.mark.parametrize( + "M,N,n_experts,topk", + [ + (16384, 2048, 32, 2), + (8192, 4096, 32, 2), + ], +) +def test_batch_transpose_smooth_fused_permute_with_indices( + M, N, n_experts, topk, benchmark +): + device = "cuda:0" if True: - logits = torch.randn((M, n_experts), dtype=torch.float32, - device=device) ** 3 - logits[:,0] -= 1000 - logits[:,2] -= 100 - probs, mask_map, token_count_per_expert, indices, row_id_map = torch_make_indices( - logits, topk=topk, bias=-0.01) + logits = torch.randn((M, n_experts), dtype=torch.float32, device=device) ** 3 + logits[:, 0] -= 1000 + logits[:, 2] -= 100 + probs, mask_map, token_count_per_expert, indices, row_id_map = ( + torch_make_indices(logits, topk=topk, bias=-0.01) + ) token_count_per_expert_list = token_count_per_expert.tolist() out_tokens = sum(token_count_per_expert_list) x = torch.randn((M, N), dtype=torch.bfloat16, device=device).to( - torch.float8_e4m3fn) + torch.float8_e4m3fn + ) scale = torch.rand((M,), dtype=torch.float32, device=device) + 0.1 org_smooth_scale = torch.rand((N,), dtype=torch.float32, device=device) + 0.1 - smooth_scales = torch.rand((out_tokens, ), dtype=torch.float32, device=device) + 0.1 + smooth_scales = ( + torch.rand((out_tokens,), dtype=torch.float32, device=device) + 0.1 + ) else: - # torch.save({"x":x, "scale":scale, "org_smooth_scale":org_smooth_scale,"smooth_scales":smooth_scales, "indices":indices, "token_count_per_expert":token_count_per_expert,"splits":splits}, '/tmp/debug.bin') - state = torch.load('/tmp/debug.bin') - x = state['x'] - scale = state['scale'] - org_smooth_scale = state['org_smooth_scale'] - smooth_scales = state['smooth_scales'] - indices = state['indices'] - token_count_per_expert = state['token_count_per_expert'] - token_count_per_expert_list = state['splits'] + state = torch.load("/tmp/debug.bin") + x = state["x"] + scale = state["scale"] + org_smooth_scale = state["org_smooth_scale"] + smooth_scales = state["smooth_scales"] + indices = state["indices"] + token_count_per_expert = state["token_count_per_expert"] + token_count_per_expert_list = state["splits"] out_tokens = sum(token_count_per_expert_list) + x_q_ref, x_scale_ref = torch_batch_transpose_smooth_fused_permute_with_indices( + x, + scale, + org_smooth_scale, + smooth_scales, + indices, + token_count_per_expert_list, + round_scale=True, + ) + + x_q, x_scale = triton_batch_transpose_smooth_fused_permute_with_indices( + x, + scale, + org_smooth_scale, + smooth_scales, + indices, + token_count_per_expert, + token_count_per_expert_list, + round_scale=True, + ) + output_check(x_q_ref, x_q, name="smoothed.data", rtol=0.125) + output_check(x_scale_ref.float(), x_scale.float(), "smoothed.scale") + + benchmark( + torch_batch_transpose_smooth_fused_permute_with_indices, + x, + scale, + org_smooth_scale, + smooth_scales, + indices, + token_count_per_expert_list, + round_scale=True, + ref_bytes=out_tokens * N * 2, + ) + benchmark( + triton_batch_transpose_smooth_fused_permute_with_indices, + x, + scale, + org_smooth_scale, + smooth_scales, + indices, + token_count_per_expert, + token_count_per_expert_list, + round_scale=True, + ref_bytes=out_tokens * N * 2, + ) + + +@pytest.mark.parametrize( + "M,N,n_experts,topk", + [ + (8192 * 2, 2048, 32, 2), + (0, 2048, 32, 2), + (8192, 1536, 32, 2), + ], +) +def test_batch_block_pad_permute_with_indices(M, N, n_experts, topk, benchmark): + device = "cuda:0" + logits = torch.randn((M, n_experts), dtype=torch.float32, device=device) ** 3 + logits[:, 0] -= 1000 + logits[:, 2] -= 100 + probs, mask_map, token_count_per_expert, indices, row_id_map = torch_make_indices( + logits, topk=topk, bias=-0.01 + ) + token_count_per_expert_list = token_count_per_expert.tolist() + + 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 = torch.randn((M, N), dtype=torch.bfloat16, device=device) + + x_q_ref, x_s_ref, xt_q_ref, xt_s_ref, p_ref = ( + torch_batch_block_pad_permute_with_indices( + x, indices, probs, token_count_per_expert_list, round_scale=True + ) + ) + + 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, + ) + output_check(x_q_ref.float(), x_q.view(-1).float(), "data") + output_check(x_s_ref.float(), x_s.view(-1).float(), "scale") + output_check(xt_q_ref.float(), xt_q.view(-1).float(), "t.data") + output_check(xt_s_ref.float(), xt_s.view(-1).float(), "t.scale") + output_check(p_ref.float(), p.view(-1).float(), "prob") + + benchmark( + triton_batch_block_pad_permute_with_indices, + x, + token_count_per_expert, + pad_indices, + token_count_per_expert_list, + probs=probs, + round_scale=True, + ref_bytes=num_out_tokens * N * 4, + ) + + benchmark( + triton_permute_with_mask_map, + x, + None, + probs, + row_id_map, + num_out_tokens, + contiguous=False, + tokens_per_expert=token_count_per_expert, + ref_bytes=num_out_tokens * N * 4, + ) + xs = x[indices] + benchmark( + triton_batch_blockwise_quant, + xs, + token_count_per_expert, + token_count_per_expert_list, + round_scale=True, + ref_bytes=num_out_tokens * N * 4, + ) + + +@pytest.mark.parametrize( + "M,N,n_experts,topk", + [ + (3095, 2048, 32, 2), + (0, 2048, 32, 2), + (1024, 1536, 32, 2), + ], +) +def test_batch_mxfp8_permute_with_indices(M, N, n_experts, topk, benchmark): + device = "cuda:0" + logits = torch.randn((M, n_experts), dtype=torch.float32, device=device) ** 3 + logits[:, 0] -= 1000 + logits[:, 2] -= 100 + probs, mask_map, token_count_per_expert, indices, row_id_map = torch_make_indices( + logits, topk=topk, bias=-0.01 + ) + token_count_per_expert_list = token_count_per_expert.tolist() - x_q_ref, x_scale_ref = torch_batch_transpose_smooth_permute_with_indices(x, scale, org_smooth_scale, smooth_scales, - indices, - token_count_per_expert_list, - round_scale=True) - - x_q, x_scale = triton_batch_transpose_smooth_permute_with_indices(x, scale, org_smooth_scale, smooth_scales, - indices, - token_count_per_expert, token_count_per_expert_list, - round_scale=True) - output_check(x_q_ref.float(), x_q.float(), 'smoothed.data') - output_check(x_scale_ref.float(), x_scale.float(), 'smoothed.scale') - - - x_q_ref, x_scale_ref = torch_batch_transpose_smooth_permute_with_indices(x, None, None, smooth_scales, - indices, - token_count_per_expert_list, - round_scale=True) - - x_q, x_scale = triton_batch_transpose_smooth_permute_with_indices(x, None, None, smooth_scales, - indices, - token_count_per_expert, token_count_per_expert_list, - round_scale=True) - output_check(x_q_ref.float(), x_q.float(), 'bf16.data') - output_check(x_scale_ref.float(), x_scale.float(), 'bf16.scale') - - - if bench: - benchmark_func(torch_batch_transpose_smooth_permute_with_indices, x, scale, org_smooth_scale, smooth_scales, - indices, - token_count_per_expert_list, - round_scale=True, - ref_bytes=out_tokens * N * 2) - benchmark_func(triton_batch_transpose_smooth_permute_with_indices, x, scale, org_smooth_scale, smooth_scales, - indices, - token_count_per_expert, token_count_per_expert_list, - round_scale=True, - ref_bytes=out_tokens * N * 2) - - -if __name__ == '__main__': - test_make_id_map(M=4098, n_experts=32, topk=2, bias=0.0, bench=False) - test_triton_permute_with_mask_map(M=16384, N=2048, n_experts=32, topk=8, bench=False) - test_triton_permute_with_mask_map(M=8192, N=4096, n_experts=32, topk=8, bench=False) - test_triton_permute_with_mask_map(M=7628, N=2048, n_experts=32, topk=8, bench=False) - - test_triton_smooth_permute_with_mask_map(M=4096, N=4096, n_experts=32, - topk=8) - test_triton_smooth_permute_with_mask_map(M=7628, N=2048, n_experts=32, - topk=8) - - test_triton_batch_transpose_smooth_permute_with_indices(M=16384, N=2048, n_experts=32, topk=2, bench=False) - test_triton_batch_transpose_smooth_permute_with_indices(M=8192, N=4096, n_experts=32, topk=2, bench=False) + num_out_tokens = sum([(x + 31) // 32 * 32 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=32 + ) + + x = torch.randn((M, N), dtype=torch.bfloat16, device=device) + + x_q_ref, x_s_ref, xt_q_ref, xt_s_ref, p_ref = ( + torch_batch_mxfp8_permute_with_indices( + x, indices, probs, token_count_per_expert_list + ) + ) + + x_q, x_s, xt_q, xt_s, p = triton_batch_mxfp8_permute_with_indices( + x, + token_count_per_expert, + pad_indices, + token_count_per_expert_list, + probs=probs, + dispatch_type="deepep", + ) + output_check(x_q_ref.float(), x_q.float(), "data") + output_check(x_s_ref.float(), x_s.float(), "scale") + output_check(xt_q_ref.float(), xt_q.float(), "t.data") + output_check(xt_s_ref.float(), xt_s.float(), "t.scale") + output_check(p_ref.float(), p.float(), "prob") + + num_out_tokens = sum(token_count_per_expert_list) + benchmark( + triton_batch_mxfp8_permute_with_indices, + x, + token_count_per_expert, + indices, + token_count_per_expert_list, + probs=probs, + dispatch_type="deepep", + ref_bytes=num_out_tokens * N * 4, + ) + + benchmark( + triton_permute_with_mask_map, + x, + None, + probs, + row_id_map, + num_out_tokens, + contiguous=False, + tokens_per_expert=token_count_per_expert, + ref_bytes=num_out_tokens * N * 4, + ) + xs = x[indices] + benchmark( + triton_batch_mxfp8_quant, + xs, + token_count_per_expert, + token_count_per_expert_list, + ref_bytes=num_out_tokens * N * 4, + ) diff --git a/tests/test_group_quant.py b/tests/test_group_quant.py index d8c68ec..27672e2 100644 --- a/tests/test_group_quant.py +++ b/tests/test_group_quant.py @@ -3,27 +3,31 @@ Copyright (c) Ant Financial Service Group and its affiliates. """ +import pytest import torch from linghe.quant.group import triton_group_quant -from linghe.tools.benchmark import benchmark_func -from linghe.tools.util import (output_check, - torch_group_quant) +from linghe.tools.check import output_check +from linghe.tools.util import torch_group_quant -def test_group_quant(M=4096, N=4096, B=128, round_scale=False, bench=False): - x = torch.randn((M, N), dtype=torch.bfloat16, device='cuda:0') ** 3 +@pytest.mark.parametrize( + "M,N", + [ + (4096, 4096), + (4096, 8192), + (2049, 8192), + (2049, 1536), + ], +) +def test_group_quant(M, N, benchmark, B=128, round_scale=False): + x = torch.randn((M, N), dtype=torch.bfloat16, device="cuda:0") ** 3 xq_ref, x_scale_ref = torch_group_quant(x, B, round_scale=round_scale) xq, x_scale = triton_group_quant(x, group_size=B, round_scale=round_scale) - output_check(xq_ref.float(), xq.float(), mode='data') - output_check(x_scale_ref.float(), x_scale.float(), mode='scale') + output_check(xq_ref, xq, name="data") + output_check(x_scale_ref, x_scale, name="scale") - if bench: - n_repeat = 100 - benchmark_func(triton_group_quant, x, group_size=B, - n_repeat=n_repeat, ref_bytes=M * N * 3) - -if __name__ == '__main__': - test_group_quant(M=4096, N=4096, B=128) - test_group_quant(M=4096, N=8192, B=128) - test_group_quant(M=2049, N=8192, B=128) + n_repeat = 100 + benchmark( + triton_group_quant, x, group_size=B, n_repeat=n_repeat, ref_bytes=M * N * 3 + ) diff --git a/tests/test_hadamard_quant.py b/tests/test_hadamard_quant.py index f508e2f..12880e9 100644 --- a/tests/test_hadamard_quant.py +++ b/tests/test_hadamard_quant.py @@ -5,96 +5,95 @@ import torch -from linghe.quant.hadamard import triton_hadamard_quant -from linghe.tools.benchmark import benchmark_func -from linghe.tools.util import (output_check, - make_hadamard_matrix, - torch_hadamard_transform, - torch_row_quant, - ) from linghe.facade.hadamard_quant_linear import HadamardQuantLinear - +from linghe.quant.hadamard import triton_hadamard_quant +from linghe.tools.check import output_check +from linghe.tools.util import ( + make_hadamard_matrix, + torch_hadamard_transform, + torch_row_quant, +) # apply hadamard transformation and quantization for x def torch_hadamard_quant(x, hm, round_scale=False): - xh = torch_hadamard_transform(x, hm, side='right') - q, s = torch_row_quant(xh, round_scale=round_scale) - xht = torch_hadamard_transform(x.t().contiguous(), hm, side='right') - qt, st = torch_row_quant(xht, round_scale=round_scale) + dtype = x.dtype + x = x.float() + hm = hm.float() + xh = torch_hadamard_transform(x, hm, side="right") + q, s = torch_row_quant(xh, round_scale=round_scale) + xht = torch_hadamard_transform(x.t().contiguous(), hm, side="right") + qt, st = torch_row_quant(xht, round_scale=round_scale) - return xh,xht,q,s,qt,st + return xh.to(dtype), xht.to(dtype), q, s, qt, st def test_hadamard_quant(M=8192, N=1024, K=2048, B=64, bench=False): - dtype = torch.bfloat16 - device = 'cuda:0' + dtype = torch.bfloat16 + device = "cuda:0" x = torch.randn((M, K), dtype=dtype, device=device) w = torch.randn((N, K), dtype=dtype, device=device) dy = torch.randn((M, N), dtype=dtype, device=device) hm = make_hadamard_matrix(B, dtype=dtype, device=device, norm=True) + y_ref = x @ w.t() + dx_ref = dy @ w + dw_ref = dy.t() @ x - y_ref = x@w.t() - dx_ref = dy@w - dw_ref = dy.t()@x + xh, xht, xq, xs, xqt, xst = torch_hadamard_quant(x, hm, round_scale=False) + wh, wht, wq, ws, wqt, wst = torch_hadamard_quant(w, hm, round_scale=False) + dyh, dyht, dyq, dys, dyqt, dyst = torch_hadamard_quant(dy, hm, round_scale=False) - xh,xht,xq,xs,xqt,xst = torch_hadamard_quant(x, hm, round_scale=False) - wh,wht,wq,ws,wqt,wst = torch_hadamard_quant(w, hm, round_scale=False) - dyh,dyht,dyq,dys,dyqt,dyst = torch_hadamard_quant(dy, hm, round_scale=False) + y = xh @ wh.t() + dx = dyh @ wht.t() + dw = dyht @ xht.t() - y = xh@wh.t() - dx = dyh@wht.t() - dw = dyht@xht.t() - - output_check(y_ref,y,'bf16.y') - output_check(dx_ref,dx,'bf16.dx') - output_check(dw_ref,dw,'bf16.dw') + output_check(y_ref, y, "bf16.y", atol=2) + output_check(dx_ref, dx, "bf16.dx", atol=2) + output_check(dw_ref, dw, "bf16.dw", atol=2) x_q, x_scale, xt_q, xt_scale = triton_hadamard_quant(x, hm) - output_check(xq, x_q, 'x.data') - output_check(xs, x_scale, 'x.scale') - output_check(xqt, xt_q, 'xt.data') - output_check(xst, xt_scale, 'xt.scale') + output_check(xq, x_q, "x.data", rtol=0.125) + output_check(xs, x_scale, "x.scale") + output_check(xqt, xt_q, "xt.data", rtol=0.125) + output_check(xst, xt_scale, "xt.scale") - w_q, w_scale, wt_q, wt_scale = triton_hadamard_quant(w, hm) - output_check(wq, w_q, 'w.data') - output_check(ws, w_scale, 'w.scale') - output_check(wqt, wt_q, 'wt.data') - output_check(wst, wt_scale, 'wt.scale') + output_check(wq, w_q, "w.data", rtol=0.125) + output_check(ws, w_scale, "w.scale") + output_check(wqt, wt_q, "wt.data", rtol=0.125) + output_check(wst, wt_scale, "wt.scale") - dy_q, dy_scale, dyt_q, dyt_scale = triton_hadamard_quant(dy, hm) - output_check(dyq, dy_q, 'dy.data') - output_check(dys, dy_scale, 'dy.scale') - output_check(dyqt, dyt_q, 'dyt.data') - output_check(dyst, dyt_scale, 'dyt.scale') + output_check(dyq, dy_q, "dy.data", rtol=0.125) + output_check(dys, dy_scale, "dy.scale") + output_check(dyqt, dyt_q, "dyt.data", rtol=0.125) + output_check(dyst, dyt_scale, "dyt.scale") def test_hadamard_quant_linear(M=8192, N=1024, K=2048, B=64): - - dtype = torch.bfloat16 - device = 'cuda:0' + dtype = torch.bfloat16 + device = "cuda:0" linear = HadamardQuantLinear(K, N, bias=False, dtype=dtype, device=device) x = torch.randn((M, K), dtype=dtype, device=device).requires_grad_() w = torch.randn((N, K), dtype=dtype, device=device) dy = torch.randn((M, N), dtype=dtype, device=device) linear.weight.data.copy_(w) - y_ref = x@w.t() + y_ref = x @ w.t() y = linear(x) - output_check(y_ref, y, mode='y') + output_check(y_ref, y, name="y", rtol=-1) - dx_ref = dy@w - dw_ref = dy.t()@x + dx_ref = dy @ w + dw_ref = dy.t() @ x y.backward(dy) - dw = linear.weight.grad + dw = linear.weight.grad dx = x.grad - output_check(dx_ref, dx, mode='dx') - output_check(dw_ref, dw, mode='dw') + output_check(dx_ref, dx, name="dx", rtol=-1) + output_check(dw_ref, dw, name="dw", rtol=-1) + -if __name__ == '__main__': +if __name__ == "__main__": test_hadamard_quant(M=8192, N=1024, K=2048, B=64, bench=False) - test_hadamard_quant_linear(M=8192, N=1024, K=2048, B=64) \ No newline at end of file + test_hadamard_quant_linear(M=8192, N=1024, K=2048, B=64) diff --git a/tests/test_la.py b/tests/test_la.py new file mode 100644 index 0000000..be2fa46 --- /dev/null +++ b/tests/test_la.py @@ -0,0 +1,167 @@ +# -*- coding: utf-8 -*- +""" +Copyright (c) Ant Financial Service Group and its affiliates. +""" + +import math + +import pytest +import torch + +from linghe.attn.la import ( + triton_lightning_attention_forward, + triton_lightning_attention_backward, + triton_fused_lightning_attention_backward, +) +from linghe.tools.check import output_check + + +def torch_la(q, k, v, s, decay_scales): + dtype = q.dtype + + q = q.double() + k = k.double() + v = v.double() + s = s.double() + bs, q_len, q_heads, head_dim = q.shape + k_heads = k.shape[2] + k_len = k.shape[1] + assert q_len == k_len + softmax_scale = 1.0 / math.sqrt(head_dim) + query = q.transpose(1, 2) # [bs, head, len, dim] + key = torch.permute(k, (0, 2, 3, 1)) # [bs, head, dim, len] + value = v.transpose(1, 2) # [bs, head, len, dim] + if k_heads != q_heads: + g = q_heads // k_heads + key = torch.repeat_interleave(key, g, dim=1) + value = torch.repeat_interleave(value, g, dim=1) + + arr = torch.arange(q_len, dtype=torch.float64, device=q.device) + decay_matrix = arr.view(-1, 1) - arr.view(1, -1) + decay_matrix = torch.exp(-decay_scales[:, None, None] * decay_matrix[None]) + decay_matrix = torch.tril(decay_matrix, 0) + + score = torch.matmul(query, key) * softmax_scale + score *= decay_matrix[None] + att = torch.matmul(score, value) + + decay_arr = torch.exp(-decay_scales[:, None, None] * (arr[:, None] + 1)) + att = att + torch.matmul(query * decay_arr, s) + + att = torch.reshape( + att.transpose(1, 2), [bs, q_len, q_heads, head_dim] + ).contiguous() + + decay_key = key * torch.exp(-decay_scales[:, None, None] * (q_len - 1 - arr)) + state = decay_key @ value + s * torch.exp(-decay_scales[:, None, None]) + + return att.to(dtype), state.to(torch.float32) + + +def torch_varlen_torch_la(q, k, v, s, cu_seqlens, padded_cu_seqlens, decay_scales): + pass + + +def make_varlen_input( + qo_heads=16, kv_heads=16, dim=128, qls=[1024, 1024], kls=[1024, 1024] +): + device = torch.device("cuda:0") + dtype = torch.bfloat16 + qs = [] + ks = [] + vs = [] + for i, ql in enumerate(qls): + kvl = kls[i] + q = torch.randn(ql, qo_heads, dim, dtype=dtype, device=device) + k = torch.randn(kvl, kv_heads, dim, dtype=dtype, device=device) + v = torch.randn(kvl, kv_heads, dim, dtype=dtype, device=device) + qs.append(q) + ks.append(k) + vs.append(v) + + q = torch.cat(qs, 0) + k = torch.cat(ks, 0) + v = torch.cat(vs, 0) + + return q, k, v + + +@pytest.mark.parametrize( + "bs,length,qo_heads,kv_heads,dim", + [ + (1, 8192, 64, 64, 128), + ], +) +def test_la(bs, length, qo_heads, kv_heads, dim, benchmark): + device = torch.device("cuda:0") + dtype = torch.bfloat16 + + q = torch.randn(bs, length, qo_heads, dim, dtype=dtype, device=device) ** 3 * 0.1 + q = q.requires_grad_() + + k = torch.randn(bs, length, kv_heads, dim, dtype=dtype, device=device) ** 3 * 0.1 + k = k.requires_grad_() + + v = torch.randn(bs, length, kv_heads, dim, dtype=dtype, device=device) ** 3 * 0.1 + v = v.requires_grad_() + + g = torch.randn(bs, length, qo_heads, dim, dtype=dtype, device=device) + + s = torch.zeros(bs, kv_heads, dim, dim, dtype=torch.float32, device=device) + + decay_scales = 2 ** ( + -0.5 * torch.arange(1, qo_heads + 1, dtype=torch.float32, device=device) + ) + # decay_scales = 0.0 * torch.ones(qo_heads, dtype=torch.float32, device=device) + output_ref, state_ref = torch_la(q, k, v, s, decay_scales) + output_ref.backward(g) + dq_ref = q.grad + dk_ref = k.grad + dv_ref = v.grad + q.grad = None + k.grad = None + v.grad = None + + output, state = triton_lightning_attention_forward(q, k, v, decay_scales) + + output_check(output_ref, output, name="output", rtol=0.1, atol=0.2) + output_check(state_ref, state, name="state", rtol=0.1, atol=0.2) + + dq, dk, dv = triton_lightning_attention_backward(g, q, k, v, decay_scales) + + output_check(dq_ref, dq, name="dq", rtol=-0.1, atol=1.0) + output_check(dk_ref, dk, name="dk", rtol=-0.1, atol=1.0) + output_check(dv_ref, dv, name="dv", rtol=-0.1, atol=1.0) + + max_decay_scale = decay_scales.max().item() + if max_decay_scale < 0.1: + dq, dk, dv = triton_fused_lightning_attention_backward( + g, q, k, v, state, decay_scales + ) + output_check(dq_ref, dq, name="dq", rtol=-0.1, atol=0.1) + output_check(dk_ref, dk, name="dk", rtol=-0.1, atol=0.1) + output_check(dv_ref, dv, name="dv", rtol=-0.1, atol=0.1) + + ref_bytes = bs * length * qo_heads * dim * 8 + bs * qo_heads * dim * dim * 8 + benchmark( + triton_lightning_attention_forward, q, k, v, decay_scales, ref_bytes=ref_bytes + ) + benchmark( + triton_lightning_attention_backward, + g, + q, + k, + v, + decay_scales, + ref_bytes=ref_bytes, + ) + benchmark( + triton_fused_lightning_attention_backward, + g, + q, + k, + v, + state, + decay_scales, + ref_bytes=ref_bytes, + ) diff --git a/tests/test_loss.py b/tests/test_loss.py index 0ba2631..e503eda 100644 --- a/tests/test_loss.py +++ b/tests/test_loss.py @@ -5,60 +5,175 @@ import random +import pytest import torch -from linghe.tools.benchmark import benchmark_func -from linghe.tools.util import output_check -from linghe.utils.loss import triton_softmax_cross_entropy_forward, \ - triton_softmax_cross_entropy_backward +from linghe.facade.loss import moe_z_loss, softmax_cross_entropy +from linghe.tools.check import output_check +from linghe.utils.loss import ( + triton_softmax_cross_entropy_forward, + triton_softmax_cross_entropy_backward, + triton_moe_z_loss_forward, + triton_moe_z_loss_backward, +) -def torch_cross_entropy(logits, targets): - float_logits = logits.float() +def torch_cross_entropy(logits, targets, grad, ignore_index=-100): + float_logits = logits.to(torch.float32) losses = torch.nn.functional.cross_entropy( float_logits.view(-1, logits.size()[-1]), targets.view(-1), - reduction='none') - loss = losses.sum() + reduction="none", + ignore_index=ignore_index, + ) + loss = (losses * grad).sum() + loss.backward() + return losses.to(torch.float32), logits.grad + + +def torch_z_loss(logits, coef=1e-6): + float_logits = logits.float() + loss = torch.mean(torch.square(torch.logsumexp(float_logits, dim=-1))) * coef loss.backward() - return losses, logits.grad + return loss, logits.grad + + +@pytest.mark.parametrize( + "M,N,coef,grad_coef,fill,ignore_index", + [ + (8192, 157184, 1.0, 1.0, False, None), + (8192, 157184, 1.0, 1e-6, False, None), + (8192, 157184, 10000.0, 100.0, True, None), + (8192, 157184, 1.0, 1.0, True, -100), + (8192, 157184, 1.0, 1.0, True, 0), + (8192, 157184 - 16, 10000.0, 100.0, True, None), + (8192, 175175, 1.0, 1.0, False, None), + (8192, 157184, 0.0, 0.0, False, None), + (8192, 157184, 0.0, 100.0, False, None), + (8192, 157184, 1000.0, 0.0, False, None), + (8192, 157184, 100.0, 100.0, True, None), + (4096, 157184, 0.1, 1.0, False, None), + ], +) +def test_triton_softmax_cross_entropy( + M, N, coef, grad_coef, ignore_index, fill, benchmark, inplace=True +): + device = "cuda:0" + dtype = torch.bfloat16 + logits = torch.randn((M, N), dtype=dtype, device=device, requires_grad=False) + + select = True + if select: + top_indices = torch.topk(logits, 1)[1].tolist() + targets = [] + for i, idx in enumerate(top_indices): + targets.append(random.choice(idx)) + + targets = torch.tensor(targets, dtype=torch.long, device=device) + else: + targets = torch.randint(0, N, (M,), dtype=torch.long, device=device) + + if ignore_index is not None: + targets[:10000] = ignore_index + if fill: + logits[:, :8192] = -10000 + # logits[:,0] = -100000000 + # logits[:,1:] = 1000 + # logits[:,100:] = 1000 + # targets = 0 * torch.ones((M,), dtype=torch.long, device=device) -def test_triton_softmax_cross_entropy(M=4096, N=157184, coef=1.0, bench=False): - device = 'cuda:0' - logits = torch.randn((M, N), dtype=torch.bfloat16, device=device, - requires_grad=False) + ignore_index = -100 if ignore_index is None else ignore_index logits = (logits * coef).detach().clone().requires_grad_() - top_indices = torch.topk(logits, 16)[1].tolist() - targets = [] - for i, idx in enumerate(top_indices): - targets.append(random.choice(idx)) - targets = torch.tensor(targets, dtype=torch.long, device=device) - input_grad = torch.ones((M,), dtype=torch.bfloat16, device=device) - loss_ref, grad_ref = torch_cross_entropy(logits, targets) - sum_exp_ref = torch.sum(torch.exp(logits.float()), dim=-1) - max_logit_ref = 0.0 * torch.amax(logits, dim=-1).float() - - loss, sum_exp, max_logit = triton_softmax_cross_entropy_forward(logits, - targets) - output_check(loss_ref, loss, mode='loss') - # output_check(sum_exp_ref, sum_exp, mode='sum_exp') - # output_check(max_logit_ref, max_logit, mode='max_logit') - - grad = triton_softmax_cross_entropy_backward(logits, targets, sum_exp, - max_logit, - input_grad, output_grad=logits) - output_check(grad_ref.float(), grad.float(), mode='grad') - if bench: - benchmark_func(torch_cross_entropy, logits, targets, - ref_bytes=M * N * 6) - benchmark_func(triton_softmax_cross_entropy_forward, logits, targets, - ref_bytes=M * N * 2) - benchmark_func(triton_softmax_cross_entropy_backward, logits, targets, - sum_exp, max_logit, input_grad, ref_bytes=M * N * 4) - - -if __name__ == '__main__': - test_triton_softmax_cross_entropy(M=8192, N=157184, coef=1.0, bench=False) - test_triton_softmax_cross_entropy(M=8192, N=157184, coef=10.0, bench=False) - test_triton_softmax_cross_entropy(M=4096, N=157175, coef=10.0, bench=False) + output_grad = torch.randn((M,), dtype=torch.float32, device=device) * grad_coef + loss_ref, grad_ref = torch_cross_entropy( + logits, targets, output_grad, ignore_index=ignore_index + ) + + loss, sum_exp, max_logit = triton_softmax_cross_entropy_forward( + logits.detach().clone(), targets, ignore_index=ignore_index + ) + output_check(loss_ref, loss, name="loss", atol=1e-4, rtol=1e-5) + + grad = triton_softmax_cross_entropy_backward( + logits.detach().clone(), + targets, + sum_exp, + max_logit, + output_grad, + ignore_index=ignore_index, + inplace=inplace, + ) + output_check(grad_ref, grad, name="grad", digest=10) + + logits_ = logits.detach().clone().requires_grad_() + loss = softmax_cross_entropy( + logits_, targets, ignore_index=ignore_index, inplace=True + ) + loss.backward(output_grad) + grad = logits_.grad + output_check(loss_ref, loss, name="loss", atol=1e-4, rtol=1e-5) + output_check(grad_ref, grad, name="grad", digest=10) + + logits_clone = logits.detach().clone() + benchmark( + torch_cross_entropy, + logits.requires_grad_(), + targets, + output_grad, + ref_bytes=M * N * 2, + ) + benchmark( + triton_softmax_cross_entropy_forward, + logits, + targets, + ignore_index=ignore_index, + ref_bytes=M * N * 2, + ) + benchmark( + triton_softmax_cross_entropy_backward, + logits_clone, + targets, + sum_exp, + max_logit, + output_grad, + ignore_index=ignore_index, + ref_bytes=M * N * 4, + ) + + +@pytest.mark.parametrize( + "L,B,N,coef", + [ + (4096, 2, 256, 1e-6), + ], +) +def test_z_loss(L, B, N, coef, benchmark): + device = "cuda:0" + logits = torch.randn( + (L, B, N), dtype=torch.float32, device=device, requires_grad=False + ) + logits = (logits * 1).detach().clone().requires_grad_() + input_grad = torch.ones((1,), dtype=torch.float32, device=device) + loss_ref, grad_ref = torch_z_loss(logits, coef=coef) + + loss = triton_moe_z_loss_forward(logits, coef=coef) + grad = triton_moe_z_loss_backward(input_grad, logits, coef=coef) + output_check(loss_ref, loss, name="loss") + output_check(grad_ref.float(), grad.float(), name="grad") + + loss = moe_z_loss(logits, coef=coef) + loss.backward(gradient=input_grad[0]) + grad = logits.grad + output_check(loss_ref, loss, name="loss") + output_check(grad_ref.float(), grad.float(), name="grad") + + benchmark(torch_z_loss, logits, coef=coef, ref_bytes=L * B * N * 4) + benchmark(triton_moe_z_loss_forward, logits, coef=coef, ref_bytes=L * B * N * 4) + benchmark( + triton_moe_z_loss_backward, + input_grad, + logits, + coef=coef, + ref_bytes=L * B * N * 8, + ) diff --git a/tests/test_mla.py b/tests/test_mla.py new file mode 100644 index 0000000..dd235b5 --- /dev/null +++ b/tests/test_mla.py @@ -0,0 +1,434 @@ +# -*- coding: utf-8 -*- +""" +Copyright (c) Ant Financial Service Group and its affiliates. +""" + +import math + +import pytest +import torch + +from linghe.attn.mla import ( + triton_mla_forward, + triton_mla_backward, + triton_fp8_mla_forward, + triton_varlen_mla_forward, + triton_varlen_mla_backward, +) +from linghe.facade.mla import multi_latend_attention +from linghe.tools.check import output_check + + +def torch_attn(q, k, v, causal=True, mask=None, clip_value=None, hp=False): + dtype = q.dtype + if hp: + q = q.float() + k = k.float() + v = v.float() + bs, q_len, q_head, q_head_dim = q.shape + v_head_dim = v.shape[-1] + k_head = k.shape[2] + k_len = k.shape[1] + if mask is None: + if causal: + mask = -10000 * torch.triu( + torch.ones((q_len, k_len), dtype=q.dtype, device="cuda:0"), + k_len - q_len + 1, + ) + else: + mask = torch.zeros((q_len, k_len), dtype=q.dtype, device="cuda:0") + + query = q.transpose(1, 2) + key = torch.permute(k, (0, 2, 3, 1)) + value = v.transpose(1, 2) + if k_head != q_head: + g = q_head // k_head + key = torch.repeat_interleave(key, g, dim=1) + value = torch.repeat_interleave(value, g, dim=1) + qk = torch.matmul(query, key) + if clip_value is not None: + qk = torch.clamp_max_(qk, clip_value).detach() + qk - qk.detach() + score = qk / math.sqrt(v_head_dim) + mask + lse = torch.sum(torch.exp(score), -1) + max_logits = torch.amax(score, -1) + prob = torch.softmax(score, dim=-1, dtype=torch.float32) + if not hp: + prob = prob.to(dtype) + att = torch.matmul(prob, value) + att = torch.reshape( + att.transpose(1, 2), [bs, q_len, q_head, v_head_dim] + ).contiguous() + return att.to(dtype), lse, max_logits + + +def torch_varlen_attn( + qs, ks, vs, cu_seqlens, padded_cu_seqlens=None, causal=True, hp=False +): + cu_seqlens = cu_seqlens.tolist() + if padded_cu_seqlens is not None: + padded_cu_seqlens = padded_cu_seqlens.tolist() + else: + padded_cu_seqlens = None + outputs = [] + lses = [] + logits = [] + for i in range(len(cu_seqlens) - 1): + if padded_cu_seqlens is None: + s = cu_seqlens[i] + e = cu_seqlens[i + 1] + q = qs[s:e][None] + k = ks[s:e][None] + v = vs[s:e][None] + else: + s = padded_cu_seqlens[i] + e = s + cu_seqlens[i + 1] - cu_seqlens[i] + q = qs[s:e][None] + k = ks[s:e][None] + v = vs[s:e][None] + out, lse, logit = torch_attn(q, k, v, causal=causal, hp=hp) + outputs.append(out[0]) + lses.append(lse[0]) + logits.append(logit[0]) + if padded_cu_seqlens is not None: + gap = (padded_cu_seqlens[i + 1] - padded_cu_seqlens[i]) - ( + cu_seqlens[i + 1] - cu_seqlens[i] + ) + outputs.append(torch.zeros_like(out[0][:gap])) + lses.append(torch.zeros_like(lse[0][:, :gap])) + logits.append(torch.zeros_like(logit[0][:, :gap])) + + outputs = torch.cat(outputs, 0) + lses = torch.cat(lses, 1) + logits = torch.cat(logits, 1) + return outputs, lses, logits + + +def torch_softmax(x): + prob = torch.softmax(x, dim=-1, dtype=torch.float32) + return prob.to(x.dtype) + + +def torch_softmax_backward(x, g): + p = torch.softmax(x, dim=-1, dtype=torch.float32) + # (dp * p - p * tl.sum(p * dp, 1)[:,None]) + gi = g * p - p * torch.sum(p * g, 1)[:, None] + return gi.to(x.dtype) + + +def head_wise_quant(x): + x = x.float() + maxs = x.abs().amax(-1) + scales = torch.maximum(maxs / 448, maxs * 0.0 + 1e-30) + x_q = (x / scales[..., None]).to(torch.float8_e4m3fn) + x_s = scales.permute(0, 2, 1).contiguous() + return x_q, x_s + + +@pytest.mark.parametrize( + "M,N", + [ + (128, 128), + ], +) +def test_softmax(M, N, benchmark): + x = torch.randn((N, N), dtype=torch.bfloat16, device="cuda:0", requires_grad=True) + g = torch.randn((N, N), dtype=torch.bfloat16, device="cuda:0") + y_ref = torch_softmax(x) + y_ref.backward(g) + grad_ref = x.grad + grad = torch_softmax_backward(x, g) + output_check(grad_ref, grad, atol=10, name="grad") + + +@pytest.mark.parametrize( + "M,N,D", + [ + (128, 128, 128), + ], +) +def test_dot_sum(M, N, D, benchmark): + p = torch.randn((M, N), dtype=torch.float32) + v = torch.randn((N, D), dtype=torch.float32) + g = torch.randn((M, D), dtype=torch.float32) + ds_ref = ((g @ v.T) * p).sum(1) + ds = ((p @ v) * g).sum(1) + output_check(ds_ref, ds, atol=10, name="dot_sum") + + +@pytest.mark.parametrize( + "L,H,causal,hpc,safe,coef,clip_value", + [ + (8192, 64, True, False, False, 1.0, None), + (8192, 64, True, False, False, 1.0, 500.0), + (8192, 64, True, True, False, 1.0, None), + (8192, 64, True, False, True, 100.0, None), + (4096, 64, True, False, False, 1.0, None), + (4096, 64, False, False, False, 1.0, None), + (8192, 64, False, False, False, 1.0, None), + (8192, 1, False, False, False, 1.0, None), + ], +) +def test_mla(L, H, causal, hpc, safe, coef, clip_value, benchmark, B=1): + dtype = torch.bfloat16 + device = "cuda:0" + q = ( + torch.randn((B, L, H, 192), device=device, dtype=dtype) * coef + ).requires_grad_() + k = torch.randn((B, L, H, 192), device=device, dtype=dtype) + k[:, :, :, 128:] = k[:, :, :1, 128:] + k = k.requires_grad_() + v = torch.randn((B, L, H, 128), device=device, dtype=dtype, requires_grad=True) + g = torch.randn((B, L, H, 128), device=device, dtype=dtype, requires_grad=True) + + output_ref, lse_ref, max_logits_ref = torch_attn(q, k, v, causal=causal, hp=True) + output_ref.backward(g, retain_graph=False) + gq_ref = q.grad + gk_ref = k.grad + gv_ref = v.grad + + q.grad = None + k.grad = None + v.grad = None + + output, lse, max_logits = triton_mla_forward( + q, k, v, causal=causal, safe=safe, clip_value=clip_value + ) + output_check(output_ref, output, atol=0.05, rtol=0.05, name="output") + # output_check(lse_ref.float(), lse, atol=0.05, rtol=0.05, name='lse') + # output_check(max_logits_ref, max_logits, atol=0.01, rtol=0.03, name='max_logits') + + gq, gk, gv = triton_mla_backward( + g, + output, + q, + k, + v, + lse, + max_logits, + causal=causal, + hpc=hpc, + safe=safe, + clip_value=clip_value, + ) + if clip_value is None: + output_check(gv_ref, gv, atol=0.05, rtol=0.05, name="gv") + output_check(gk_ref, gk, atol=0.05 * coef, rtol=0.05, name="gk") + output_check(gq_ref, gq, atol=0.05 * coef, rtol=0.05, name="gq") + + ref_flops = B * L * L * H * (192 + 128) * (1 if causal else 2) + benchmark( + triton_mla_forward, + q, + k, + v, + causal=causal, + safe=safe, + clip_value=clip_value, + ref_flops=ref_flops, + ) + ref_flops = B * L * L * H * (192 + 128 * 2 + 192 * 2) * (1 if causal else 2) + benchmark( + triton_mla_backward, + g, + output, + q, + k, + v, + lse, + max_logits, + causal=causal, + hpc=hpc, + safe=safe, + clip_value=clip_value, + ref_flops=ref_flops, + n_profile=0, + ) + + +@pytest.mark.parametrize( + "LS,H,causal,hpc,safe,coef,clip_value,pad", + [ + ([8192], 64, True, False, False, 1.0, None, False), + ([8192], 64, True, False, True, 1.0, 100.0, False), + ([8192], 64, True, False, True, 1.0, None, True), + ([4096, 4096], 64, True, False, True, 1.0, None, False), + ([2048, 2048, 4096], 64, True, True, True, 1.0, None, False), + ([127, 873, 3096], 64, False, False, False, 1.0, None, False), + ([127, 873, 3456], 16, False, False, True, 1.0, 100.0, False), + ([127, 873, 3456], 16, False, False, True, 1.0, None, True), + ([127, 873, 3456], 1, True, False, True, 1.0, None, False), + ], +) +def test_varlen_mla(LS, H, causal, hpc, safe, coef, clip_value, pad, benchmark): + dtype = torch.bfloat16 + device = "cuda:0" + if pad: + cu_seqlens = torch.cumsum(torch.tensor([0] + LS, device=device), 0) + LS = [x + 7 for x in LS] + padded_cu_seqlens = torch.cumsum(torch.tensor([0] + LS, device=device), 0) + else: + cu_seqlens = torch.cumsum(torch.tensor([0] + LS, device=device), 0) + padded_cu_seqlens = None + + L = sum(LS) + q = (torch.randn((L, H, 192), device=device, dtype=dtype) * coef).requires_grad_() + k = torch.randn((L, H, 192), device=device, dtype=dtype) + k[:, :, 128:] = k[:, :1, 128:] + k = k.requires_grad_() + v = torch.randn((L, H, 128), device=device, dtype=dtype, requires_grad=True) + g = torch.randn((L, H, 128), device=device, dtype=dtype, requires_grad=True) + cu_seqlens = torch.cumsum(torch.tensor([0] + LS, device=device), 0) + max_q_length = max(LS) + + output_ref, lse_ref, max_logit_ref = torch_varlen_attn( + q, k, v, cu_seqlens, causal=causal, hp=True + ) + output_ref.backward(g, retain_graph=False) + gq_ref = q.grad + gk_ref = k.grad + gv_ref = v.grad + + q.grad = None + k.grad = None + v.grad = None + + output, lse, max_logits = triton_varlen_mla_forward( + q, + k, + v, + cu_seqlens, + max_q_length, + causal=causal, + safe=safe, + clip_value=clip_value, + ) + output_check(output_ref, output, atol=0.05, rtol=0.05, name="output") + if clip_value is None and not safe: + output_check(lse_ref.float(), lse, atol=0.05, rtol=0.05, name="lse") + if clip_value is None and safe: + output_check(max_logit_ref, max_logits, atol=0.01, rtol=0.03, name="max_logits") + + gq, gk, gv = triton_varlen_mla_backward( + g, + output, + q, + k, + v, + lse, + max_logits, + cu_seqlens, + max_q_length, + causal=causal, + hpc=hpc, + safe=safe, + clip_value=clip_value, + ) + if clip_value is None: + output_check(gv_ref, gv, atol=0.05, rtol=0.05, name="gv") + output_check(gk_ref, gk, atol=0.05 * coef, rtol=0.05, name="gk") + output_check(gq_ref, gq, atol=0.05 * coef, rtol=0.05, name="gq") + + output = multi_latend_attention( + q, + k, + v, + cu_seqlens=cu_seqlens, + padded_cu_seqlens=padded_cu_seqlens, + max_q_length=max_q_length, + causal=causal, + safe=safe, + clip_value=0.0 if clip_value is None else clip_value, + ) + output.backward(g) + gq = q.grad + gk = k.grad + gv = v.grad + if clip_value is None and not safe: + output_check(lse_ref.float(), lse, atol=0.05, rtol=0.05, name="lse") + if clip_value is None and safe: + output_check(max_logit_ref, max_logits, atol=0.01, rtol=0.03, name="max_logits") + if clip_value is None: + output_check(gv_ref, gv, atol=0.05, rtol=0.05, name="gv") + output_check(gk_ref, gk, atol=0.05 * coef, rtol=0.05, name="gk") + output_check(gq_ref, gq, atol=0.05 * coef, rtol=0.05, name="gq") + + ref_flops = sum([L * L * H * (192 + 128) * (1 if causal else 2) for L in LS]) + benchmark( + triton_varlen_mla_forward, + q, + k, + v, + cu_seqlens, + max_q_length, + causal=causal, + safe=safe, + clip_value=clip_value, + ref_flops=ref_flops, + ) + ref_flops = sum( + [L * L * H * (192 + 128 * 2 + 192 * 2) * (1 if causal else 2) for L in LS] + ) + benchmark( + triton_varlen_mla_backward, + g, + output, + q, + k, + v, + lse, + max_logits, + cu_seqlens, + max_q_length, + padded_cu_seqlens=padded_cu_seqlens, + causal=causal, + hpc=hpc, + safe=safe, + clip_value=clip_value, + ref_flops=ref_flops, + n_profile=0, + ) + + +@pytest.mark.parametrize( + "B,L,H,causal,hpc,quant_value", + [ + (1, 8192, 64, True, False, False), + ], +) +def test_fp8_mla(B, L, H, causal, hpc, quant_value, benchmark): + dtype = torch.bfloat16 + device = "cuda:0" + q = torch.randn((B, L, H, 192), device=device, dtype=dtype, requires_grad=True) + k = torch.randn((B, L, H, 192), device=device, dtype=dtype, requires_grad=True) + v = torch.randn((B, L, H, 128), device=device, dtype=dtype, requires_grad=True) + + q_q, q_s = head_wise_quant(q) + k_q, k_s = head_wise_quant(k) + v_q, v_s = head_wise_quant(v) + + output_ref, lse_ref, max_logits_ref = torch_attn(q, k, v, causal=causal, hp=True) + + output, lse, max_logits = triton_fp8_mla_forward( + q_q, + k_q, + v_q if quant_value else v, + q_s, + k_s, + vs=v_s if quant_value else None, + causal=causal, + ) + output_check(output_ref, output, atol=0.2, rtol=0.5, name="fp8.output") + output_check(lse_ref.float(), lse, atol=0.2, rtol=0.5, name="fp8.lse") + + ref_flops = B * L * L * H * (192 + 128) * (1 if causal else 2) + benchmark( + triton_fp8_mla_forward, + q_q, + k_q, + v_q if quant_value else v, + q_s, + k_s, + vs=v_s if quant_value else None, + causal=causal, + ref_flops=ref_flops, + ) diff --git a/tests/test_mul.py b/tests/test_mul.py new file mode 100644 index 0000000..670a5b1 --- /dev/null +++ b/tests/test_mul.py @@ -0,0 +1,114 @@ +# -*- coding: utf-8 -*- +""" +Copyright (c) Ant Financial Service Group and its affiliates. +""" + +import random + +import pytest +import torch + +from linghe.tools.check import output_check +from linghe.utils.mul import triton_dot, triton_batch_scale, triton_inplace_scale + + +def torch_fp16_dot(x, y): + return (x * y).sum(1) + + +def torch_inplace_scale(x, scale): + x *= scale + return x + + +def torch_batch_scale(xs, scale): + return [x * scale for x in xs] + + +@pytest.mark.parametrize( + "M,N", + [ + (4096, 4096), + ], +) +def test_dot(M, N, benchmark): + dtype = torch.bfloat16 + device = "cuda:0" + + n_repeat = 100 + + x = torch.randn(M, N, dtype=dtype, device=device) + y = torch.randn(M, N, dtype=dtype, device=device) + q = torch.randn(M, N, dtype=dtype, device=device).to(torch.float8_e4m3fn) + quant_scale = torch.randn(M, dtype=torch.float32, device=device).abs() + smooth_scale = torch.randn(N, dtype=torch.float32, device=device).abs() + + sums_ref = torch_fp16_dot(x, q.float().to(dtype)) + sums = triton_dot(x, q) + output_check(sums_ref, sums, "dot", atol=1.0) + + sums_ref = ( + x.float() * (q.to(torch.float32) * quant_scale[:, None] * smooth_scale[None, :]) + ).sum(dim=1) + + ref_time = benchmark(torch_fp16_dot, x, y, n_repeat=n_repeat) + benchmark(triton_dot, x, q, n_repeat=n_repeat, ref_time=ref_time) + + +@pytest.mark.parametrize( + "M", + [ + 2**28 + 1, + ], +) +def test_inplace_scale(M, benchmark): + x = torch.randn((M,), device="cuda:0", dtype=torch.float32) + scale = 7.86 + sum_ref = torch_inplace_scale(x, scale) + sums = triton_inplace_scale(x, scale) + output_check(sum_ref, sums, "sum") + + ref_bytes = M * 8 + + ref_time = benchmark(torch_inplace_scale, x, scale, ref_bytes=ref_bytes) + benchmark(triton_inplace_scale, x, scale, ref_bytes=ref_bytes, ref_time=ref_time) + + +@pytest.mark.parametrize( + "M,N,k,scale", + [ + (2048, 1024, 128, 2.0), + (2048, 1024, 128, 0.0), + ], +) +def test_batch_scale(M, N, k, scale, benchmark): + dtype = torch.float32 + xs = [ + torch.randn( + random.randint(1, int(M**0.5)) ** 2, + random.randint(1, int(N**0.5)) ** 2, + dtype=dtype, + device="cuda:0", + ) + for i in range(k) + ] + xs1 = [x.clone().detach() for x in xs] + xs2 = [x.clone().detach() for x in xs] + if scale == 0.0: + xs1[0][:10] = float("inf") + xs1[0][10:20] = -float("inf") + xs2[0][:10] = float("inf") + xs2[0][10:20] = -float("inf") + + sum_ref = torch_batch_scale(xs1, scale) + sums = triton_batch_scale(xs2, scale) + output_check( + torch.cat([x.view(-1) for x in sum_ref], 0), + torch.cat([x.view(-1) for x in sums], 0), + "batch_clip", + ) + + ref_bytes = sum([x.numel() for x in xs]) * 8 + + ref_time = benchmark(torch_batch_scale, xs, scale, ref_bytes=ref_bytes) + benchmark(triton_batch_scale, xs, scale, ref_bytes=ref_bytes, ref_time=ref_time) diff --git a/tests/test_norm.py b/tests/test_norm.py index c94f455..6745e77 100644 --- a/tests/test_norm.py +++ b/tests/test_norm.py @@ -4,72 +4,66 @@ """ import torch -import torch.nn.functional as F - -from linghe.utils.norm import (triton_rms_norm_and_smooth_quant_forward, - triton_rms_norm_and_block_quant_forward, - triton_rms_norm_backward, - triton_rms_norm_forward, - triton_group_rms_norm_gate_forward, - triton_group_rms_norm_gate_backward) -from linghe.tools.util import (output_check, - torch_smooth_quant, - torch_group_quant) + +from linghe.gemm.fp32_gemm import triton_fp32_gemm +from linghe.quant.group import triton_group_quant from linghe.tools.benchmark import benchmark_func +from linghe.tools.check import output_check +from linghe.tools.util import torch_smooth_quant, torch_group_quant +from linghe.utils.norm import ( + triton_rms_norm_and_smooth_quant_forward, + triton_rms_norm_and_block_quant_forward, + triton_rms_norm_fp32_gemm_block_quant_forward, + triton_rms_norm_backward, + triton_rms_norm_forward, +) def torch_rms_forward(x, weight): + dtype = x.dtype x = x.float() weight = weight.float() N = x.shape[-1] rmsnorm = torch.nn.RMSNorm( - normalized_shape=N, - eps=1e-6, - dtype=torch.float32, - device=x.device + normalized_shape=N, eps=1e-6, dtype=torch.float32, device=x.device ) with torch.no_grad(): rmsnorm.weight.copy_(weight) - return rmsnorm(x) + return rmsnorm(x).to(dtype) def torch_rms_backward(x, weight, dy): + dtype = x.dtype x = x.float() weight = weight.float() dy = dy.float() N = x.shape[-1] rmsnorm = torch.nn.RMSNorm( - normalized_shape=N, - eps=1e-6, - dtype=torch.float32, - device=x.device + normalized_shape=N, eps=1e-6, dtype=torch.float32, device=x.device ) with torch.no_grad(): rmsnorm.weight.copy_(weight) x = x.clone().detach().requires_grad_() y = rmsnorm(x) y.backward(gradient=dy) - return x.grad, rmsnorm.weight.grad + return x.grad.to(dtype), rmsnorm.weight.grad.to(dtype) -def torch_rms_and_smooth_quant_forward(x, weight, smooth_scale=None, - round_scale=False): +def torch_rms_and_smooth_quant_forward(x, weight, smooth_scale=None, round_scale=False): x = x.float() weight = weight.float() smooth_scale = smooth_scale.float() N = x.shape[-1] rmsnorm = torch.nn.RMSNorm( - normalized_shape=N, - eps=1e-6, - dtype=torch.float32, - device=x.device + normalized_shape=N, eps=1e-6, dtype=torch.float32, device=x.device ) with torch.no_grad(): rmsnorm.weight.copy_(weight) y = rmsnorm(x) # smooth - y_q, y_scale, y_maxs = torch_smooth_quant(y, smooth_scale, reverse=False, - round_scale=round_scale) + y_q, y_scale, y_maxs = torch_smooth_quant( + y, smooth_scale, reverse=False, round_scale=round_scale + ) return y_q, y_scale, y_maxs @@ -78,233 +72,266 @@ def torch_rms_and_block_quant_forward(x, weight, round_scale=False): weight = weight.float() N = x.shape[-1] rmsnorm = torch.nn.RMSNorm( - normalized_shape=N, - eps=1e-6, - dtype=torch.float32, - device=x.device + normalized_shape=N, eps=1e-6, dtype=torch.float32, device=x.device ) with torch.no_grad(): rmsnorm.weight.copy_(weight) y = rmsnorm(x) + rms = torch.rsqrt(torch.sum(x**2, 1) / N + 1e-6) # blockwise y_q, y_scale = torch_group_quant(y, round_scale=round_scale) yt_q, yt_scale = torch_group_quant(y.t(), round_scale=round_scale) - return y_q, y_scale, yt_q, yt_scale + return y_q, y_scale.t(), rms, yt_q, yt_scale.t() -@torch.compile -def torch_group_rms_norm_gate_forward(x, gate, weight, eps=1e-6, group_size=4, transpose=True): +def torch_rms_gemm_block_quant_forward(x, norm_weight, route_weight, round_scale=False): + dtype = x.dtype x = x.float() - gate = gate.float() - weight = weight.float() - if transpose: - length, bs, dim = gate.shape - else: - bs, length, dim = gate.shape - d = dim // group_size - attn_output = x.view(bs, length, group_size, d).transpose(0, 1) - outputs = [] - for i in range(group_size): - outputs.append(F.rms_norm(attn_output[:, :, i], [d], - weight=weight[i * d:(i + 1) * d], eps=eps)) - outputs = torch.stack(outputs, 2) - if transpose: - outputs = outputs.view(length, bs, dim) - else: - outputs = outputs.view(bs, length, dim) - gate = F.sigmoid(gate) - return outputs * gate - - -def torch_group_rms_norm_gate_backward(grad_output, x, gate, weight, eps=1e-6, - group_size=4, - transpose=True): - grad_output = grad_output.float() - x = x.float().clone().detach().requires_grad_() - gate = gate.float().clone().detach().requires_grad_() - weight = weight.float().clone().detach().requires_grad_() - y = torch_group_rms_norm_gate_forward(x, gate, weight, eps=eps, - group_size=group_size, - transpose=transpose) - y.backward(gradient=grad_output) - return x.grad, gate.grad, weight.grad + norm_weight = norm_weight.float() + route_weight = route_weight.float() + N = x.shape[-1] + rmsnorm = torch.nn.RMSNorm( + normalized_shape=N, eps=1e-6, dtype=torch.float32, device=x.device + ) + with torch.no_grad(): + rmsnorm.weight.copy_(norm_weight) + y = rmsnorm(x) + logits = y @ route_weight.t() + # blockwise + y_q, y_scale = torch_group_quant(y, round_scale=round_scale) + yt_q, yt_scale = torch_group_quant(y.t(), round_scale=round_scale) + + return y.to(dtype), logits, y_q, y_scale, yt_q, yt_scale + + +def split_rms_gemm_block_quant_forward(x, norm_weight, route_weight, round_scale=False): + y, _ = triton_rms_norm_forward(x, norm_weight) + logit = triton_fp32_gemm(y, route_weight) + q, s = triton_group_quant(y, round_scale=round_scale) + return y, logit, q, s def test_rmsnorm(M=4096, N=4096, bench=False): dtype = torch.bfloat16 - device = 'cuda:0' + device = "cuda:0" x = torch.randn(M, N, dtype=dtype, requires_grad=True, device=device) ** 3 weight = torch.randn(N, dtype=dtype, requires_grad=True, device=device) dy = torch.randn(M, N, dtype=dtype, device=device) y_ref = torch_rms_forward(x, weight) - y = triton_rms_norm_forward(x, weight) - output_check(y_ref.float(), y.float(), 'y') + y, rms = triton_rms_norm_forward(x, weight) + output_check(y_ref, y, "y") + + y_with_rms, _ = triton_rms_norm_forward(x, weight, rms=rms) + output_check(y_ref, y_with_rms, "y_with_rms") dx_ref, dw_ref = torch_rms_backward(x, weight, dy) dx, dw = triton_rms_norm_backward(dy, x, weight) - output_check(dx_ref, dx, mode="dx") - output_check(dw_ref, dw, mode='dw') + output_check(dx_ref, dx, name="dx") + output_check(dw_ref, dw.to(dtype), name="dw") + + dx_with_rms, dw_with_rms = triton_rms_norm_backward(dy, x, weight, rms=rms) + output_check(dx_ref, dx_with_rms, name="dx_with_rms") + output_check(dw_ref, dw_with_rms.to(dtype), name="dw_with_rms") if bench: benchmark_func(triton_rms_norm_forward, x, weight, ref_bytes=M * N * 3) - benchmark_func(triton_rms_norm_backward, dy, x, weight, - ref_bytes=M * N * 3) + benchmark_func(triton_rms_norm_forward, x, weight, rms=rms, ref_bytes=M * N * 3) + benchmark_func(triton_rms_norm_backward, dy, x, weight, ref_bytes=M * N * 3) + benchmark_func( + triton_rms_norm_backward, dy, x, weight, rms=rms, ref_bytes=M * N * 3 + ) def test_rmsnorm_and_smooth_quant(M=4096, N=4096, bench=False): dtype = torch.bfloat16 - device = 'cuda:0' + device = "cuda:0" x = torch.randn(M, N, dtype=dtype, requires_grad=True, device=device) weight = torch.randn(N, dtype=dtype, requires_grad=True, device=device) - smooth_scale = torch.rand(N, dtype=torch.float32, requires_grad=False, - device=device) + 0.1 + smooth_scale = ( + torch.rand(N, dtype=torch.float32, requires_grad=False, device=device) + 0.1 + ) calibrate = True # smooth - q_ref, scale_ref, maxs_ref = torch_rms_and_smooth_quant_forward(x, weight, - smooth_scale=smooth_scale, - round_scale=True) - - q, scale, maxs, rms = triton_rms_norm_and_smooth_quant_forward(x, weight, - smooth_scale=smooth_scale, - calibrate=calibrate, - output_rms=True, - round_scale=True) - output_check(q_ref, q, mode="smooth.data") - output_check(scale_ref, scale, mode='smooth.scale') + q_ref, scale_ref, maxs_ref = torch_rms_and_smooth_quant_forward( + x, weight, smooth_scale=smooth_scale, round_scale=True + ) + + q, scale, maxs, rms = triton_rms_norm_and_smooth_quant_forward( + x, + weight, + smooth_scale=smooth_scale, + calibrate=calibrate, + output_rms=True, + round_scale=True, + ) + output_check(q_ref, q, name="smooth.data", atol=-1) + output_check(scale_ref, scale, name="smooth.scale") if calibrate: - output_check(maxs_ref, maxs, mode="smooth.maxs") + output_check(maxs_ref, maxs, name="smooth.maxs") if bench: - benchmark_func(triton_rms_norm_and_smooth_quant_forward, x, weight, - smooth_scale=smooth_scale, - calibrate=True, - round_scale=True, - output_rms=True, - ref_bytes=M * N * 3) + benchmark_func( + triton_rms_norm_and_smooth_quant_forward, + x, + weight, + smooth_scale=smooth_scale, + calibrate=True, + round_scale=True, + output_rms=True, + ref_bytes=M * N * 3, + ) def test_rmsnorm_and_block_quant(M=4096, N=4096, bench=False): dtype = torch.bfloat16 - device = 'cuda:0' + device = "cuda:0" x = torch.randn(M, N, dtype=dtype, requires_grad=True, device=device) ** 2 weight = torch.randn(N, dtype=dtype, requires_grad=True, device=device) # blockwise - q_ref, scale_ref, qt_ref, scale_t_ref = torch_rms_and_block_quant_forward(x, - weight, - round_scale=True) - q, scale, rms, q_t, scale_t = triton_rms_norm_and_block_quant_forward(x, - weight, - round_scale=True, - output_mode=2) - output_check(q_ref, q, mode="2.block.data") - output_check(scale_ref.t(), scale, mode='2.block.scale') - output_check(qt_ref, q_t, mode='2.block.t_data') - output_check(scale_t_ref.t(), scale_t, mode="2.block.t_scale") - - q, scale, _, _, _ = triton_rms_norm_and_block_quant_forward(x, weight, - round_scale=True, - output_mode=0) - output_check(q_ref, q, mode="0.block.data") - output_check(scale_ref.t(), scale, mode='0.block.scale') - - _, _, _, q_t, scale_t = triton_rms_norm_and_block_quant_forward(x, weight, - round_scale=True, - rms=rms, - output_mode=1) - output_check(qt_ref, q_t, mode='0.block.t_data') - output_check(scale_t_ref.t(), scale_t, mode="0.block.t_scale") - - if bench: - benchmark_func(triton_rms_norm_and_block_quant_forward, x, weight, - round_scale=True, - output_mode=0, - ref_bytes=M * N * 3) - - benchmark_func(triton_rms_norm_and_block_quant_forward, x, weight, - round_scale=True, - output_mode=1, - ref_bytes=M * N * 3) + q_ref, scale_ref, rms_ref, qt_ref, scale_t_ref = torch_rms_and_block_quant_forward( + x, weight, round_scale=True + ) - benchmark_func(triton_rms_norm_and_block_quant_forward, x, weight, - round_scale=True, - output_mode=2, - ref_bytes=M * N * 4) + q, scale, rms, _, _ = triton_rms_norm_and_block_quant_forward( + x, weight, round_scale=True, output_mode=0 + ) + output_check(q_ref, q, name="0.block.data", rtol=0.125) + output_check(scale_ref, scale, name="0.block.scale") + _, _, _, q_t, scale_t = triton_rms_norm_and_block_quant_forward( + x, weight, round_scale=True, rms=rms, output_mode=1 + ) + output_check(qt_ref, q_t, name="1.block.t_data", rtol=0.125) + output_check(scale_t_ref, scale_t, name="1.block.t_scale") -def test_group_rms_norm_gate(bs=1, length=4096, dim=4096, group_size=4, - transpose=True, - bench=False): - dtype = torch.bfloat16 - device = 'cuda:0' - x = torch.randn(bs, length, dim, dtype=dtype, requires_grad=True, - device=device) ** 2 - weight = torch.randn(dim, dtype=dtype, requires_grad=True, device=device) - if transpose: - gate = torch.randn(length, bs, dim, dtype=dtype, requires_grad=True, - device=device) - grad_output = torch.randn(length, bs, dim, dtype=dtype, requires_grad=True, - device=device) - else: - gate = torch.randn(bs, length, dim, dtype=dtype, requires_grad=True, - device=device) - grad_output = torch.randn(bs, length, dim, dtype=dtype, requires_grad=True, - device=device) - - output_ref = torch_group_rms_norm_gate_forward(x, gate, weight, - group_size=group_size, - transpose=transpose) - output = triton_group_rms_norm_gate_forward(x, gate, weight, - group_size=group_size, - transpose=transpose) - output_check(output_ref, output.float(), mode='group_norm_gate.y') - - dx_ref, dg_ref, dw_ref = torch_group_rms_norm_gate_backward(grad_output, x, - gate, weight, - group_size=group_size, - transpose=transpose) - dx, dg, dw = triton_group_rms_norm_gate_backward(grad_output, x, gate, weight, - group_size=group_size, - transpose=transpose) - output_check(dx_ref, dx.float(), mode='group_norm_gate.dx') - output_check(dg_ref, dg.float(), mode='group_norm_gate.dg') - output_check(dw_ref, dw.float(), mode='group_norm_gate.dw') + q, scale, rms, q_t, scale_t = triton_rms_norm_and_block_quant_forward( + x, weight, round_scale=True, output_mode=2 + ) + output_check(q_ref, q, name="2.block.data", rtol=0.125) + output_check(scale_ref, scale, name="2.block.scale") + output_check(qt_ref, q_t, name="2.block.t_data", rtol=0.125) + output_check(scale_t_ref, scale_t, name="2.block.t_scale") if bench: - benchmark_func(torch_group_rms_norm_gate_forward, x, gate, weight, - group_size=group_size, - ref_bytes=bs * length * dim * 6) - - benchmark_func(triton_group_rms_norm_gate_forward, x, gate, weight, - group_size=group_size, - ref_bytes=bs * length * dim * 6) + benchmark_func( + triton_rms_norm_and_block_quant_forward, + x, + weight, + round_scale=True, + output_mode=0, + ref_bytes=M * N * 3, + ) + + benchmark_func( + triton_rms_norm_and_block_quant_forward, + x, + weight, + round_scale=True, + output_mode=1, + rms=rms, + ref_bytes=M * N * 3, + ) + + benchmark_func( + triton_rms_norm_and_block_quant_forward, + x, + weight, + round_scale=True, + output_mode=2, + ref_bytes=M * N * 6, + ) + + +def test_rms_norm_fp32_gemm_block_quant_forward(M=8192, N=256, K=2048, bench=False): + dtype = torch.bfloat16 + device = "cuda:0" - benchmark_func(triton_group_rms_norm_gate_backward, grad_output, x, gate, - weight, group_size=group_size, - ref_bytes=bs * length * dim * 10) + round_scale = False + x = torch.randn(M, K, dtype=dtype, requires_grad=True, device=device) ** 2 + norm_weight = torch.randn(K, dtype=dtype, requires_grad=True, device=device) + route_weight = torch.randn(N, K, dtype=dtype, requires_grad=True, device=device) + # blockwise + y_ref, logit_ref, q_ref, scale_ref, qt_ref, scale_t_ref = ( + torch_rms_gemm_block_quant_forward( + x, norm_weight, route_weight, round_scale=round_scale + ) + ) + y, rms, logit, q, scale, q_t, scale_t = ( + triton_rms_norm_fp32_gemm_block_quant_forward( + x, norm_weight, route_weight, round_scale=round_scale, output_mode=0 + ) + ) + output_check(y_ref, y, name="0.y") + output_check(logit_ref, logit, name="0.logit", atol=0.1, rtol=0.001) + output_check(q_ref, q, name="0.block.data") + output_check(scale_ref.t(), scale, name="0.block.scale") + + y, rms, logit, q, scale, q_t, scale_t = ( + triton_rms_norm_fp32_gemm_block_quant_forward( + x, + norm_weight, + route_weight, + rms=rms, + round_scale=round_scale, + output_mode=1, + ) + ) + output_check(qt_ref, q_t, name="1.block.data") + output_check(scale_t_ref.t(), scale_t, name="1.block.scale") -if __name__ == '__main__': + if bench: + benchmark_func( + triton_rms_norm_fp32_gemm_block_quant_forward, + x, + norm_weight, + route_weight, + round_scale=round_scale, + output_mode=0, + ref_bytes=M * K * 5, + ref_flops=M * K * N * 2, + ) + + benchmark_func( + triton_rms_norm_fp32_gemm_block_quant_forward, + x, + norm_weight, + route_weight, + rms=rms, + round_scale=round_scale, + output_mode=1, + ref_bytes=M * K * 3, + ) + + benchmark_func( + split_rms_gemm_block_quant_forward, + x, + norm_weight, + route_weight, + round_scale=round_scale, + ref_bytes=M * K * 9, + ) + + +if __name__ == "__main__": test_rmsnorm(M=16384, N=2048, bench=False) + test_rmsnorm(M=16384, N=1664, bench=False) + test_rmsnorm(M=1664, N=1664, bench=False) test_rmsnorm(M=8192, N=4096, bench=False) test_rmsnorm(M=4096, N=8192, bench=False) + + test_rmsnorm_and_block_quant(M=4096, N=2048, bench=False) + test_rmsnorm_and_block_quant(M=8192, N=1536, bench=False) + test_rmsnorm_and_block_quant(M=16384, N=1536, bench=False) + test_rmsnorm_and_smooth_quant(M=16384, N=2048, bench=False) test_rmsnorm_and_smooth_quant(M=8192, N=4096, bench=False) test_rmsnorm_and_smooth_quant(M=4096, N=8192, bench=False) - test_rmsnorm_and_block_quant(M=128, N=2048, bench=False) - test_rmsnorm_and_block_quant(M=8192, N=4096, bench=False) - test_group_rms_norm_gate(bs=2, length=4096, dim=2048, group_size=4, - transpose=True, - bench=False) - test_group_rms_norm_gate(bs=2, length=4096, dim=2048, group_size=4, - transpose=False, - bench=False) - test_group_rms_norm_gate(bs=1, length=4096, dim=4096, group_size=4, - bench=False) - + test_rms_norm_fp32_gemm_block_quant_forward(M=8192 * 2, N=256, K=2048, bench=False) diff --git a/tests/test_rearange.py b/tests/test_rearange.py index 384ba72..c063fbd 100644 --- a/tests/test_rearange.py +++ b/tests/test_rearange.py @@ -3,14 +3,14 @@ Copyright (c) Ant Financial Service Group and its affiliates. """ +import pytest import torch -from linghe.tools.benchmark import benchmark_func -from linghe.tools.util import output_check -from linghe.utils.rearange import triton_split_and_cat +from linghe.tools.check import output_check +from linghe.utils.rearange import triton_sort_chunks_by_index -def torch_split_and_cat(x, scales, counts, indices): +def torch_sort_chunks_by_index(x, scales, counts, indices): n = len(counts) chunks = torch.split(x, counts) chunks = [chunks[indices[i]] for i in range(n)] @@ -22,18 +22,15 @@ def torch_split_and_cat(x, scales, counts, indices): return output_data, output_scale -def test_triton_split_and_cat(M=4096, N=4096, bench=False): - # M, N, K = 8192, 10240, 8192 # max qkv - # M, N, K = 8192, 8192, 8192 # max out - # M, N, K = 2048, 4096, 8192 # max gate_up - # M, N, K = 2048, 8192, 2048 # max down - # M, N, K = 8192, 8192, 8192 - # M, N, K = 2048, 8192, 8192 - - # M, N, K = M-1, N-1, K-1 - +@pytest.mark.parametrize( + "M,N", + [ + (4096, 4096), + ], +) +def test_sort_chunks_by_index(M, N, benchmark): dtype = torch.bfloat16 - device = 'cuda:0' + device = "cuda:0" n_repeat = 100 x = torch.randn(M, N, dtype=dtype, device=device) @@ -49,29 +46,32 @@ def test_triton_split_and_cat(M=4096, N=4096, bench=False): chunks = torch.split(x_q.view(torch.float8_e4m3fn), split_size_list) scale_chunks = torch.split(x_scales, split_size_list) - data_ref, scale_ref = torch_split_and_cat(x_q.view(torch.float8_e4m3fn), - x_scales, split_size_list, - sorted_indices_list) - - data, scale = triton_split_and_cat(x_q, counts, indices, scales=x_scales) - - output_check(data_ref.view(torch.float8_e4m3fn).float(), data.float(), - mode='data') - output_check(scale_ref, scale, mode='scale') - - if bench: - benchmark_func(torch.split, x_q.view(torch.uint8), split_size_list, - n_repeat=n_repeat) - benchmark_func(torch.cat, chunks, dim=0, n_repeat=n_repeat) - benchmark_func(torch.split, x_scales, split_size_list, - n_repeat=n_repeat) - benchmark_func(torch.cat, scale_chunks, dim=0, n_repeat=n_repeat) - benchmark_func(torch_split_and_cat, x_q.view(torch.float8_e4m3fn), - x_scales, split_size_list, sorted_indices_list, - n_repeat=n_repeat) - benchmark_func(triton_split_and_cat, x_q, counts, indices, - scales=x_scales, n_repeat=n_repeat) - - -if __name__ == '__main__': - test_triton_split_and_cat(M=4096, N=4096) + data_ref, scale_ref = torch_sort_chunks_by_index( + x_q.view(torch.float8_e4m3fn), x_scales, split_size_list, sorted_indices_list + ) + + data, scale = triton_sort_chunks_by_index(x_q, counts, indices, scales=x_scales) + + output_check(data_ref.view(torch.float8_e4m3fn), data, name="data") + output_check(scale_ref, scale, name="scale") + + benchmark(torch.split, x_q.view(torch.uint8), split_size_list, n_repeat=n_repeat) + benchmark(torch.cat, chunks, dim=0, n_repeat=n_repeat) + benchmark(torch.split, x_scales, split_size_list, n_repeat=n_repeat) + benchmark(torch.cat, scale_chunks, dim=0, n_repeat=n_repeat) + benchmark( + torch_sort_chunks_by_index, + x_q.view(torch.float8_e4m3fn), + x_scales, + split_size_list, + sorted_indices_list, + n_repeat=n_repeat, + ) + benchmark( + triton_sort_chunks_by_index, + x_q, + counts, + indices, + scales=x_scales, + n_repeat=n_repeat, + ) diff --git a/tests/test_reduce.py b/tests/test_reduce.py index bff2b65..37d7947 100644 --- a/tests/test_reduce.py +++ b/tests/test_reduce.py @@ -3,42 +3,70 @@ Copyright (c) Ant Financial Service Group and its affiliates. """ +import random + +import pytest import torch -from linghe.tools.benchmark import benchmark_func -from linghe.tools.util import output_check -from linghe.utils.reduce import (triton_abs_max, - triton_batch_count_zero, - triton_batch_sum_with_ord) +from linghe.tools.check import output_check +from linghe.utils.reduce import ( + triton_abs_max, + triton_batch_count_zero, + triton_norm, + triton_batch_norm, +) -def torch_sum(xs): - return sum([x.square().sum() for x in xs]) +def torch_sum(xs, ord=2, norm=True): + if ord == 2: + output = sum([x.square().sum() for x in xs]) + if norm: + output = torch.sqrt(output) + return output + elif ord == 1: + return sum([x.abs().sum() for x in xs]) + elif ord == -1: + return max([x.abs().max() for x in xs]) def torch_count_zero(xs): - count = torch.tensor([0], dtype=torch.int64, device='cuda') + count = torch.tensor([0], dtype=torch.int64, device="cuda") for x in xs: count += x.numel() - torch.count_nonzero(x) return count -def test_triton_abs_max(M=4096, N=4096, bench=False): - x = 100 * torch.randn(M, 1, N, dtype=torch.bfloat16, device='cuda:0') +@pytest.mark.parametrize( + "M,N", + [ + (4096, 4096), + ], +) +def test_triton_abs_max(M, N, benchmark): + x = 100 * torch.randn(M, 1, N, dtype=torch.bfloat16, device="cuda:0") # scales = 1.0/torch.sqrt(torch.maximum(x[:,0].abs().float().amax(0), torch.ones(M,N,dtype=dtype,device=device)) ) # smooth_scale_ref = torch.exp2(torch.ceil(torch.log2(scales))) maxs_ref = x.abs().amax(0).float().view(N) maxs = triton_abs_max(x) - output_check(maxs_ref, maxs, 'abs_max') - if bench: - benchmark_func(triton_abs_max, x, n_repeat=100, ref_bytes=M * N * 2) - - -def test_count_zero(M=4096, N=8192, k=32, bench=False): - xs = [torch.randn(M, N, dtype=torch.float32, device='cuda:0').to( - torch.float8_e4m3fn).to(torch.float32) for i in range(k)] + output_check(maxs_ref, maxs, "abs_max") + benchmark(triton_abs_max, x, n_repeat=100, ref_bytes=M * N * 2) + + +@pytest.mark.parametrize( + "M,N,k", + [ + (4096, 8192, 32), + ], +) +def test_count_zero(M, N, k, benchmark): + xs = [ + torch.randn(M, N, dtype=torch.float32, device="cuda:0") + .to(torch.float8_e4m3fn) + .to(torch.float32) + for i in range(k) + ] ref_bytes = sum([x.numel() for x in xs]) * 4 @@ -47,37 +75,81 @@ def test_count_zero(M=4096, N=8192, k=32, bench=False): # print(f'{count_ref=} {count=}') assert count_ref.item() - count.item() == 0 - sum_ref = torch_sum(xs) - sums = triton_batch_sum_with_ord(xs) - output_check(sum_ref, sums, 'count_zero') - - if bench: - n_repeat = 100 - ref_time = benchmark_func(torch_count_zero, xs, n_repeat=n_repeat, - ref_bytes=ref_bytes) - benchmark_func(triton_batch_count_zero, xs, n_repeat=n_repeat, - ref_bytes=ref_bytes, ref_time=ref_time) - - -def test_ord_sum(M=4096, N=8192, k=32, bench=False): - xs = [torch.randn(M, N, dtype=torch.float32, device='cuda:0').to( - torch.float8_e4m3fn).to(torch.float32) for i in range(k)] + n_repeat = 100 + ref_time = benchmark(torch_count_zero, xs, n_repeat=n_repeat, ref_bytes=ref_bytes) + benchmark( + triton_batch_count_zero, + xs, + n_repeat=n_repeat, + ref_bytes=ref_bytes, + ref_time=ref_time, + ) + + +@pytest.mark.parametrize( + "M,N", + [ + (100000, 8192), + ], +) +def test_norm(M, N, benchmark, coef=1.0): + x = torch.randn(M, N, dtype=torch.float32, device="cuda:0") * 1.0 + + sum_ref = x.norm(p=2) + sums = triton_norm(x, ord=2, norm=True, scalar=True) + output_check(sum_ref, sums, "l2_norm") + + sum_ref = x.norm(p=1) + sums = triton_norm(x, ord=1, norm=True, scalar=True) + output_check(sum_ref, sums, "l1_norm") + + ref_bytes = M * N * 4 + n_repeat = 100 + ref_time = benchmark( + lambda x: x.norm(p=2), x, n_repeat=n_repeat, ref_bytes=ref_bytes + ) + benchmark( + triton_norm, + x, + ord=2, + norm=True, + scalar=True, + n_repeat=n_repeat, + ref_bytes=ref_bytes, + ref_time=ref_time, + ) + + +@pytest.mark.parametrize( + "N,k,coef", + [ + (1024, 16, 1.0), + (1024, 64, 1.0), + (2048, 1024, 1e12), + ], +) +def test_batch_norm(N, k, coef, benchmark, M=4096): + bs = [random.randint(1, int(M**0.5)) ** 2 for i in range(k)] + xs = [ + torch.randn(bs[i], N, dtype=torch.float32, device="cuda:0") * coef + for i in range(k) + ] + + sum_ref = torch_sum(xs, ord=2, norm=False) + sums = triton_batch_norm(xs, ord=2, norm=False) + output_check(sum_ref, sums, "l2_norm") + + sum_ref = torch_sum(xs, ord=1, norm=False) + sums = triton_batch_norm(xs, ord=1, norm=False) + output_check(sum_ref, sums, "l1_norm") + + sum_ref = torch_sum(xs, ord=-1, norm=False) + sums = triton_batch_norm(xs, ord=-1, norm=False) + output_check(sum_ref, sums, "inf_norm") ref_bytes = sum([x.numel() for x in xs]) * 4 - - sum_ref = torch_sum(xs) - sums = triton_batch_sum_with_ord(xs) - output_check(sum_ref, sums, 'norm') - - if bench: - n_repeat = 100 - ref_time = benchmark_func(torch_sum, xs, n_repeat=n_repeat, - ref_bytes=ref_bytes) - benchmark_func(triton_batch_sum_with_ord, xs, n_repeat=n_repeat, - ref_bytes=ref_bytes, ref_time=ref_time) - - -if __name__ == '__main__': - test_triton_abs_max(M=4096, N=4096) - test_count_zero(M=4096, N=8192, k=32) - test_ord_sum(M=4096, N=8192, k=32) + n_repeat = 100 + ref_time = benchmark(torch_sum, xs, n_repeat=n_repeat, ref_bytes=ref_bytes) + benchmark( + triton_batch_norm, xs, n_repeat=n_repeat, ref_bytes=ref_bytes, ref_time=ref_time + ) diff --git a/tests/test_rope.py b/tests/test_rope.py index f3ffd31..21852ec 100644 --- a/tests/test_rope.py +++ b/tests/test_rope.py @@ -3,55 +3,72 @@ Copyright (c) Ant Financial Service Group and its affiliates. """ +import pytest import torch -from linghe.tools.benchmark import benchmark_func -from linghe.tools.util import output_check -from linghe.utils.rope import triton_half_rope_forward, \ - triton_half_rope_backward, triton_qk_norm_and_half_rope_forward, \ - triton_qk_norm_and_half_rope_backward +from linghe.facade.rope import qk_norm_half_rope +from linghe.tools.check import output_check +from linghe.utils.rope import ( + triton_half_rope_forward, + triton_half_rope_backward, + triton_qk_norm_and_half_rope_forward, + triton_qk_norm_and_half_rope_backward, + triton_mla_rope_forward, + triton_mla_rope_backward, + triton_varlen_qk_norm_and_half_rope_forward, + triton_varlen_qk_norm_and_half_rope_backward, +) def rotate_half(x): """Rotates half the hidden dims of the input.""" x1 = x[..., : x.shape[-1] // 2] - x2 = x[..., x.shape[-1] // 2:] + x2 = x[..., x.shape[-1] // 2 :] return torch.cat((-x2, x1), dim=-1) def apply_rotary_pos_emb(q, k, cos, sin, position_ids): - cos = cos[position_ids][:, :, None] - sin = sin[position_ids][:, :, None] + if cos.ndim == 2: + cos = cos[position_ids][:, :, None] + sin = sin[position_ids][:, :, None] + elif cos.ndim == 4: + cos = cos[:, 0, 0][position_ids][:, :, None] + sin = sin[:, 0, 0][position_ids][:, :, None] + else: + raise ValueError("unsupported ndim=3") q_embed = (q * cos) + (rotate_half(q) * sin) k_embed = (k * cos) + (rotate_half(k) * sin) return q_embed, k_embed 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() + 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 torch_half_rope(q, k, freqs, rope_theta=10000.0, transposed=True): +def torch_half_rope(q, k, freqs, transposed=True): + dtype = q.dtype if transposed: L, B, H, D = q.shape else: B, L, H, D = q.shape d = D // 2 - cos = freqs.cos().to(q.dtype) - sin = freqs.sin().to(q.dtype) + cos = freqs.cos() + sin = freqs.sin() if transposed: - position_ids = torch.arange(L, device='cuda:0')[:, None].expand(-1, B) + position_ids = torch.arange(L, device="cuda:0")[:, None].expand(-1, B) else: - position_ids = torch.arange(L, device='cuda:0')[None, :].expand(B, -1) - qr, kr = apply_rotary_pos_emb(q[:, :, :, :d], k[:, :, :, :d], cos, sin, - position_ids) + position_ids = torch.arange(L, device="cuda:0")[None, :].expand(B, -1) + qr, kr = apply_rotary_pos_emb( + q[:, :, :, :d], k[:, :, :, :d], cos, sin, position_ids + ) qo = torch.cat([qr, q[:, :, :, d:]], dim=-1) ko = torch.cat([kr, k[:, :, :, d:]], dim=-1) - return qo, ko + return qo.to(dtype), ko.to(dtype) def torch_qk_norm(q, k, qw, kw, eps=1e-6, transposed=True): @@ -69,15 +86,132 @@ def torch_qk_norm(q, k, qw, kw, eps=1e-6, transposed=True): return q.to(dtype), k.to(dtype) -def torch_qk_norm_and_half_rope(qkv, qw, kw, freqs, rope_theta=10000.0, H=32, - h=4, eps=1e-6, interleaved=True, transposed=True): +def torch_mla_rope(q, kv, k_pos_emb, freqs, mscale=1.0, transpose=False): + dtype = q.dtype + q = q.float() + kv = kv.float() + k_pos_emb = k_pos_emb.float() + L, B, H, _ = q.shape + q_no_pe, q_pos_emb = torch.split(q, [128, 64], dim=-1) + + k_no_pe, value = torch.split(kv, [128, 128], dim=-1) + + cos = freqs.cos() * mscale + sin = freqs.sin() * mscale + position_ids = torch.arange(L, device="cuda:0")[:, None].expand(-1, B) + + q_pos_emb = torch.cat([q_pos_emb[:, :, :, 0::2], q_pos_emb[:, :, :, 1::2]], -1) + k_pos_emb = torch.cat([k_pos_emb[:, :, :, 0::2], k_pos_emb[:, :, :, 1::2]], -1) + + q_pos_emb, k_pos_emb = apply_rotary_pos_emb( + q_pos_emb, k_pos_emb, cos, sin, position_ids + ) + + query = torch.cat([q_no_pe, q_pos_emb], dim=-1) + + k_pos_emb = k_pos_emb.expand(-1, -1, H, -1) + + key = torch.cat([k_no_pe, k_pos_emb], dim=-1) + + value = value.contiguous() + + if transpose: + query = query.transpose(0, 1) + key = key.transpose(0, 1) + value = value.transpose(0, 1) + + return query.to(dtype), key.to(dtype), value.to(dtype) + + +def torch_varlen_mla_rope( + qs, kvs, k_pos_embs, freqs, lengths, mscale=1.0, cp_size=1, cp_rank=0 +): + ls = [x // cp_size for x in lengths] + + B = len(lengths) + N, H, _ = qs.shape + dtype = qs.dtype + + qss = qs.split(ls, 0) + kvss = kvs.split(ls, 0) + k_pos_embss = k_pos_embs.split(ls, 0) + + seg_size = 2 * cp_size + qoss = [] + koss = [] + voss = [] + for i in range(B): + q01 = qss[i][:, None].split([ls[i] // 2] * 2, 0) + kv01 = kvss[i][:, None].split([ls[i] // 2] * 2, 0) + pos01 = k_pos_embss[i][:, None].split([ls[i] // 2] * 2, 0) + + for j in range(2): + q_no_pe, q_pos_emb = torch.split(q01[j], [128, 64], dim=-1) + + k_no_pe, value = torch.split(kv01[j], [128, 128], dim=-1) + + cos = freqs.cos().to(dtype) * mscale + sin = freqs.sin().to(dtype) * mscale + if j == 0: + p = cp_rank * lengths[i] // seg_size + else: + p = (cp_size * 2 - cp_rank - 1) * lengths[i] // seg_size + position_ids = ( + p + torch.arange(lengths[i] // seg_size, device="cuda:0")[:, None] + ) + + q_pos_emb = torch.cat( + [q_pos_emb[:, :, :, 0::2], q_pos_emb[:, :, :, 1::2]], -1 + ) + k_pos_emb = torch.cat( + [pos01[j][:, :, :, 0::2], pos01[j][:, :, :, 1::2]], -1 + ) + + q_pos_emb, k_pos_emb = apply_rotary_pos_emb( + q_pos_emb, k_pos_emb, cos, sin, position_ids + ) + + query = torch.cat([q_no_pe, q_pos_emb], dim=-1) + + k_pos_emb = k_pos_emb.expand(-1, -1, H, -1) + + key = torch.cat([k_no_pe, k_pos_emb], dim=-1) + + qoss.append(query) + koss.append(key) + voss.append(value) + + qoss = torch.cat(qoss, 0)[:, 0] + koss = torch.cat(koss, 0)[:, 0] + voss = torch.cat(voss, 0)[:, 0] + + return qoss, koss, voss + + +def torch_qk_norm_and_half_rope( + qkv, + qw, + kw, + freqs, + H=32, + h=4, + eps=1e-6, + interleaved=True, + transposed=True, + silu=False, +): if transposed: length, bs, dim = qkv.shape else: bs, length, dim = qkv.shape + dtype = qkv.dtype qkv = qkv.float() qw = qw.float() kw = kw.float() + + if silu: + qkv = torch.nn.functional.silu(qkv) + D = dim // (H + 2 * h) if interleaved: if transposed: @@ -96,132 +230,606 @@ def torch_qk_norm_and_half_rope(qkv, qw, kw, freqs, rope_theta=10000.0, H=32, qkv = qkv.view(bs, length, H + 2 * h, D) q, k, v = torch.split(qkv, [H, h, h], dim=2) q, k = torch_qk_norm(q, k, qw, kw, eps=eps, transposed=transposed) - q, k = torch_half_rope(q, k, freqs, rope_theta=rope_theta, transposed=transposed) + q, k = torch_half_rope(q, k, freqs, transposed=transposed) if transposed: q = q.transpose(0, 1) k = k.transpose(0, 1) v = v.transpose(0, 1) - return q, k, v + return q.to(dtype), k.to(dtype), v.to(dtype) -def test_half_rope(B=2, L=4096, H=32, h=8, D=128, rope_theta=10000.0, - transposed=True, - bench=False): - dtype = torch.bfloat16 - device = 'cuda:0' +def torch_varlen_qk_norm_and_half_rope( + qkvs, + qw, + kw, + freqs, + lengths, + H=32, + h=4, + interleaved=True, + silu=False, + eps=1e-6, + mscale=1.0, + cp_size=1, + cp_rank=0, +): + ls = [x // cp_size for x in lengths] + D = qkvs.size(1) + + B = len(lengths) + + qkvss = qkvs.split(ls, 0) + + seg_size = 2 * cp_size + qoss = [] + koss = [] + voss = [] + for i in range(B): + qkv01 = qkvss[i].split([ls[i] // 2] * 2, 0) + + for j in range(2): + qkv = qkv01[j].view(1, ls[i] // 2, D) + + if j == 0: + p = cp_rank * lengths[i] // seg_size + else: + p = (cp_size * 2 - cp_rank - 1) * lengths[i] // seg_size + position_ids = p + torch.arange(lengths[i] // seg_size, device="cuda:0") + fs = freqs[position_ids] + query, key, value = torch_qk_norm_and_half_rope( + qkv, + qw, + kw, + fs, + H=H, + h=h, + eps=eps, + interleaved=interleaved, + transposed=False, + silu=silu, + ) + qoss.append(query) + koss.append(key) + voss.append(value) + + qoss = torch.cat(qoss, 1)[0] + koss = torch.cat(koss, 1)[0] + voss = torch.cat(voss, 1)[0] + + return qoss, koss, voss + + +@pytest.mark.parametrize( + "B,L,H,h,D,rope_theta,transposed", + [ + (2, 4096, 32, 8, 128, 10000.0, True), + (2, 4096, 32, 8, 128, 10000.0, False), + ], +) +def test_half_rope(B, L, H, h, D, rope_theta, transposed, benchmark): + dtype = torch.float32 + device = "cuda:0" q = torch.randn(L, B, H, D, dtype=dtype, device=device) k = torch.randn(L, B, h, D, dtype=dtype, device=device) freqs = rope_freqs(L, D // 2, rope_theta=rope_theta) freqs = torch.cat([freqs, freqs], -1) - q_ref, k_ref = torch_half_rope(q, k, freqs, rope_theta=rope_theta, transposed=transposed) + q_ref, k_ref = torch_half_rope(q, k, freqs, transposed=transposed) qo, ko = triton_half_rope_forward(q, k, freqs, transposed=transposed) - output_check(q_ref, qo, mode='q') - output_check(k_ref, ko, mode='k') + output_check(q_ref, qo, name="q") + output_check(k_ref, ko, name="k") q_grad = torch.randn(L, B, H, D, dtype=dtype, device=device) k_grad = torch.randn(L, B, h, D, dtype=dtype, device=device) q_ref = q.detach().clone().requires_grad_() k_ref = k.detach().clone().requires_grad_() - qo_ref, ko_ref = torch_half_rope(q_ref, k_ref, freqs, rope_theta=rope_theta, transposed=transposed) + qo_ref, ko_ref = torch_half_rope(q_ref, k_ref, freqs, transposed=transposed) qo_ref.backward(gradient=q_grad) ko_ref.backward(gradient=k_grad) dq_ref = q_ref.grad dk_ref = k_ref.grad - dq, dk = triton_half_rope_backward(q_grad, k_grad, freqs, inplace=True, transposed=transposed) - output_check(dq_ref, dq, mode='dq') - output_check(dk_ref, dk, mode='dk') + dq, dk = triton_half_rope_backward( + q_grad, k_grad, freqs, inplace=True, transposed=transposed + ) + output_check(dq_ref, dq, name="dq") + output_check(dk_ref, dk, name="dk", rtol=0.05, atol=0.1) - if bench: - benchmark_func(triton_half_rope_forward, q, k, freqs, - ref_bytes=L * B * (H + h) * D * 4, - n_profile=0) + benchmark( + triton_half_rope_forward, + q, + k, + freqs, + ref_bytes=L * B * (H + h) * D * 4, + n_profile=0, + ) -def test_qk_norm_and_half_rope(B=2, L=4096, H=32, h=8, D=128, - rope_theta=10000.0, - interleaved=True, - transposed=True, - bench=False): +@pytest.mark.parametrize( + "B,L,H,h,D,rope_theta,eps,interleaved,transposed,silu", + [ + (2, 4096, 16, 16, 128, 10000.0, 1e-6, True, True, True), + (2, 4096, 16, 16, 128, 10000.0, 1e-6, True, True, False), + (4, 4096, 16, 4, 128, 10000.0, 1e-6, True, False, True), + (4, 4096, 16, 4, 128, 10000.0, 1e-6, True, False, False), + (4, 4096, 32, 4, 128, 10000.0, 1e-6, False, True, True), + (4, 4096, 24, 6, 128, 10000.0, 1e-6, True, True, False), + (4, 4096, 32, 32, 128, 10000.0, 1e-6, False, False, True), + (1, 4096, 32, 32, 128, 10000.0, 1e-6, False, False, False), + (256, 4, 32, 32, 128, 10000.0, 1e-6, False, False, False), + ], +) +def test_qk_norm_and_half_rope( + B, L, H, h, D, rope_theta, eps, interleaved, transposed, silu, benchmark +): dtype = torch.bfloat16 - device = 'cuda:0' + device = "cuda:0" if transposed: qkv = torch.randn(L, B, (H + 2 * h) * D, dtype=dtype, device=device) else: qkv = torch.randn(B, L, (H + 2 * h) * D, dtype=dtype, device=device) - qw = torch.randn(D, dtype=dtype, device=device) - kw = torch.randn(D, dtype=dtype, device=device) + qkv = (qkv * qkv.abs()).requires_grad_() + qw = torch.nn.Parameter( + torch.randn(D, dtype=dtype, device=device), requires_grad=True + ) + kw = torch.nn.Parameter( + torch.randn(D, dtype=dtype, device=device), requires_grad=True + ) freqs = rope_freqs(L, D // 2, rope_theta=rope_theta) freqs = torch.cat([freqs, freqs], -1) - q_ref, k_ref, v_ref = torch_qk_norm_and_half_rope(qkv, qw, kw, freqs, - rope_theta, H=H, h=h, - eps=1e-6, - transposed=transposed, - interleaved=interleaved) - qo, ko, vo = triton_qk_norm_and_half_rope_forward(qkv, qw, kw, freqs, H=H, - h=h, eps=1e-6, - transposed=transposed, - interleaved=interleaved) - output_check(q_ref, qo, mode='q') - output_check(k_ref, ko, mode='k') - output_check(v_ref, vo, mode='v') - - q_grad = torch.randn(B, L, H, D, dtype=dtype, device=device) - k_grad = torch.randn(B, L, h, D, dtype=dtype, device=device) - v_grad = torch.randn(B, L, h, D, dtype=dtype, device=device) - qkv_ref = qkv.detach().clone().requires_grad_() - qw_ref = qw.detach().clone().requires_grad_() - kw_ref = kw.detach().clone().requires_grad_() - qo_ref, ko_ref, vo_ref = torch_qk_norm_and_half_rope(qkv_ref, qw_ref, - kw_ref, freqs, - rope_theta=rope_theta, - H=H, h=h, eps=1e-6, - transposed=transposed, - interleaved=interleaved) - qo_ref.backward(gradient=q_grad) - ko_ref.backward(gradient=k_grad) - vo_ref.backward(gradient=v_grad) - - dqkv_ref = qkv_ref.grad - dqw_ref = qw_ref.grad - dkw_ref = kw_ref.grad - - dqkv, dqw, dkw = triton_qk_norm_and_half_rope_backward(q_grad, k_grad, - v_grad, qkv, qw, kw, - freqs, eps=1e-6, - transposed=transposed, - interleaved=interleaved) - output_check(dqkv_ref, dqkv, mode='dqkv') - output_check(dqw_ref, dqw, mode='dqw') - output_check(dkw_ref, dkw, mode='dkw') - - if bench: - benchmark_func(triton_qk_norm_and_half_rope_forward, qkv, qw, kw, freqs, - H=H, h=h, eps=1e-6, transposed=transposed, interleaved=interleaved, - ref_bytes=L * B * (H + 2 * h) * D * 4, - n_profile=0) - benchmark_func(triton_qk_norm_and_half_rope_backward, q_grad, k_grad, - v_grad, qkv, qw, kw, freqs, eps=1e-6, transpose=transposed, interleaved=interleaved, - ref_bytes=L * B * (H + 2 * h) * D * 6, - n_profile=0) - - -if __name__ == '__main__': - test_half_rope(B=2, L=4096, H=32, h=8, D=128, rope_theta=10000.0, transposed=True, - bench=False) - test_half_rope(B=2, L=4096, H=32, h=8, D=128, rope_theta=10000.0, transposed=False, - bench=False) - test_qk_norm_and_half_rope(B=1, L=4096, H=16, h=16, D=128, - rope_theta=10000.0, interleaved=False, transposed=True, bench=False) - test_qk_norm_and_half_rope(B=2, L=4096, H=16, h=4, D=128, - rope_theta=10000.0, interleaved=False, transposed=True, bench=False) - - - test_qk_norm_and_half_rope(B=4, L=4096, H=16, h=4, D=128, - rope_theta=10000.0, interleaved=True, transposed=True, bench=False) - test_qk_norm_and_half_rope(B=4, L=4096, H=32, h=8, D=128, - rope_theta=10000.0, interleaved=True, transposed=True, bench=False) - test_qk_norm_and_half_rope(B=4, L=4096, H=32, h=8, D=128, - rope_theta=10000.0, interleaved=False, transposed=True, - bench=False) + q_grad = torch.randn(B, L, H, D, dtype=dtype, device=device) * 0.1 + k_grad = torch.randn(B, L, h, D, dtype=dtype, device=device) * 0.1 + v_grad = torch.randn(B, L, h, D, dtype=dtype, device=device) * 0.1 + + q_ref, k_ref, v_ref = torch_qk_norm_and_half_rope( + qkv, + qw, + kw, + freqs, + H=H, + h=h, + eps=eps, + transposed=transposed, + interleaved=interleaved, + silu=silu, + ) + q_ref.backward(gradient=q_grad, retain_graph=True) + k_ref.backward(gradient=k_grad, retain_graph=True) + v_ref.backward(gradient=v_grad, retain_graph=True) + dqkv_ref = qkv.grad + dqw_ref = qw.grad + dkw_ref = kw.grad + + qo, ko, vo = triton_qk_norm_and_half_rope_forward( + qkv, + qw, + kw, + freqs, + H=H, + h=h, + eps=eps, + transposed=transposed, + interleaved=interleaved, + silu=silu, + ) + output_check(q_ref, qo, name="q") + output_check(k_ref, ko, name="k") + output_check(v_ref, vo, name="v") + + dqkv, dqw, dkw = triton_qk_norm_and_half_rope_backward( + q_grad, + k_grad, + v_grad, + qkv, + qw, + kw, + freqs, + eps=eps, + transposed=transposed, + interleaved=interleaved, + silu=silu, + ) + output_check(dqkv_ref, dqkv, name="dqkv") + output_check(dqw_ref, dqw.to(dtype), name="dqw", amp=10) + output_check(dkw_ref, dkw.to(dtype), name="dkw", amp=10) + + if transposed and interleaved and not silu: + qkv.grad = None + qw.grad = None + kw.grad = None + q, k, v = qk_norm_half_rope(qkv, qw, kw, freqs, H=H, h=h, eps=eps) + # q.backward(gradient=q_grad, retain_graph=True) + # k.backward(gradient=k_grad, retain_graph=True) + # v.backward(gradient=v_grad, retain_graph=True) + loss = (q * q_grad).sum() + (k * k_grad).sum() + (v * v_grad).sum() + loss.backward() + dqkv = qkv.grad + dqw = qw.grad + dkw = kw.grad + output_check(dqkv_ref, dqkv, name="dqkv") + output_check(dqw_ref, dqw, name="dqw") + output_check(dkw_ref, dkw, name="dkw") + + benchmark( + triton_qk_norm_and_half_rope_forward, + qkv, + qw, + kw, + freqs, + H=H, + h=h, + eps=1e-6, + transposed=transposed, + interleaved=interleaved, + silu=silu, + ref_bytes=L * B * (H + 2 * h) * D * 4, + n_profile=0, + ) + benchmark( + triton_qk_norm_and_half_rope_backward, + q_grad, + k_grad, + v_grad, + qkv, + qw, + kw, + freqs, + eps=1e-6, + transposed=transposed, + interleaved=interleaved, + silu=silu, + ref_bytes=L * B * (H + 2 * h) * D * 6, + n_profile=0, + ) + + +@pytest.mark.parametrize( + "lengths,H,h,dim,rope_theta,silu,interleaved,cp_size,cp_rank", + [ + ([2048], 24, 6, 128, 10000.0, False, True, 1, 0), + ([1024, 4096, 4096, 568], 32, 4, 128, 10000.0, False, True, 1, 0), + ([2048, 4096, 4096], 32, 4, 128, 10000.0, False, True, 1, 0), + ([2048, 3072, 4096], 32, 4, 128, 10000.0, False, True, 4, 0), + ([2048, 4096, 4096], 32, 4, 128, 10000.0, True, False, 4, 0), + ([16] * 512, 32, 4, 128, 10000.0, False, False, 4, 0), + ], +) +def test_varlen_qk_norm_and_half_rope( + lengths, H, h, dim, rope_theta, silu, interleaved, cp_size, cp_rank, benchmark +): + # weight grad of torch impl has great error with large seq nums + dtype = torch.bfloat16 if len(lengths) < 8 else torch.float32 + device = "cuda:0" + N = sum(lengths) // cp_size + qkv = torch.randn(N, (H + 2 * h) * dim, dtype=dtype, device=device) + 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), dim // 2, rope_theta=rope_theta) + freqs = torch.cat([freqs, freqs], -1) + + mscale = 1.0 + qw = torch.randn(dim, dtype=dtype, device=device).requires_grad_() + kw = torch.randn(dim, dtype=dtype, device=device).requires_grad_() + + q_grad = torch.randn(sum(lengths) // cp_size, H, dim, dtype=dtype, device=device) + k_grad = torch.randn(sum(lengths) // cp_size, h, dim, dtype=dtype, device=device) + v_grad = torch.randn(sum(lengths) // cp_size, h, dim, dtype=dtype, device=device) + qkv = qkv.detach().clone().requires_grad_() + qo_ref, ko_ref, vo_ref = torch_varlen_qk_norm_and_half_rope( + qkv, + qw, + kw, + freqs, + lengths, + H=H, + h=h, + interleaved=interleaved, + silu=silu, + mscale=mscale, + cp_size=cp_size, + cp_rank=cp_rank, + ) + + qo_ref.backward(gradient=q_grad, retain_graph=True) + ko_ref.backward(gradient=k_grad, retain_graph=True) + vo_ref.backward(gradient=v_grad, retain_graph=True) + + dqkv_ref = qkv.grad + dqw_ref = qw.grad + dkw_ref = kw.grad + + qo, ko, vo = triton_varlen_qk_norm_and_half_rope_forward( + qkv, + qw, + kw, + freqs, + cu_seqlens_q, + cu_seqlens_kv, + H=H, + h=h, + interleaved=interleaved, + silu=silu, + mscale=mscale, + cp_size=cp_size, + cp_rank=cp_rank, + ) + output_check(qo_ref, qo, name="q") + output_check(ko_ref, ko, name="k") + output_check(vo_ref, vo, name="v") + + dqkv, dqw, dkw = triton_varlen_qk_norm_and_half_rope_backward( + q_grad, + k_grad, + v_grad, + qkv, + qw, + kw, + freqs, + cu_seqlens_q, + cu_seqlens_kv, + mscale=mscale, + interleaved=interleaved, + silu=silu, + cp_size=cp_size, + cp_rank=cp_rank, + ) + output_check(dqkv_ref, dqkv, name="dqkv", atol=0.1, rtol=0.02) + output_check(dqw_ref, dqw.to(dtype), name="dqw", atol=2.0 * len(lengths), rtol=0.02) + output_check(dkw_ref, dkw.to(dtype), name="dkw", atol=2.0 * len(lengths), rtol=0.02) + + lbh = sum(lengths) // cp_size * H + benchmark( + triton_varlen_qk_norm_and_half_rope_forward, + qkv, + qw, + kw, + freqs, + cu_seqlens_q, + cu_seqlens_kv, + interleaved=interleaved, + H=H, + h=h, + silu=silu, + mscale=mscale, + 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( + triton_varlen_qk_norm_and_half_rope_backward, + q_grad, + k_grad, + v_grad, + qkv, + qw, + kw, + freqs, + cu_seqlens_q, + cu_seqlens_kv, + mscale=mscale, + interleaved=interleaved, + silu=silu, + cp_size=cp_size, + cp_rank=cp_rank, + ref_bytes=lbh * (64 * 2 + 256 * 2 + 64 * 2 + 192 * 2 + 128 * 2), + n_profile=0, + ) + + +@pytest.mark.parametrize( + "B,L,H,rope_theta,transpose", + [ + (4, 4096, 16, 10000.0, False), + (4, 4096, 16, 10000.0, True), + ], +) +def test_mla_rope(B, L, H, rope_theta, transpose, benchmark): + dtype = torch.bfloat16 + device = "cuda:0" + q = torch.randn(L, B, H, 192, dtype=dtype, device=device, requires_grad=True) + kv = torch.randn(L, B, H, 256, dtype=dtype, device=device, requires_grad=True) + 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 + + q_ref, k_ref, v_ref = torch_mla_rope( + q, kv, k_pos_emb, freqs, mscale=mscale, transpose=transpose + ) + q_ref.backward(gradient=q_grad, retain_graph=True) + k_ref.backward(gradient=k_grad, retain_graph=True) + v_ref.backward(gradient=v_grad, retain_graph=True) + dq_ref = q.grad + dkv_ref = kv.grad + dp_ref = k_pos_emb.grad + + qo, ko, vo = triton_mla_rope_forward( + q.clone().detach(), kv, k_pos_emb, freqs, mscale=mscale, transpose=transpose + ) + output_check(q_ref, qo, name="q") + output_check(k_ref, ko, name="k") + output_check(v_ref, vo, name="v") + + dq, dkv, dp = triton_mla_rope_backward( + q_grad.clone().detach(), + k_grad, + v_grad, + freqs, + mscale=mscale, + transposed=transpose, + ) + output_check(dq_ref, dq, name="dq") + output_check(dkv_ref, dkv, name="dkv") + output_check(dp_ref, dp, name="dp") + + lbh = L * B * H + benchmark( + triton_mla_rope_forward, + q, + kv, + k_pos_emb, + freqs, + ref_bytes=lbh * (64 * 2 + 256 * 2 + 64 * 2 + 192 * 2 + 128 * 2), + n_profile=0, + ) + benchmark( + triton_mla_rope_backward, + q_grad, + k_grad, + v_grad, + freqs, + ref_bytes=lbh * (64 * 2 + 256 * 2 + 64 * 2 + 192 * 2 + 128 * 2), + n_profile=0, + ) + + +@pytest.mark.parametrize( + "lengths,H,rope_theta,cp_size,cp_rank", + [ + ([8192], 64, 10000.0, 1, 0), + ([4096, 4096], 16, 10000.0, 1, 0), + ([4096 * 4, 2048 * 4, 2048 * 4], 32, 10000.0, 4, 0), + ([4096 * 4, 2048 * 4, 2048 * 4], 32, 10000.0, 4, 1), + ([4096 * 4, 2048 * 4, 2048 * 4], 32, 10000.0, 4, 2), + ([4096 * 4, 2048 * 4, 2048 * 4], 32, 10000.0, 4, 3), + ([16] * 236, 32, 10000.0, 4, 1), + ], +) +def test_varlen_mla_rope(lengths, H, rope_theta, cp_size, cp_rank, benchmark): + # weight grad of torch impl has great error with large seq nums + dtype = torch.bfloat16 if len(lengths) < 8 else torch.float32 + device = "cuda:0" + qc = torch.randn( + sum(lengths) // cp_size, H, 192, dtype=dtype, device=device + ).requires_grad_() + kvc = torch.randn( + sum(lengths) // cp_size, H, 256, dtype=dtype, device=device + ).requires_grad_() + k_pos_emb = torch.randn(sum(lengths) // cp_size, 576, dtype=dtype, device=device) + k_pos_embc = ( + k_pos_emb[:, 512:].view(sum(lengths) // cp_size, 1, 64).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), 64, rope_theta=rope_theta) + freqs = torch.cat([freqs, freqs], -1) + + mscale = 1.0 + q_ref, k_ref, v_ref = torch_varlen_mla_rope( + qc, + kvc, + k_pos_embc, + freqs, + lengths, + mscale=mscale, + cp_size=cp_size, + cp_rank=cp_rank, + ) + qo, ko, vo = triton_mla_rope_forward( + qc.clone().detach(), + kvc, + k_pos_embc, + freqs, + mscale=mscale, + cu_seqlens_q=cu_seqlens_q, + cu_seqlens_kv=cu_seqlens_kv, + transpose=False, + cp_size=cp_size, + cp_rank=cp_rank, + ) + output_check(q_ref, qo, name="q", amp=20) + output_check(k_ref, ko, name="k", amp=20) + output_check(v_ref, vo, name="v", amp=20) + + 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) + q_i = qc.detach().clone().requires_grad_() + kv_i = kvc.detach().clone().requires_grad_() + k_pos_emb_i = k_pos_embc.detach().clone().requires_grad_() + qo_ref, ko_ref, vo_ref = torch_varlen_mla_rope( + q_i, + kv_i, + k_pos_emb_i, + freqs, + lengths, + mscale=mscale, + cp_size=cp_size, + cp_rank=cp_rank, + ) + qo_ref.backward(gradient=q_grad.clone().detach(), retain_graph=True) + ko_ref.backward(gradient=k_grad, retain_graph=True) + vo_ref.backward(gradient=v_grad, retain_graph=True) + dq_ref = q_i.grad + dkv_ref = kv_i.grad + dp_ref = k_pos_emb_i.grad + dq, dkv, dp = triton_mla_rope_backward( + q_grad.clone().detach(), + k_grad, + v_grad, + freqs, + mscale=mscale, + cu_seqlens_q=cu_seqlens_q, + cu_seqlens_kv=cu_seqlens_kv, + cp_size=cp_size, + cp_rank=cp_rank, + transposed=False, + ) + output_check(dq_ref, dq, name="dq", atol=0.1, rtol=0.02) + output_check(dkv_ref, dkv, name="dkv", atol=0.1, rtol=0.02) + output_check(dp_ref, dp, name="dp", atol=0.2, rtol=0.02) + + lbh = sum(lengths) // cp_size * H + benchmark( + triton_mla_rope_forward, + qc, + kvc, + k_pos_embc, + freqs, + mscale=mscale, + cu_seqlens_q=cu_seqlens_q, + cu_seqlens_kv=cu_seqlens_kv, + cp_size=cp_size, + cp_rank=cp_rank, + transpose=False, + ref_bytes=lbh * (64 * 2 + 256 * 2 + 64 * 2 + 192 * 2 + 128 * 2), + n_profile=0, + ) + benchmark( + triton_mla_rope_backward, + q_grad, + k_grad, + v_grad, + freqs, + mscale=mscale, + cu_seqlens_q=cu_seqlens_q, + cu_seqlens_kv=cu_seqlens_kv, + cp_size=cp_size, + cp_rank=cp_rank, + transposed=False, + ref_bytes=lbh * (64 * 2 + 256 * 2 + 64 * 2 + 192 * 2 + 128 * 2), + n_profile=0, + ) diff --git a/tests/test_scatter.py b/tests/test_scatter.py index fa9fbd6..a518274 100644 --- a/tests/test_scatter.py +++ b/tests/test_scatter.py @@ -3,34 +3,44 @@ Copyright (c) Ant Financial Service Group and its affiliates. """ +import pytest import torch -from linghe.tools.benchmark import benchmark_func -from linghe.tools.util import (output_check, - torch_make_indices) -from linghe.utils.scatter import (triton_scatter_add, - triton_unpermute_with_mask_map - ) +from linghe.tools.check import output_check +from linghe.tools.util import torch_make_indices +from linghe.utils.scatter import triton_scatter_add, triton_unpermute_with_mask_map # os.environ["CUDA_LAUNCH_BLOCKING"] = "1" -def torch_fp16_scatter_add(x, outputs, indices, weights): +def torch_scatter_add(x, outputs, indices, weights): + dtype = x.dtype + x = x.float() if weights is not None: x = x * weights[:, None] dim = x.size(1) + outputs = outputs.float() outputs.scatter_add_(0, indices.unsqueeze(1).expand(-1, dim), x) - return outputs - - -def test_scatter(M=4098, N=4096, n_experts=32, topk=2, bias=0.0, bench=False): + return outputs.to(dtype) + + +@pytest.mark.parametrize( + "M,N,bias", + [ + (4098, 4096, 0.0), + (2467, 4096, -0.1), + (2467, 1536, -0.1), + ], +) +def test_scatter(M, N, bias, benchmark, n_experts=32, topk=2): dtype = torch.bfloat16 - device = 'cuda:0' + device = "cuda:0" 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=bias) + logits, topk=topk, bias=bias + ) token_count_per_expert_list = token_count_per_expert.tolist() out_tokens = sum(token_count_per_expert_list) @@ -39,26 +49,22 @@ def test_scatter(M=4098, N=4096, n_experts=32, topk=2, bias=0.0, bench=False): outputs = torch.zeros((M, N), dtype=dtype, device=device) - sums_ref = torch_fp16_scatter_add(x, outputs.clone(), indices, None) - counts = mask_map.sum(1) - unpermuted_prob = probs.T.contiguous().masked_select( - mask_map.T.contiguous()) - - sums_unpermute, output_prob = triton_unpermute_with_mask_map(x, row_id_map, - unpermuted_prob) - output_check(sums_ref, sums_unpermute, 'unpermute_data') - output_check(probs, output_prob, 'unpermute_prob') - - if bench: - n_repeat = 100 - # ref_time = benchmark_func(torch_fp16_scatter_add,x, outputs, indices, weights,n_repeat=n_repeat) - # benchmark_func(triton_aligned_scatter_add,x, outputs, indices, weights=weights, n_repeat=n_repeat,ref_time=ref_time) - 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) - - -if __name__ == '__main__': - test_scatter(M=4098, N=4096, n_experts=32, topk=2, bias=0.0) - test_scatter(M=2467, N=4096, n_experts=32, topk=2, bias=-0.1) + sums_ref = torch_scatter_add(x, outputs.clone(), indices, None) + unpermuted_prob = probs.T.contiguous().masked_select(mask_map.T.contiguous()) + + sums_unpermute, output_prob = triton_unpermute_with_mask_map( + x, row_id_map, unpermuted_prob + ) + output_check(sums_ref, sums_unpermute, "unpermute_data") + output_check(probs, output_prob, "unpermute_prob") + + n_repeat = 100 + ref_time = benchmark(triton_scatter_add, x, outputs, indices, n_repeat=n_repeat) + benchmark( + triton_unpermute_with_mask_map, + x, + row_id_map, + probs, + n_repeat=n_repeat, + ref_time=ref_time, + ) diff --git a/tests/test_silu.py b/tests/test_silu.py index 2375b90..ab9aa48 100644 --- a/tests/test_silu.py +++ b/tests/test_silu.py @@ -7,68 +7,82 @@ random.seed(7) +import pytest import torch -from linghe.tools.benchmark import benchmark_func -from linghe.utils.silu import (triton_weighted_silu_forward, - triton_weighted_silu_backward, - triton_batch_weighted_silu_and_smooth_quant_backward, - triton_batch_weighted_silu_and_smooth_quant_forward, - triton_batch_weighted_silu_and_block_quant_backward, - triton_batch_weighted_silu_and_block_quant_forward, - triton_silu_and_smooth_quant_backward, - triton_silu_and_smooth_quant_forward, - triton_silu_and_block_quant_backward, - triton_silu_and_block_quant_forward) -from linghe.tools.util import output_check, torch_smooth_quant, \ - torch_group_quant - - -def torch_silu(x): +from linghe.utils.silu import ( + triton_weighted_silu_forward, + triton_weighted_silu_backward, + triton_batch_weighted_silu_and_smooth_quant_backward, + triton_batch_weighted_silu_and_smooth_quant_forward, + triton_batch_weighted_silu_and_block_quant_backward, + triton_batch_weighted_silu_and_block_quant_forward, + triton_batch_weighted_silu_and_mxfp8_quant_backward, + triton_batch_weighted_silu_and_mxfp8_quant_forward, + triton_silu_and_smooth_quant_backward, + triton_silu_and_smooth_quant_forward, + triton_silu_and_block_quant_backward, + triton_silu_and_block_quant_forward, + triton_silu_and_mxfp8_quant_backward, + triton_silu_and_mxfp8_quant_forward, +) +from linghe.tools.util import torch_smooth_quant, torch_group_quant, torch_mxfp8_quant +from linghe.tools.check import output_check + + +def torch_silu(x, limit=None): + x = x.float() M, N = x.shape x1, x2 = torch.split(x, N // 2, dim=1) - y = torch.sigmoid(x1) * x1 * x2 + if limit is None: + y = torch.sigmoid(x1) * x1 * x2 + else: + y = torch.clamp(torch.sigmoid(x1) * x1, max=limit) * torch.clamp( + x2, max=limit, min=-limit + ) return y -def torch_weighted_silu(x, weight): +def torch_weighted_silu(x, weight, limit=None): + dtype = x.dtype + x = x.float() + weight = weight.float() M, N = x.shape x1, x2 = torch.split(x, N // 2, dim=1) - y = torch.sigmoid(x1) * x1 * x2 * weight - return y + if limit is None: + y = torch.sigmoid(x1) * x1 * x2 * weight + else: + y = ( + torch.clamp(torch.sigmoid(x1) * x1, max=limit) + * torch.clamp(x2, max=limit, min=-limit) + * weight + ) + return y.to(dtype) -def torch_weighted_silu_backward(dy, x, weight): +def torch_weighted_silu_backward(dy, x, weight, limit=None): + dtype = x.dtype + x = x.float() x = x.clone().detach().requires_grad_() weight = weight.clone().detach().requires_grad_() - y = torch_weighted_silu(x, weight) + y = torch_weighted_silu(x, weight, limit=limit) y.backward(gradient=dy) - return x.grad, weight.grad - + return x.grad.to(dtype), weight.grad -def torch_silu_and_smooth_quant_forward(x, smooth_scale=None, round_scale=True): - M, N = x.shape - x = x.float() - x1, x2 = torch.split(x, N // 2, dim=1) - y = torch.sigmoid(x1) * x1 * x2 +def torch_silu_and_smooth_quant_forward( + x, limit=None, smooth_scale=None, round_scale=True +): + y = torch_silu(x, limit=limit) # smooth - y_q, y_scale, x_maxs = torch_smooth_quant(y, smooth_scale, reverse=False, - round_scale=round_scale) - # y_smooth = y / smooth_scale - # x_maxs = y.abs().float().amax(0) - # y_scale = y_smooth.abs().amax(1) / 448 - # if round_scale: - # y_scale = torch.exp2(torch.ceil(torch.log2(y_scale))) - # y_q = (y_smooth / y_scale[:, None]).to(torch.float8_e4m3fn) - return y_q, y_scale, x_maxs - - -def torch_silu_and_block_quant_forward(x, round_scale=True): - M, N = x.shape - x = x.float() - x1, x2 = torch.split(x, N // 2, dim=1) - y = torch.sigmoid(x1) * x1 * x2 + y_q, y_scale = torch_smooth_quant( + y, smooth_scale, reverse=False, round_scale=round_scale + ) + return y_q, y_scale + + +def torch_silu_and_block_quant_forward(x, limit=None, round_scale=True): + y = torch_silu(x, limit=limit) # blockwise y_q, y_scale = torch_group_quant(y, round_scale=round_scale) yt_q, yt_scale = torch_group_quant(y.t(), round_scale=round_scale) @@ -76,28 +90,44 @@ def torch_silu_and_block_quant_forward(x, round_scale=True): return y_q, y_scale, yt_q, yt_scale -def torch_silu_and_smooth_quant_backward(grad, x, smooth_scale=None, - transpose_smooth_scale=None, - round_scale=True, reverse=True): +def torch_silu_and_mxfp8_quant_forward(x, limit=None): + y = torch_silu(x, limit=limit) + y_q, y_scale, yt_q, yt_scale = torch_mxfp8_quant(y) + + return y_q, y_scale, yt_q, yt_scale + + +def torch_silu_and_smooth_quant_backward( + grad, + x, + smooth_scale=None, + transpose_smooth_scale=None, + limit=None, + round_scale=True, + reverse=True, +): grad = grad.float() x = x.float().detach().clone().requires_grad_() - y = torch_silu(x) + y = torch_silu(x, limit=limit) y.backward(gradient=grad) dx = x.grad - q, dx_scale, ms = torch_smooth_quant(dx, smooth_scale, reverse=reverse, - round_scale=round_scale) - yt_q, yt_scale, ms = torch_smooth_quant(dx.t().contiguous(), - transpose_smooth_scale, - reverse=reverse, - round_scale=round_scale) + q, dx_scale = torch_smooth_quant( + dx, smooth_scale, reverse=reverse, round_scale=round_scale + ) + yt_q, yt_scale = torch_smooth_quant( + dx.t().contiguous(), + transpose_smooth_scale, + reverse=reverse, + round_scale=round_scale, + ) return q, dx_scale, yt_q, yt_scale -def torch_silu_and_block_quant_backward(grad, x, round_scale=True): +def torch_silu_and_block_quant_backward(grad, x, limit=None, round_scale=True): grad = grad.float() x = x.float().detach().clone().requires_grad_() - y = torch_silu(x) + y = torch_silu(x, limit=limit) y.backward(gradient=grad) dx = x.grad # blockwise @@ -107,19 +137,27 @@ def torch_silu_and_block_quant_backward(grad, x, round_scale=True): return q, dx_scale, yt_q, yt_scale -def torch_batch_weighted_silu_and_smooth_quant_forward(xs, weight, - counts, - smooth_scales=None, - round_scale=True, - reverse=False): +def torch_silu_and_mxfp8_quant_backward(grad, x, limit=None): + grad = grad.float() + x = x.float().detach().clone().requires_grad_() + y = torch_silu(x, limit=limit) + y.backward(gradient=grad) + dx = x.grad + y_q, y_scale, yt_q, yt_scale = torch_mxfp8_quant(dx) + + return y_q, y_scale, yt_q, yt_scale + + +def torch_batch_weighted_silu_and_smooth_quant_forward( + xs, weight, counts, smooth_scales=None, limit=None, round_scale=True, reverse=False +): counts = counts.tolist() N = xs.shape[1] if sum(counts) == 0: device = xs.device qs = torch.empty((0, N // 2), device=device, dtype=torch.float8_e4m3fn) scales = torch.empty((0,), device=device, dtype=torch.float32) - maxs = torch.zeros((len(counts), N), device=device, dtype=torch.float32) - return qs, scales, maxs + return qs, scales xs = xs.float() weight = weight.float() @@ -127,27 +165,25 @@ def torch_batch_weighted_silu_and_smooth_quant_forward(xs, weight, qs = [] scales = [] - maxs = [] s = 0 for i, c in enumerate(counts): - x = xs[s:s + c] - y = torch_weighted_silu(x, weight[s:s + c]) - q, scale, ms = torch_smooth_quant(y, smooth_scales[i], reverse=reverse, - round_scale=round_scale) + x = xs[s : s + c] + y = torch_weighted_silu(x, weight[s : s + c], limit=limit) + q, scale = torch_smooth_quant( + y, smooth_scales[i], reverse=reverse, round_scale=round_scale + ) qs.append(q) scales.append(scale) - maxs.append(ms) s += c qs = torch.cat(qs, 0) scales = torch.cat(scales, 0) - maxs = torch.cat(maxs, 0) - return qs, scales, maxs + return qs, scales -def torch_batch_weighted_silu_and_block_quant_forward(xs, weight, - counts, - round_scale=True): +def torch_batch_weighted_silu_and_block_quant_forward( + xs, weight, counts, limit=None, round_scale=True +): counts = counts.tolist() N = xs.shape[1] if sum(counts) == 0: @@ -167,8 +203,8 @@ def torch_batch_weighted_silu_and_block_quant_forward(xs, weight, qtscales = [] s = 0 for i, c in enumerate(counts): - x = xs[s:s + c] - y = torch_weighted_silu(x, weight[s:s + c]) + x = xs[s : s + c] + y = torch_weighted_silu(x, weight[s : s + c], limit=limit) q, scale = torch_group_quant(y, round_scale=round_scale) qt, qtscale = torch_group_quant(y.t(), round_scale=round_scale) qs.append(q) @@ -184,12 +220,54 @@ def torch_batch_weighted_silu_and_block_quant_forward(xs, weight, return qs, scales, qts, qtscales -def torch_batch_weighted_silu_and_smooth_quant_backward(grad_output, x, weight, - counts, - smooth_scales=None, - transpose_smooth_scale=None, - round_scale=True, - reverse=False): +def torch_batch_weighted_silu_and_mxfp8_quant_forward(xs, weight, counts, limit=None): + counts = counts.tolist() + N = xs.shape[1] + if sum(counts) == 0: + device = xs.device + qs = torch.empty((0, N // 2), device=device, dtype=torch.float8_e4m3fn) + scales = torch.empty((0, N // 64), device=device, dtype=torch.uint8) + qts = torch.empty((0, N // 2), device=device, dtype=torch.float8_e4m3fn) + qtscales = torch.zeros((0, N // 2), device=device, dtype=torch.uint8) + return qs, scales, qts, qtscales + + xs = xs.float() + weight = weight.float() + + qs = [] + scales = [] + qts = [] + qtscales = [] + s = 0 + for i, c in enumerate(counts): + x = xs[s : s + c] + y = torch_weighted_silu(x, weight[s : s + c], limit=limit) + + y_q, y_scale, yt_q, yt_scale = torch_mxfp8_quant(y) + qs.append(y_q) + scales.append(y_scale) + qts.append(yt_q) + qtscales.append(yt_scale) + + s += c + qs = torch.cat(qs, 0) + scales = torch.cat(scales, 0) + qts = torch.cat(qts, 0) + qtscales = torch.cat(qtscales, 0) + return qs, scales, qts, qtscales + + +def torch_batch_weighted_silu_and_smooth_quant_backward( + grad_output, + x, + weight, + counts, + smooth_scales=None, + transpose_smooth_scale=None, + limit=None, + round_scale=True, + reverse=False, +): if sum(counts) == 0: device = x.device N = x.shape[1] @@ -197,8 +275,7 @@ def torch_batch_weighted_silu_and_smooth_quant_backward(grad_output, x, weight, dx_scale = torch.empty((0,), device=device, dtype=torch.float32) dw = torch.empty_like(weight) qts = torch.empty((0,), device=device, dtype=torch.float8_e4m3fn) - qtscales = torch.zeros((N * len(counts),), device=device, - dtype=torch.float32) + qtscales = torch.zeros((N * len(counts),), device=device, dtype=torch.float32) return dx_q, dx_scale, dw, qts, qtscales grad_output = grad_output.float() @@ -207,25 +284,25 @@ def torch_batch_weighted_silu_and_smooth_quant_backward(grad_output, x, weight, smooth_scales = smooth_scales.float() transpose_smooth_scale = transpose_smooth_scale.float() - dx, dw = torch_weighted_silu_backward(grad_output, x, weight) + dx, dw = torch_weighted_silu_backward(grad_output, x, weight, limit=limit) qs = [] scales = [] qts = [] qtscales = [] s = 0 for i, c in enumerate(counts): - q, scale, dx_max = torch_smooth_quant(dx[s:s + c], smooth_scales[i], - reverse=reverse, - round_scale=round_scale) - dxt = dx[s:s + c].t().contiguous() - dxt_s = transpose_smooth_scale[s:s + c] + q, scale = torch_smooth_quant( + dx[s : s + c], smooth_scales[i], reverse=reverse, round_scale=round_scale + ) + dxt = dx[s : s + c].t().contiguous() + dxt_s = transpose_smooth_scale[s : s + c] padding_size = (c + 31) // 32 * 32 - c if padding_size > 0: dxt = torch.nn.functional.pad(dxt, (0, padding_size, 0, 0)) dxt_s = torch.nn.functional.pad(dxt_s, (0, padding_size)) - qt, t_scale, dx_max = torch_smooth_quant(dxt, dxt_s, - reverse=reverse, - round_scale=round_scale) + qt, t_scale = torch_smooth_quant( + dxt, dxt_s, reverse=reverse, round_scale=round_scale + ) qs.append(q) scales.append(scale) @@ -239,9 +316,9 @@ def torch_batch_weighted_silu_and_smooth_quant_backward(grad_output, x, weight, return dx_q, dx_scale, dw, qts, qtscales -def torch_batch_weighted_silu_and_block_quant_backward(grad_output, x, weight, - counts, - round_scale=True): +def torch_batch_weighted_silu_and_block_quant_backward( + grad_output, x, weight, counts, limit=None, round_scale=True +): if sum(counts) == 0: device = x.device N = x.shape[1] @@ -249,24 +326,22 @@ def torch_batch_weighted_silu_and_block_quant_backward(grad_output, x, weight, dx_scale = torch.empty((0,), device=device, dtype=torch.float32) dw = torch.empty_like(weight) qts = torch.empty((0,), device=device, dtype=torch.float8_e4m3fn) - qtscales = torch.zeros((N * len(counts),), device=device, - dtype=torch.float32) + qtscales = torch.zeros((0,), device=device, dtype=torch.float32) return dx_q, dx_scale, dw, qts, qtscales grad_output = grad_output.float() x = x.float() weight = weight.float() - dx, dw = torch_weighted_silu_backward(grad_output, x, weight) + dx, dw = torch_weighted_silu_backward(grad_output, x, weight, limit=limit) qs = [] scales = [] qts = [] qtscales = [] s = 0 for i, c in enumerate(counts): - q, scale = torch_group_quant(dx[s:s + c], round_scale=round_scale) - qt, qtscale = torch_group_quant(dx[s:s + c].t(), - round_scale=round_scale) + q, scale = torch_group_quant(dx[s : s + c], round_scale=round_scale) + qt, qtscale = torch_group_quant(dx[s : s + c].t(), round_scale=round_scale) qs.append(q) scales.append(scale.t().contiguous().view(-1)) qts.append(qt.view(-1)) @@ -280,313 +355,669 @@ def torch_batch_weighted_silu_and_block_quant_backward(grad_output, x, weight, return dx_q, dx_scale, dw, qts, qtscales -def test_weighted_silu(M=4096, N=4096, bench=False): - x = torch.randn((M, N), dtype=torch.bfloat16, device='cuda:0') - x = (x ** 3 // 10).clone().detach().requires_grad_() - weight = torch.randn((M, 1), dtype=torch.bfloat16, device='cuda:0') - grad_output = torch.randn((M, N // 2), dtype=torch.bfloat16, - device='cuda:0') +def torch_batch_weighted_silu_and_mxfp8_quant_backward( + grad_output, x, weight, counts, limit=None +): + if sum(counts) == 0: + device = x.device + N = x.shape[1] + dx_q = torch.empty((0, N), device=device, dtype=torch.float8_e4m3fn) + dx_scale = torch.empty((0, N // 32), device=device, dtype=torch.uint8) + dw = torch.empty_like(weight) + qts = torch.empty((0, N), device=device, dtype=torch.float8_e4m3fn) + qtscales = torch.zeros((0, N), device=device, dtype=torch.uint8) + return dx_q, dx_scale, dw, qts, qtscales + + grad_output = grad_output.float() + x = x.float() + weight = weight.float() + + dx, dw = torch_weighted_silu_backward(grad_output, x, weight, limit=limit) + qs = [] + scales = [] + qts = [] + qtscales = [] + s = 0 + for i, c in enumerate(counts): + q, scale, qt, qtscale = torch_mxfp8_quant(dx[s : s + c]) + + qs.append(q) + scales.append(scale) + qts.append(qt) + qtscales.append(qtscale) + + s += c + dx_q = torch.cat(qs, 0) + dx_scale = torch.cat(scales, 0) + qts = torch.cat(qts, 0) + qtscales = torch.cat(qtscales, 0) + return dx_q, dx_scale, dw, qts, qtscales + + +@pytest.mark.parametrize( + "M,N,coef,asm", + [ + (32 * 2048, 2048, 1.0, False), + (16384, 4096, 1.0, True), + (8192, 1536, 1.0, False), + (0, 1536, 1.0, False), + ], +) +def test_weighted_silu(M, N, coef, asm, benchmark): + x = torch.randn((M, N * 2), dtype=torch.bfloat16, device="cuda:0") + x = (x * coef).clone().detach().requires_grad_() + weight = torch.randn((M, 1), dtype=torch.float32, device="cuda:0") + grad_output = torch.randn((M, N), dtype=torch.bfloat16, device="cuda:0") ref_y = torch_weighted_silu(x, weight) - y = triton_weighted_silu_forward(x, weight) - output_check(ref_y, y, 'y') + y = triton_weighted_silu_forward(x, weight, asm=asm) + output_check(ref_y, y, "y") dx_ref, dw_ref = torch_weighted_silu_backward(grad_output, x, weight) dx, dw = triton_weighted_silu_backward(grad_output, x, weight) - output_check(dx_ref, dx, 'dx') - output_check(dw_ref, dw, 'dw') - - if bench: - benchmark_func(triton_weighted_silu_forward, x, weight, n_repeat=100, - ref_bytes=M * N * 3) - benchmark_func(triton_weighted_silu_backward, grad_output, x, weight, - n_repeat=100, ref_bytes=M * N * 5) - - -def test_silu_and_smooth_quant(M=4096, N=4096, bench=False): - if True: - x = torch.randn((M, N), dtype=torch.bfloat16, device='cuda:0') - x = (x * 10).clone().detach().requires_grad_() - grad_output = torch.randn((M, N // 2), dtype=torch.bfloat16, - device='cuda:0') - smooth_scale = 1 + torch.rand((N // 2,), dtype=torch.float32, - device='cuda:0') - grad_smooth_scale = 1 + torch.rand((N,), dtype=torch.float32, - device='cuda:0') - transpose_grad_smooth_scale = 1 + torch.rand((M,), dtype=torch.float32, - device='cuda:0') - else: - d = torch.load('/ossfs/workspace/tmp/vis/silu.bin') - x = d['x'].clone().detach().to('cuda:0').requires_grad_() - grad_output = d['g'].to('cuda:0') - grad_smooth_scale = d['smooth_scale'].to('cuda:0') - N = x.shape[-1] - M = x.shape[0] - smooth_scale = 1 + torch.rand((N // 2,), dtype=torch.float32, - device='cuda:0') - - y_q_ref, y_scale_ref, y_maxs_ref = torch_silu_and_smooth_quant_forward(x, - smooth_scale=smooth_scale) - y_q, y_scale, y_maxs = triton_silu_and_smooth_quant_forward(x, - smooth_scale=smooth_scale, - round_scale=True, - calibrate=True) - output_check(y_q_ref.float(), y_q.float(), 'smooth.y_q') - output_check(y_scale_ref, y_scale, 'smooth.y_scale') - output_check(y_maxs_ref, y_maxs, 'smooth.y_max') - - dx_q_ref, dx_scale_ref, dxt_q_ref, dxt_scale_ref = torch_silu_and_smooth_quant_backward( - grad_output, x, - smooth_scale=grad_smooth_scale, - transpose_smooth_scale=transpose_grad_smooth_scale, - reverse=True, - round_scale=True) + output_check(dx_ref, dx, "dx") + output_check(dw_ref, dw, "dw", rtol=3e-3, atol=3e-3) + + benchmark( + triton_weighted_silu_forward, + x, + weight, + asm=asm, + n_repeat=100, + ref_bytes=M * N * 6, + ) + benchmark( + triton_weighted_silu_backward, + grad_output, + x, + weight, + n_repeat=100, + ref_bytes=M * N * 10, + ) + + +@pytest.mark.parametrize( + "M,N", + [ + (16384, 1024), + (8192, 2048), + (4096, 10240), + (4096, 5120), + ], +) +def test_silu_and_smooth_quant(M, N, benchmark, coef=1.0, grad_coef=1.0): + x = torch.randn((M, 2 * N), dtype=torch.bfloat16, device="cuda:0") + x = (x * coef).clone().detach().requires_grad_() + grad_output = torch.randn((M, N), dtype=torch.bfloat16, device="cuda:0") * grad_coef + smooth_scale = 1 + torch.rand((N,), dtype=torch.float32, device="cuda:0") + grad_smooth_scale = 1 + torch.rand((2 * N,), dtype=torch.float32, device="cuda:0") + transpose_grad_smooth_scale = 1 + torch.rand( + (M,), dtype=torch.float32, device="cuda:0" + ) + + round_scale = False + y_q_ref, y_scale_ref = torch_silu_and_smooth_quant_forward( + x, smooth_scale=smooth_scale, round_scale=round_scale + ) + y_q, y_scale = triton_silu_and_smooth_quant_forward( + x, smooth_scale=smooth_scale, round_scale=round_scale + ) + output_check(y_q_ref, y_q, "smooth.y_q", rtol=0.125) + output_check(y_scale_ref, y_scale, "smooth.y_scale") + + dx_q_ref, dx_scale_ref, dxt_q_ref, dxt_scale_ref = ( + torch_silu_and_smooth_quant_backward( + grad_output, + x, + smooth_scale=grad_smooth_scale, + transpose_smooth_scale=transpose_grad_smooth_scale, + reverse=True, + round_scale=True, + ) + ) dx_q, dx_scale, dxt_q, dxt_scale = triton_silu_and_smooth_quant_backward( - grad_output, x, + grad_output, + x, smooth_scale=grad_smooth_scale, transpose_smooth_scale=transpose_grad_smooth_scale, reverse=True, - round_scale=True) - - output_check(dx_q_ref.float(), dx_q.float(), 'smooth.dx_data') - output_check(dx_scale_ref, dx_scale, 'smooth.dx_scale') - output_check(dxt_q_ref.float(), dxt_q.float(), 'smooth.dxt_data') - output_check(dxt_scale_ref, dxt_scale, 'smooth.dxt_scale') - - if bench: - benchmark_func(torch_silu_and_smooth_quant_forward, x, - smooth_scale=smooth_scale, - n_repeat=100, ref_bytes=M * N * 2.5) - benchmark_func(triton_silu_and_smooth_quant_forward, x, - smooth_scale=smooth_scale, - n_repeat=100, ref_bytes=M * N * 2.5) - benchmark_func(triton_silu_and_smooth_quant_backward, grad_output, x, - smooth_scale=grad_smooth_scale, - transpose_smooth_scale=transpose_grad_smooth_scale, - n_repeat=100, ref_bytes=M * N * 5) - - -def test_silu_and_block_quant(M=4096, N=4096, bench=False): - if True: - x = torch.randn((M, N), dtype=torch.bfloat16, device='cuda:0') - x = (x * 10).clone().detach().requires_grad_() - grad_output = torch.randn((M, N // 2), dtype=torch.bfloat16, - device='cuda:0') - else: - d = torch.load('/ossfs/workspace/tmp/vis/silu.bin') - x = d['x'].clone().detach().to('cuda:0').requires_grad_() - grad_output = d['g'].to('cuda:0') - N = x.shape[-1] - M = x.shape[0] + round_scale=True, + ) + output_check(dx_q_ref, dx_q, "smooth.dx_data", rtol=0.125) + output_check(dx_scale_ref, dx_scale, "smooth.dx_scale") + output_check(dxt_q_ref, dxt_q, "smooth.dxt_data", rtol=0.125) + output_check(dxt_scale_ref, dxt_scale, "smooth.dxt_scale") + + benchmark( + torch_silu_and_smooth_quant_forward, + x, + smooth_scale=smooth_scale, + n_repeat=100, + ref_bytes=M * N * 5, + ) + benchmark( + triton_silu_and_smooth_quant_forward, + x, + smooth_scale=smooth_scale, + n_repeat=100, + ref_bytes=M * N * 5, + ) + benchmark( + triton_silu_and_smooth_quant_backward, + grad_output, + x, + smooth_scale=grad_smooth_scale, + transpose_smooth_scale=transpose_grad_smooth_scale, + n_repeat=100, + ref_bytes=M * N * 10, + ) + + +@pytest.mark.parametrize( + "M,N,coef,grad_coef,limit", + [ + (16384, 1024, 1.0, 1.0, None), + (16384, 1024, 1.0, 1.0, 4.0), + (16384, 2048, 1.0, 1.0, 4.0), + (8192, 4096, 1.0, 1.0, None), + (16384, 1536, 1.0, 1.0, None), + (4096, 1536 * 8, 1.0, 1.0, None), + (4096, 1536 * 8, 100.0, 100.0, None), + (4096, 1536 * 8, 0.0, 0.0, None), + ], +) +def test_silu_and_block_quant(M, N, coef, grad_coef, limit, benchmark): + x = torch.randn((M, 2 * N), dtype=torch.bfloat16, device="cuda:0") + x = (x * coef).clone().detach().requires_grad_() + grad_output = torch.randn((M, N), dtype=torch.bfloat16, device="cuda:0") * grad_coef + + round_scale = False y_q_ref, y_scale_ref, yt_q_ref, yt_scale_ref = torch_silu_and_block_quant_forward( - x, round_scale=True) - y_q, y_scale, yt_q, yt_scale = triton_silu_and_block_quant_forward(x, - round_scale=True, - output_mode=2) - output_check(y_q_ref.float(), y_q.float(), 'block.2.y_q') - output_check(y_scale_ref, y_scale.t(), 'block.2.y_scale') - output_check(yt_q_ref, yt_q, 'block.2.yt_q') - output_check(yt_scale_ref, yt_scale.t(), 'block.2.yt_scale') - - y_q, y_scale, yt_q, yt_scale = triton_silu_and_block_quant_forward(x, - round_scale=True, - output_mode=0) - output_check(y_q_ref.float(), y_q.float(), 'block.0.y_q') - output_check(y_scale_ref, y_scale.t(), 'block.0.y_scale') - - y_q, y_scale, yt_q, yt_scale = triton_silu_and_block_quant_forward(x, - round_scale=True, - output_mode=1) - output_check(yt_q_ref.float(), yt_q.float(), 'block.1.yt_q') - output_check(yt_scale_ref, yt_scale.t(), 'block.1.yt_scale') - - dx_q_ref, dx_scale_ref, dxt_q_ref, dxt_scale_ref = torch_silu_and_block_quant_backward( - grad_output, x, - round_scale=True) + x, round_scale=round_scale, limit=limit + ) + + y_q, y_scale, yt_q, yt_scale = triton_silu_and_block_quant_forward( + x, round_scale=round_scale, limit=limit, output_mode=0 + ) + output_check(y_q_ref, y_q, "block.0.y_q", rtol=0.125) + output_check(y_scale_ref, y_scale.t(), "block.0.y_scale") + + y_q, y_scale, yt_q, yt_scale = triton_silu_and_block_quant_forward( + x, round_scale=round_scale, limit=limit, output_mode=1 + ) + output_check(yt_q_ref, yt_q, "block.1.yt_q", rtol=0.125) + output_check(yt_scale_ref, yt_scale.t(), "block.1.yt_scale") + + y_q, y_scale, yt_q, yt_scale = triton_silu_and_block_quant_forward( + x, round_scale=round_scale, limit=limit, output_mode=2 + ) + output_check(y_q_ref, y_q, "block.2.y_q", rtol=0.125) + output_check(y_scale_ref, y_scale.t(), "block.2.y_scale") + output_check(yt_q_ref, yt_q, "block.2.yt_q", rtol=0.125) + output_check(yt_scale_ref, yt_scale.t(), "block.2.yt_scale") + + dx_q_ref, dx_scale_ref, dxt_q_ref, dxt_scale_ref = ( + torch_silu_and_block_quant_backward( + grad_output, x, round_scale=round_scale, limit=limit + ) + ) dx_q, dx_scale, dxt_q, dxt_scale = triton_silu_and_block_quant_backward( - grad_output, x, - round_scale=True) - output_check(dx_q_ref.float(), dx_q.float(), 'block.dx_q') - output_check(dx_scale_ref.t(), dx_scale, 'block.dx_scale') - output_check(dxt_q_ref.float(), dxt_q.float(), 'block.dxt_q') - output_check(dxt_scale_ref.t(), dxt_scale, 'block.dxt_scale') - - if bench: - benchmark_func(triton_silu_and_block_quant_forward, x, - n_repeat=100, ref_bytes=M * N * 3) - benchmark_func(triton_silu_and_block_quant_backward, grad_output, x, - n_repeat=100, ref_bytes=M * N * 5) - - -def test_triton_batch_weighted_silu_and_smooth_quant(M=1024, N=4096, - n_experts=32, - bench=False): - if True: - count_list = [random.randint(M // 2, M // 2 * 3) // 16 * 16 for _ in - range(n_experts)] - counts = torch.tensor(count_list, device='cuda:0', dtype=torch.int32) - bs = sum(count_list) - - x = torch.randn((bs, N), dtype=torch.bfloat16, device='cuda:0') ** 3 / 4 - weight = torch.randn((bs, 1), dtype=torch.float32, device='cuda:0') - smooth_scales = 1 + torch.rand((n_experts, N // 2), dtype=torch.float32, - device='cuda:0') * 10 - else: - d = torch.load('/ossfs/workspace/Megatron-LM/silu.bin') - counts = d['counts'].cuda() - x = d['x'].cuda() - weight = d['weight'].cuda() - smooth_scales = d['smooth_scale'].cuda() - bs = sum(counts.tolist()) - N = x.shape[-1] - n_experts = counts.shape[0] - - grad_output = torch.randn((bs, N // 2), dtype=torch.bfloat16, - device='cuda:0') ** 3 - grad_smooth_scales = 1 + torch.rand((n_experts, N), dtype=torch.float32, - device='cuda:0') * 10 - transpose_grad_smooth_scales = 1 + torch.rand((bs,), dtype=torch.float32, - device='cuda:0') * 10 + grad_output, x, round_scale=round_scale, limit=limit + ) + output_check(dx_q_ref, dx_q, "block.dx_q", rtol=0.125) + output_check(dx_scale_ref.t(), dx_scale, "block.dx_scale") + output_check(dxt_q_ref, dxt_q, "block.dxt_q", rtol=0.125) + output_check(dxt_scale_ref.t(), dxt_scale, "block.dxt_scale") + + n_profile = 0 + benchmark( + triton_silu_and_block_quant_forward, + x, + round_scale=round_scale, + output_mode=0, + limit=limit, + n_repeat=100, + ref_bytes=M * N * 6, + n_profile=n_profile, + ) + benchmark( + triton_silu_and_block_quant_forward, + x, + round_scale=round_scale, + output_mode=1, + limit=limit, + n_repeat=100, + ref_bytes=M * N * 6, + n_profile=n_profile, + ) + benchmark( + triton_silu_and_block_quant_forward, + x, + round_scale=round_scale, + output_mode=2, + limit=limit, + n_repeat=100, + ref_bytes=M * N * 6, + n_profile=n_profile, + ) + benchmark( + triton_silu_and_block_quant_backward, + grad_output, + x, + limit=limit, + n_repeat=100, + ref_bytes=M * N * 10, + n_profile=n_profile, + ) + + +@pytest.mark.parametrize( + "M,N,coef,grad_coef,limit", + [ + (16384, 1024, 1.0, 1e-8, 4.0), + (2345, 1024, 10.0, 1.0, None), + (2345, 1536, 10.0, 1e-8, None), + ], +) +def test_silu_and_mxfp8_quant(M, N, coef, grad_coef, limit, benchmark): + x = torch.randn((M, 2 * N), dtype=torch.bfloat16, device="cuda:0") + x = (x * coef).clone().detach().requires_grad_() + grad_output = torch.randn((M, N), dtype=torch.bfloat16, device="cuda:0") * grad_coef + + y_q_ref, y_scale_ref, yt_q_ref, yt_scale_ref = torch_silu_and_mxfp8_quant_forward( + x, limit + ) + y_q, y_scale, yt_q, yt_scale = triton_silu_and_mxfp8_quant_forward( + x, limit=limit, output_mode=2 + ) + output_check(y_q_ref.float(), y_q.float(), "block.2.y_q") + output_check(y_scale_ref, y_scale, "block.2.y_scale") + output_check(yt_q_ref, yt_q, "block.2.yt_q") + output_check(yt_scale_ref, yt_scale, "block.2.yt_scale") + + y_q, y_scale, yt_q, yt_scale = triton_silu_and_mxfp8_quant_forward( + x, limit=limit, output_mode=0 + ) + output_check(y_q_ref.float(), y_q.float(), "block.0.y_q") + output_check(y_scale_ref, y_scale, "block.0.y_scale") + + y_q, y_scale, yt_q, yt_scale = triton_silu_and_mxfp8_quant_forward( + x, limit=limit, output_mode=1 + ) + output_check(yt_q_ref.float(), yt_q.float(), "block.1.yt_q") + output_check(yt_scale_ref, yt_scale, "block.1.yt_scale") + + dx_q_ref, dx_scale_ref, dxt_q_ref, dxt_scale_ref = ( + torch_silu_and_mxfp8_quant_backward(grad_output, x, limit=limit) + ) + dx_q, dx_scale, dxt_q, dxt_scale = triton_silu_and_mxfp8_quant_backward( + grad_output, x, limit=limit + ) + output_check(dx_q_ref, dx_q, "block.dx_q", rtol=0.125) + output_check(dx_scale_ref, dx_scale, "block.dx_scale") + output_check(dxt_q_ref, dxt_q, "block.dxt_q", rtol=0.125) + output_check(dxt_scale_ref, dxt_scale, "block.dxt_scale") + + benchmark(triton_silu_and_mxfp8_quant_forward, x, n_repeat=100, ref_bytes=M * N * 6) + benchmark( + triton_silu_and_mxfp8_quant_backward, + grad_output, + x, + n_repeat=100, + ref_bytes=M * N * 10, + ) + + +@pytest.mark.parametrize( + "M,N,n_experts", + [ + (0, 2048, 32), + (2048, 2048, 32), + ], +) +def test_triton_batch_weighted_silu_and_smooth_quant( + M, N, n_experts, benchmark, coef=1.0, grad_coef=1.0 +): + count_list = [ + random.randint(M // 2, M // 2 * 3) // 16 * 16 for _ in range(n_experts) + ] + counts = torch.tensor(count_list, device="cuda:0", dtype=torch.int32) + bs = sum(count_list) + + x = torch.randn((bs, 2 * N), dtype=torch.bfloat16, device="cuda:0") * coef + weight = torch.randn((bs, 1), dtype=torch.float32, device="cuda:0") + smooth_scales = ( + 1 + torch.rand((n_experts, N), dtype=torch.float32, device="cuda:0") * 10 + ) + + grad_output = ( + torch.randn((bs, N), dtype=torch.bfloat16, device="cuda:0") * grad_coef + ) + grad_smooth_scales = ( + 1 + torch.rand((n_experts, 2 * N), dtype=torch.float32, device="cuda:0") * 10 + ) + transpose_grad_smooth_scales = ( + 1 + torch.rand((bs,), dtype=torch.float32, device="cuda:0") * 10 + ) round_scale = True - - x_q_ref, x_scale_ref, x_max_ref = torch_batch_weighted_silu_and_smooth_quant_forward( + rtol = 2 if round_scale else 0.125 + x_q_ref, x_scale_ref = torch_batch_weighted_silu_and_smooth_quant_forward( x, weight, counts, smooth_scales=smooth_scales, round_scale=round_scale, - reverse=False) - x_q, x_scale, maxs = triton_batch_weighted_silu_and_smooth_quant_forward(x, - weight, - counts, - smooth_scale=smooth_scales, - round_scale=round_scale, - reverse=False) - output_check(x_q_ref.float(), x_q.float(), 'smooth.data') - output_check(x_scale_ref, x_scale, 'smooth.scale') - - dx_ref, dx_scale_ref, dw_ref, dxt_ref, dxt_scale_ref = torch_batch_weighted_silu_and_smooth_quant_backward( - grad_output, x, weight, count_list, - smooth_scales=grad_smooth_scales, - transpose_smooth_scale=transpose_grad_smooth_scales, - round_scale=round_scale, reverse=False) - dx, dx_scale, dw, dxt, dxt_scale = triton_batch_weighted_silu_and_smooth_quant_backward( - grad_output, x, weight, counts, - smooth_scale=grad_smooth_scales, + reverse=False, + ) + x_q, x_scale = triton_batch_weighted_silu_and_smooth_quant_forward( + x, + weight, + counts, + smooth_scale=smooth_scales, + round_scale=round_scale, + reverse=False, + ) + output_check(x_q_ref, x_q, "smooth.data", rtol=rtol) + output_check(x_scale_ref, x_scale, "smooth.scale") + + dx_ref, dx_scale_ref, dw_ref, dxt_ref, dxt_scale_ref = ( + torch_batch_weighted_silu_and_smooth_quant_backward( + grad_output, + x, + weight, + count_list, + smooth_scales=grad_smooth_scales, + transpose_smooth_scale=transpose_grad_smooth_scales, + round_scale=round_scale, + reverse=False, + ) + ) + dx, dx_scale, dw, dxt, dxt_scale = ( + triton_batch_weighted_silu_and_smooth_quant_backward( + grad_output, + x, + weight, + counts, + smooth_scale=grad_smooth_scales, + transpose_smooth_scale=transpose_grad_smooth_scales, + splits=count_list, + round_scale=round_scale, + reverse=False, + ) + ) + output_check(dx_ref, dx, "smooth.dx", rtol=rtol) + output_check(dx_scale_ref, dx_scale, "smooth.dx_scale") + rate = coef**0.75 if coef > 1 else 1 + output_check(dw_ref, dw, "smooth.dw", rtol=1e-3 * rate, atol=1e-3 * rate) + output_check(dxt_ref, dxt, "smooth.dxt", rtol=rtol) + output_check(dxt_scale_ref, dxt_scale.view(-1), "smooth.dxt_scale") + + ref_time = None + benchmark( + triton_batch_weighted_silu_and_smooth_quant_forward, + x, + weight, + counts, + smooth_scale=smooth_scales, + round_scale=True, + ref_bytes=n_experts * M * N * 5, + ref_time=ref_time, + ) + benchmark( + triton_batch_weighted_silu_and_smooth_quant_backward, + grad_output, + x, + weight, + counts, + smooth_scale=smooth_scales, transpose_smooth_scale=transpose_grad_smooth_scales, splits=count_list, + round_scale=True, + ref_bytes=n_experts * M * N * 8, + ref_time=ref_time, + ) + + +@pytest.mark.parametrize( + "M,N,n_experts,limit,coef,grad_coef", + [ + (0, 1536, 32, None, 1.0, 1.0), + (2048, 2048, 32, None, 1.0, 1.0), + (2048, 2048, 32, 4.0, 1.0, 1.0), + (2048, 1536, 32, None, 100.0, 100.0), + (2048, 1536, 32, None, 0.0, 0.0), + ], +) +def test_triton_batch_weighted_silu_and_block_quant( + M, N, n_experts, limit, coef, grad_coef, benchmark +): + count_list = [ + random.randint(M // 2, M // 2 * 3) // 16 * 16 for _ in range(n_experts) + ] + counts = torch.tensor(count_list, device="cuda:0", dtype=torch.int32) + bs = sum(count_list) + + x = torch.randn((bs, 2 * N), dtype=torch.bfloat16, device="cuda:0") * coef + if bs > 3: + x[:3] = 0.0 + weight = torch.randn((bs, 1), dtype=torch.float32, device="cuda:0") + + grad_output = ( + torch.randn((bs, N), dtype=torch.bfloat16, device="cuda:0") * grad_coef + ) + round_scale = False + rtol = 2 if round_scale else 0.125 + + x_q_ref, x_scale_ref, xt_q_ref, xt_scale_ref = ( + torch_batch_weighted_silu_and_block_quant_forward( + x, weight, counts, limit=limit, round_scale=round_scale + ) + ) + x_q, x_scale, xt_q, xt_scale = triton_batch_weighted_silu_and_block_quant_forward( + x, + weight, + counts, + count_list, + limit=limit, round_scale=round_scale, - reverse=False) - output_check(dx_ref.float(), dx.float(), 'smooth.dx') - output_check(dx_scale_ref, dx_scale, 'smooth.dx_scale') - output_check(dw_ref, dw, 'smooth.dw') - output_check(dxt_ref.float(), dxt.float(), 'smooth.dxt') - output_check(dxt_scale_ref, dxt_scale.view(-1), 'smooth.dxt_scale') - - if bench: - ref_time = None - benchmark_func(triton_batch_weighted_silu_and_smooth_quant_forward, x, - weight, - counts, smooth_scale=smooth_scales, round_scale=True, - ref_bytes=n_experts * M * N * 2.5, ref_time=ref_time) - benchmark_func(triton_batch_weighted_silu_and_smooth_quant_backward, - grad_output, x, weight, counts, - smooth_scale=smooth_scales, - transpose_smooth_scale=transpose_grad_smooth_scales, - splits=count_list, - round_scale=True, - ref_bytes=n_experts * M * N * 4, ref_time=ref_time) - - -def test_triton_batch_weighted_silu_and_block_quant(M=1024, N=4096, - n_experts=32, - bench=False): - if True: - count_list = [random.randint(M // 2, M // 2 * 3) // 16 * 16 for _ in - range(n_experts)] - counts = torch.tensor(count_list, device='cuda:0', dtype=torch.int32) - bs = sum(count_list) - - x = torch.randn((bs, N), dtype=torch.bfloat16, - device='cuda:0') ** 3 / 10 - weight = torch.randn((bs, 1), dtype=torch.float32, device='cuda:0') - else: - d = torch.load('/ossfs/workspace/Megatron-LM/silu.bin') - counts = d['counts'].cuda() - x = d['x'].cuda() - weight = d['weight'].cuda() - smooth_scales = d['smooth_scale'].cuda() - bs = sum(counts.tolist()) - N = x.shape[-1] - n_experts = counts.shape[0] - - grad_output = torch.randn((bs, N // 2), dtype=torch.bfloat16, - device='cuda:0') ** 3 - round_scale = True + output_mode=2, + ) + + output_check(x_q_ref, x_q, "block.q", rtol=rtol) + output_check(x_scale_ref, x_scale.view(-1), "block.scale") + output_check(xt_q_ref, xt_q.view(-1), "block.qt", rtol=rtol) + output_check(xt_scale_ref, xt_scale.view(-1), "block.t_scale", atol=1e-3) - x_q_ref, x_scale_ref, xt_q_ref, xt_scale_ref = torch_batch_weighted_silu_and_block_quant_forward( + x_q, x_scale, xt_q, xt_scale = triton_batch_weighted_silu_and_block_quant_forward( x, weight, counts, - round_scale=round_scale) + count_list, + limit=limit, + round_scale=round_scale, + output_mode=0, + ) + output_check(x_q_ref, x_q, "block.q", rtol=rtol) + output_check(x_scale_ref, x_scale.view(-1), "block.scale") + x_q, x_scale, xt_q, xt_scale = triton_batch_weighted_silu_and_block_quant_forward( x, weight, counts, count_list, + limit=limit, round_scale=round_scale, - output_mode=2) - - output_check(x_q_ref.float(), x_q.float(), 'block.q') - output_check(x_scale_ref, x_scale, 'block.scale') - output_check(xt_q_ref.float(), xt_q.float(), 'block.qt') - output_check(xt_scale_ref, xt_scale, 'block.t_scale') - - dx_ref, dx_scale_ref, dw_ref, dxt_ref, dxt_scale_ref = torch_batch_weighted_silu_and_block_quant_backward( - grad_output, x, weight, counts, - round_scale=round_scale) - dx, dx_scale, dw, dxt, dxt_scale = triton_batch_weighted_silu_and_block_quant_backward( - grad_output, x, weight, counts, splits=count_list, - round_scale=round_scale) - output_check(dx_ref.float(), dx.float(), 'block.dx') - output_check(dx_scale_ref, dx_scale, 'block.dx_scale') - output_check(dw_ref, dw, 'block.dw') - output_check(dxt_ref.float(), dxt.float(), 'block.dxt') - output_check(dxt_scale_ref, dxt_scale, 'block.dxt_scale') - - if bench: - ref_time = None - benchmark_func(triton_batch_weighted_silu_and_block_quant_forward, x, - weight, - counts, round_scale=True, splits=count_list, - output_mode=0, n_repeat=100, - ref_bytes=n_experts * M * N * 2.5, ref_time=ref_time) - benchmark_func(triton_batch_weighted_silu_and_block_quant_forward, x, - weight, - counts, round_scale=True, splits=count_list, - output_mode=2, n_repeat=100, - ref_bytes=n_experts * M * N * 3, ref_time=ref_time) - benchmark_func(triton_batch_weighted_silu_and_block_quant_backward, - grad_output, x, weight, counts, - round_scale=True, splits=count_list, n_repeat=100, - ref_bytes=n_experts * M * N * 4, ref_time=ref_time) - - -if __name__ == '__main__': - test_weighted_silu(M=16384, N=1024, bench=True) - - test_silu_and_smooth_quant(M=16384, N=1024, bench=False) - test_silu_and_smooth_quant(M=8192, N=2048, bench=False) - test_silu_and_smooth_quant(M=4096, N=10240, bench=False) - test_silu_and_smooth_quant(M=4096, N=5120, bench=False) - - test_silu_and_block_quant(M=16384, N=1024, bench=True) - - test_triton_batch_weighted_silu_and_smooth_quant(M=2048, N=2048, - n_experts=32, bench=False) - test_triton_batch_weighted_silu_and_smooth_quant(M=800, N=2048, n_experts=32, bench=False) - test_triton_batch_weighted_silu_and_smooth_quant(M=0, N=2048, n_experts=32, bench=False) - - test_triton_batch_weighted_silu_and_block_quant(M=4096, N=2048, - n_experts=32, bench=True) - test_triton_batch_weighted_silu_and_block_quant(M=1008, N=2048, n_experts=32, bench=False) + output_mode=1, + ) + output_check(xt_q_ref, xt_q.view(-1), "block.qt", rtol=rtol) + output_check(xt_scale_ref, xt_scale.view(-1), "block.t_scale") + + dx_ref, dx_scale_ref, dw_ref, dxt_ref, dxt_scale_ref = ( + torch_batch_weighted_silu_and_block_quant_backward( + grad_output, x, weight, counts, round_scale=round_scale, limit=limit + ) + ) + dx, dx_scale, dw, dxt, dxt_scale = ( + triton_batch_weighted_silu_and_block_quant_backward( + grad_output, + x, + weight, + counts, + splits=count_list, + limit=limit, + round_scale=round_scale, + ) + ) + output_check(dx_ref, dx, "block.dx", rtol=rtol) + output_check(dx_scale_ref, dx_scale.view(-1), "block.dx_scale") + rate = (coef * grad_coef) ** 0.75 if coef * grad_coef > 1 else 1 + output_check(dw_ref, dw, "block.dw", rtol=1e-3 * rate, atol=1e-3 * rate) + output_check(dxt_ref, dxt.view(-1), "block.dxt", rtol=rtol) + output_check(dxt_scale_ref, dxt_scale.view(-1), "block.dxt_scale") + + ref_time = None + benchmark( + triton_batch_weighted_silu_and_block_quant_forward, + x, + weight, + counts, + round_scale=True, + splits=count_list, + output_mode=0, + limit=limit, + n_repeat=100, + ref_bytes=n_experts * M * N * 5, + ref_time=ref_time, + ) + benchmark( + triton_batch_weighted_silu_and_block_quant_forward, + x, + weight, + counts, + round_scale=True, + splits=count_list, + output_mode=1, + limit=limit, + n_repeat=100, + ref_bytes=n_experts * M * N * 6, + ref_time=ref_time, + ) + benchmark( + triton_batch_weighted_silu_and_block_quant_forward, + x, + weight, + counts, + round_scale=True, + splits=count_list, + output_mode=2, + limit=limit, + n_repeat=100, + ref_bytes=n_experts * M * N * 6, + ref_time=ref_time, + ) + benchmark( + triton_batch_weighted_silu_and_block_quant_backward, + grad_output, + x, + weight, + counts, + round_scale=True, + limit=limit, + splits=count_list, + n_repeat=100, + ref_bytes=n_experts * M * N * 8, + ref_time=ref_time, + ) + + +@pytest.mark.parametrize( + "M,N,n_experts,limit,coef,grad_coef", + [ + (0, 2048, 32, None, 1.0, 1.0), + (2048, 2048, 32, 4.0, 1.0, 1.0), + (2048, 1536, 32, None, 1.0, 1.0), + (2048, 1536, 32, None, 10000.0, 10000.0), + ], +) +def test_triton_batch_weighted_silu_and_mxfp8_quant( + M, N, n_experts, limit, coef, grad_coef, benchmark +): + count_list = [ + random.randint(M // 2, M // 2 * 3) // 16 * 16 for _ in range(n_experts) + ] + counts = torch.tensor(count_list, device="cuda:0", dtype=torch.int32) + bs = sum(count_list) + + x = torch.randn((bs, 2 * N), dtype=torch.bfloat16, device="cuda:0") * coef + weight = torch.randn((bs, 1), dtype=torch.float32, device="cuda:0") + + grad_output = ( + torch.randn((bs, N), dtype=torch.bfloat16, device="cuda:0") * grad_coef + ) + + x_q_ref, x_scale_ref, xt_q_ref, xt_scale_ref = ( + torch_batch_weighted_silu_and_mxfp8_quant_forward( + x, weight, counts, limit=limit + ) + ) + x_q, x_scale, xt_q, xt_scale = triton_batch_weighted_silu_and_mxfp8_quant_forward( + x, weight, counts, count_list, limit=limit, output_mode=2 + ) + + rtol = 2 + output_check(x_q_ref, x_q, "mxfp8.q", rtol=rtol) + output_check(x_scale_ref, x_scale, "mxfp8.scale", itol=1) + output_check(xt_q_ref, xt_q, "mxfp8.qt", rtol=rtol) + output_check(xt_scale_ref, xt_scale, "mxfp8.t_scale", itol=1) + + dx_ref, dx_scale_ref, dw_ref, dxt_ref, dxt_scale_ref = ( + torch_batch_weighted_silu_and_mxfp8_quant_backward( + grad_output, x, weight, counts, limit=limit + ) + ) + dx, dx_scale, dw, dxt, dxt_scale = ( + triton_batch_weighted_silu_and_mxfp8_quant_backward( + grad_output, x, weight, counts, splits=count_list, limit=limit + ) + ) + output_check(dx_ref, dx, "mxfp8.dx", rtol=rtol) + output_check(dx_scale_ref, dx_scale, "mxfp8.dx_scale", itol=1) + rate = coef**0.75 if coef > 1 else 1 + output_check(dw_ref, dw, "mxfp8.dw", rtol=1e-3 * rate, atol=1e-3 * rate) + output_check(dxt_ref, dxt, "mxfp8.dxt", rtol=rtol) + output_check(dxt_scale_ref, dxt_scale, "mxfp8.dxt_scale", itol=1) + + ref_time = None + benchmark( + triton_batch_weighted_silu_and_mxfp8_quant_forward, + x, + weight, + counts, + splits=count_list, + output_mode=0, + n_repeat=100, + ref_bytes=n_experts * M * N * 5, + ref_time=ref_time, + ) + benchmark( + triton_batch_weighted_silu_and_mxfp8_quant_forward, + x, + weight, + counts, + splits=count_list, + output_mode=2, + n_repeat=100, + ref_bytes=n_experts * M * N * 6, + ref_time=ref_time, + ) + benchmark( + triton_batch_weighted_silu_and_mxfp8_quant_backward, + grad_output, + x, + weight, + counts, + splits=count_list, + n_repeat=100, + ref_bytes=n_experts * M * N * 8, + ref_time=ref_time, + ) diff --git a/tests/test_smooth_quant.py b/tests/test_smooth_quant.py index b0011d9..563c2df 100644 --- a/tests/test_smooth_quant.py +++ b/tests/test_smooth_quant.py @@ -3,19 +3,21 @@ Copyright (c) Ant Financial Service Group and its affiliates. """ +import pytest import torch -from linghe.quant.smooth import (triton_batch_smooth_quant, - triton_subrow_smooth_quant, - triton_transpose_rescale_smooth_quant, - triton_smooth_quant, - triton_transpose_smooth_quant) -from linghe.tools.benchmark import benchmark_func -from linghe.tools.util import (output_check, - torch_make_indices, - torch_smooth_quant, - round_up) -from linghe.facade.smooth_quant_linear import SmoothQuantLinear +from linghe.facade.linear import SmoothQuantLinear +from linghe.quant.smooth import ( + triton_smooth_quant, + triton_transpose_smooth_quant, + triton_batch_smooth_quant, + triton_batch_transpose_smooth_quant, + triton_transpose_rescale_smooth_quant, + triton_subrow_smooth_quant, +) +from linghe.tools.check import output_check +from linghe.tools.util import torch_make_indices, torch_smooth_quant, round_up + def torch_split_smooth_quant(x_split, smooth_scales, round_scale=False): x_qs = [] @@ -31,14 +33,55 @@ def torch_split_smooth_quant(x_split, smooth_scales, round_scale=False): x_qs.append(x_q_) x_scales.append(x_scale_) x_maxs = torch.stack(x_maxs, 0) - return x_qs, x_scales, x_maxs + return x_qs, x_scales + + +def torch_split_transpose_smooth_quant( + x_split, smooth_scale_split, round_scale=False, reverse=True +): + assert reverse + x_qs = [] + x_scales = [] + for i, x_ in enumerate(x_split): + smooth_scale = smooth_scale_split[i] + M, N = x_.shape + if M % 32 != 0: + x_ = torch.cat( + [x_, torch.zeros((32 - M % 32, N), dtype=x_.dtype, device=x_.device)], 0 + ) + smooth_scale = torch.cat( + [ + smooth_scale, + torch.zeros((32 - M % 32,), dtype=torch.float32, device=x_.device), + ], + 0, + ) + M = (M + 31) // 32 * 32 + x_smooth = x_ * smooth_scale[:, None] + x_scale_ = x_smooth.float().abs().amax(0) / 448 + x_scale_ = torch.maximum(x_scale_, 1e-30 + x_scale_ * 0.0) + if round_scale: + x_scale_ = torch.exp2(torch.ceil(torch.log2(x_scale_))) + x_q_ = (x_smooth / x_scale_).t().contiguous().to(torch.float8_e4m3fn) + x_qs.append(x_q_.view(M, N)) + x_scales.append(x_scale_) + x_qs = torch.cat(x_qs, 0) + x_scales = torch.stack(x_scales, 0) + return x_qs, x_scales -def torch_subrow_smooth_quant(x, smooth_scale, x_q, x_scale, subrow_scales, - offset, size, - reverse=False, round_scale=False): - limit = 448 * torch.ones((1,), dtype=smooth_scale.dtype, - device=smooth_scale.device) +def torch_subrow_smooth_quant( + x, + smooth_scale, + x_q, + x_scale, + subrow_scales, + offset, + size, + reverse=False, + round_scale=False, +): + limit = 448 * torch.ones((1,), dtype=smooth_scale.dtype, device=smooth_scale.device) # subrow_scales is saved as 448/max M, N = x_q.shape @@ -46,7 +89,7 @@ def torch_subrow_smooth_quant(x, smooth_scale, x_q, x_scale, subrow_scales, si = offset % N k = N - si x_slice = x.view(-1)[0:k] - smooth_scale_slice = smooth_scale[si: N] + smooth_scale_slice = smooth_scale[si:N] if not reverse: smooth_scale_slice = 1 / smooth_scale_slice x_smooth = x_slice * smooth_scale_slice @@ -55,32 +98,41 @@ def torch_subrow_smooth_quant(x, smooth_scale, x_q, x_scale, subrow_scales, if round_scale: scale = torch.exp2(torch.floor(torch.log2(scale))) - x_q_slice = torch.minimum(torch.maximum(x_smooth / scale, -limit), - limit).to(torch.float8_e4m3fn) - x_q.view(-1)[offset:offset + k] = x_q_slice + x_q_slice = torch.minimum(torch.maximum(x_smooth / scale, -limit), limit).to( + torch.float8_e4m3fn + ) + x_q.view(-1)[offset : offset + k] = x_q_slice if (offset + size) % N > 0: k = (offset + size) % N x_slice = x.view(-1)[-k:] - smooth_scale_slice = smooth_scale[0: k] + smooth_scale_slice = smooth_scale[0:k] if not reverse: smooth_scale_slice = 1 / smooth_scale_slice x_smooth = x_slice * smooth_scale_slice scale = subrow_scales[1:2] if round_scale: scale = torch.exp2(torch.floor(torch.log2(scale))) - x_q_slice = torch.minimum(torch.maximum(x_smooth / scale, -limit), - limit).to(torch.float8_e4m3fn) - x_q.view(-1)[(offset + size - k):(offset + size)] = x_q_slice + x_q_slice = torch.minimum(torch.maximum(x_smooth / scale, -limit), limit).to( + torch.float8_e4m3fn + ) + x_q.view(-1)[(offset + size - k) : (offset + size)] = x_q_slice x_scale[(offset + size) // N] = scale -def torch_rescale_quant(y_q, org_smooth_scale, y_scale, transpose_smooth_scale, - reverse=True, round_scale=True): +def torch_rescale_quant( + y_q, + org_smooth_scale, + y_scale, + transpose_smooth_scale, + reverse=True, + round_scale=True, +): assert reverse y = y_q.float() / org_smooth_scale * y_scale[:, None] - y_q, y_scale, _ = torch_smooth_quant(y.t(), transpose_smooth_scale, - reverse=True, round_scale=round_scale) + y_q, y_scale = torch_smooth_quant( + y.t(), transpose_smooth_scale, reverse=True, round_scale=round_scale + ) return y_q, y_scale @@ -88,268 +140,378 @@ def triton_split_smooth_quant(x_split, smooth_scales): x_qs = [] x_scales = [] for i, x_ in enumerate(x_split): - x_q_, x_scale_, _ = triton_smooth_quant(x_, smooth_scales[i]) + x_q_, x_scale_ = triton_smooth_quant(x_, smooth_scales[i]) x_qs.append(x_q_) x_scales.append(x_scale_) return x_qs, x_scales -def test_triton_smooth_quant(M=4096, N=4096, bench=False): - device = 'cuda:0' +@pytest.mark.parametrize( + "M,N", + [ + (16384, 2048), + (8192, 4096), + (4096, 8192), + (8192, 3072), + (8192, 6144), + (16384, 512), + (3457, 512), + ], +) +def test_triton_smooth_quant(M, N, benchmark): + device = "cuda:0" x = torch.randn((M, N), dtype=torch.bfloat16, device=device) - smooth_scale = torch.randn((N,), device=device, dtype=torch.float32).abs() - x_q_ref, scales_ref, x_maxs_ref = torch_smooth_quant(x, smooth_scale, - reverse=False, - round_scale=True) - - x_q, x_scale, x_maxs = triton_smooth_quant(x, smooth_scale, - reverse=False, - round_scale=True, - calibrate=True) - output_check(x_q_ref.float(), x_q.float(), - 'triton_smooth_quant.data') - output_check(scales_ref, x_scale, 'triton_smooth_quant.scale') - output_check(x_maxs_ref, x_maxs, 'triton_smooth_quant.x_maxs') - - if bench: - benchmark_func(triton_smooth_quant, x, - smooth_scale, - reverse=False, - round_scale=True, - calibrate=False, - ref_bytes=M * N * 3) - - -def test_triton_subrow_smooth_quant(M=4096, N=5120, offset=4096, - size=16384): - device = 'cuda:0' + smooth_scale = torch.randn((N,), device=device, dtype=torch.float32).abs() + 1.0 + round_scale = False + rtol = 2 if round_scale else 0.125 + x_q_ref, scales_ref = torch_smooth_quant( + x, smooth_scale, reverse=False, round_scale=round_scale + ) + + x_q, x_scale = triton_smooth_quant( + x, smooth_scale, reverse=False, round_scale=round_scale + ) + output_check(x_q_ref, x_q, "triton_smooth_quant.data", rtol=rtol) + output_check(scales_ref, x_scale, "triton_smooth_quant.scale") + + benchmark( + triton_smooth_quant, + x, + smooth_scale, + reverse=False, + round_scale=True, + ref_bytes=M * N * 3, + ) + + +@pytest.mark.parametrize( + "M,N,offset,size", + [ + (4096, 5120, 5120, 2048), + (4096, 5120, 4096, 5120), + (4096, 5120, 5120, 5120 * 10 - 1024), + ], +) +def test_triton_subrow_smooth_quant(M, N, offset, size, benchmark): + device = "cuda:0" x = torch.randn((size,), dtype=torch.float32, device=device) x_q = torch.zeros((M, N), dtype=torch.bfloat16, device=device).to( - torch.float8_e4m3fn) + torch.float8_e4m3fn + ) x_scale = torch.zeros((M,), dtype=torch.float32, device=device).abs() - smooth_scale = torch.randn((N,), device=device, - dtype=torch.float32).abs() + 1 - subrow_scales = torch.randn((2,), device=device, - dtype=torch.float32).abs() + 1 + smooth_scale = torch.randn((N,), device=device, dtype=torch.float32).abs() + 1 + subrow_scales = torch.randn((2,), device=device, dtype=torch.float32).abs() + 1 x_ref = x.clone() x_q_ref = x_q.clone() x_scale_ref = x_scale.clone() subrow_scales_ref = subrow_scales.clone() - torch_subrow_smooth_quant(x_ref, smooth_scale, x_q_ref, x_scale_ref, - subrow_scales_ref, offset, size, - reverse=False, round_scale=False) - - triton_subrow_smooth_quant(x, smooth_scale, x_q, x_scale, - subrow_scales, offset, size, - reverse=False, round_scale=False) - - output_check(x_q_ref.float(), x_q.float(), 'subrow.data') - output_check(x_scale_ref, x_scale, 'subrow.scale') + torch_subrow_smooth_quant( + x_ref, + smooth_scale, + x_q_ref, + x_scale_ref, + subrow_scales_ref, + offset, + size, + reverse=False, + round_scale=False, + ) + + triton_subrow_smooth_quant( + x, + smooth_scale, + x_q, + x_scale, + subrow_scales, + offset, + size, + reverse=False, + round_scale=False, + ) + + output_check(x_q_ref.float(), x_q.float(), "subrow.data") + output_check(x_scale_ref, x_scale, "subrow.scale") if offset % N > 0: k = N - offset % N - output_check(x_q_ref.float().view(-1)[offset:offset + k], - x_q.float().view(-1)[offset:offset + k], - 'subrow.data.tail') + output_check( + x_q_ref.float().view(-1)[offset : offset + k], + x_q.float().view(-1)[offset : offset + k], + "subrow.data.tail", + ) if (offset + size) % N > 0: k = (offset + size) % N - output_check(x_q_ref.float().view(-1)[offset + size - k:offset + size], - x_q.float().view(-1)[offset + size - k:offset + size], - 'subrow.data.head') + output_check( + x_q_ref.float().view(-1)[offset + size - k : offset + size], + x_q.float().view(-1)[offset + size - k : offset + size], + "subrow.data.head", + ) row_id = (offset + size) // N - output_check(x_scale_ref[row_id], x_scale[row_id], 'subrow.scale.slice') - - -def test_triton_transpose_smooth_quant(M=4096, N=4096, bench=False): - device = 'cuda:0' + output_check(x_scale_ref[row_id], x_scale[row_id], "subrow.scale.slice") + + +@pytest.mark.parametrize( + "M,N", + [ + (16384, 2048), + (8192, 4096), + (4096, 8192), + (4096, 3072), + ], +) +def test_triton_transpose_smooth_quant(M, N, benchmark): + device = "cuda:0" P = round_up(M, b=32) y = torch.randn((M, N), dtype=torch.bfloat16, device=device) ** 3 * 1e-10 - transpose_smooth_scale = torch.randn((M,), device=device, - dtype=torch.float32).abs() * 10 + 1 - yt_q, yt_scale = triton_transpose_smooth_quant(y, - transpose_smooth_scale, - reverse=True, - pad=True, - round_scale=True) - q_ref, scale_ref, maxs_ref = torch_smooth_quant(y.T.contiguous(), - transpose_smooth_scale, - reverse=True, - round_scale=True) + transpose_smooth_scale = ( + torch.randn((M,), device=device, dtype=torch.float32).abs() * 10 + 1 + ) + yt_q, yt_scale = triton_transpose_smooth_quant( + y, transpose_smooth_scale, reverse=True, pad=True, round_scale=True + ) + q_ref, scale_ref = torch_smooth_quant( + y.T.contiguous(), transpose_smooth_scale, reverse=True, round_scale=True + ) assert yt_q.shape[1] == P if P > M: assert yt_q.float()[:, M:].abs().sum().item() == 0 - output_check(q_ref, yt_q[:, :M], - 'triton_transpose_smooth_quant.data') - output_check(scale_ref, yt_scale, - 'triton_transpose_smooth_quant.scale') - - if bench: - benchmark_func(triton_transpose_smooth_quant, y, - transpose_smooth_scale, - reverse=True, - pad=True, - round_scale=True, - ref_bytes=M * N * 3) - - -def test_triton_transpose_rescale_smooth_quant(M=4096, N=4096, - round_scale=False): - device = 'cuda:0' + output_check(q_ref, yt_q[:, :M], "triton_transpose_smooth_quant.data") + output_check(scale_ref, yt_scale, "triton_transpose_smooth_quant.scale") + + benchmark( + triton_transpose_smooth_quant, + y, + transpose_smooth_scale, + reverse=True, + pad=True, + round_scale=True, + ref_bytes=M * N * 3, + ) + + +@pytest.mark.parametrize( + "M,N,round_scale", + [ + (4096, 4096, True), + (3895, 4096, True), + (4096, 3072, True), + (395, 2048, True), + ], +) +def test_triton_transpose_rescale_smooth_quant(M, N, round_scale, benchmark): + device = "cuda:0" P = round_up(M, b=32) y = torch.randn((M, N), dtype=torch.bfloat16, device=device) ** 3 - org_smooth_scale = torch.randn((N,), device=device, - dtype=torch.float32).abs() * 10 + 1 + org_smooth_scale = ( + torch.randn((N,), device=device, dtype=torch.float32).abs() * 10 + 1 + ) if round_scale: org_smooth_scale = torch.exp2(torch.ceil(torch.log2(org_smooth_scale))) - transpose_smooth_scale = torch.randn((M,), device=device, - dtype=torch.float32).abs() + 0.1 + transpose_smooth_scale = ( + torch.randn((M,), device=device, dtype=torch.float32).abs() + 0.1 + ) if round_scale: transpose_smooth_scale = torch.exp2( - torch.ceil(torch.log2(transpose_smooth_scale))) - - y_q, y_scale, y_maxs = triton_smooth_quant(y, org_smooth_scale, - reverse=True, - round_scale=round_scale) - - yt_gt, yt_scale_gt, yt_maxs_gt = torch_smooth_quant(y.t(), - transpose_smooth_scale, - reverse=True, - round_scale=round_scale) - - yt_q_ref, yt_scale_ref = torch_rescale_quant(y_q, org_smooth_scale, y_scale, - transpose_smooth_scale, - reverse=True, - round_scale=round_scale) - - yt_q, yt_scale = triton_transpose_rescale_smooth_quant(y_q, - org_smooth_scale, - y_scale, - transpose_smooth_scale, - reverse=True, - pad=True, - round_scale=round_scale) + torch.ceil(torch.log2(transpose_smooth_scale)) + ) + + y_q, y_scale = triton_smooth_quant( + y, org_smooth_scale, reverse=True, round_scale=round_scale + ) + + yt_gt, yt_scale_gt = torch_smooth_quant( + y.t(), transpose_smooth_scale, reverse=True, round_scale=round_scale + ) + + yt_q_ref, yt_scale_ref = torch_rescale_quant( + y_q, + org_smooth_scale, + y_scale, + transpose_smooth_scale, + reverse=True, + round_scale=round_scale, + ) + + yt_q, yt_scale = triton_transpose_rescale_smooth_quant( + y_q, + org_smooth_scale, + y_scale, + transpose_smooth_scale, + reverse=True, + pad=True, + round_scale=round_scale, + ) if P > M: assert yt_q.shape[1] == P yt_q.float()[:, M:].abs().sum().item() == 0 - output_check(yt_q_ref, yt_q[:, :M], - 'triton_transpose_rescale_smooth_quant.data') - output_check(yt_scale_ref, yt_scale, - 'triton_transpose_rescale_smooth_quant.scale') + output_check(yt_q_ref, yt_q[:, :M], "triton_transpose_rescale_smooth_quant.data") + output_check(yt_scale_ref, yt_scale, "triton_transpose_rescale_smooth_quant.scale") # should dequant and compare with gt # output_check(yt_gt, yt_q[:, :M], # 'triton_transpose_rescale_smooth_quant.data.gt') # output_check(yt_scale_gt, yt_scale, # 'triton_transpose_rescale_smooth_quant.scale.gt') - - -def test_triton_batch_smooth_quant(M=4096, N=4096, n_experts=32, topk=8, - round_scale=False, bench=False): - device = 'cuda:0' - - smooth_scales = 1 + 10 * torch.rand((n_experts, N), device=device, - dtype=torch.float32) + ref_bytes = M * N * 3 + benchmark( + triton_transpose_rescale_smooth_quant, + y_q, + org_smooth_scale, + y_scale, + transpose_smooth_scale, + reverse=True, + pad=True, + round_scale=round_scale, + ref_bytes=ref_bytes, + ) + + +@pytest.mark.parametrize( + "M,N,n_experts,topk,round_scale", + [ + (4096, 4096, 32, 8, False), + ], +) +def test_batch_smooth_quant(M, N, n_experts, topk, round_scale, benchmark): + device = "cuda:0" + + smooth_scales = 1 + 10 * torch.rand( + (n_experts, N), device=device, dtype=torch.float32 + ) 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() - x = torch.randn((sum(token_count_per_expert_list), N), dtype=torch.bfloat16, - device=device) + x = torch.randn( + (sum(token_count_per_expert_list), N), dtype=torch.bfloat16, device=device + ) - x_q, x_scale, x_maxs = triton_batch_smooth_quant(x, smooth_scales, - token_count_per_expert, - reverse=False, - round_scale=round_scale, - calibrate=True) + x_q, x_scale = triton_batch_smooth_quant( + x, smooth_scales, token_count_per_expert, reverse=False, round_scale=round_scale + ) x_split = torch.split(x, token_count_per_expert_list) - x_q_ref, x_scale_ref, x_maxs_ref = torch_split_smooth_quant(x_split, - smooth_scales) - x_q_ref = torch.cat([x.view(torch.uint8) for x in x_q_ref], 0).view( - torch.float8_e4m3fn) + x_q_ref, x_scale_ref = torch_split_smooth_quant(x_split, smooth_scales) + x_q_ref = torch.cat(x_q_ref, 0) x_scale_ref = torch.cat(x_scale_ref, 0) - output_check(x_q_ref.float(), x_q.float(), 'triton_batch_smooth_quant.data') - output_check(x_scale_ref.float(), x_scale.float(), - 'triton_batch_smooth_quant.scale') - output_check(x_maxs_ref.float(), x_maxs.float(), - 'triton_batch_smooth_quant.maxs') - - if bench: - n_repeat = 100 - ref_time = benchmark_func(triton_split_smooth_quant, x_split, - smooth_scales, n_repeat=n_repeat) - benchmark_func(triton_batch_smooth_quant, x, smooth_scales, - token_count_per_expert, reverse=False, - round_scale=round_scale, n_repeat=n_repeat, - ref_time=ref_time) - benchmark_func(triton_batch_smooth_quant, x, smooth_scales, - token_count_per_expert, reverse=False, - round_scale=round_scale, calibrate=True, - n_repeat=n_repeat, ref_time=ref_time) - + rtol = 2 if round_scale else 0.125 + output_check(x_q_ref, x_q, "triton_batch_smooth_quant.data", rtol=rtol) + output_check( + x_scale_ref.float(), x_scale.float(), "triton_batch_smooth_quant.scale" + ) + + ref_bytes = sum(token_count_per_expert_list) * N * 3 + ref_time = benchmark( + triton_split_smooth_quant, x_split, smooth_scales, ref_bytes=ref_bytes + ) + benchmark( + triton_batch_smooth_quant, + x, + smooth_scales, + token_count_per_expert, + reverse=False, + round_scale=round_scale, + ref_bytes=ref_bytes, + ref_time=ref_time, + ) + benchmark( + triton_batch_smooth_quant, + x, + smooth_scales, + token_count_per_expert, + reverse=False, + round_scale=round_scale, + ref_bytes=ref_bytes, + ref_time=ref_time, + ) + + +@pytest.mark.parametrize( + "M,N,n_experts,topk,round_scale", + [ + (4096, 4096, 32, 8, False), + ], +) +def test_batch_transpose_smooth_quant(M, N, n_experts, topk, round_scale, benchmark): + device = "cuda:0" + 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() + x = torch.randn( + (sum(token_count_per_expert_list), N), dtype=torch.bfloat16, device=device + ) + smooth_scales = 1 + torch.rand( + (sum(token_count_per_expert_list),), device=device, dtype=torch.float32 + ) + + x_q, x_scale = triton_batch_transpose_smooth_quant( + x, + smooth_scales, + token_count_per_expert, + token_count_per_expert_list, + reverse=True, + round_scale=round_scale, + pad=True, + ) - -def test_smooth_quant_linear(M=8192, N=1024, K=2048): - - dtype = torch.bfloat16 - device = 'cuda:0' + x_split = torch.split(x, token_count_per_expert_list) + smooth_scale_split = torch.split(smooth_scales, token_count_per_expert_list) + x_q_ref, x_scale_ref = torch_split_transpose_smooth_quant( + x_split, smooth_scale_split, round_scale=round_scale, reverse=True + ) + rtol = 2 if round_scale else 0.125 + output_check(x_q_ref, x_q, "batch_transpose_smooth_quant.data", rtol=rtol) + output_check( + x_scale_ref.float(), x_scale.float(), "batch_transpose_smooth_quant.scale" + ) + + ref_bytes = sum(token_count_per_expert_list) * N * 3 + benchmark( + triton_batch_transpose_smooth_quant, + x, + smooth_scales, + token_count_per_expert, + token_count_per_expert_list, + reverse=True, + round_scale=round_scale, + pad=True, + ref_bytes=ref_bytes, + ) + + +@pytest.mark.parametrize( + "M,N,K", + [ + (8192, 1024, 2048), + ], +) +def test_smooth_quant_linear(M, N, K, benchmark): + dtype = torch.bfloat16 + device = "cuda:0" linear = SmoothQuantLinear(K, N, bias=False, dtype=dtype, device=device) - x = (10*torch.randn((M, K), dtype=dtype, device=device)).requires_grad_() - w = 0.1*torch.randn((N, K), dtype=dtype, device=device) - dy = 1e-6*torch.randn((M, N), dtype=dtype, device=device) + x = (10 * torch.randn((M, K), dtype=dtype, device=device)).requires_grad_() + w = 0.1 * torch.randn((N, K), dtype=dtype, device=device) + dy = 1e-6 * torch.randn((M, N), dtype=dtype, device=device) linear.weight.data.copy_(w) - y_ref = x@w.t() + y_ref = x @ w.t() y = linear(x) - output_check(y_ref, y, mode='y') + output_check(y_ref, y, name="y", atol=-1) - dx_ref = dy@w - dw_ref = dy.t()@x + dx_ref = dy @ w + dw_ref = dy.t() @ x y.backward(dy) - dw = linear.weight.grad + dw = linear.weight.grad dx = x.grad - output_check(dx_ref, dx, mode='dx') - output_check(dw_ref, dw, mode='dw') - - - -if __name__ == '__main__': - test_triton_smooth_quant(M=16384, N=2048, bench=False) - test_triton_smooth_quant(M=8192, N=4096, bench=False) - test_triton_smooth_quant(M=4096, N=8192, bench=False) - test_triton_smooth_quant(M=8192, N=3072, bench=False) - test_triton_smooth_quant(M=8192, N=6144, bench=False) - test_triton_smooth_quant(M=16384, N=512, bench=False) - test_triton_smooth_quant(M=3457, N=512, bench=False) - - test_triton_subrow_smooth_quant(M=4096, N=5120, offset=5120, - size=2048) - test_triton_subrow_smooth_quant(M=4096, N=5120, offset=4096, - size=5120) - test_triton_subrow_smooth_quant(M=4096, N=5120, offset=5120, - size=5120 * 10 - 1024) - - test_triton_transpose_smooth_quant(M=16384, N=2048, bench=False) - test_triton_transpose_smooth_quant(M=8192, N=4096, bench=False) - test_triton_transpose_smooth_quant(M=4096, N=8192, bench=False) - test_triton_transpose_smooth_quant(M=4096, N=3072, bench=False) - - test_triton_transpose_rescale_smooth_quant(M=4096, N=4096, - round_scale=True) - test_triton_transpose_rescale_smooth_quant(M=3895, N=4096, - round_scale=True) - test_triton_transpose_rescale_smooth_quant(M=4096, N=3072, - round_scale=True) - test_triton_transpose_rescale_smooth_quant(M=395, N=2048, - round_scale=True) - - test_triton_batch_smooth_quant(M=4096, N=4096, n_experts=32, topk=8, - round_scale=False) - test_smooth_quant_linear(M=8192, N=1024, K=2048) \ No newline at end of file + output_check(dx_ref, dx, name="dx", atol=-1) + output_check(dw_ref, dw, name="dw", atol=-1) diff --git a/tests/test_topk.py b/tests/test_topk.py new file mode 100644 index 0000000..6eeaddd --- /dev/null +++ b/tests/test_topk.py @@ -0,0 +1,284 @@ +# -*- coding: utf-8 -*- +""" +Copyright (c) Ant Financial Service Group and its affiliates. +""" + +import torch + +from linghe.facade.topk import fused_topk, group_topk_score +from linghe.tools.benchmark import benchmark_func +from linghe.tools.check import output_check +from linghe.utils.topk import ( + triton_topk_forward, + triton_topk_backward, + triton_group_topk_score_forward, + triton_group_topk_score_backward, +) + + +def group_limited_topk( + scores: torch.Tensor, + topk: int, + num_tokens: int, + num_experts: int, + num_groups: int, + group_topk: int, +): + # Organize the experts into groups + # Select groups based on sum of top-(topk/group_topk) routing scores within each group + group_scores = ( + scores.view(num_tokens, num_groups, -1) + .topk(topk // group_topk, dim=-1)[0] + .sum(dim=-1) + ) + group_idx = torch.topk(group_scores, k=group_topk, dim=-1, sorted=False)[1] + group_mask = torch.zeros_like(group_scores) + group_mask.scatter_(1, group_idx, 1) + + # Mask the experts based on selection groups + score_mask = ( + group_mask.unsqueeze(-1) + .expand(num_tokens, num_groups, num_experts // num_groups) + .reshape(num_tokens, -1) + ) + + masked_scores = scores.masked_fill(~score_mask.bool(), float("-inf")) + probs, top_indices = torch.topk(masked_scores, k=topk, dim=-1) + + return probs, top_indices + + +def torch_group_topk_score( + logits, + expert_bias=None, + num_experts=256, + topk=8, + num_groups=32, + group_topk=4, + scaling_factor=1.0, + eps=1e-20, +): + num_tokens, num_experts = logits.shape + scores = torch.sigmoid(logits).to(torch.float64) + if expert_bias is not None: + expert_bias = expert_bias.to(torch.float64) + scores_for_routing = ( + scores + + expert_bias + - torch.arange(0, num_experts, device=logits.device).to(torch.float64) + * 1e-12 + ) + _, top_indices = group_limited_topk( + scores_for_routing, topk, num_tokens, num_experts, num_groups, group_topk + ) + scores = torch.gather(scores, dim=1, index=top_indices) + else: + scores = ( + scores + - torch.arange(0, num_experts, device=logits.device).to(torch.float64) + * 1e-12 + ) + scores, top_indices = group_limited_topk( + scores, topk, num_tokens, num_experts, num_groups, group_topk + ) + probs = scores / (scores.sum(dim=-1, keepdim=True) + eps) if topk > 1 else scores + + if scaling_factor: + probs = probs * scaling_factor + + # TODO Try using element-wise operations instead of scatter? + topk_masked_gates = torch.zeros_like(logits, dtype=torch.float64).scatter( + 1, top_indices, probs + ) + topk_map = ( + torch.zeros_like(logits, dtype=torch.float64) + .int() + .scatter(1, top_indices, 1) + .bool() + ) + + tokens_per_expert = topk_map.sum(dim=0) + return topk_masked_gates.float(), topk_map, tokens_per_expert + + +def test_topk(M=4096, B=1, N=256, k=8, equal=False, bench=False): + dtype = torch.float32 + device = "cuda:0" + + if B == 0: + x = torch.randn(M, N, dtype=dtype, device=device) + else: + x = torch.randn(M, B, N, dtype=dtype, device=device) + if equal: + x[..., 0] = x[..., -1] + + x = x.requires_grad_() + + if equal: + xd = x.to(torch.float64) * ( + 1 - torch.arange(0, N, device=device).to(torch.float64) * 1e-12 + ) + value_ref, index_ref = torch.topk(xd, k) + value_ref = value_ref.float() + else: + value_ref, index_ref = torch.topk(x, k) + + loss_ref = (value_ref * index_ref.float()).sum() + loss_ref.backward() + grad_ref = x.grad + x.grad = None + + value, index = triton_topk_forward(x, k) + grad = triton_topk_backward(index_ref.float(), index, N) + output_check(value_ref, value, "value") + output_check(index_ref, index.to(torch.int64), "index") + output_check(grad_ref, grad, "grad") + + value, index = fused_topk(x, k) + value.backward(index_ref.float()) + grad = x.grad + output_check(value_ref, value, "value") + output_check(index_ref, index.to(torch.int64), "index") + output_check(grad_ref, grad, "grad") + + if bench: + ref_time = benchmark_func(torch.topk, x, k, ref_bytes=M * N * 4) + benchmark_func( + triton_topk_forward, x, k, ref_bytes=M * N * 4, ref_time=ref_time + ) + benchmark_func( + triton_topk_backward, index_ref.float(), index, N, ref_bytes=M * N * 4 + ) + + +def test_group_topk_score( + M=4096, + N=256, + k=8, + num_groups=32, + group_topk=4, + scaling_factor=1.0, + equal=False, + bias=True, + bench=False, +): + dtype = torch.float32 + device = "cuda:0" + + x = torch.randn(M, N, dtype=dtype, device=device) + + dy = torch.randn(M, N, dtype=dtype, device=device) + if bias: + expert_bias = torch.randn(N, dtype=dtype, device=device) * 10.0 + else: + expert_bias = None + if equal: + x[..., 0] = x[..., 1] + x = x.requires_grad_() + + prob_ref, map_ref, count_ref = torch_group_topk_score( + x, + expert_bias=expert_bias, + num_experts=N, + topk=k, + num_groups=num_groups, + group_topk=group_topk, + scaling_factor=scaling_factor, + ) + loss_ref = (prob_ref * map_ref.float() * dy).sum() + loss_ref.backward() + grad_ref = x.grad + x.grad = None + + prob, maps, count = triton_group_topk_score_forward( + x, + k, + expert_bias=expert_bias, + num_groups=num_groups, + group_topk=group_topk, + scaling_factor=scaling_factor, + ) + grad = triton_group_topk_score_backward( + map_ref.float() * dy, x, maps, scaling_factor=scaling_factor + ) + output_check(prob_ref, prob, "prob", atol=-1) # may have mismatched results + err = output_check(map_ref, maps, "maps") + output_check(count_ref, count, "count") + output_check(grad_ref, grad, "grad") + + prob, maps, count = group_topk_score( + x, + k, + expert_bias=expert_bias, + num_groups=num_groups, + group_topk=group_topk, + scaling_factor=scaling_factor, + ) + prob.backward(gradient=map_ref.float() * dy) + grad = x.grad + output_check(prob_ref, prob, "prob") + output_check(map_ref, maps, "maps") + output_check(count_ref, count, "count") + output_check(grad_ref, grad, "grad") + + if bench: + ref_time = benchmark_func( + torch_group_topk_score, + x, + expert_bias, + num_experts=N, + topk=k, + num_groups=num_groups, + group_topk=group_topk, + scaling_factor=scaling_factor, + ) + benchmark_func( + triton_group_topk_score_forward, + x, + k, + expert_bias=expert_bias, + num_groups=num_groups, + group_topk=group_topk, + scaling_factor=scaling_factor, + ref_time=ref_time, + ) + benchmark_func(triton_group_topk_score_backward, map_ref.float(), x, maps) + + +if __name__ == "__main__": + test_topk(M=8192, B=0, N=256, k=8, equal=False, bench=False) + test_topk(M=4096, B=2, N=256, k=8, equal=False, bench=False) + test_topk(M=4096, B=2, N=256, k=8, equal=True, bench=False) + test_group_topk_score( + M=8192, + N=256, + k=8, + num_groups=32, + group_topk=4, + scaling_factor=1.0, + equal=False, + bias=True, + bench=False, + ) + test_group_topk_score( + M=8192, + N=256, + k=8, + num_groups=32, + group_topk=4, + scaling_factor=1.0, + equal=False, + bias=False, + bench=False, + ) + test_group_topk_score( + M=8192, + N=256, + k=8, + num_groups=32, + group_topk=4, + scaling_factor=1.0, + equal=True, + bias=True, + bench=False, + ) diff --git a/tests/test_transpose.py b/tests/test_transpose.py index b542ad0..317f617 100644 --- a/tests/test_transpose.py +++ b/tests/test_transpose.py @@ -8,16 +8,16 @@ import torch from linghe.tools.benchmark import benchmark_func -from linghe.tools.util import output_check -from linghe.utils.transpose import (round_up, - triton_batch_transpose, - triton_batch_transpose_and_pad, - triton_transpose, - triton_transpose_and_pad) +from linghe.tools.check import output_check +from linghe.utils.transpose import ( + round_up, + triton_batch_transpose, + triton_batch_transpose_and_pad, + triton_transpose, + triton_transpose_and_pad, +) -# from torch.profiler import profile, record_function, ProfilerActivity - def torch_nd_transpose(x, dim0, dim1): return x.transpose(dim0, dim1).contiguous() @@ -34,7 +34,7 @@ def triton_split_transpose(xs, count_list): s = 0 outputs = [] for i, c in enumerate(count_list): - x = xs[s:s + c] + x = xs[s : s + c] output = triton_transpose_and_pad(x, pad=True) outputs.append(output) s += c @@ -47,7 +47,7 @@ def test_transpose(M=4096, N=4096, bench=False): # M, N, K = 4096, 4096, 4096 dtype = torch.bfloat16 - device = 'cuda:0' + device = "cuda:0" n_repeat = 100 @@ -57,10 +57,9 @@ def test_transpose(M=4096, N=4096, bench=False): ref_output = x_q.t().contiguous() opt_output = triton_transpose(x_q) - output_check(ref_output.float(), opt_output.float(), 'transpose') + output_check(ref_output.float(), opt_output.float(), "transpose") if bench: - benchmark_func(triton_transpose, x_q, n_repeat=n_repeat, - ref_bytes=M * N * 2) + benchmark_func(triton_transpose, x_q, n_repeat=n_repeat, ref_bytes=M * N * 2) def test_nd_transpose(B=4096, M=4, N=4096, bench=False): @@ -69,32 +68,55 @@ def test_nd_transpose(B=4096, M=4, N=4096, bench=False): # M, N, K = 4096, 4096, 4096 dtype = torch.bfloat16 - device = 'cuda:0' + device = "cuda:0" n_repeat = 100 x = torch.randn(B, M, N, dtype=dtype, device=device) t_ref = torch_nd_transpose(x, 0, 1) - t = triton_transpose(x, dim0=0, dim1=1) - output_check(t_ref, t, '3d_transpose') + t = triton_transpose(x, inner=True) + output_check(t_ref, t, "3d_transpose") - x = torch.randn(B, M, N, dtype=dtype, device=device)[:, :M // 2] + x = torch.randn(B, M, N, dtype=dtype, device=device)[:, : M // 2] t_ref = torch_nd_transpose(x, 0, 1) - t = triton_transpose(x, dim0=0, dim1=1) - output_check(t_ref, t, '3d_transpose_stride') + t = triton_transpose(x, inner=True) + output_check(t_ref, t, "3d_transpose_stride") - x = torch.randn(B, M, N // 128, 128, dtype=dtype, device=device)[:, :M // 2] + x = torch.randn(B, M, N // 128, 128, dtype=dtype, device=device)[:, : M // 2] t_ref = torch_nd_transpose(x, 0, 1) - t = triton_transpose(x, dim0=0, dim1=1) - output_check(t_ref, t, '4d_transpose') + t = triton_transpose(x, inner=True) + output_check(t_ref, t, "4d_transpose") + + x = torch.randn(B, M, N, dtype=dtype, device=device) + t_ref = torch_nd_transpose(x, 1, 2) + t = triton_transpose(x, inner=False) + output_check(t_ref, t, "3d_outer_transpose") if bench: x = torch.randn(B, M, N, dtype=dtype, device=device) - ref_time = benchmark_func(torch_nd_transpose, x, 0, 1, - n_repeat=n_repeat, - ref_bytes=B * M * N * 4) - benchmark_func(triton_transpose, x, dim0=0, dim1=1, n_repeat=n_repeat, - ref_bytes=B * M * N * 4, ref_time=ref_time) + ref_time = benchmark_func( + torch_nd_transpose, x, 0, 1, n_repeat=n_repeat, ref_bytes=B * M * N * 4 + ) + benchmark_func( + triton_transpose, + x, + inner=True, + n_repeat=n_repeat, + ref_bytes=B * M * N * 4, + ref_time=ref_time, + ) + x = torch.randn(M, B, N, dtype=dtype, device=device) + ref_time = benchmark_func( + torch_nd_transpose, x, 1, 2, n_repeat=n_repeat, ref_bytes=B * M * N * 4 + ) + benchmark_func( + triton_transpose, + x, + inner=False, + n_repeat=n_repeat, + ref_bytes=B * M * N * 4, + ref_time=ref_time, + ) def test_transpose_and_pad(M=4095, N=4096, bench=False): @@ -103,80 +125,89 @@ def test_transpose_and_pad(M=4095, N=4096, bench=False): # M, N, K = 4096, 4096, 4096 dtype = torch.bfloat16 - device = 'cuda:0' - - n_repeat = 100 + device = "cuda:0" - if True: - x = torch.randn(M, N, dtype=dtype, device=device) - else: - x = torch.load('/ossfs/workspace/tmp/vis/backward.bin')['w'] - M, N = x.shape + x = torch.randn(M, N, dtype=dtype, device=device) P = round_up(M, b=32) tail = P - M x_q = x.to(torch.float8_e4m3fn) ref_output = x_q.t().contiguous() - opt_output = torch.randn((N, P), dtype=dtype, device=device).to( - torch.float8_e4m3fn) + opt_output = torch.randn((N, P), dtype=dtype, device=device).to(torch.float8_e4m3fn) opt_output = triton_transpose_and_pad(x_q, out=opt_output, pad=True) - output_check(ref_output.float(), opt_output[:, :M].float(), - 'transpose_and_pad') + output_check(ref_output.float(), opt_output[:, :M].float(), "transpose_and_pad") if tail > 0: assert opt_output[:, -tail:].float().abs().sum().item() == 0 if bench: - benchmark_func(triton_transpose_and_pad, x_q, n_repeat=n_repeat, - ref_bytes=M * N * 2) + benchmark_func(triton_transpose_and_pad, x_q, ref_bytes=M * N * 2) def test_batch_transpose(M=4096, N=4096, k=32, bench=False): dtype = torch.bfloat16 - device = 'cuda:0' + device = "cuda:0" xs = [ torch.randn((M, N), dtype=dtype, device=device).to(torch.float8_e4m3fn) - for _ in range(k)] + for _ in range(k) + ] xts = triton_batch_transpose(xs) + xts = torch.cat([x.view(-1) for x in xts]) x_t_ref = triton_sequence_transpose(xs) - for i in range(len(xs)): - output_check(x_t_ref[i].float(), xts[i].float(), f'batch_transpose_{i}') + x_t_ref = torch.cat([x.view(-1) for x in x_t_ref]) + + output_check(x_t_ref, xts, f"batch_transpose") if bench: n_repeat = 100 - ref_time = benchmark_func(triton_sequence_transpose, xs, - n_repeat=n_repeat, ref_bytes=M * M * 2 * k) - benchmark_func(triton_batch_transpose, xs, n_repeat=n_repeat, - ref_bytes=M * N * 2 * k, ref_time=ref_time) + ref_time = benchmark_func( + triton_sequence_transpose, xs, n_repeat=n_repeat, ref_bytes=M * M * 2 * k + ) + benchmark_func( + triton_batch_transpose, + xs, + n_repeat=n_repeat, + ref_bytes=M * N * 2 * k, + ref_time=ref_time, + ) def test_batch_transpose_and_pad(M=4096, N=4096, k=32, bench=False): dtype = torch.bfloat16 - device = 'cuda:0' + device = "cuda:0" count_list = [random.randint(1500, 2600) for x in range(k)] xs = torch.randn((sum(count_list), N), dtype=dtype, device=device).to( - torch.float8_e4m3fn) + torch.float8_e4m3fn + ) x_t = triton_batch_transpose_and_pad(xs, count_list, x_t=None, pad=True) + x_t = torch.cat([x.view(-1) for x in x_t]) x_t_ref = triton_split_transpose(xs, count_list) + x_t_ref = torch.cat([x.view(-1) for x in x_t_ref]) - for i in range(len(count_list)): - output_check(x_t_ref[i].float(), x_t[i].float(), - f'batch_transpose_and_pad_{i}') + output_check(x_t_ref, x_t, f"batch_transpose_and_pad") if bench: n_repeat = 100 - ref_time = benchmark_func(triton_split_transpose, xs, count_list, - n_repeat=n_repeat) - benchmark_func(triton_batch_transpose_and_pad, xs, count_list, x_t=None, - pad=True, n_repeat=n_repeat, ref_time=ref_time) - - -if __name__ == '__main__': + ref_time = benchmark_func( + triton_split_transpose, xs, count_list, n_repeat=n_repeat + ) + benchmark_func( + triton_batch_transpose_and_pad, + xs, + count_list, + x_t=None, + pad=True, + n_repeat=n_repeat, + ref_time=ref_time, + ) + + +if __name__ == "__main__": test_transpose(M=4096, N=4096) test_transpose_and_pad(M=4095, N=4096) - test_nd_transpose(B=4096, M=4, N=4096) - test_batch_transpose(M=4096,N=4096,k=32) - test_batch_transpose_and_pad(M=4096,N=4096,k=32) + test_nd_transpose(B=4096, M=4, N=2048, bench=False) + test_batch_transpose(M=4096, N=4096, k=32, bench=False) + test_batch_transpose_and_pad(M=4096, N=4096, k=32) diff --git a/tests/test_unary.py b/tests/test_unary.py index fb00756..3d142fc 100644 --- a/tests/test_unary.py +++ b/tests/test_unary.py @@ -3,39 +3,109 @@ Copyright (c) Ant Financial Service Group and its affiliates. """ +import random + +import pytest import torch -from linghe.utils.unary import triton_calculate_smooth_scale -from linghe.tools.benchmark import benchmark_func -from linghe.tools.util import output_check +from linghe.tools.check import output_check +from linghe.utils.unary import triton_calculate_smooth_scale, triton_clip, triton_batch_clip -def torch_calculate_smooth_scale(x, min_value=1.0, smooth_coef=0.5): +def torch_calculate_smooth_scale(x, min_value=1.0, smooth_coef=0.5, + round_scale=False): one = torch.ones([1], dtype=torch.float32, device=x.device) - input_smooth_scales = torch.pow(torch.maximum(x, min_value*one), smooth_coef) - weight_smooth_scales = 1/input_smooth_scales - weight_smooth_scales = torch.exp2(torch.ceil(torch.log2(weight_smooth_scales))) + input_smooth_scales = torch.pow(torch.maximum(x, min_value * one), + smooth_coef) + weight_smooth_scales = 1 / input_smooth_scales + if round_scale: + weight_smooth_scales = torch.exp2( + torch.ceil(torch.log2(weight_smooth_scales))) return weight_smooth_scales -def test_calculate_smooth_scale(N=4096, bench=False): +def torch_clip(x, clip_value): + x = torch.clamp(x, -clip_value, clip_value) + return x + +def torch_batch_clip(xs, clip_value): + torch._foreach_clamp_min_(xs, -clip_value) + torch._foreach_clamp_max_(xs, clip_value) + return xs + - x = torch.randn(N, dtype=torch.float32, device='cuda:0').abs()**3+0.1 +@pytest.mark.parametrize("N", [ + 4096 * 32, + 4096 * 32 - 1897, +]) +def test_calculate_smooth_scale(N, benchmark): + x = torch.randn(N, dtype=torch.float32, device='cuda:0').abs() ** 3 + 0.1 min_value = 0.0 - smooth_coef = 0.7 - out_ref = torch_calculate_smooth_scale(x, min_value=min_value, smooth_coef=smooth_coef) - out = triton_calculate_smooth_scale(x, min_value=min_value, smooth_coef=smooth_coef) + smooth_coef = 0.5 + out_ref = torch_calculate_smooth_scale(x, min_value=min_value, + smooth_coef=smooth_coef, + round_scale=True) + out = triton_calculate_smooth_scale(x, min_value=min_value, + smooth_coef=smooth_coef, + round_scale=True) output_check(out_ref, out, 'torch_calculate_smooth_scale') n_repeat = 100 - if bench: - ref_time = benchmark_func(torch_calculate_smooth_scale, x, - n_repeat=n_repeat) - benchmark_func(torch_calculate_smooth_scale, x, n_repeat=n_repeat, - ref_time=ref_time, ref_bytes=N * 8) + ref_time = benchmark(torch_calculate_smooth_scale, x, n_repeat=n_repeat) + benchmark(torch_calculate_smooth_scale, x, n_repeat=n_repeat, ref_time=ref_time, ref_bytes=N * 8) + + +@pytest.mark.parametrize("M,N,clip_value,inf", [ + (2048, 8192, 0.1, False), + (10000, 8192, 100.0, True), +]) +def test_clip(M, N, clip_value, inf, benchmark): + x = torch.randn(M, N, dtype=torch.float32, device='cuda:0') + x1 = x.clone().detach() + x2 = x.clone().detach() + + if inf: + x1[0][:100] = float('inf') + x2[0][:100] = float('inf') + + out_ref = torch_clip(x1, clip_value) + out = triton_clip(x2, clip_value) + output_check(out_ref, out, 'clip') + + ref_bytes = M * N * 8 + x3 = x.clone().detach() + n_repeat = 1 # inplace update will speedup our triton op + ref_time = benchmark(torch_clip, x3, clip_value, ref_bytes=ref_bytes, n_repeat=n_repeat, n_warmup=0) + benchmark(triton_clip, x3, clip_value, ref_bytes=ref_bytes, ref_time=ref_time, n_repeat=n_repeat, n_warmup=0) + + +@pytest.mark.parametrize("M,N,k,clip_value,inf", [ + (2048, 8192, 128, 0.1, False), + (2048, 1024, 128, 1.0, False), + (2048, 1024, 128, 100.0, True), +]) +def test_batch_clip(M, N, k, clip_value, inf, benchmark): + shapes1 = [random.randint(1, int(M ** 0.5)) ** 2 for i in range(k)] + shapes2 = [random.randint(1, int(N ** 0.5)) ** 2 for i in range(k)] + xs = [torch.randn(shapes1[i], shapes2[i], dtype=torch.float32, + device='cuda:0') for i in range(k)] + xs1 = [x.clone().detach() for x in xs] + xs2 = [x.clone().detach() for x in xs] + + if inf: + xs1[0][:100] = float('inf') + xs2[0][:100] = float('inf') + + sum_ref = torch_batch_clip(xs1, clip_value) + sums = triton_batch_clip(xs2, clip_value) + output_check(torch.cat([x.view(-1) for x in sum_ref], 0), + torch.cat([x.view(-1) for x in sums], 0), 'batch_clip') -if __name__ == '__main__': - test_calculate_smooth_scale(N=4096*32) - test_calculate_smooth_scale(N=4096*32-1897) + ref_bytes = sum([x.numel() for x in xs]) * 8 + xs3 = [x.clone().detach() for x in xs] + n_repeat = 1 # inplace update will speedup our triton op + ref_time = benchmark(torch_batch_clip, xs3, clip_value, ref_bytes=ref_bytes, n_repeat=n_repeat, n_warmup=0) + xs4 = [x.clone().detach() for x in xs] + benchmark(triton_batch_clip, xs4, clip_value, ref_bytes=ref_bytes, ref_time=ref_time, n_repeat=n_repeat, n_warmup=0)